| """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() |
|
|