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