Compare commits

...
3 Commits
Author SHA1 Message Date
dela 47c72e5bb8 Warn and reset chunk_index when resume changes batch or seq_len
Trying a larger micro-batch on an existing 0.5b run re-packs wiki
chunks; keep tokens/opt_step and restart the data cursor.
2026-08-25 22:00:28 +08:00
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
dela 53d0f4b17a Cut 1B-run I/O: rarer SwanLab, ckpt, and generate
Every micro-step was hitting SwanLab, and every 100 steps wrote a 5GB
ckpt plus greedy decode. 0.5b now logs every 20, eval/held-out every 500,
saves _last every 1000, samples every 2000.
2026-08-25 20:59:29 +08:00
4 changed files with 267 additions and 37 deletions
+108
View File
@@ -89,6 +89,17 @@ def pretrain_dir() -> Path:
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")
@@ -287,6 +298,103 @@ def encode_sft_row(
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)
+38 -1
View File
@@ -1,4 +1,11 @@
from kda.training.data import IGNORE_INDEX, collate_sft, encode_sft_row, load_sft_rows
from kda.training.data import (
IGNORE_INDEX,
collate_sft,
encode_sft_row,
fetch_opus100_enzh,
load_sft_rows,
resolve_sft_rows,
)
from kda.training.prompts import instruction_prompt
@@ -42,6 +49,36 @@ def test_collate_and_jsonl(tmp_path):
assert (y == IGNORE_INDEX).any()
def test_resolve_sft_rows_reads_local_jsonl(tmp_path):
path = tmp_path / "bitext.jsonl"
path.write_text(
'{"src": "你好", "tgt": "Hello", "target_lang": "en"}\n',
encoding="utf-8",
)
rows = resolve_sft_rows(str(path))
assert rows == [{"src": "你好", "tgt": "Hello", "target_lang": "en"}]
def test_fetch_opus_skips_eval_sentences(tmp_path, monkeypatch):
eval_dir = tmp_path / "eval"
eval_dir.mkdir()
(eval_dir / "zh2en.src.txt").write_text("禁止句\n", encoding="utf-8")
class _DS:
def __iter__(self):
yield {"translation": {"en": "Hello", "zh": "你好"}}
yield {"translation": {"en": "skip", "zh": "禁止句"}}
yield {"translation": {"en": "Thanks", "zh": "谢谢"}}
fake = type("datasets", (), {"load_dataset": staticmethod(lambda *a, **k: _DS())})
monkeypatch.setitem(__import__("sys").modules, "datasets", fake)
rows = fetch_opus100_enzh(10, cache_dir=tmp_path / "sft", eval_dir=eval_dir)
srcs = {r["src"] for r in rows}
assert "禁止句" not in srcs
assert "你好" in srcs and "Hello" in srcs
assert "谢谢" in srcs and "Thanks" in srcs
def test_toy_sft_file_parses():
from pathlib import Path
+95 -29
View File
@@ -35,6 +35,9 @@ _TOY_TRAIN = {
"warmup": 50,
"grad_acc": 1,
"eval_every": 100,
"log_every": 10,
"ckpt_every": 100,
"gen_every": 200,
}
_B500M_TRAIN = {
"tokenizer": "Qwen/Qwen3-8B",
@@ -46,7 +49,10 @@ _B500M_TRAIN = {
"lr": 3e-4,
"warmup": 64,
"grad_acc": 8,
"eval_every": 100,
"eval_every": 500,
"log_every": 20,
"ckpt_every": 1000,
"gen_every": 2000,
}
@@ -93,6 +99,10 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace):
"gradient_checkpointing": cfg.gradient_checkpointing,
"moe_aux_loss_coef": cfg.moe_aux_loss_coef,
"moe_z_loss_coef": cfg.moe_z_loss_coef,
"eval_every": args.eval_every,
"log_every": args.log_every,
"ckpt_every": args.ckpt_every,
"gen_every": args.gen_every,
},
)
except Exception as exc:
@@ -123,6 +133,9 @@ def _payload(
"tokens": tokens,
"chunk_index": chunk_index,
"best_heldout": best_heldout,
"batch": args.batch,
"seq_len": args.seq_len,
"grad_acc": args.grad_acc,
}
@@ -201,6 +214,24 @@ def main() -> None:
)
p.add_argument("--grad-acc", type=int, default=train_defaults["grad_acc"])
p.add_argument("--eval-every", type=int, default=train_defaults["eval_every"])
p.add_argument(
"--log-every",
type=int,
default=train_defaults["log_every"],
help="swanlab scalar period in micro-steps",
)
p.add_argument(
"--ckpt-every",
type=int,
default=train_defaults["ckpt_every"],
help="write _last/_best this many micro-steps (1B default 1000)",
)
p.add_argument(
"--gen-every",
type=int,
default=train_defaults["gen_every"],
help="sample prefixes this often; 0 disables",
)
p.add_argument(
"--langs",
default="zh,en",
@@ -265,6 +296,10 @@ def main() -> None:
raise SystemExit(
"KDA training needs bf16; this GPU does not support it (avoid V100 fp16)"
)
if device == "cuda":
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.set_float32_matmul_precision("high")
print(f"loading tokenizer {args.tokenizer} ...")
tok = load_tokenizer(args.tokenizer)
@@ -344,7 +379,11 @@ def main() -> None:
f"({args.steps * tpm:,} tokens); pass --max-tokens for a real run"
)
else:
print(f"token budget: {args.max_tokens:,} cosine horizon {horizon} opt steps")
print(
f"token budget: {args.max_tokens:,} cosine horizon {horizon} opt steps "
f"log/{args.log_every} eval/{args.eval_every} ckpt/{args.ckpt_every} "
f"gen/{args.gen_every}"
)
train_chunks, held_chunks, n_ids = load_pretrain_chunks(
tok,
@@ -356,10 +395,23 @@ def main() -> None:
)
print(
f"packed tokens {n_ids:,} -> {train_chunks.size(0)} train / "
f"{held_chunks.size(0)} held-out chunks of [{args.batch}, {args.seq_len}]"
f"{held_chunks.size(0)} held-out chunks of [{args.batch}, {args.seq_len}] "
f"{tpm} tok/micro"
)
if train_chunks.size(0) == 0:
raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len")
if args.resume:
old_batch = payload.get("batch")
old_seq = payload.get("seq_len")
if old_batch is not None and (
int(old_batch) != args.batch or int(old_seq or args.seq_len) != args.seq_len
):
print(
f"warning: resume pack [{old_batch}, {old_seq}] -> "
f"[{args.batch}, {args.seq_len}]; reset chunk_index 0 "
f"(tokens/opt_step kept)"
)
chunk_index = 0
optim = torch.optim.AdamW(
model.parameters(),
@@ -429,13 +481,17 @@ def main() -> None:
if elapsed > 0:
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
log_now = (
micro_step % args.eval_every == 0
or micro_step == 1
or (args.max_tokens is not None and tokens >= args.max_tokens)
or (args.max_tokens is None and micro_step >= args.steps)
ended = (
args.max_tokens is not None and tokens >= args.max_tokens
) or (args.max_tokens is None and micro_step >= args.steps)
log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended
eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended
ckpt_now = micro_step % args.ckpt_every == 0 or ended
gen_now = args.gen_every > 0 and (
micro_step % args.gen_every == 0 or micro_step == 1 or ended
)
if log_now:
if eval_now:
held = _heldout_loss(model, held_chunks, device, use_bf16)
if held is not None:
metrics["heldout/loss"] = held
@@ -446,17 +502,36 @@ def main() -> None:
+ (f" held {held:.4f}" if held is not None else "")
+ f" aux {metrics['moe/aux']:.4f} z {metrics['moe/z']:.4f}"
)
if micro_step % (args.eval_every * 2) == 0 or micro_step <= args.eval_every:
for prefix in args.gen_prefix:
sample = gen_sample(prefix)
print(f" gen[{prefix[:16]}]: {sample}")
if tracker is not None:
import swanlab
if held is not None and held < best_heldout:
best_heldout = held
payload = _payload(
cfg,
model,
optim,
args,
micro_step=micro_step,
opt_step=opt_step,
tokens=tokens,
chunk_index=chunk_index + 1,
best_heldout=best_heldout,
)
_save(_sibling(args.out, "_best"), payload)
print(
f" best held-out {best_heldout:.4f} -> {_sibling(args.out, '_best')}"
)
del payload
if gen_now:
for prefix in args.gen_prefix:
sample = gen_sample(prefix)
print(f" gen[{prefix[:16]}]: {sample}")
if tracker is not None:
import swanlab
tracker.log(
{f"gen/{prefix[:24]}": swanlab.Text(sample)},
step=micro_step,
)
tracker.log(
{f"gen/{prefix[:24]}": swanlab.Text(sample)},
step=micro_step,
)
if ckpt_now:
payload = _payload(
cfg,
model,
@@ -469,17 +544,8 @@ def main() -> None:
best_heldout=best_heldout,
)
_save(_sibling(args.out, "_last"), payload)
if held is not None and held < best_heldout:
best_heldout = held
payload["best_heldout"] = best_heldout
_save(_sibling(args.out, "_best"), payload)
print(
f" best held-out {best_heldout:.4f} -> {_sibling(args.out, '_best')}"
)
del payload
if device == "cuda":
torch.cuda.empty_cache()
if tracker is not None:
if tracker is not None and (log_now or eval_now):
tracker.log(metrics, step=micro_step)
payload = _payload(
+26 -7
View File
@@ -1,9 +1,9 @@
"""Instruction SFT for zh↔en translation. Prompt template matches eval_mt.
用法:
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/train.jsonl
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data data/sft/opus.jsonl \\
--seq-len 512 --batch 4 --lr 5e-5 --epochs 2
uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/toy.jsonl
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\
--limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2
"""
from __future__ import annotations
@@ -18,7 +18,7 @@ from kda.layers.latent_moe import moe_router_losses
from kda.training.data import (
IGNORE_INDEX,
iter_sft_batches,
load_sft_rows,
resolve_sft_rows,
load_tokenizer,
)
from kda.training.eval_mt import evaluate_pairs
@@ -72,7 +72,22 @@ def _read_lines(path: str) -> list[str]:
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--ckpt", required=True)
p.add_argument("--data", required=True, help="jsonl {src,tgt,target_lang} or TSV")
p.add_argument(
"--data",
default="opus-100",
help="local jsonl/tsv, or 'opus-100' to stream Helsinki-NLP/opus-100 en-zh",
)
p.add_argument(
"--limit",
type=int,
default=100_000,
help="OPUS source pairs to pull (each becomes zh2en + en2zh unless --one-dir)",
)
p.add_argument(
"--one-dir",
action="store_true",
help="only zh→en rows when pulling OPUS",
)
p.add_argument("--out", default="ckpts/k3_sft.pt")
p.add_argument("--tokenizer", default=None)
p.add_argument("--batch", type=int, default=4)
@@ -103,9 +118,13 @@ def main() -> None:
if not tok_src:
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
tok = load_tokenizer(tok_src)
rows = load_sft_rows(args.data)
rows = resolve_sft_rows(
args.data,
limit=args.limit,
both_dirs=not args.one_dir,
)
if not rows:
raise SystemExit(f"no SFT rows in {args.data}")
raise SystemExit(f"no SFT rows from {args.data}")
print(f"SFT {len(rows)} rows from {args.data}; model {cfg.__class__.__name__}")
steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1)