| |
| """ |
| train_chat.py: Chat SFT for MetaDiffusion-150M-exp (LLaDA Algorithm 2 style). |
| |
| Checkpoints use the same format as train.py: |
| {step, model_state_dict, optimizer_state_dict, scheduler_state_dict, config} |
| |
| Usage: |
| python3 train_chat.py \ |
| --model-path ../hf_release \ |
| --data-dir data/no_robots_chatml \ |
| --output-dir checkpoints_chat \ |
| --epochs 8 |
| """ |
|
|
| import argparse |
| import glob |
| import heapq |
| import json |
| import logging |
| import math |
| import os |
| import re |
| import shutil |
| import sys |
| import time |
| from dataclasses import asdict |
| from pathlib import Path |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from safetensors.torch import load_file |
| from torch.optim import AdamW |
| from torch.optim.lr_scheduler import LambdaLR |
| from torch.utils.data import DataLoader, Dataset |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) |
| from model import MetaDiffusionLM, MetaDiffusionConfig |
|
|
| MASK_TOKEN_ID = 32000 |
| BASE_VOCAB = 32000 |
| RESERVED = "<|reserved|>" |
| BASE_VOCAB_WITH_RESERVED = BASE_VOCAB + 1 |
| CHAT_TOKENS = ["<|im_start|>", "<|im_end|>"] + [f"<|r{i}|>" for i in range(1, 8)] |
| CHAT_VOCAB = BASE_VOCAB + 1 + len(CHAT_TOKENS) |
|
|
|
|
| def ensure_chat_tokens(tokenizer): |
| """Make sure chat tokens live at ids 32001..32009 |
| |
| Handles both a fresh base tokenizer (adds <|reserved|> at 32000 first) and |
| an already-prepared one (no-op). |
| """ |
| if tokenizer.convert_tokens_to_ids("<|im_start|>") == tokenizer.unk_token_id: |
| if len(tokenizer) == BASE_VOCAB: |
| tokenizer.add_special_tokens({"additional_special_tokens": [RESERVED]}) |
| tokenizer.add_special_tokens({"additional_special_tokens": CHAT_TOKENS}) |
| return tokenizer |
|
|
| logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") |
| logger = logging.getLogger(__name__) |
|
|
| def format_bytes(b): |
| for unit in ["B", "KB", "MB", "GB", "TB"]: |
| if b < 1024: |
| return f"{b:.1f} {unit}" |
| b /= 1024 |
| return f"{b:.1f} PB" |
|
|
|
|
| def format_duration(seconds): |
| seconds = max(0, int(seconds)) |
| hours, rem = divmod(seconds, 3600) |
| minutes, secs = divmod(rem, 60) |
| if hours > 0: |
| return f"{hours}h{minutes:02d}m{secs:02d}s" |
| if minutes > 0: |
| return f"{minutes}m{secs:02d}s" |
| return f"{secs}s" |
|
|
|
|
| def get_gpu_memory_info(): |
| if not torch.cuda.is_available(): |
| return None, None |
| torch.cuda.synchronize() |
| free, total = torch.cuda.mem_get_info() |
| return total, free |
|
|
|
|
| def detect_max_batch_size(model, seq_len, device, keep_free_fraction=0.1, |
| amp_dtype=None): |
| """Largest batch size that fits in VRAM with headroom (from train.py).""" |
| if not torch.cuda.is_available(): |
| return 8 |
| total_mem, free_mem = get_gpu_memory_info() |
| if total_mem is None: |
| return 8 |
| logger.info(f"GPU memory: {format_bytes(total_mem)} total, {format_bytes(free_mem)} free") |
| model = model.to(device).train() |
| mem_limit = total_mem - int(total_mem * keep_free_fraction) |
| last_working, first_oom = 1, None |
| for bs in [1, 2, 4, 8, 16, 32, 64, 128, 256]: |
| torch.cuda.synchronize() |
| if torch.cuda.memory_allocated() >= mem_limit: |
| first_oom = bs |
| break |
| try: |
| input_ids = torch.randint(0, 32000, (bs, seq_len), device=device) |
| labels = torch.randint(0, 32000, (bs, seq_len), device=device) |
| mask_positions = torch.rand(bs, seq_len, device=device) < 0.5 |
| timesteps = torch.rand(bs, device=device) |
| with torch.autocast("cuda", dtype=amp_dtype or torch.float16): |
| logits = model(input_ids, timesteps) |
| loss, num_masked = model.compute_loss(logits, labels, mask_positions) |
| if num_masked > 0: |
| (loss / 4).backward() |
| torch.cuda.synchronize() |
| peak_mem = torch.cuda.max_memory_allocated() |
| logger.info(f" batch_size={bs:>3d}: peak VRAM={format_bytes(peak_mem)} " |
| f"(limit={format_bytes(mem_limit)})") |
| if peak_mem >= mem_limit: |
| first_oom = bs |
| break |
| last_working = bs |
| except RuntimeError as e: |
| if "out of memory" in str(e).lower(): |
| first_oom = bs |
| break |
| raise |
| finally: |
| model.zero_grad(set_to_none=True) |
| torch.cuda.empty_cache() |
| if first_oom is not None and last_working < first_oom - 1: |
| lo, hi = last_working, first_oom |
| while lo + 1 < hi: |
| mid = (lo + hi) // 2 |
| try: |
| input_ids = torch.randint(0, 32000, (mid, seq_len), device=device) |
| labels = torch.randint(0, 32000, (mid, seq_len), device=device) |
| mask_positions = torch.rand(mid, seq_len, device=device) < 0.5 |
| timesteps = torch.rand(mid, device=device) |
| with torch.autocast("cuda", dtype=amp_dtype or torch.float16): |
| logits = model(input_ids, timesteps) |
| loss, num_masked = model.compute_loss(logits, labels, mask_positions) |
| if num_masked > 0: |
| (loss / 4).backward() |
| torch.cuda.synchronize() |
| if torch.cuda.max_memory_allocated() < mem_limit: |
| lo = mid |
| else: |
| hi = mid |
| except RuntimeError as e: |
| if "out of memory" in str(e).lower(): |
| hi = mid |
| else: |
| raise |
| finally: |
| model.zero_grad(set_to_none=True) |
| torch.cuda.empty_cache() |
| last_working = lo |
| model.zero_grad(set_to_none=True) |
| torch.cuda.empty_cache() |
| logger.info(f"Detected max batch_size: {last_working}") |
| return last_working |
|
|
|
|
| def get_step_from_filename(filename): |
| basename = os.path.basename(filename) |
| m = re.match(r"step_(\d+)(?:_\w+)?\.pt$", basename) |
| return int(m.group(1)) if m else None |
|
|
|
|
| def load_best_steps(stats_path, max_n): |
| if not os.path.exists(stats_path): |
| return set() |
| entries = [] |
| with open(stats_path) as f: |
| for line in f: |
| line = line.strip() |
| if not line: |
| continue |
| try: |
| entry = json.loads(line) |
| if "step" in entry and "loss" in entry: |
| entries.append((entry["loss"], entry["step"])) |
| except json.JSONDecodeError: |
| continue |
| return {step for _, step in heapq.nsmallest(max_n, entries)} |
|
|
|
|
| def cleanup_checkpoints(output_dir, keep_first_n, keep_last_n, keep_best_n, stats_path): |
| all_ckpts = sorted(glob.glob(os.path.join(output_dir, "step_*.pt"))) |
| if len(all_ckpts) <= keep_first_n + keep_last_n + keep_best_n: |
| return |
| first_steps = {get_step_from_filename(c) for c in all_ckpts[:keep_first_n]} |
| last_steps = {get_step_from_filename(c) for c in all_ckpts[-keep_last_n:]} |
| best_steps = load_best_steps(stats_path, keep_best_n) |
| keep_steps = (first_steps | last_steps | best_steps) - {None} |
| for ckpt in all_ckpts: |
| s = get_step_from_filename(ckpt) |
| if s is not None and s not in keep_steps: |
| try: |
| os.remove(ckpt) |
| except OSError: |
| pass |
| logger.info(f"Cleaned old checkpoints (kept {len(keep_steps)}: " |
| f"{len(first_steps)} first, {len(last_steps)} last, {len(best_steps)} best)") |
|
|
|
|
| def get_free_disk_space(path): |
| return shutil.disk_usage(path).free |
|
|
| def build_config(config_dict): |
| """Build MetaDiffusionConfig, ignoring non-dataclass keys (model_type, ...).""" |
| valid = {k: v for k, v in config_dict.items() if k in MetaDiffusionConfig.__dataclass_fields__} |
| config = MetaDiffusionConfig(**valid) |
| config.tie_word_embeddings = False |
| return config |
|
|
|
|
| def load_model(model_path, device): |
| """Load MetaDiffusionLM from a dir (config.json + model.safetensors) or a step_*.pt.""" |
| path = Path(model_path) |
| if path.is_dir(): |
| with open(path / "config.json") as f: |
| config = build_config(json.load(f)) |
| model = MetaDiffusionLM(config).to(device) |
| sd = load_file(path / "model.safetensors") |
| sd = {k[len("model."):] if k.startswith("model.") else k: v for k, v in sd.items()} |
| missing, unexpected = model.load_state_dict(sd, strict=False) |
| if missing or unexpected: |
| logger.warning(f"missing={missing[:5]} unexpected={unexpected[:5]}") |
| else: |
| ckpt = torch.load(path, map_location=device, weights_only=False) |
| config = build_config(ckpt["config"]) |
| model = MetaDiffusionLM(config).to(device) |
| model.load_state_dict(clean_state_dict(ckpt["model_state_dict"])) |
| return model, config |
|
|
|
|
| def clean_state_dict(state_dict): |
| """Strip torch.compile's _orig_mod. prefix from checkpoint keys.""" |
| return {k.replace("_orig_mod.", "", 1) if k.startswith("_orig_mod.") else k: v |
| for k, v in state_dict.items()} |
|
|
|
|
| def expand_embeddings(model, new_vocab): |
| """Mean-init new rows (ChatML + rainbow tokens) in embed_tokens and lm_head. |
| |
| New modules are created on the model's device/dtype: nn.Embedding/nn.Linear |
| default to CPU, which would crash the first forward ("Tensor device |
| mismatch") unless something else moves the model afterwards. |
| """ |
| old_vocab = model.config.mask_vocab_size |
| if new_vocab <= old_vocab: |
| return |
| device = model.embed_tokens.weight.device |
| dtype = model.embed_tokens.weight.dtype |
| mean_emb = model.embed_tokens.weight.data.mean(dim=0, keepdim=True) |
| n_new = new_vocab - old_vocab |
|
|
| emb = torch.cat([model.embed_tokens.weight.data, mean_emb.expand(n_new, -1)], dim=0) |
| model.embed_tokens = nn.Embedding(new_vocab, model.config.hidden_size, |
| padding_idx=model.config.pad_token_id).to(device, dtype) |
| model.embed_tokens.weight.data.copy_(emb) |
|
|
| head = torch.cat([model.lm_head.weight.data, mean_emb.expand(n_new, -1)], dim=0) |
| model.lm_head = nn.Linear(model.config.hidden_size, new_vocab, bias=False).to(device, dtype) |
| model.lm_head.weight.data.copy_(head) |
|
|
| model.config.mask_vocab_size = new_vocab |
| logger.info(f"Expanded embeddings {old_vocab} -> {new_vocab} (mean init)") |
|
|
|
|
| class ChatDataset(Dataset): |
| """no_robots ChatML examples; masks ONLY the last assistant response. |
| |
| Examples are stored as plain int lists (NOT tensors): with forkserver |
| workers (torch's default once CUDA is initialized), every tensor in the |
| dataset is transferred through shared memory at worker spawn, and 9000 |
| tensors blows the open-file limit. Lists pickle as bytes. |
| """ |
|
|
| def __init__(self, data_path, tokenizer, seq_len, seed=42): |
| raw = torch.load(data_path, weights_only=True)["examples"] |
| self.examples = [ |
| { |
| "ids": ex["input_ids"].tolist(), |
| "a0": int(ex["assistant_start"]), |
| "a1": int(ex["assistant_end"]), |
| } |
| for ex in raw |
| ] |
| self.seq_len = seq_len |
| self.mask_id = MASK_TOKEN_ID |
| self.rainbow_ids = [ |
| tokenizer.convert_tokens_to_ids(f"<|r{i}|>") for i in range(1, 8) |
| ] |
| self.seed = seed |
| logger.info(f"Loaded {len(self.examples)} examples from {data_path}") |
|
|
| def __len__(self): |
| return len(self.examples) |
|
|
| def __getitem__(self, idx): |
| ex = self.examples[idx] |
| ids = ex["ids"] |
| a0, a1 = ex["a0"], ex["a1"] |
|
|
| |
| if len(ids) > self.seq_len: |
| resp = ids[a0:a1] |
| if len(resp) > self.seq_len: |
| resp = resp[: self.seq_len] |
| room = self.seq_len - len(resp) |
| hist = ids[:a0] |
| hist = hist[len(hist) - room:] if room > 0 else [] |
| ids = hist + resp |
| a0, a1 = len(hist), len(ids) |
|
|
| |
| n = len(ids) |
| pad = self.seq_len - n |
| full = ids + [self.rainbow_ids[j % 7] for j in range(pad)] |
|
|
| can_mask = torch.zeros(self.seq_len, dtype=torch.bool) |
| can_mask[a0:a1] = True |
|
|
| t = torch.rand(1).item() |
| rand = torch.rand(self.seq_len) |
| mask_pos = (rand < t) & can_mask |
|
|
| input_ids = torch.tensor(full, dtype=torch.long) |
| input_ids[mask_pos] = self.mask_id |
| attention = torch.ones(self.seq_len, dtype=torch.long) |
| attention[n:] = 0 |
|
|
| return { |
| "input_ids": input_ids, |
| "labels": torch.tensor(full, dtype=torch.long), |
| "mask_positions": mask_pos, |
| "timesteps": torch.tensor(t, dtype=torch.float32), |
| "attention_mask": attention, |
| "resp_len": torch.tensor(max(a1 - a0, 1), dtype=torch.float32), |
| } |
|
|
|
|
| def collate_fn(batch): |
| return { |
| k: torch.stack([b[k] for b in batch]) for k in batch[0] |
| } |
|
|
|
|
| def worker_init_fn(worker_id): |
| torch.manual_seed(42 + worker_id) |
|
|
| def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps, |
| min_lr_ratio=0.1): |
| def lr_lambda(current_step): |
| if current_step < num_warmup_steps: |
| return float(current_step) / float(max(1, num_warmup_steps)) |
| progress = float(current_step - num_warmup_steps) / float( |
| max(1, num_training_steps - num_warmup_steps) |
| ) |
| return max(min_lr_ratio, 0.5 * (1.0 + math.cos(math.pi * progress))) |
| return LambdaLR(optimizer, lr_lambda) |
|
|
|
|
| @torch.no_grad() |
| def evaluate(model, val_dataset, batch_size, device, dtype, n_max=100): |
| model.eval() |
| losses, n_seen = [], 0 |
| for start in range(0, min(len(val_dataset), n_max), batch_size): |
| idxs = list(range(start, min(start + batch_size, n_max))) |
| batch = collate_fn([val_dataset[i] for i in idxs]) |
| input_ids = batch["input_ids"].to(device) |
| labels = batch["labels"].to(device) |
| mask_positions = batch["mask_positions"].to(device) |
| timesteps = batch["timesteps"].to(device) |
| attention = batch["attention_mask"].to(device) |
| resp_len = batch["resp_len"].to(device) |
| with torch.autocast("cuda", dtype=dtype): |
| logits = model(input_ids, timesteps, attention_mask=attention) |
| if mask_positions.any(): |
| ce = F.cross_entropy(logits.float()[mask_positions], |
| labels[mask_positions], reduction="none") |
| w = (1.0 / (timesteps * resp_len)).unsqueeze(1).expand_as(labels) |
| loss = (ce * w[mask_positions]).sum() / input_ids.shape[0] |
| losses.append(loss.item()) |
| n_seen += 1 |
| model.train() |
| return sum(losses) / len(losses) if losses else float("nan"), n_seen |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Chat SFT for MetaDiffusion (LLaDA Algorithm 2)") |
| parser.add_argument("--model-path", default="../hf_release", |
| help="Dir with config.json + model.safetensors, or a step_*.pt") |
| parser.add_argument("--data-dir", default="data/no_robots_chatml") |
| parser.add_argument("--output-dir", default="checkpoints_chat") |
| parser.add_argument("--seq-len", type=int, default=512) |
| parser.add_argument("--batch-size", type=int, default=0, help="0 = auto-detect") |
| parser.add_argument("--grad-accum-steps", type=int, default=4) |
| parser.add_argument("--num-workers", type=int, default=4, |
| help="DataLoader workers (0 if forkserver shm issues)") |
| parser.add_argument("--lr", type=float, default=3e-5) |
| parser.add_argument("--min-lr-ratio", type=float, default=0.1) |
| parser.add_argument("--warmup-steps", type=int, default=100) |
| parser.add_argument("--weight-decay", type=float, default=0.1) |
| parser.add_argument("--epochs", type=int, default=8) |
| parser.add_argument("--max-steps", type=int, default=0, help="0 = epochs only") |
| parser.add_argument("--transferred-lr-mult", type=float, default=0.33) |
| parser.add_argument("--new-lr-mult", type=float, default=1.0) |
| parser.add_argument("--max-grad-norm", type=float, default=1.0) |
| parser.add_argument("--save-every", type=int, default=500) |
| parser.add_argument("--log-every", type=int, default=50) |
| parser.add_argument("--val-every", type=int, default=200) |
| parser.add_argument("--patience", type=int, default=3, |
| help="Early stop after N val checks without improvement (0 = off)") |
| parser.add_argument("--min-delta", type=float, default=0.001, |
| help="Relative val-loss improvement required to count as progress") |
| parser.add_argument("--resume-from", type=str, default=None) |
| parser.add_argument("--seed", type=int, default=42) |
| parser.add_argument("--device", default="cuda", help="cuda, cuda:1, cpu") |
| parser.add_argument("--bf16", action="store_true", help="Use bf16 instead of fp16") |
| parser.add_argument("--no-compile", action="store_true") |
| parser.add_argument("--keep-free-vram", type=float, default=0.1) |
| parser.add_argument("--keep-first-n", type=int, default=2) |
| parser.add_argument("--keep-last-n", type=int, default=2) |
| parser.add_argument("--keep-best-n", type=int, default=2) |
| parser.add_argument("--disk-min-gb", type=float, default=5.0) |
| parser.add_argument("--export-dir", type=str, default=None, |
| help="Export final dir (config+safetensors+tokenizer)") |
| args = parser.parse_args() |
|
|
| device = torch.device(args.device if torch.cuda.is_available() else "cpu") |
| torch.manual_seed(args.seed) |
|
|
| model, config = load_model(args.model_path, device) |
| logger.info(f"Loaded: {config.num_hidden_layers}L x {config.hidden_size}W, " |
| f"vocab={config.mask_vocab_size}") |
|
|
| if args.resume_from is None: |
| expand_embeddings(model, CHAT_VOCAB) |
| else: |
| logger.info(f"Resuming: keeping expanded vocab {config.mask_vocab_size}") |
|
|
| from transformers import AutoTokenizer |
| tokenizer = AutoTokenizer.from_pretrained(os.path.join(args.data_dir, "tokenizer")) |
| ensure_chat_tokens(tokenizer) |
| im_end = tokenizer.convert_tokens_to_ids("<|im_end|>") |
| assert im_end == 32002, ( |
| f"Tokenizer has im_end={im_end}, expected 32002. " |
| f"Data dir is stale (pre-fix ids): re-run prepare_data.py first." |
| ) |
| logger.info(f"Tokenizer vocab: {len(tokenizer)} | im_end={im_end}") |
|
|
| if args.bf16: |
| model = model.to(torch.bfloat16) |
| amp_dtype = torch.bfloat16 |
| else: |
| |
| |
| |
| amp_dtype = torch.float16 |
|
|
| train_ds = ChatDataset(os.path.join(args.data_dir, "train.pt"), tokenizer, |
| args.seq_len, seed=args.seed) |
| val_ds = ChatDataset(os.path.join(args.data_dir, "val.pt"), tokenizer, |
| args.seq_len, seed=args.seed) |
|
|
| transferred_names, new_names = set(), set() |
| for name, p in model.named_parameters(): |
| if any(k in name for k in ["timestep_emb", "timestep_residual", "lm_head", |
| "embed_tokens.weight"]): |
| new_names.add(name) |
| else: |
| transferred_names.add(name) |
| param_groups = [ |
| {"params": [p for n, p in model.named_parameters() if n in transferred_names], |
| "lr": args.lr * args.transferred_lr_mult, "name": "transferred"}, |
| {"params": [p for n, p in model.named_parameters() if n in new_names], |
| "lr": args.lr * args.new_lr_mult, "name": "new"}, |
| ] |
| for pg in param_groups: |
| logger.info(f" {pg['name']}: {sum(p.numel() for p in pg['params']):,} params, " |
| f"lr={pg['lr']:.2e}") |
|
|
| optimizer = AdamW(param_groups, weight_decay=args.weight_decay) |
|
|
| batch_size = args.batch_size |
| if batch_size <= 0 and torch.cuda.is_available(): |
| batch_size = detect_max_batch_size(model, args.seq_len, device, |
| args.keep_free_vram, amp_dtype) |
| if batch_size <= 0: |
| batch_size = 8 |
| eff_batch = batch_size * args.grad_accum_steps |
| steps_per_epoch = max(1, math.ceil(len(train_ds) / eff_batch)) |
| total_steps = args.max_steps if args.max_steps > 0 else steps_per_epoch * args.epochs |
| logger.info(f"batch={batch_size} accum={args.grad_accum_steps} " |
| f"eff={eff_batch} steps/epoch={steps_per_epoch} total={total_steps}") |
|
|
| |
| |
| scheduler = get_cosine_schedule_with_warmup( |
| optimizer, args.warmup_steps, total_steps, args.min_lr_ratio |
| ) |
|
|
| dataloader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, |
| num_workers=args.num_workers, pin_memory=True, |
| drop_last=False, collate_fn=collate_fn, |
| worker_init_fn=worker_init_fn) |
|
|
| global_step = 0 |
| if args.resume_from: |
| ckpt = torch.load(args.resume_from, map_location=device, weights_only=False) |
| model.load_state_dict(clean_state_dict(ckpt["model_state_dict"])) |
| optim_state = ckpt.get("optimizer_state_dict", {}) |
| if optim_state and "param_groups" in optim_state and "state" in optim_state: |
| try: |
| optimizer.load_state_dict(optim_state) |
| except (ValueError, KeyError) as e: |
| logger.warning(f"Optimizer state not loaded: {e}") |
| sched_state = ckpt.get("scheduler_state_dict", {}) |
| if sched_state and sched_state.get("last_epoch", 0) == ckpt.get("step", 0): |
| try: |
| scheduler.load_state_dict(sched_state) |
| except (ValueError, KeyError) as e: |
| logger.warning(f"Scheduler state not loaded: {e}") |
| global_step = ckpt.get("step", 0) |
| logger.info(f"Resumed from step {global_step}") |
|
|
| if not args.no_compile: |
| logger.info("Compiling model...") |
| model = torch.compile(model) |
|
|
| os.makedirs(args.output_dir, exist_ok=True) |
| stats_path = os.path.join(args.output_dir, "stats.jsonl") |
| stats_file = open(stats_path, "a") |
| with open(os.path.join(args.output_dir, "config.json"), "w") as f: |
| json.dump(asdict(model.config), f, indent=2, default=str) |
|
|
| scaler = torch.amp.GradScaler("cuda", enabled=not args.bf16) |
| free_disk = get_free_disk_space(args.output_dir) |
| if free_disk < args.disk_min_gb * 1e9: |
| stats_file.close() |
| raise RuntimeError(f"Insufficient disk space: {format_bytes(free_disk)}") |
|
|
| model.train() |
| optimizer.zero_grad() |
| loss_total, loss_count = 0.0, 0 |
| start_time = time.time() |
| last_log_time = start_time |
| data_iter = iter(dataloader) |
| epoch = 0 |
| best_val = float("inf") |
| no_improve = 0 |
| early_stopped = False |
|
|
| while global_step < total_steps: |
| if global_step % steps_per_epoch == 0 and global_step > 0: |
| epoch += 1 |
| try: |
| batch = next(data_iter) |
| except StopIteration: |
| epoch += 1 |
| data_iter = iter(dataloader) |
| batch = next(data_iter) |
|
|
| input_ids = batch["input_ids"].to(device) |
| labels = batch["labels"].to(device) |
| mask_positions = batch["mask_positions"].to(device) |
| timesteps = batch["timesteps"].to(device) |
| attention = batch["attention_mask"].to(device) |
| resp_len = batch["resp_len"].to(device) |
|
|
| with torch.autocast("cuda", dtype=amp_dtype): |
| logits = model(input_ids, timesteps, attention_mask=attention) |
|
|
| num_masked = mask_positions.sum().item() |
| if num_masked > 0: |
| |
| |
| ce = F.cross_entropy(logits.float()[mask_positions], |
| labels[mask_positions], reduction="none") |
| w = (1.0 / (timesteps * resp_len)).unsqueeze(1).expand_as(labels) |
| loss = (ce * w[mask_positions]).sum() / input_ids.shape[0] |
| scaler.scale(loss / args.grad_accum_steps).backward() |
| loss_total += loss.item() |
| else: |
| loss = torch.tensor(0.0, device=device) |
|
|
| global_step += 1 |
|
|
| if global_step % args.grad_accum_steps == 0: |
| scaler.unscale_(optimizer) |
| torch.nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm) |
| skipped = scaler.step(optimizer) |
| scaler.update() |
| if not skipped: |
| scheduler.step() |
| optimizer.zero_grad() |
|
|
| loss_count += 1 |
|
|
| if global_step % args.log_every == 0: |
| now = time.time() |
| elapsed = now - start_time |
| avg_loss = loss_total / max(1, loss_count) |
| ppl = math.exp(min(avg_loss, 20)) |
| lr = scheduler.get_last_lr()[0] |
| steps_per_sec = args.log_every / max(now - last_log_time, 1e-6) |
| eta = format_duration((total_steps - global_step) / steps_per_sec) |
| logger.info(f"Step {global_step:>6d} | epoch {epoch:.1f} | loss={avg_loss:.4f} | " |
| f"ppl={ppl:.1f} | lr={lr:.2e} | {steps_per_sec:.1f} steps/s | " |
| f"elapsed={elapsed:.0f}s | eta={eta}") |
| loss_total, loss_count = 0.0, 0 |
| last_log_time = now |
|
|
| stats_entry = {"step": global_step, "epoch": round(epoch, 2), |
| "loss": round(avg_loss, 4), "ppl": round(ppl, 1), "lr": lr} |
| stats_file.write(json.dumps(stats_entry) + "\n") |
| stats_file.flush() |
|
|
| |
| if global_step % args.val_every == 0: |
| val_loss, _ = evaluate(model, val_ds, batch_size, device, amp_dtype) |
| logger.info(f" val_loss={val_loss:.4f}") |
| stats_file.write(json.dumps({"step": global_step, |
| "val_loss": round(val_loss, 4)}) + "\n") |
| stats_file.flush() |
| if val_loss < best_val * (1.0 - args.min_delta): |
| best_val = val_loss |
| no_improve = 0 |
| ckpt_path = os.path.join(args.output_dir, "best.pt") |
| torch.save({"step": global_step, |
| "model_state_dict": clean_state_dict(model.state_dict()), |
| "optimizer_state_dict": optimizer.state_dict(), |
| "scheduler_state_dict": scheduler.state_dict(), |
| "config": asdict(model.config)}, ckpt_path) |
| logger.info(f" Best val loss, saved {ckpt_path}") |
| else: |
| no_improve += 1 |
| logger.info(f" No val improvement ({no_improve}/{args.patience} checks, " |
| f"best={best_val:.4f})") |
| if args.patience > 0 and no_improve >= args.patience: |
| logger.info(f"Early stopping at step {global_step}: no val loss " |
| f"improvement for {args.patience} checks " |
| f"(best={best_val:.4f})") |
| ckpt_path = os.path.join(args.output_dir, f"step_{global_step}.pt") |
| torch.save({"step": global_step, |
| "model_state_dict": clean_state_dict(model.state_dict()), |
| "optimizer_state_dict": optimizer.state_dict(), |
| "scheduler_state_dict": scheduler.state_dict(), |
| "config": asdict(model.config)}, ckpt_path) |
| logger.info(f"Saved final checkpoint: {ckpt_path}") |
| stats_file.write(json.dumps( |
| {**stats_entry, "best_val": round(best_val, 4), |
| "early_stopped": True}) + "\n") |
| stats_file.flush() |
| stats_file.close() |
| early_stopped = True |
| break |
|
|
| if global_step % args.save_every == 0: |
| ckpt_path = os.path.join(args.output_dir, f"step_{global_step}.pt") |
| torch.save({"step": global_step, |
| "model_state_dict": clean_state_dict(model.state_dict()), |
| "optimizer_state_dict": optimizer.state_dict(), |
| "scheduler_state_dict": scheduler.state_dict(), |
| "config": asdict(model.config)}, ckpt_path) |
| logger.info(f"Saved checkpoint: {ckpt_path}") |
| cleanup_checkpoints(args.output_dir, args.keep_first_n, |
| args.keep_last_n, args.keep_best_n, stats_path) |
| if get_free_disk_space(args.output_dir) < args.disk_min_gb * 1e9: |
| logger.warning("Low disk after save; stopping") |
| stats_file.close() |
| return |
|
|
| stats_file.close() |
| if early_stopped: |
| logger.info(f"Early stopping triggered; best val loss {best_val:.4f} " |
| f"saved as best.pt") |
| else: |
| logger.info(f"Training complete at step {global_step}") |
|
|
| if args.export_dir: |
| from export_hf import export |
| export(os.path.join(args.output_dir, f"step_{global_step}.pt") |
| if not os.path.exists(os.path.join(args.output_dir, "best.pt")) |
| else os.path.join(args.output_dir, "best.pt"), |
| os.path.join(args.data_dir, "tokenizer"), |
| args.export_dir) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|