Compare commits
3
Commits
8442f92c58
...
47c72e5bb8
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
47c72e5bb8 | ||
|
|
e7185cbf49 | ||
|
|
53d0f4b17a |
@@ -89,6 +89,17 @@ def pretrain_dir() -> Path:
|
|||||||
return Path("data/pretrain")
|
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]:
|
def _wiki_files(lang: str, n_shards: int) -> list[str]:
|
||||||
if lang not in WIKI_SHARD_TOTAL:
|
if lang not in WIKI_SHARD_TOTAL:
|
||||||
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
|
raise ValueError(f"unsupported wiki lang {lang!r}; expected zh or en")
|
||||||
@@ -287,6 +298,103 @@ def encode_sft_row(
|
|||||||
return ids, labels
|
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]:
|
def load_sft_rows(path: str | Path) -> list[dict]:
|
||||||
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
|
"""jsonl ``{src,tgt,target_lang}`` or TSV ``src\\ttgt\\ttarget_lang``."""
|
||||||
p = Path(path)
|
p = Path(path)
|
||||||
|
|||||||
@@ -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
|
from kda.training.prompts import instruction_prompt
|
||||||
|
|
||||||
|
|
||||||
@@ -42,6 +49,36 @@ def test_collate_and_jsonl(tmp_path):
|
|||||||
assert (y == IGNORE_INDEX).any()
|
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():
|
def test_toy_sft_file_parses():
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|||||||
+95
-29
@@ -35,6 +35,9 @@ _TOY_TRAIN = {
|
|||||||
"warmup": 50,
|
"warmup": 50,
|
||||||
"grad_acc": 1,
|
"grad_acc": 1,
|
||||||
"eval_every": 100,
|
"eval_every": 100,
|
||||||
|
"log_every": 10,
|
||||||
|
"ckpt_every": 100,
|
||||||
|
"gen_every": 200,
|
||||||
}
|
}
|
||||||
_B500M_TRAIN = {
|
_B500M_TRAIN = {
|
||||||
"tokenizer": "Qwen/Qwen3-8B",
|
"tokenizer": "Qwen/Qwen3-8B",
|
||||||
@@ -46,7 +49,10 @@ _B500M_TRAIN = {
|
|||||||
"lr": 3e-4,
|
"lr": 3e-4,
|
||||||
"warmup": 64,
|
"warmup": 64,
|
||||||
"grad_acc": 8,
|
"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,
|
"gradient_checkpointing": cfg.gradient_checkpointing,
|
||||||
"moe_aux_loss_coef": cfg.moe_aux_loss_coef,
|
"moe_aux_loss_coef": cfg.moe_aux_loss_coef,
|
||||||
"moe_z_loss_coef": cfg.moe_z_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:
|
except Exception as exc:
|
||||||
@@ -123,6 +133,9 @@ def _payload(
|
|||||||
"tokens": tokens,
|
"tokens": tokens,
|
||||||
"chunk_index": chunk_index,
|
"chunk_index": chunk_index,
|
||||||
"best_heldout": best_heldout,
|
"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("--grad-acc", type=int, default=train_defaults["grad_acc"])
|
||||||
p.add_argument("--eval-every", type=int, default=train_defaults["eval_every"])
|
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(
|
p.add_argument(
|
||||||
"--langs",
|
"--langs",
|
||||||
default="zh,en",
|
default="zh,en",
|
||||||
@@ -265,6 +296,10 @@ def main() -> None:
|
|||||||
raise SystemExit(
|
raise SystemExit(
|
||||||
"KDA training needs bf16; this GPU does not support it (avoid V100 fp16)"
|
"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} ...")
|
print(f"loading tokenizer {args.tokenizer} ...")
|
||||||
tok = load_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"
|
f"({args.steps * tpm:,} tokens); pass --max-tokens for a real run"
|
||||||
)
|
)
|
||||||
else:
|
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(
|
train_chunks, held_chunks, n_ids = load_pretrain_chunks(
|
||||||
tok,
|
tok,
|
||||||
@@ -356,10 +395,23 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
print(
|
print(
|
||||||
f"packed tokens {n_ids:,} -> {train_chunks.size(0)} train / "
|
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:
|
if train_chunks.size(0) == 0:
|
||||||
raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len")
|
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(
|
optim = torch.optim.AdamW(
|
||||||
model.parameters(),
|
model.parameters(),
|
||||||
@@ -429,13 +481,17 @@ def main() -> None:
|
|||||||
if elapsed > 0:
|
if elapsed > 0:
|
||||||
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
|
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
|
||||||
|
|
||||||
log_now = (
|
ended = (
|
||||||
micro_step % args.eval_every == 0
|
args.max_tokens is not None and tokens >= args.max_tokens
|
||||||
or micro_step == 1
|
) or (args.max_tokens is None and micro_step >= args.steps)
|
||||||
or (args.max_tokens is not None and tokens >= args.max_tokens)
|
log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended
|
||||||
or (args.max_tokens is None and micro_step >= args.steps)
|
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)
|
held = _heldout_loss(model, held_chunks, device, use_bf16)
|
||||||
if held is not None:
|
if held is not None:
|
||||||
metrics["heldout/loss"] = held
|
metrics["heldout/loss"] = held
|
||||||
@@ -446,17 +502,36 @@ def main() -> None:
|
|||||||
+ (f" held {held:.4f}" if held is not None else "")
|
+ (f" held {held:.4f}" if held is not None else "")
|
||||||
+ f" aux {metrics['moe/aux']:.4f} z {metrics['moe/z']:.4f}"
|
+ 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:
|
if held is not None and held < best_heldout:
|
||||||
for prefix in args.gen_prefix:
|
best_heldout = held
|
||||||
sample = gen_sample(prefix)
|
payload = _payload(
|
||||||
print(f" gen[{prefix[:16]}]: {sample}")
|
cfg,
|
||||||
if tracker is not None:
|
model,
|
||||||
import swanlab
|
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(
|
tracker.log(
|
||||||
{f"gen/{prefix[:24]}": swanlab.Text(sample)},
|
{f"gen/{prefix[:24]}": swanlab.Text(sample)},
|
||||||
step=micro_step,
|
step=micro_step,
|
||||||
)
|
)
|
||||||
|
if ckpt_now:
|
||||||
payload = _payload(
|
payload = _payload(
|
||||||
cfg,
|
cfg,
|
||||||
model,
|
model,
|
||||||
@@ -469,17 +544,8 @@ def main() -> None:
|
|||||||
best_heldout=best_heldout,
|
best_heldout=best_heldout,
|
||||||
)
|
)
|
||||||
_save(_sibling(args.out, "_last"), payload)
|
_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
|
del payload
|
||||||
if device == "cuda":
|
if tracker is not None and (log_now or eval_now):
|
||||||
torch.cuda.empty_cache()
|
|
||||||
if tracker is not None:
|
|
||||||
tracker.log(metrics, step=micro_step)
|
tracker.log(metrics, step=micro_step)
|
||||||
|
|
||||||
payload = _payload(
|
payload = _payload(
|
||||||
|
|||||||
+26
-7
@@ -1,9 +1,9 @@
|
|||||||
"""Instruction SFT for zh↔en translation. Prompt template matches eval_mt.
|
"""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_wiki.pt --data data/sft/toy.jsonl
|
||||||
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data data/sft/opus.jsonl \\
|
uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\
|
||||||
--seq-len 512 --batch 4 --lr 5e-5 --epochs 2
|
--limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -18,7 +18,7 @@ from kda.layers.latent_moe import moe_router_losses
|
|||||||
from kda.training.data import (
|
from kda.training.data import (
|
||||||
IGNORE_INDEX,
|
IGNORE_INDEX,
|
||||||
iter_sft_batches,
|
iter_sft_batches,
|
||||||
load_sft_rows,
|
resolve_sft_rows,
|
||||||
load_tokenizer,
|
load_tokenizer,
|
||||||
)
|
)
|
||||||
from kda.training.eval_mt import evaluate_pairs
|
from kda.training.eval_mt import evaluate_pairs
|
||||||
@@ -72,7 +72,22 @@ def _read_lines(path: str) -> list[str]:
|
|||||||
def main() -> None:
|
def main() -> None:
|
||||||
p = argparse.ArgumentParser(description=__doc__)
|
p = argparse.ArgumentParser(description=__doc__)
|
||||||
p.add_argument("--ckpt", required=True)
|
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("--out", default="ckpts/k3_sft.pt")
|
||||||
p.add_argument("--tokenizer", default=None)
|
p.add_argument("--tokenizer", default=None)
|
||||||
p.add_argument("--batch", type=int, default=4)
|
p.add_argument("--batch", type=int, default=4)
|
||||||
@@ -103,9 +118,13 @@ def main() -> None:
|
|||||||
if not tok_src:
|
if not tok_src:
|
||||||
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
|
raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint")
|
||||||
tok = load_tokenizer(tok_src)
|
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:
|
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__}")
|
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)
|
steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1)
|
||||||
|
|||||||
Reference in New Issue
Block a user