"""Fine-tune T5-small as a text humanizer (one language per run). English uses ``google-t5/t5-small``; Chinese uses a Chinese T5-small (``uer/t5-small-chinese-cluecorpussmall``). Supervised pairs come from ``scripts/build_dataset.py`` (rule-generated), filtered to one language. MPS notes: avoid per-step host sync; accumulate the loss tensor and sync at evaluation. Save a checkpoint at the end of every epoch. """ import argparse import json import time from pathlib import Path import torch from torch.optim import AdamW from torch.utils.data import DataLoader, Dataset from transformers import BertTokenizer, T5ForConditionalGeneration, T5Tokenizer BASE_MODELS = { "en": "google-t5/t5-small", "zh": "uer/t5-small-chinese-cluecorpussmall", } class PairsDataset(Dataset): def __init__(self, path: Path, tokenizer: T5Tokenizer, max_len: int, lang: str): self.pairs = [] with open(path, encoding="utf-8") as f: for line in f: row = json.loads(line) if row["lang"] != lang: continue self.pairs.append((row["input_text"], row["output_text"])) self.tokenizer = tokenizer self.max_len = max_len def __len__(self): return len(self.pairs) def __getitem__(self, idx): src, tgt = self.pairs[idx] enc = self.tokenizer( src, max_length=self.max_len, padding="max_length", truncation=True ) dec = self.tokenizer( tgt, max_length=self.max_len, padding="max_length", truncation=True ) labels = torch.tensor(dec["input_ids"]) labels[labels == self.tokenizer.pad_token_id] = -100 return { "input_ids": torch.tensor(enc["input_ids"]), "attention_mask": torch.tensor(enc["attention_mask"]), "labels": labels, } def decode(ids, tokenizer): ids = torch.where( (ids == -100) | (ids == tokenizer.pad_token_id), torch.tensor(tokenizer.pad_token_id), ids, ) return tokenizer.decode(ids, skip_special_tokens=True) def main(): parser = argparse.ArgumentParser() parser.add_argument("--lang", choices=["en", "zh"], default="en") parser.add_argument("--data", default="data") parser.add_argument("--out", default=None) parser.add_argument("--base-model", default=None) parser.add_argument("--max-len", type=int, default=128) parser.add_argument("--batch-size", type=int, default=32) parser.add_argument("--lr", type=float, default=3e-4) parser.add_argument("--epochs", type=int, default=3) parser.add_argument("--warmup-steps", type=int, default=100) parser.add_argument("--eval-every", type=int, default=200) parser.add_argument("--seed", type=int, default=42) args = parser.parse_args() base_model = args.base_model or BASE_MODELS[args.lang] out_dir = Path(args.out or f"checkpoints/humanize-text-model/{args.lang}") torch.manual_seed(args.seed) device = "mps" if torch.backends.mps.is_available() else ( "cuda" if torch.cuda.is_available() else "cpu" ) print(f"lang={args.lang} base={base_model} device={device}", flush=True) tokenizer = ( BertTokenizer.from_pretrained(base_model) if args.lang == "zh" else T5Tokenizer.from_pretrained(base_model) ) if (out_dir / "pytorch_model.bin").exists() or (out_dir / "model.safetensors").exists(): print(f"resuming from existing checkpoint in {out_dir}", flush=True) model = T5ForConditionalGeneration.from_pretrained(out_dir) else: model = T5ForConditionalGeneration.from_pretrained(base_model) model.to(device) train_ds = PairsDataset(Path(args.data) / "train.jsonl", tokenizer, args.max_len, args.lang) val_ds = PairsDataset(Path(args.data) / "val.jsonl", tokenizer, args.max_len, args.lang) val_samples = [ json.loads(l) for l in open(Path(args.data) / "val.jsonl", encoding="utf-8") if json.loads(l)["lang"] == args.lang ][:3] train_loader = DataLoader(train_ds, batch_size=args.batch_size, shuffle=True, num_workers=0) optimizer = AdamW(model.parameters(), lr=args.lr) total_steps = args.epochs * len(train_loader) scheduler = torch.optim.lr_scheduler.LinearLR( optimizer, start_factor=1 / (args.warmup_steps + 1), total_iters=args.warmup_steps ) out_dir.mkdir(parents=True, exist_ok=True) step = 0 t0 = time.time() for epoch in range(1, args.epochs + 1): model.train() epoch_loss = torch.zeros(()) for batch in train_loader: batch = {k: v.to(device) for k, v in batch.items()} loss = model(**batch).loss loss.backward() optimizer.step() scheduler.step() optimizer.zero_grad() epoch_loss = epoch_loss + loss.detach() step += 1 if step % 100 == 0: print(f"[hb] step {step}/{total_steps} ({time.time() - t0:.0f}s)", flush=True) if step % args.eval_every == 0: model.eval() vloss = 0.0 with torch.no_grad(): for i in range(16): vb = val_ds[i] vb = {k: v.unsqueeze(0).to(device) for k, v in vb.items()} vloss += model(**vb).loss.item() avg_train = (epoch_loss / step).item() model.train() print( f"[step {step}] epoch={epoch} train_loss={avg_train:.4f} " f"val_loss={vloss / 16:.4f} ({time.time() - t0:.0f}s)", flush=True, ) sample = val_samples[0] model.eval() with torch.no_grad(): ids = model.generate( input_ids=tokenizer(sample["input_text"], return_tensors="pt").input_ids.to(device), max_length=args.max_len, ) out = decode(ids[0], tokenizer) model.train() print(" IN :", sample["input_text"][:90]) print(" OUT:", out[:90], flush=True) avg = (epoch_loss / len(train_loader)).item() print(f"epoch {epoch} done: avg_loss={avg:.4f}", flush=True) model.save_pretrained(out_dir) tokenizer.save_pretrained(out_dir) print(f"checkpoint saved to {out_dir}", flush=True) print(f"final model saved to {out_dir}", flush=True) if __name__ == "__main__": main()