Compare commits
2
Commits
5a7d949b01
...
24c9d56b72
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
24c9d56b72 | ||
|
|
5cc0555563 |
@@ -138,7 +138,7 @@ PYTHONPATH=. python train.py
|
|||||||
|
|
||||||
### 复现训练
|
### 复现训练
|
||||||
|
|
||||||
目标是 **~0.5B zh↔en 指令翻译模型**(`K3Config.preset("0.5b")` = 482M,tied Qwen3 词表)。本机 RTX 3060 6GB 只跑 8M 全流程孪生;0.5B 预训练需要 **32–40GB Ampere bf16**。
|
目标是 **~0.5B zh↔en 指令翻译模型**(`K3Config.preset("0.5b")` ≈ 415M,tied Yi-6B 64k 词表)。本机 RTX 3060 6GB 只跑 8M 全流程孪生;0.5B 预训练需要 **32–40GB Ampere bf16**。
|
||||||
|
|
||||||
成功标准是冻结集上的 `translation_success()`,**不是** wiki train loss。wiki 预训练没见过 `Translate to English:\n...`,预训练阶段 `eval_mt` 的 success_rate 预期 ≈0。
|
成功标准是冻结集上的 `translation_success()`,**不是** wiki train loss。wiki 预训练没见过 `Translate to English:\n...`,预训练阶段 `eval_mt` 的 success_rate 预期 ≈0。
|
||||||
|
|
||||||
@@ -160,7 +160,7 @@ uv run python train_k3.py --preset 0.5b --attnres block \
|
|||||||
--max-tokens 1000000000 --warmup 2000
|
--max-tokens 1000000000 --warmup 2000
|
||||||
```
|
```
|
||||||
|
|
||||||
`0.5b` 预设:`d=768`,`L=24`(6×3 KDA + 1 MLA),`H=12`,`head=64`,`chunk=64`,LatentMoE `ℓ=384` / 16 routed / Top-2 / shared 2,tied Qwen3 embedding,activation checkpoint 默认开。训练默认 seq 2048、micro-batch 2、grad-acc 8、lr 3e-4、warmup **64 optimizer steps**。checkpoint:`ckpts/k3_0.5b.pt`,另写 `_last` / `_best`(best 按 held-out CE)。
|
`0.5b` 预设:`d=768`,`L=24`(6×3 KDA + 1 MLA),`H=12`,`head=64`,`chunk=64`,LatentMoE `ℓ=384` / 16 routed / Top-2 / shared 2,tied Yi-6B embedding(64k,有 EOS),activation checkpoint 默认开。训练默认 seq 2048、micro-batch 2、grad-acc 8、lr 3e-4、warmup **64 optimizer steps**。checkpoint:`ckpts/k3_0.5b.pt`,另写 `_last` / `_best`(best 按 held-out CE)。**不能**从 Qwen3 词表的旧 ckpt `--resume`。
|
||||||
|
|
||||||
Token 会计:`step` 仍是 micro-batch;有效 token = `batch × seq_len × micro_steps`。默认 0.5b 冒烟是 **8.2M token ≈ 0.017 tok/param**。翻译前置 LM 的最低有意义预算是 **1B token**(`--max-tokens`),不是 2000 step。
|
Token 会计:`step` 仍是 micro-batch;有效 token = `batch × seq_len × micro_steps`。默认 0.5b 冒烟是 **8.2M token ≈ 0.017 tok/param**。翻译前置 LM 的最低有意义预算是 **1B token**(`--max-tokens`),不是 2000 step。
|
||||||
|
|
||||||
|
|||||||
@@ -9,13 +9,15 @@ Hybrid Attention (K3): 每 4 层 1 次 Gated MLA, 末层强制 MLA.
|
|||||||
|
|
||||||
Presets:
|
Presets:
|
||||||
toy — ~8M, 自训 8k SP, 本地过拟合
|
toy — ~8M, 自训 8k SP, 本地过拟合
|
||||||
0.5b — ~482M, Qwen3 词表, 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens
|
0.5b — ~415M, Yi-6B 词表 (64k), 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
# Qwen3 config.json; train_k3 overrides with len(tokenizer).
|
# 01-ai/Yi-6B config.json; train_k3 overrides with len(tokenizer).
|
||||||
|
YI6B_VOCAB_SIZE = 64000
|
||||||
|
# Kept for old Qwen3 checkpoints / docs.
|
||||||
QWEN3_VOCAB_SIZE = 151936
|
QWEN3_VOCAB_SIZE = 151936
|
||||||
|
|
||||||
|
|
||||||
@@ -24,7 +26,7 @@ class K3Config:
|
|||||||
# 主干
|
# 主干
|
||||||
hidden_size: int = 256
|
hidden_size: int = 256
|
||||||
num_hidden_layers: int = 4
|
num_hidden_layers: int = 4
|
||||||
vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Qwen3
|
vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Yi-6B 64k
|
||||||
initializer_range: float = 0.02
|
initializer_range: float = 0.02
|
||||||
norm_eps: float = 1e-6
|
norm_eps: float = 1e-6
|
||||||
tie_word_embeddings: bool = False
|
tie_word_embeddings: bool = False
|
||||||
@@ -76,11 +78,11 @@ class K3Config:
|
|||||||
return cls()
|
return cls()
|
||||||
if name in {"0.5b", "500m"}:
|
if name in {"0.5b", "500m"}:
|
||||||
# H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm.
|
# H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm.
|
||||||
# ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA).
|
# ~415M with tied Yi-6B embeddings. 6×(3 KDA + 1 MLA).
|
||||||
return cls(
|
return cls(
|
||||||
hidden_size=768,
|
hidden_size=768,
|
||||||
num_hidden_layers=24,
|
num_hidden_layers=24,
|
||||||
vocab_size=QWEN3_VOCAB_SIZE,
|
vocab_size=YI6B_VOCAB_SIZE,
|
||||||
tie_word_embeddings=True,
|
tie_word_embeddings=True,
|
||||||
max_position_embeddings=2048,
|
max_position_embeddings=2048,
|
||||||
num_heads=12,
|
num_heads=12,
|
||||||
|
|||||||
@@ -67,14 +67,18 @@ def evaluate_pairs(
|
|||||||
want = "zh" if target_lang.startswith("zh") else "en"
|
want = "zh" if target_lang.startswith("zh") else "en"
|
||||||
lang_ok += int(_detect_lang(hyp) == want)
|
lang_ok += int(_detect_lang(hyp) == want)
|
||||||
chrf_sum += _chrf(hyp, ref)
|
chrf_sum += _chrf(hyp, ref)
|
||||||
corpus = {}
|
corpus = {"chrf": chrf_sum / max(n, 1), "bleu": None}
|
||||||
try:
|
try:
|
||||||
from sacrebleu.metrics import BLEU, CHRF
|
from sacrebleu.metrics import CHRF
|
||||||
|
|
||||||
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
|
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
from sacrebleu.metrics import BLEU
|
||||||
|
|
||||||
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
|
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
|
||||||
except Exception:
|
except Exception:
|
||||||
corpus["chrf"] = chrf_sum / max(n, 1)
|
|
||||||
corpus["bleu"] = None
|
corpus["bleu"] = None
|
||||||
return {
|
return {
|
||||||
"n": n,
|
"n": n,
|
||||||
@@ -131,6 +135,8 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
printable = {k: v for k, v in out.items() if k != "hyps"}
|
printable = {k: v for k, v in out.items() if k != "hyps"}
|
||||||
print(json.dumps(printable, ensure_ascii=False, indent=2))
|
print(json.dumps(printable, ensure_ascii=False, indent=2))
|
||||||
|
for i, hyp in enumerate(out["hyps"][: min(5, out["n"])]):
|
||||||
|
print(f" [{i}] {hyp}")
|
||||||
elif not args.prefix:
|
elif not args.prefix:
|
||||||
raise SystemExit("pass --prefix and/or --src + --ref")
|
raise SystemExit("pass --prefix and/or --src + --ref")
|
||||||
|
|
||||||
|
|||||||
+27
-11
@@ -28,8 +28,33 @@ def _detect_lang(text: str) -> str | None:
|
|||||||
return tag[:2]
|
return tag[:2]
|
||||||
|
|
||||||
|
|
||||||
|
def _chrf_ngram(hyp: str, ref: str, max_n: int = 4) -> float:
|
||||||
|
"""Count-based char n-gram F (β=2), 0–100. Not set-overlap unigrams."""
|
||||||
|
from collections import Counter
|
||||||
|
|
||||||
|
hyp, ref = hyp.strip(), ref.strip()
|
||||||
|
if not hyp or not ref:
|
||||||
|
return 0.0
|
||||||
|
scores: list[float] = []
|
||||||
|
for n in range(1, max_n + 1):
|
||||||
|
if len(hyp) < n or len(ref) < n:
|
||||||
|
scores.append(0.0)
|
||||||
|
continue
|
||||||
|
hc = Counter(hyp[i : i + n] for i in range(len(hyp) - n + 1))
|
||||||
|
rc = Counter(ref[i : i + n] for i in range(len(ref) - n + 1))
|
||||||
|
overlap = sum((hc & rc).values())
|
||||||
|
prec = overlap / max(sum(hc.values()), 1)
|
||||||
|
rec = overlap / max(sum(rc.values()), 1)
|
||||||
|
if prec + rec == 0:
|
||||||
|
scores.append(0.0)
|
||||||
|
continue
|
||||||
|
beta2 = 4.0
|
||||||
|
scores.append((1.0 + beta2) * prec * rec / (beta2 * prec + rec))
|
||||||
|
return 100.0 * sum(scores) / max(len(scores), 1)
|
||||||
|
|
||||||
|
|
||||||
def _chrf(hyp: str, ref: str) -> float:
|
def _chrf(hyp: str, ref: str) -> float:
|
||||||
"""chrF++ in 0–100. Falls back to char unigram F if sacrebleu is missing."""
|
"""chrF++ in 0–100. Falls back to count-based char n-grams if sacrebleu is missing."""
|
||||||
if not hyp or not ref:
|
if not hyp or not ref:
|
||||||
return 0.0
|
return 0.0
|
||||||
try:
|
try:
|
||||||
@@ -37,16 +62,7 @@ def _chrf(hyp: str, ref: str) -> float:
|
|||||||
|
|
||||||
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
|
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
|
||||||
except Exception:
|
except Exception:
|
||||||
hyp_c, ref_c = list(hyp), list(ref)
|
return _chrf_ngram(hyp, 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(
|
def translation_success(
|
||||||
|
|||||||
@@ -28,6 +28,16 @@ def test_good_zh2en_passes():
|
|||||||
assert translation_success(src, hyp, ref, target_lang="en") is True
|
assert translation_success(src, hyp, ref, target_lang="en") is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_english_wiki_garbage_does_not_pass_zh2en():
|
||||||
|
from kda.training.success import _chrf
|
||||||
|
|
||||||
|
src = "今天天气很好。"
|
||||||
|
hyp = "The first one's the time."
|
||||||
|
ref = "The weather is very nice today."
|
||||||
|
assert _chrf(hyp, ref) < 40.0
|
||||||
|
assert translation_success(src, hyp, ref, target_lang="en") is False
|
||||||
|
|
||||||
|
|
||||||
def test_container_help_exits_2():
|
def test_container_help_exits_2():
|
||||||
import importlib.util
|
import importlib.util
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|||||||
@@ -78,6 +78,7 @@ def test_preset_0_5b_schedule():
|
|||||||
assert cfg.hidden_size == 768
|
assert cfg.hidden_size == 768
|
||||||
assert cfg.num_heads * cfg.head_dim == cfg.hidden_size
|
assert cfg.num_heads * cfg.head_dim == cfg.hidden_size
|
||||||
assert cfg.num_hidden_layers == 24
|
assert cfg.num_hidden_layers == 24
|
||||||
|
assert cfg.vocab_size == 64000
|
||||||
assert cfg.tie_word_embeddings
|
assert cfg.tie_word_embeddings
|
||||||
assert cfg.chunk_size == 64
|
assert cfg.chunk_size == 64
|
||||||
assert cfg.gradient_checkpointing is True
|
assert cfg.gradient_checkpointing is True
|
||||||
|
|||||||
@@ -0,0 +1,81 @@
|
|||||||
|
"""Resume must not clobber attnres or skip a chunk at budget exit."""
|
||||||
|
from argparse import Namespace
|
||||||
|
from dataclasses import asdict
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from kda.models.causal_lm import CausalLM
|
||||||
|
from kda.models.k3_config import K3Config
|
||||||
|
from kda.training.toy import load_ckpt
|
||||||
|
from train_k3 import _apply_cli_overrides
|
||||||
|
|
||||||
|
|
||||||
|
def _tiny_block():
|
||||||
|
return K3Config(
|
||||||
|
hidden_size=32,
|
||||||
|
num_hidden_layers=4,
|
||||||
|
num_heads=4,
|
||||||
|
head_dim=8,
|
||||||
|
chunk_size=4,
|
||||||
|
vocab_size=64,
|
||||||
|
moe_latent_size=16,
|
||||||
|
moe_d_ff=16,
|
||||||
|
n_routed=4,
|
||||||
|
top_k=2,
|
||||||
|
n_shared=1,
|
||||||
|
kv_lora_rank=8,
|
||||||
|
q_lora_rank=16,
|
||||||
|
qk_nope_head_dim=8,
|
||||||
|
v_head_dim=8,
|
||||||
|
attnres="block",
|
||||||
|
attnres_block_size=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _cli(**kwargs):
|
||||||
|
base = dict(
|
||||||
|
attnres=None,
|
||||||
|
attnres_block_size=None,
|
||||||
|
grad_checkpoint=None,
|
||||||
|
moe_aux_coef=None,
|
||||||
|
moe_z_coef=None,
|
||||||
|
)
|
||||||
|
base.update(kwargs)
|
||||||
|
return Namespace(**base)
|
||||||
|
|
||||||
|
|
||||||
|
def test_resume_without_attnres_flag_keeps_block():
|
||||||
|
cfg = _tiny_block()
|
||||||
|
_apply_cli_overrides(cfg, _cli())
|
||||||
|
assert cfg.attnres == "block"
|
||||||
|
assert cfg.attnres_block_size == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_explicit_attnres_overrides_resume():
|
||||||
|
cfg = _tiny_block()
|
||||||
|
_apply_cli_overrides(cfg, _cli(attnres="full", attnres_block_size=1))
|
||||||
|
assert cfg.attnres == "full"
|
||||||
|
assert cfg.attnres_block_size == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_block_state_with_off_config_cannot_load(tmp_path):
|
||||||
|
cfg = _tiny_block()
|
||||||
|
model = CausalLM(cfg)
|
||||||
|
payload = {"config": asdict(cfg), "model_state": model.state_dict()}
|
||||||
|
payload["config"]["attnres"] = "off"
|
||||||
|
payload["config"]["attnres_block_size"] = None
|
||||||
|
path = str(tmp_path / "polluted.pt")
|
||||||
|
torch.save(payload, path)
|
||||||
|
with pytest.raises(RuntimeError, match="Unexpected key"):
|
||||||
|
load_ckpt(path)
|
||||||
|
|
||||||
|
|
||||||
|
def test_budget_break_does_not_skip_yielded_chunk():
|
||||||
|
next_chunk = 10
|
||||||
|
for chunk_index in (10, 11, 12):
|
||||||
|
tokens = 100
|
||||||
|
if tokens >= 100:
|
||||||
|
break
|
||||||
|
next_chunk = chunk_index + 1
|
||||||
|
assert next_chunk == 10
|
||||||
+39
-23
@@ -41,7 +41,7 @@ _TOY_TRAIN = {
|
|||||||
"gen_every": 200,
|
"gen_every": 200,
|
||||||
}
|
}
|
||||||
_B500M_TRAIN = {
|
_B500M_TRAIN = {
|
||||||
"tokenizer": "Qwen/Qwen3-8B",
|
"tokenizer": "01-ai/Yi-6B",
|
||||||
"out": "ckpts/k3_0.5b.pt",
|
"out": "ckpts/k3_0.5b.pt",
|
||||||
"limit": 20000,
|
"limit": 20000,
|
||||||
"batch": 2,
|
"batch": 2,
|
||||||
@@ -184,6 +184,20 @@ def _apply_moe_coefs(model, cfg: K3Config) -> None:
|
|||||||
module.z_loss_coef = cfg.moe_z_loss_coef
|
module.z_loss_coef = cfg.moe_z_loss_coef
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_cli_overrides(cfg: K3Config, args: argparse.Namespace) -> None:
|
||||||
|
"""Copy only flags the user actually passed. CLI defaults must not clobber a resume."""
|
||||||
|
if args.attnres is not None:
|
||||||
|
cfg.attnres = args.attnres
|
||||||
|
if args.attnres_block_size is not None:
|
||||||
|
cfg.attnres_block_size = args.attnres_block_size
|
||||||
|
if args.grad_checkpoint is not None:
|
||||||
|
cfg.gradient_checkpointing = args.grad_checkpoint
|
||||||
|
if args.moe_aux_coef is not None:
|
||||||
|
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
||||||
|
if args.moe_z_coef is not None:
|
||||||
|
cfg.moe_z_loss_coef = args.moe_z_coef
|
||||||
|
|
||||||
|
|
||||||
def _moe_log(model) -> dict:
|
def _moe_log(model) -> dict:
|
||||||
frac = moe_route_frac(model)
|
frac = moe_route_frac(model)
|
||||||
if frac is None:
|
if frac is None:
|
||||||
@@ -267,9 +281,10 @@ def main() -> None:
|
|||||||
p.add_argument("--device", default="auto")
|
p.add_argument("--device", default="auto")
|
||||||
p.add_argument(
|
p.add_argument(
|
||||||
"--attnres",
|
"--attnres",
|
||||||
default="off",
|
default=None,
|
||||||
choices=["off", "full", "block"],
|
choices=["off", "full", "block"],
|
||||||
help="depth mixer: off=standard residual, block=K3 AttnRes, full=per-layer AttnRes",
|
help="depth mixer: off=standard residual (preset default), block=K3 AttnRes, "
|
||||||
|
"full=per-layer AttnRes. Omit on --resume to keep the checkpoint value",
|
||||||
)
|
)
|
||||||
p.add_argument(
|
p.add_argument(
|
||||||
"--attnres-block-size",
|
"--attnres-block-size",
|
||||||
@@ -329,14 +344,7 @@ def main() -> None:
|
|||||||
tok = load_tokenizer(args.tokenizer)
|
tok = load_tokenizer(args.tokenizer)
|
||||||
cfg = K3Config.preset(args.preset)
|
cfg = K3Config.preset(args.preset)
|
||||||
cfg.vocab_size = tok.vocab_size
|
cfg.vocab_size = tok.vocab_size
|
||||||
cfg.attnres = args.attnres
|
_apply_cli_overrides(cfg, args)
|
||||||
cfg.attnres_block_size = args.attnres_block_size
|
|
||||||
if args.grad_checkpoint is not None:
|
|
||||||
cfg.gradient_checkpointing = args.grad_checkpoint
|
|
||||||
if args.moe_aux_coef is not None:
|
|
||||||
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
|
||||||
if args.moe_z_coef is not None:
|
|
||||||
cfg.moe_z_loss_coef = args.moe_z_coef
|
|
||||||
|
|
||||||
langs = [part.strip() for part in args.langs.split(",") if part.strip()]
|
langs = [part.strip() for part in args.langs.split(",") if part.strip()]
|
||||||
tpm = tokens_per_micro(args.batch, args.seq_len)
|
tpm = tokens_per_micro(args.batch, args.seq_len)
|
||||||
@@ -364,19 +372,16 @@ def main() -> None:
|
|||||||
f"{type(loaded_cfg).__name__}"
|
f"{type(loaded_cfg).__name__}"
|
||||||
)
|
)
|
||||||
cfg = loaded_cfg
|
cfg = loaded_cfg
|
||||||
cfg.attnres = args.attnres
|
_apply_cli_overrides(cfg, args)
|
||||||
cfg.attnres_block_size = args.attnres_block_size
|
|
||||||
if args.grad_checkpoint is not None:
|
|
||||||
cfg.gradient_checkpointing = args.grad_checkpoint
|
|
||||||
if args.moe_aux_coef is not None:
|
|
||||||
cfg.moe_aux_loss_coef = args.moe_aux_coef
|
|
||||||
if args.moe_z_coef is not None:
|
|
||||||
cfg.moe_z_loss_coef = args.moe_z_coef
|
|
||||||
model.gradient_checkpointing = cfg.gradient_checkpointing
|
model.gradient_checkpointing = cfg.gradient_checkpointing
|
||||||
model.to(device)
|
model.to(device)
|
||||||
payload = torch.load(args.resume, map_location="cpu", weights_only=False)
|
payload = torch.load(args.resume, map_location="cpu", weights_only=False)
|
||||||
if payload.get("tokenizer") and payload["tokenizer"] != args.tokenizer:
|
if payload.get("tokenizer") and payload["tokenizer"] != args.tokenizer:
|
||||||
print(f"warning: ckpt tokenizer {payload['tokenizer']} != {args.tokenizer}")
|
raise SystemExit(
|
||||||
|
f"tokenizer mismatch: ckpt {payload['tokenizer']!r} vs "
|
||||||
|
f"CLI {args.tokenizer!r}; embeddings are not interchangeable "
|
||||||
|
f"(do not resume a Qwen ckpt with Yi)"
|
||||||
|
)
|
||||||
micro_step = int(payload.get("micro_step", 0))
|
micro_step = int(payload.get("micro_step", 0))
|
||||||
opt_step = int(payload.get("opt_step", 0))
|
opt_step = int(payload.get("opt_step", 0))
|
||||||
tokens = int(payload.get("tokens", 0))
|
tokens = int(payload.get("tokens", 0))
|
||||||
@@ -469,6 +474,10 @@ def main() -> None:
|
|||||||
model.train()
|
model.train()
|
||||||
t0 = time.perf_counter()
|
t0 = time.perf_counter()
|
||||||
tokens_at_t0 = tokens
|
tokens_at_t0 = tokens
|
||||||
|
# Index of the next untrained chunk. Mid-loop saves use last_trained+1.
|
||||||
|
# The final save must NOT +1 again: the loop may break on a yielded chunk
|
||||||
|
# that was never trained (budget check is at the top).
|
||||||
|
next_chunk = chunk_index
|
||||||
for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index):
|
for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index):
|
||||||
if args.max_tokens is not None:
|
if args.max_tokens is not None:
|
||||||
if tokens >= args.max_tokens:
|
if tokens >= args.max_tokens:
|
||||||
@@ -497,6 +506,7 @@ def main() -> None:
|
|||||||
lr_now = optim.param_groups[0]["lr"]
|
lr_now = optim.param_groups[0]["lr"]
|
||||||
if raw_loss < best_train:
|
if raw_loss < best_train:
|
||||||
best_train = raw_loss
|
best_train = raw_loss
|
||||||
|
next_chunk = chunk_index + 1
|
||||||
|
|
||||||
metrics = {
|
metrics = {
|
||||||
"train/loss": raw_loss,
|
"train/loss": raw_loss,
|
||||||
@@ -542,7 +552,7 @@ def main() -> None:
|
|||||||
micro_step=micro_step,
|
micro_step=micro_step,
|
||||||
opt_step=opt_step,
|
opt_step=opt_step,
|
||||||
tokens=tokens,
|
tokens=tokens,
|
||||||
chunk_index=chunk_index + 1,
|
chunk_index=next_chunk,
|
||||||
best_heldout=best_heldout,
|
best_heldout=best_heldout,
|
||||||
)
|
)
|
||||||
_save(_sibling(args.out, "_best"), payload)
|
_save(_sibling(args.out, "_best"), payload)
|
||||||
@@ -570,7 +580,7 @@ def main() -> None:
|
|||||||
micro_step=micro_step,
|
micro_step=micro_step,
|
||||||
opt_step=opt_step,
|
opt_step=opt_step,
|
||||||
tokens=tokens,
|
tokens=tokens,
|
||||||
chunk_index=chunk_index + 1,
|
chunk_index=next_chunk,
|
||||||
best_heldout=best_heldout,
|
best_heldout=best_heldout,
|
||||||
)
|
)
|
||||||
_save(_sibling(args.out, "_last"), payload)
|
_save(_sibling(args.out, "_last"), payload)
|
||||||
@@ -586,7 +596,7 @@ def main() -> None:
|
|||||||
micro_step=micro_step,
|
micro_step=micro_step,
|
||||||
opt_step=opt_step,
|
opt_step=opt_step,
|
||||||
tokens=tokens,
|
tokens=tokens,
|
||||||
chunk_index=chunk_index + 1,
|
chunk_index=next_chunk,
|
||||||
best_heldout=best_heldout,
|
best_heldout=best_heldout,
|
||||||
)
|
)
|
||||||
_save(args.out, payload)
|
_save(args.out, payload)
|
||||||
@@ -596,6 +606,12 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
if tracker is not None:
|
if tracker is not None:
|
||||||
tracker.finish()
|
tracker.finish()
|
||||||
|
if device == "cuda":
|
||||||
|
try:
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -339,6 +339,8 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
printable = {k: v for k, v in out.items() if k != "hyps"}
|
printable = {k: v for k, v in out.items() if k != "hyps"}
|
||||||
print(printable)
|
print(printable)
|
||||||
|
for i, hyp in enumerate((out.get("hyps") or [])[:2]):
|
||||||
|
print(f" hyp[{i}] {hyp}")
|
||||||
if tracker is not None:
|
if tracker is not None:
|
||||||
tracker.log(
|
tracker.log(
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user