Files
K3/kda/training/data.py
T
dela e7185cbf49 Pull OPUS-100 en-zh for SFT instead of a checked-in jsonl
train_sft --data opus-100 streams Helsinki-NLP/opus-100, writes both
directions, and skips frozen eval sentences. Runtime cache stays under
data/sft/ (gitignored).
2026-08-25 21:39:58 +08:00

467 lines
14 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 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