"""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_BASE = ( "https://huggingface.co/datasets/wikimedia/wikipedia/resolve/main/20231101.{lang}" ) IGNORE_INDEX = -100 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 _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 = WIKI_BASE.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 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