Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
"""Training and checkpoint helpers."""
|
||||
|
||||
from .toy import load_ckpt, make_toy_data, save_ckpt, train_one_batch
|
||||
|
||||
__all__ = ["load_ckpt", "make_toy_data", "save_ckpt", "train_one_batch"]
|
||||
@@ -0,0 +1,355 @@
|
||||
"""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
|
||||
@@ -0,0 +1,139 @@
|
||||
"""Greedy translation eval on line-aligned src/ref files.
|
||||
|
||||
python -m kda.training.eval_mt \\
|
||||
--ckpt ckpts/k3_wiki.pt --src /data/eval/zh2en.src.txt \\
|
||||
--ref /data/eval/zh2en.ref.txt --target-lang en
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from kda.training.data import eos_id, load_tokenizer
|
||||
from kda.training.prompts import instruction_prompt
|
||||
from kda.training.success import _chrf, _detect_lang, translation_success
|
||||
from kda.training.toy import load_ckpt
|
||||
|
||||
|
||||
def _read_lines(path: str) -> list[str]:
|
||||
return [ln.strip() for ln in Path(path).read_text(encoding="utf-8").splitlines() if ln.strip()]
|
||||
|
||||
|
||||
def _instruction(src: str, target_lang: str) -> str:
|
||||
return instruction_prompt(src, target_lang)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def decode_one(model, tok, prompt: str, device: str, max_new: int) -> str:
|
||||
ids = tok.encode(prompt)
|
||||
if not ids:
|
||||
return ""
|
||||
inp = torch.tensor([ids], dtype=torch.long, device=device)
|
||||
out = model.generate(inp, max_new, eos_token_id=eos_id(tok))
|
||||
gen = out[0, inp.size(1) :].tolist()
|
||||
return tok.decode(gen).strip()
|
||||
|
||||
|
||||
def evaluate_pairs(
|
||||
model,
|
||||
tok,
|
||||
srcs: list[str],
|
||||
refs: list[str],
|
||||
*,
|
||||
target_lang: str,
|
||||
device: str,
|
||||
max_new: int,
|
||||
limit: int | None,
|
||||
) -> dict:
|
||||
n = len(srcs)
|
||||
if limit is not None:
|
||||
n = min(n, limit)
|
||||
hyps: list[str] = []
|
||||
wins = 0
|
||||
copies = 0
|
||||
lang_ok = 0
|
||||
chrf_sum = 0.0
|
||||
for i in range(n):
|
||||
src, ref = srcs[i], refs[i]
|
||||
hyp = decode_one(model, tok, _instruction(src, target_lang), device, max_new)
|
||||
hyps.append(hyp)
|
||||
ok = translation_success(src, hyp, ref, target_lang=target_lang)
|
||||
wins += int(ok)
|
||||
copies += int(_chrf(hyp, src) >= 80.0 or hyp == src)
|
||||
want = "zh" if target_lang.startswith("zh") else "en"
|
||||
lang_ok += int(_detect_lang(hyp) == want)
|
||||
chrf_sum += _chrf(hyp, ref)
|
||||
corpus = {}
|
||||
try:
|
||||
from sacrebleu.metrics import BLEU, CHRF
|
||||
|
||||
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
|
||||
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
|
||||
except Exception:
|
||||
corpus["chrf"] = chrf_sum / max(n, 1)
|
||||
corpus["bleu"] = None
|
||||
return {
|
||||
"n": n,
|
||||
"success_rate": wins / max(n, 1),
|
||||
"copy_rate": copies / max(n, 1),
|
||||
"lang_ok": lang_ok / max(n, 1),
|
||||
"chrf": corpus["chrf"],
|
||||
"bleu": corpus["bleu"],
|
||||
"hyps": hyps,
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--ckpt", required=True)
|
||||
p.add_argument("--tokenizer", default=None, help="override ckpt tokenizer field")
|
||||
p.add_argument("--src", default=None, help="one source sentence per line")
|
||||
p.add_argument("--ref", default=None, help="one reference sentence per line")
|
||||
p.add_argument("--target-lang", default="en", choices=["en", "zh"])
|
||||
p.add_argument("--max-new", type=int, default=64)
|
||||
p.add_argument("--limit", type=int, default=None)
|
||||
p.add_argument("--prefix", default=None, help="single-prompt smoke decode")
|
||||
p.add_argument("--device", default="auto")
|
||||
args = p.parse_args()
|
||||
|
||||
device = args.device
|
||||
if device == "auto":
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
model, _config = load_ckpt(args.ckpt)
|
||||
model.to(device).eval()
|
||||
payload = torch.load(args.ckpt, map_location="cpu", weights_only=False)
|
||||
tok_src = args.tokenizer or payload.get("tokenizer")
|
||||
if not tok_src:
|
||||
raise SystemExit("need --tokenizer or a 'tokenizer' field in the checkpoint")
|
||||
tok = load_tokenizer(tok_src)
|
||||
|
||||
if args.prefix:
|
||||
print(decode_one(model, tok, args.prefix, device, args.max_new))
|
||||
|
||||
if args.src and args.ref:
|
||||
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
|
||||
if len(srcs) != len(refs):
|
||||
raise SystemExit(f"src/ref length mismatch: {len(srcs)} vs {len(refs)}")
|
||||
out = evaluate_pairs(
|
||||
model,
|
||||
tok,
|
||||
srcs,
|
||||
refs,
|
||||
target_lang=args.target_lang,
|
||||
device=device,
|
||||
max_new=args.max_new,
|
||||
limit=args.limit,
|
||||
)
|
||||
printable = {k: v for k, v in out.items() if k != "hyps"}
|
||||
print(json.dumps(printable, ensure_ascii=False, indent=2))
|
||||
elif not args.prefix:
|
||||
raise SystemExit("pass --prefix and/or --src + --ref")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,7 @@
|
||||
"""Instruction strings shared by SFT and eval. Do not drift."""
|
||||
|
||||
|
||||
def instruction_prompt(src: str, target_lang: str) -> str:
|
||||
if target_lang.startswith("zh"):
|
||||
return f"Translate to Chinese:\n{src}"
|
||||
return f"Translate to English:\n{src}"
|
||||
@@ -0,0 +1,47 @@
|
||||
"""LR scale and token-horizon helpers for train_k3 / train_sft."""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
|
||||
def lr_scale(
|
||||
opt_step: int,
|
||||
warmup: int,
|
||||
total_opt: int,
|
||||
min_ratio: float = 0.1,
|
||||
) -> float:
|
||||
"""Linear warmup (optimizer steps) then cosine down to ``min_ratio``.
|
||||
|
||||
``opt_step`` is 0-indexed at the optimizer update that is about to run.
|
||||
"""
|
||||
if warmup > 0 and opt_step < warmup:
|
||||
return (opt_step + 1) / warmup
|
||||
denom = max(total_opt - warmup - 1, 1)
|
||||
progress = min(max(opt_step - warmup, 0) / denom, 1.0)
|
||||
cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
|
||||
return min_ratio + (1.0 - min_ratio) * cosine
|
||||
|
||||
|
||||
def tokens_per_micro(batch: int, seq_len: int) -> int:
|
||||
return batch * seq_len
|
||||
|
||||
|
||||
def total_opt_steps(
|
||||
*,
|
||||
max_tokens: int | None,
|
||||
max_micro: int | None,
|
||||
batch: int,
|
||||
seq_len: int,
|
||||
grad_acc: int,
|
||||
) -> int:
|
||||
"""Optimizer-step horizon used by cosine. At least 1."""
|
||||
acc = max(grad_acc, 1)
|
||||
candidates: list[int] = []
|
||||
if max_tokens is not None and max_tokens > 0:
|
||||
tpm = max(tokens_per_micro(batch, seq_len), 1)
|
||||
candidates.append(math.ceil(max_tokens / (tpm * acc)))
|
||||
if max_micro is not None and max_micro > 0:
|
||||
candidates.append(math.ceil(max_micro / acc))
|
||||
if not candidates:
|
||||
return 1
|
||||
return max(min(candidates), 1)
|
||||
@@ -0,0 +1,86 @@
|
||||
"""Frozen translation success() — SFT eval and RL reward must call this."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
CHRF_MIN = 40.0
|
||||
COPY_CHRF_MAX = 80.0
|
||||
_CJK = re.compile(r"[\u4e00-\u9fff]")
|
||||
|
||||
|
||||
def _detect_lang(text: str) -> str | None:
|
||||
sample = text.strip()
|
||||
if not sample:
|
||||
return None
|
||||
try:
|
||||
from langdetect import detect
|
||||
|
||||
tag = detect(sample)
|
||||
except Exception:
|
||||
if _CJK.search(sample):
|
||||
return "zh"
|
||||
if any(c.isascii() and c.isalpha() for c in sample):
|
||||
return "en"
|
||||
return None
|
||||
if tag.startswith("zh"):
|
||||
return "zh"
|
||||
return tag[:2]
|
||||
|
||||
|
||||
def _chrf(hyp: str, ref: str) -> float:
|
||||
"""chrF++ in 0–100. Falls back to char unigram F if sacrebleu is missing."""
|
||||
if not hyp or not ref:
|
||||
return 0.0
|
||||
try:
|
||||
from sacrebleu.metrics import CHRF
|
||||
|
||||
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
|
||||
except Exception:
|
||||
hyp_c, ref_c = list(hyp), list(ref)
|
||||
if not hyp_c:
|
||||
return 0.0
|
||||
ref_set = set(ref_c)
|
||||
overlap = sum(1 for c in hyp_c if c in ref_set)
|
||||
prec = overlap / len(hyp_c)
|
||||
rec = overlap / max(len(ref_c), 1)
|
||||
if prec + rec == 0:
|
||||
return 0.0
|
||||
return 100.0 * 2 * prec * rec / (prec + rec)
|
||||
|
||||
|
||||
def translation_success(
|
||||
src: str,
|
||||
hyp: str,
|
||||
ref: str | None = None,
|
||||
*,
|
||||
target_lang: str,
|
||||
chrf_min: float = CHRF_MIN,
|
||||
copy_chrf_max: float = COPY_CHRF_MAX,
|
||||
) -> bool:
|
||||
"""Binary task success for zh↔en instruction translation.
|
||||
|
||||
1. non-empty hyp, no instruction leak prefix
|
||||
2. langid(hyp) matches target_lang (zh / en)
|
||||
3. hyp is not a copy of src
|
||||
4. if ref is given, chrF(hyp, ref) >= chrf_min
|
||||
"""
|
||||
hyp = hyp.strip()
|
||||
src = src.strip()
|
||||
if not hyp:
|
||||
return False
|
||||
leak = ("翻译如下", "translate to", "translation:", "译文:")
|
||||
head = hyp[:40].lower()
|
||||
if any(p in head or p in hyp[:20] for p in leak):
|
||||
return False
|
||||
want = "zh" if target_lang.startswith("zh") else "en"
|
||||
got = _detect_lang(hyp)
|
||||
if got != want:
|
||||
return False
|
||||
if src and _chrf(hyp, src) >= copy_chrf_max:
|
||||
return False
|
||||
if hyp == src:
|
||||
return False
|
||||
if ref is not None and _chrf(hyp, ref.strip()) < chrf_min:
|
||||
return False
|
||||
return True
|
||||
@@ -0,0 +1,110 @@
|
||||
"""L7: toy training loop — overfit 起步.
|
||||
|
||||
target:
|
||||
端到端验证模型 + 数据流 + optimizer + ckpt + generate.
|
||||
|
||||
toy data:
|
||||
建一份 256-token vocab 的小数据集: e.g. 1000 个长度 32 随机 token 序列
|
||||
起步只取 batch=4, 看能否在 ~320 steps 内把 loss 压到 < 0.1 (overfit 单 batch).
|
||||
|
||||
step:
|
||||
optimizer = AdamW(lr=1e-3, wd=0.01)
|
||||
loss.backward(); optimizer.step(); optimizer.zero_grad()
|
||||
every N steps: 打印 loss
|
||||
end: 保存 ckpt to ckpts/kda_toy.pt
|
||||
|
||||
ckpt:
|
||||
save:
|
||||
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
|
||||
load:
|
||||
torch.load -> model.load_state_dict
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import asdict, fields
|
||||
|
||||
import torch
|
||||
|
||||
from ..models.causal_lm import CausalLM
|
||||
from ..models.config import KDAConfig
|
||||
from ..models.k3_config import K3Config
|
||||
|
||||
|
||||
def make_toy_data(batch: int = 4, seq_len: int = 32, vocab: int = 256, seed: int = 42):
|
||||
"""单 batch overfit 数据: 同一组序列循环."""
|
||||
torch.manual_seed(seed)
|
||||
seq = torch.randint(0, vocab, (batch, seq_len), dtype=torch.long)
|
||||
return seq # 用作 input_ids 和 labels (shift one inside forward)
|
||||
|
||||
|
||||
def train_one_batch(model, optimizer, input_ids, labels):
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss = model(input_ids, labels=labels)
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
return loss.detach()
|
||||
|
||||
|
||||
def save_ckpt(model, config, path: str):
|
||||
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
|
||||
torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path)
|
||||
|
||||
|
||||
#: The feed-forward submodule was named after its contents (``mlp`` in the
|
||||
#: dense config, ``moe`` in K3) before both were unified under ``ffn``.
|
||||
#: Checkpoints saved before that rename still carry the old prefixes.
|
||||
_LEGACY_PREFIXES = {
|
||||
".mlp.": ".ffn.",
|
||||
".mlp_norm.": ".ffn_norm.",
|
||||
".moe.": ".ffn.",
|
||||
".moe_norm.": ".ffn_norm.",
|
||||
}
|
||||
|
||||
|
||||
def _rename_legacy_keys(state: dict) -> dict:
|
||||
def fix(key: str) -> str:
|
||||
for old, new in _LEGACY_PREFIXES.items():
|
||||
if old in key:
|
||||
return key.replace(old, new)
|
||||
return key
|
||||
|
||||
return {fix(k): v for k, v in state.items()}
|
||||
|
||||
|
||||
def _config_from(payload_config: dict) -> K3Config | KDAConfig:
|
||||
"""Pick the config class the checkpoint was written with.
|
||||
|
||||
``moe_latent_size`` is a K3-only field, so its presence identifies the
|
||||
hybrid K3 architecture; anything else is the dense KDA config.
|
||||
"""
|
||||
cls = K3Config if "moe_latent_size" in payload_config else KDAConfig
|
||||
known = {item.name for item in fields(cls)}
|
||||
return cls(**{k: v for k, v in payload_config.items() if k in known})
|
||||
|
||||
|
||||
def load_ckpt(path: str, model: CausalLM | None = None) -> tuple[CausalLM, K3Config | KDAConfig]:
|
||||
payload = torch.load(path, map_location="cpu", weights_only=False)
|
||||
config = _config_from(payload["config"])
|
||||
if model is None:
|
||||
model = CausalLM(config)
|
||||
model.load_state_dict(_rename_legacy_keys(payload["model_state"]))
|
||||
return model, config
|
||||
|
||||
|
||||
def main():
|
||||
"""主入口: overfit 起步. 320 steps 期望 loss < 0.1."""
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
config = KDAConfig()
|
||||
model = CausalLM(config).to(device)
|
||||
tokens = make_toy_data(seq_len=32, vocab=config.vocab_size).to(device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
|
||||
for step in range(320):
|
||||
loss = train_one_batch(model, optimizer, tokens, tokens)
|
||||
if step % 64 == 0 or step == 319:
|
||||
print(f"step {step:3d} loss {loss.item():.4f}")
|
||||
save_ckpt(model, config, "ckpts/kda_toy.pt")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Train a SentencePiece tokenizer on a Chinese Wikipedia subset.
|
||||
|
||||
用法:
|
||||
uv run python kda/training/train_tokenizer.py \
|
||||
--out data/spm_4k --vocab-size 4096 --limit 20000
|
||||
|
||||
产出:
|
||||
data/spm_4k.model / data/spm_4k.vocab (BPE/unigram, 中文小语料)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
||||
import sentencepiece as spm
|
||||
|
||||
from .data import fetch_wiki_texts
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--out", default="data/spm_4k", help="输出前缀 (model/vocab 文件)")
|
||||
p.add_argument("--vocab-size", type=int, default=8192)
|
||||
p.add_argument("--limit", type=int, default=20000, help="用于训练的 wiki 文章数")
|
||||
p.add_argument("--model-type", default="unigram", choices=["unigram", "bpe"])
|
||||
p.add_argument("--character-coverage", type=float, default=0.9995)
|
||||
args = p.parse_args()
|
||||
|
||||
texts = fetch_wiki_texts(args.limit)
|
||||
corpus = "".join(texts)
|
||||
tmp = args.out + ".corpus.txt"
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
f.write(corpus)
|
||||
print(f"corpus: {len(corpus):,} chars from {len(texts)} articles")
|
||||
|
||||
spm.SentencePieceTrainer.train(
|
||||
input=tmp,
|
||||
model_prefix=args.out,
|
||||
vocab_size=args.vocab_size,
|
||||
model_type=args.model_type,
|
||||
character_coverage=args.character_coverage,
|
||||
unk_id=0,
|
||||
pad_id=1,
|
||||
bos_id=-1,
|
||||
eos_id=-1,
|
||||
num_threads=4,
|
||||
)
|
||||
print(f"tokenizer saved: {args.out}.model / {args.out}.vocab")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user