Route with σ(W_r x), Top-k(s+b), then L1-normalize over the selected set. Add Switch/GShard aux and router z-loss into train_k3 and train_sft. Wiki parquet URLs honor HF_ENDPOINT for mirrored downloads.
359 lines
11 KiB
Python
359 lines
11 KiB
Python
"""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 _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 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
|