ViClickbait-2025 / scripts /train_gpt_oss_20b.py
minhy112's picture
Upload ViClickbait-2025 project (PhoBERT + GPT-OSS-20B LoRA SFT)
af55aed verified
Raw
History Blame Contribute Delete
15.8 kB
"""
train_gpt_oss_20b.py โ€“ Generative SFT fine-tune openai/gpt-oss-20b
----------------------------------------------------------------------
Phฦฐฦกng phรกp: Supervised Fine-Tuning (SFT) vแป›i causal LM loss
- KHร”NG dรนng classification head (random init โ†’ khรดng hแป™i tแปฅ)
- Fine-tune model SINH ra "clickbait" hoแบทc "khรดng clickbait"
- ฤรกnh giรก: so sรกnh log-probability cแปงa 2 response
Tแบกi sao approach cลฉ (SeqCls) thแบฅt bแบกi:
1. score.weight random init โ†’ cแบงn LR riรชng, vแบซn kรฉm
2. Causal LM dรนng last token rep โ†’ yแบฟu hฦกn bidirectional encoder
3. Dataset nhแป 2.7K โ†’ khรดng ฤ‘แปง ฤ‘แปƒ train head tแปซ ฤ‘แบงu
Tแบกi sao SFT tแป‘t hฦกn:
1. Dรนng lm_head ฤ‘รฃ pretrain โ†’ khรดng random init
2. Aligned vแป›i pretraining objective โ†’ hแป™i tแปฅ nhanh
3. Tแบญn dแปฅng khแบฃ nฤƒng hiแปƒu ngรดn ngแปฏ cแปงa 20B model
Dรนng:
python3 scripts/train_gpt_oss_20b.py \\
--data_dir data/splits --output_dir outputs/gpt20b \\
--epochs 10 --batch_size 4 --grad_accum 8 \\
--lr 2e-5 --lora_r 32 --lora_alpha 64
"""
import argparse
import json
import logging
import os
import warnings
from dataclasses import dataclass
from typing import Dict, List
# โ”€โ”€ Suppress warnings โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
warnings.filterwarnings("ignore")
logging.getLogger("transformers").setLevel(logging.ERROR)
logging.getLogger("datasets").setLevel(logging.ERROR)
logging.getLogger("huggingface_hub").setLevel(logging.ERROR)
import numpy as np
import pandas as pd
import torch
from datasets import Dataset
from sklearn.metrics import (
accuracy_score,
classification_report,
f1_score,
precision_score,
recall_score,
)
from tqdm.auto import tqdm
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
EarlyStoppingCallback,
Trainer,
TrainingArguments,
set_seed,
)
import transformers
transformers.logging.set_verbosity_error()
# โ”€โ”€ Constants โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
PROMPT_TEMPLATE = (
"Phรขn loแบกi bร i viแบฟt sau lร  clickbait hay khรดng.\n\n"
"Bร i viแบฟt: {text}\n\n"
"Nhรฃn:"
)
LABEL_TEXT = {0: " khรดng clickbait", 1: " clickbait"}
# โ”€โ”€ Data Collator โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
@dataclass
class PromptResponseCollator:
"""Left-pad sequences cho causal LM. Labels nhแบญn -100 แปŸ padding."""
pad_token_id: int
def __call__(self, features: List[Dict]) -> Dict[str, torch.Tensor]:
max_len = max(len(f["input_ids"]) for f in features)
batch_ids, batch_mask, batch_labels = [], [], []
for f in features:
pad_len = max_len - len(f["input_ids"])
batch_ids.append([self.pad_token_id] * pad_len + f["input_ids"])
batch_mask.append([0] * pad_len + f["attention_mask"])
batch_labels.append([-100] * pad_len + f["labels"])
return {
"input_ids": torch.tensor(batch_ids, dtype=torch.long),
"attention_mask": torch.tensor(batch_mask, dtype=torch.long),
"labels": torch.tensor(batch_labels, dtype=torch.long),
}
# โ”€โ”€ Helpers โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
def load_splits(data_dir: str):
train_df = pd.read_csv(os.path.join(data_dir, "train.csv"))
val_df = pd.read_csv(os.path.join(data_dir, "val.csv"))
test_df = pd.read_csv(os.path.join(data_dir, "test.csv"))
return train_df, val_df, test_df
def format_and_tokenize(df: pd.DataFrame, tokenizer, max_length: int) -> List[Dict]:
"""Tแบกo prompt + response cho mแป—i mแบซu, tokenize, mask label trรชn prompt."""
records = []
for _, row in df.iterrows():
prompt_text = PROMPT_TEMPLATE.format(text=row["text"])
response_text = LABEL_TEXT[int(row["label_id"])]
full_text = prompt_text + response_text
prompt_ids = tokenizer.encode(prompt_text, add_special_tokens=False)
full_ids = tokenizer.encode(full_text, add_special_tokens=False)
# Truncate nแบฟu quรก dร i (giแปฏ response nguyรชn, cแบฏt prompt)
if len(full_ids) > max_length:
resp_ids = tokenizer.encode(response_text, add_special_tokens=False)
prompt_ids = prompt_ids[: max_length - len(resp_ids)]
full_ids = prompt_ids + resp_ids
# Labels: -100 cho prompt tokens, real IDs cho response tokens
labels = [-100] * len(prompt_ids) + full_ids[len(prompt_ids):]
assert len(labels) == len(full_ids), f"len mismatch: {len(labels)} vs {len(full_ids)}"
records.append({
"input_ids": full_ids,
"attention_mask": [1] * len(full_ids),
"labels": labels,
})
return records
def predict_by_scoring(
model, tokenizer, texts: List[str], max_length: int = 512
) -> List[int]:
"""Dแปฑ ฤ‘oรกn bแบฑng so sรกnh log-probability cแปงa 2 response.
Cho mแป—i text:
1. Tแบกo full sequence: prompt + " clickbait" โ†’ tรญnh log P(response | prompt)
2. Tแบกo full sequence: prompt + " khรดng clickbait" โ†’ tรญnh log P(response | prompt)
3. Chแปn label cรณ normalized log-prob cao hฦกn
"""
model.eval()
predictions = []
for text in tqdm(texts, desc="ฤรกnh giรก"):
prompt_text = PROMPT_TEMPLATE.format(text=text)
prompt_ids = tokenizer.encode(prompt_text, add_special_tokens=False)
scores = {}
for label_id, label_text in LABEL_TEXT.items():
full_text = prompt_text + label_text
full_ids = tokenizer.encode(full_text, add_special_tokens=False)
# Truncate nแบฟu cแบงn
if len(full_ids) > max_length:
resp_ids = tokenizer.encode(label_text, add_special_tokens=False)
p_ids = prompt_ids[: max_length - len(resp_ids)]
full_ids = p_ids + resp_ids
resp_start = len(p_ids)
else:
resp_start = len(prompt_ids)
input_ids = torch.tensor([full_ids], device=model.device)
with torch.no_grad():
logits = model(input_ids).logits # (1, seq_len, vocab_size)
# Log-probability trung bรฌnh cแปงa response tokens
log_probs = torch.nn.functional.log_softmax(logits[0], dim=-1)
total_lp = 0.0
n_tokens = len(full_ids) - resp_start
for i in range(resp_start, len(full_ids)):
total_lp += log_probs[i - 1, full_ids[i]].item()
scores[label_id] = total_lp / max(n_tokens, 1) # normalize by length
predictions.append(max(scores, key=scores.get))
return predictions
# โ”€โ”€ Main โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
def main() -> None:
parser = argparse.ArgumentParser(
description="Generative SFT fine-tune GPT-OSS-20B cho clickbait detection"
)
parser.add_argument("--data_dir", default="data/splits")
parser.add_argument("--output_dir", default="outputs/gpt20b")
parser.add_argument("--model_name", default="openai/gpt-oss-20b")
parser.add_argument("--max_length", type=int, default=256)
parser.add_argument("--batch_size", type=int, default=2)
parser.add_argument("--grad_accum", type=int, default=16)
parser.add_argument("--lr", type=float, default=2e-5)
parser.add_argument("--epochs", type=int, default=10)
parser.add_argument("--warmup_ratio", type=float, default=0.1)
parser.add_argument("--weight_decay", type=float, default=0.01)
parser.add_argument("--patience", type=int, default=3)
parser.add_argument("--seed", type=int, default=42)
# LoRA
parser.add_argument("--lora_r", type=int, default=32)
parser.add_argument("--lora_alpha", type=int, default=64)
parser.add_argument("--lora_dropout", type=float, default=0.05)
args = parser.parse_args()
set_seed(args.seed)
os.makedirs(args.output_dir, exist_ok=True)
from peft import LoraConfig, TaskType, get_peft_model
# โ”€โ”€ 1. Load data โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
train_df, val_df, test_df = load_splits(args.data_dir)
print(f"Train={len(train_df)} Val={len(val_df)} Test={len(test_df)}")
# โ”€โ”€ 2. Tokenizer โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
print(f"\nNแบกp tokenizer: {args.model_name}")
tokenizer = AutoTokenizer.from_pretrained(args.model_name, use_fast=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.pad_token_id = tokenizer.eos_token_id
tokenizer.padding_side = "left"
# โ”€โ”€ 3. Format & tokenize โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
print("Tokenizing data...")
train_records = format_and_tokenize(train_df, tokenizer, args.max_length)
val_records = format_and_tokenize(val_df, tokenizer, args.max_length)
train_ds = Dataset.from_list(train_records)
val_ds = Dataset.from_list(val_records)
collator = PromptResponseCollator(pad_token_id=tokenizer.pad_token_id)
# โ”€โ”€ 4. Model โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
print(f"Nแบกp model: {args.model_name}")
model = AutoModelForCausalLM.from_pretrained(
args.model_name,
dtype=torch.bfloat16,
device_map="auto",
)
# Sync token configs
for cfg in filter(None, [model.config, getattr(model, "generation_config", None)]):
cfg.pad_token_id = tokenizer.pad_token_id
cfg.eos_token_id = tokenizer.eos_token_id
if tokenizer.bos_token_id is not None:
cfg.bos_token_id = tokenizer.bos_token_id
# โ”€โ”€ 5. LoRA โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM, # โ† dรนng CausalLM, KHร”NG phแบฃi SEQ_CLS
r=args.lora_r,
lora_alpha=args.lora_alpha,
lora_dropout=args.lora_dropout,
bias="none",
target_modules=["q_proj", "v_proj"], # attention only, bแป MoE FFN
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# โ”€โ”€ 6. Training args โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
steps_per_epoch = max(1, len(train_ds) // (args.batch_size * args.grad_accum))
warmup_steps = int(args.warmup_ratio * steps_per_epoch * args.epochs)
training_args = TrainingArguments(
output_dir=args.output_dir,
per_device_train_batch_size=args.batch_size,
per_device_eval_batch_size=args.batch_size,
gradient_accumulation_steps=args.grad_accum,
learning_rate=args.lr,
num_train_epochs=args.epochs,
warmup_steps=warmup_steps,
weight_decay=args.weight_decay,
lr_scheduler_type="cosine",
max_grad_norm=0.3,
bf16=torch.cuda.is_bf16_supported(),
fp16=not torch.cuda.is_bf16_supported() and torch.cuda.is_available(),
eval_strategy="epoch",
save_strategy="epoch",
load_best_model_at_end=True,
metric_for_best_model="eval_loss", # causal LM loss โ†’ lower is better
greater_is_better=False,
save_total_limit=2,
logging_steps=20,
report_to="none",
seed=args.seed,
dataloader_num_workers=2,
gradient_checkpointing=True,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_ds,
eval_dataset=val_ds,
data_collator=collator,
processing_class=tokenizer,
callbacks=[EarlyStoppingCallback(early_stopping_patience=args.patience)],
)
# โ”€โ”€ 7. Train โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
print("\n" + "=" * 60)
print(f"BแบฎT ฤแบฆU HUแบคN LUYแป†N (Generative SFT)")
print(f" Model: {args.model_name}")
print(f" Method: SFT (causal LM loss on response tokens only)")
print(f" Effective batch size: {args.batch_size * args.grad_accum}")
print(f" LoRA r={args.lora_r}, alpha={args.lora_alpha}")
print("=" * 60)
trainer.train()
# โ”€โ”€ 8. Evaluate on Test set โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
print("\n" + "=" * 60)
print("ฤรNH GIร TRรŠN TEST SET (log-probability scoring)")
print("=" * 60)
preds = predict_by_scoring(model, tokenizer, test_df["text"].tolist(), args.max_length)
labels = test_df["label_id"].tolist()
report = classification_report(
labels, preds,
target_names=["non-clickbait", "clickbait"],
digits=4,
)
print(report)
results = {
"model": args.model_name,
"method": "generative_sft",
"lora_r": args.lora_r,
"lora_alpha": args.lora_alpha,
"lr": args.lr,
"accuracy": float(accuracy_score(labels, preds)),
"f1_binary": float(f1_score(labels, preds, average="binary", zero_division=0)),
"f1_macro": float(f1_score(labels, preds, average="macro", zero_division=0)),
"precision": float(precision_score(labels, preds, average="binary", zero_division=0)),
"recall": float(recall_score(labels, preds, average="binary", zero_division=0)),
"classification_report": report,
"split_sizes": {"train": len(train_df), "val": len(val_df), "test": len(test_df)},
}
out_path = os.path.join(args.output_dir, "test_results.json")
with open(out_path, "w", encoding="utf-8") as f:
json.dump(results, f, ensure_ascii=False, indent=2)
print(f"\nKแบฟt quแบฃ lฦฐu tแบกi: {out_path}")
print("\nโ”€โ”€ Tรณm tแบฏt โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€")
print(f" Accuracy : {results['accuracy']:.4f}")
print(f" F1 Binary : {results['f1_binary']:.4f}")
print(f" F1 Macro : {results['f1_macro']:.4f}")
print(f" Precision : {results['precision']:.4f}")
print(f" Recall : {results['recall']:.4f}")
if __name__ == "__main__":
main()