"""Pretrain / SFT sample construction. Pretrain: Wikipedia parquet → tokenize → pack (B, T). Languages mix 1:1 by token via seq_len-sized blocks so each training chunk is monolingual. SFT: instruction-parallel rows → prompt-masked labels. Template lives in ``prompts.instruction_prompt`` (same string as eval_mt). """ from __future__ import annotations import json import os from dataclasses import dataclass from pathlib import Path from typing import Iterable, Protocol import torch from .prompts import instruction_prompt WIKI_SHARD_TOTAL = {"zh": 6, "en": 41} WIKI_PATH = "datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}" IGNORE_INDEX = -100 def _hf_endpoint() -> str: """Hub origin. OpenBayes/CN: export HF_ENDPOINT=https://hf-mirror.com""" return os.environ.get("HF_ENDPOINT", "https://huggingface.co").rstrip("/") class Tokenizer(Protocol): vocab_size: int def encode(self, text: str) -> list[int]: ... def decode(self, ids: list[int]) -> str: ... @dataclass class SentencePieceTokenizer: _sp: object @property def vocab_size(self) -> int: return int(self._sp.vocab_size()) def encode(self, text: str) -> list[int]: return list(self._sp.encode(text, out_type=int)) def decode(self, ids: list[int]) -> str: return str(self._sp.decode(ids)) @dataclass class HuggingFaceTokenizer: _tok: object @property def vocab_size(self) -> int: return int(len(self._tok)) def encode(self, text: str) -> list[int]: return list(self._tok.encode(text, add_special_tokens=False)) def decode(self, ids: list[int]) -> str: return str(self._tok.decode(ids, skip_special_tokens=True)) def load_tokenizer(source: str) -> Tokenizer: """`.model` 走 SentencePiece, 其它当作 HuggingFace 名或本地目录.""" if source.endswith(".model"): from sentencepiece import SentencePieceProcessor return SentencePieceTokenizer(SentencePieceProcessor(model_file=source)) from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained(source, trust_remote_code=True) return HuggingFaceTokenizer(tok) def pretrain_dir() -> Path: for candidate in ( os.environ.get("KDA_PRETRAIN_DIR"), "/data/pretrain", "data/pretrain", ): if candidate and Path(candidate).is_dir(): return Path(candidate) return Path("data/pretrain") def sft_dir() -> Path: for candidate in ( os.environ.get("KDA_SFT_DIR"), "/data/sft", "data/sft", ): if candidate and Path(candidate).is_dir(): return Path(candidate) return Path("data/sft") def _wiki_files(lang: str, n_shards: int) -> list[str]: if lang not in WIKI_SHARD_TOTAL: raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en") total = WIKI_SHARD_TOTAL[lang] n = min(max(n_shards, 1), total) base = f"{_hf_endpoint()}/{WIKI_PATH.format(lang=lang)}" return [f"{base}/train-{i:05d}-of-{total:05d}.parquet" for i in range(n)] def _cache_path(cache_dir: Path, lang: str, n_shards: int, limit: int) -> Path: return cache_dir / f"wiki-{lang}-n{n_shards}-limit{limit}.jsonl" def fetch_wiki_texts( limit: int, lang: str = "zh", n_shards: int = 2, cache_dir: str | Path | None = None, ) -> list[str]: """Load up to ``limit`` article bodies, caching jsonl under pretrain_dir.""" cache = Path(cache_dir) if cache_dir is not None else pretrain_dir() cache.mkdir(parents=True, exist_ok=True) path = _cache_path(cache, lang, n_shards, limit) if path.exists(): texts: list[str] = [] with path.open(encoding="utf-8") as fh: for line in fh: line = line.strip() if not line: continue texts.append(json.loads(line)["text"]) if len(texts) >= limit: break if texts: return texts from datasets import load_dataset files = _wiki_files(lang, n_shards) ds = load_dataset("parquet", data_files=files, split="train", streaming=True) texts = [] for i, row in enumerate(ds): if i >= limit: break texts.append(row["text"]) tmp = path.with_suffix(path.suffix + ".tmp") with tmp.open("w", encoding="utf-8") as fh: for text in texts: fh.write(json.dumps({"text": text}, ensure_ascii=False) + "\n") tmp.replace(path) return texts def tokenize_corpus(texts: list[str], tok: Tokenizer) -> list[int]: ids: list[int] = [] for text in texts: ids.extend(tok.encode(text)) return ids def interleave_balanced(ids_a: list[int], ids_b: list[int], block: int) -> list[int]: """1:1 by token: seq_len-sized monolingual blocks, drop the longer tail.""" if block < 1: raise ValueError(f"block must be >= 1, got {block}") n = min(len(ids_a), len(ids_b)) n = (n // block) * block out: list[int] = [] a, b = ids_a, ids_b for i in range(0, n, block): out.extend(a[i : i + block]) out.extend(b[i : i + block]) return out def chunk_ids(ids: list[int], batch: int, seq_len: int) -> torch.Tensor: """切成 (num_chunks, B, T); 末尾不足部分丢弃.""" n = (len(ids) // (batch * seq_len)) * (batch * seq_len) t = torch.tensor(ids[:n], dtype=torch.long) if n == 0: return t.view(0, batch, seq_len) return t.view(batch, -1, seq_len).transpose(0, 1) def split_heldout( chunks: torch.Tensor, frac: float = 0.01, min_heldout: int = 1, ) -> tuple[torch.Tensor, torch.Tensor]: """Last ``frac`` of packed chunks for CE only. Empty held-out if too few.""" n = int(chunks.size(0)) if n <= 1 or frac <= 0: return chunks, chunks[:0] h = max(min_heldout, int(n * frac)) h = min(h, n - 1) return chunks[:-h], chunks[-h:] def iter_chunks(chunks: torch.Tensor): """逐块产出 (input_ids, labels), labels 右移 (模型内 CE shift).""" for chunk in chunks: yield chunk, chunk.clone() def iter_indexed(chunks: torch.Tensor, start: int = 0): """Infinite cycle with a global index (for --resume).""" n = int(chunks.size(0)) if n == 0: raise ValueError("no training chunks") i = start while True: x = chunks[i % n] yield i, x, x.clone() i += 1 def load_pretrain_chunks( tok: Tokenizer, *, langs: Iterable[str], limit: int, batch: int, seq_len: int, heldout_frac: float = 0.01, n_shards: int = 2, cache_dir: str | Path | None = None, ) -> tuple[torch.Tensor, torch.Tensor, int]: """Fetch / cache / tokenize / pack. Returns train chunks, held-out, token count.""" lang_list = [lang.strip() for lang in langs if lang.strip()] if not lang_list: raise ValueError("langs must contain at least one of zh, en") streams: list[list[int]] = [] for lang in lang_list: print(f"loading {limit} wiki articles ({lang}) ...") texts = fetch_wiki_texts(limit, lang=lang, n_shards=n_shards, cache_dir=cache_dir) streams.append(tokenize_corpus(texts, tok)) print(f" {lang}: {len(streams[-1]):,} tokens from {len(texts)} articles") if len(streams) == 1: ids = streams[0] else: ids = streams[0] for extra in streams[1:]: ids = interleave_balanced(ids, extra, seq_len) chunks = chunk_ids(ids, batch, seq_len) train, held = split_heldout(chunks, heldout_frac) return train, held, len(ids) def pad_id(tok: Tokenizer) -> int: inner = getattr(tok, "_tok", None) if inner is not None: pid = getattr(inner, "pad_token_id", None) if pid is not None: return int(pid) eid = getattr(inner, "eos_token_id", None) if eid is not None: return int(eid) return 0 def eos_id(tok: Tokenizer) -> int | None: inner = getattr(tok, "_tok", None) if inner is not None: eid = getattr(inner, "eos_token_id", None) if eid is not None: return int(eid) convert = getattr(inner, "convert_tokens_to_ids", None) if convert is not None: tid = convert("<|im_end|>") if isinstance(tid, int) and tid >= 0: return tid return None def encode_sft_row( tok: Tokenizer, src: str, tgt: str, target_lang: str, max_len: int, eos: int | None = None, ) -> tuple[list[int], list[int]]: prompt_ids = tok.encode(instruction_prompt(src, target_lang)) tgt_ids = tok.encode(tgt) if eos is not None: tgt_ids = tgt_ids + [eos] ids = prompt_ids + tgt_ids labels = [IGNORE_INDEX] * len(prompt_ids) + list(tgt_ids) if len(ids) > max_len: overflow = len(ids) - max_len cut = min(overflow, max(len(prompt_ids) - 1, 0)) ids = ids[cut:] labels = labels[cut:] if len(ids) > max_len: ids = ids[:max_len] labels = labels[:max_len] return ids, labels def _eval_blocklist(eval_dir: str | Path | None = None) -> set[str]: """Frozen eval sentences must not appear in SFT bitext.""" blocked: set[str] = set() folders = [] if eval_dir is not None: folders.append(Path(eval_dir)) folders.extend( [ Path(os.environ["KDA_EVAL_DIR"]) if os.environ.get("KDA_EVAL_DIR") else None, Path("/data/eval"), Path("data/eval"), ] ) for folder in folders: if folder is None or not folder.is_dir(): continue for path in folder.glob("*.txt"): for line in path.read_text(encoding="utf-8").splitlines(): text = line.strip() if text: blocked.add(text) return blocked def fetch_opus100_enzh( limit: int, *, both_dirs: bool = True, cache_dir: str | Path | None = None, eval_dir: str | Path | None = None, ) -> list[dict]: """Stream Helsinki-NLP/opus-100 ``en-zh`` train. ``limit`` is source pairs.""" if limit < 1: raise ValueError(f"limit must be >= 1, got {limit}") cache = Path(cache_dir) if cache_dir is not None else sft_dir() cache.mkdir(parents=True, exist_ok=True) tag = "both" if both_dirs else "enzh" path = cache / f"opus100-en-zh-{tag}-limit{limit}.jsonl" if path.exists(): rows = load_sft_rows(path) if rows: return rows from datasets import load_dataset ds = load_dataset("Helsinki-NLP/opus-100", "en-zh", split="train", streaming=True) blocked = _eval_blocklist(eval_dir) rows: list[dict] = [] n_src = 0 for row in ds: trans = row.get("translation") if isinstance(row, dict) else None blob = trans if isinstance(trans, dict) else row en = str(blob.get("en") or "").strip() zh = str(blob.get("zh") or "").strip() if not en or not zh or en == zh: continue if en in blocked or zh in blocked: continue if min(len(en), len(zh)) < 2: continue n_src += 1 rows.append({"src": zh, "tgt": en, "target_lang": "en"}) if both_dirs: rows.append({"src": en, "tgt": zh, "target_lang": "zh"}) if n_src >= limit: break tmp = path.with_suffix(path.suffix + ".tmp") with tmp.open("w", encoding="utf-8") as fh: for row in rows: fh.write(json.dumps(row, ensure_ascii=False) + "\n") tmp.replace(path) return rows def resolve_sft_rows( source: str, *, limit: int = 100_000, both_dirs: bool = True, cache_dir: str | Path | None = None, eval_dir: str | Path | None = None, ) -> list[dict]: """Local jsonl/tsv, or ``opus-100`` / ``opus`` to pull OPUS-100 en-zh from HF.""" path = Path(source) if path.is_file(): return load_sft_rows(path) key = source.strip().lower().replace("_", "-") if key in {"opus", "opus-100", "opus100", "helsinki-nlp/opus-100"}: print(f"fetching OPUS-100 en-zh (limit {limit} pairs, both_dirs={both_dirs})") return fetch_opus100_enzh( limit, both_dirs=both_dirs, cache_dir=cache_dir, eval_dir=eval_dir ) raise FileNotFoundError( f"SFT source {source!r} is not a file; use a jsonl path or 'opus-100'" ) def load_sft_rows(path: str | Path) -> list[dict]: """jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``.""" p = Path(path) rows: list[dict] = [] text = p.read_text(encoding="utf-8") if p.suffix == ".jsonl" or p.suffix == ".json": for line in text.splitlines(): line = line.strip() if not line: continue obj = json.loads(line) rows.append( { "src": obj["src"], "tgt": obj["tgt"], "target_lang": obj.get("target_lang", "en"), } ) return rows for line in text.splitlines(): line = line.strip() if not line or line.startswith("#"): continue parts = line.split("\t") if len(parts) < 2: raise ValueError(f"SFT TSV needs src, tgt [, target_lang]: {line[:80]!r}") lang = parts[2] if len(parts) > 2 else "en" rows.append({"src": parts[0], "tgt": parts[1], "target_lang": lang}) return rows def collate_sft( rows: list[dict], tok: Tokenizer, max_len: int, ) -> tuple[torch.Tensor, torch.Tensor]: pad = pad_id(tok) eos = eos_id(tok) encoded = [ encode_sft_row(tok, r["src"], r["tgt"], r["target_lang"], max_len, eos) for r in rows ] width = min(max(len(ids) for ids, _ in encoded), max_len) width = max(width, 2) bsz = len(encoded) input_ids = torch.full((bsz, width), pad, dtype=torch.long) labels = torch.full((bsz, width), IGNORE_INDEX, dtype=torch.long) for i, (ids, lab) in enumerate(encoded): n = min(len(ids), width) input_ids[i, :n] = torch.tensor(ids[:n], dtype=torch.long) labels[i, :n] = torch.tensor(lab[:n], dtype=torch.long) return input_ids, labels def iter_sft_batches( rows: list[dict], tok: Tokenizer, batch: int, max_len: int, start: int = 0, ): n = len(rows) if n == 0: raise ValueError("no SFT rows") i = start while True: sl = [rows[j % n] for j in range(i, i + batch)] yield i, *collate_sft(sl, tok, max_len) i += batch