Danny
Bilingual humanize-text-model (EN T5-small + ZH Chinese T5-small)
7cb8aac verified
Raw
History Blame Contribute Delete
6.61 kB
"""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()