| """ |
| 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 |
|
|
| |
| 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() |
|
|
|
|
| |
|
|
| 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"} |
|
|
|
|
| |
|
|
| @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), |
| } |
|
|
|
|
| |
|
|
| 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) |
|
|
| |
| 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] * 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) |
|
|
| |
| 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 |
|
|
| |
| 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) |
|
|
| predictions.append(max(scores, key=scores.get)) |
|
|
| return predictions |
|
|
|
|
| |
|
|
| 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) |
| |
| 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 |
|
|
| |
| 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)}") |
|
|
| |
| 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" |
|
|
| |
| 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) |
|
|
| |
| print(f"Nแบกp model: {args.model_name}") |
| model = AutoModelForCausalLM.from_pretrained( |
| args.model_name, |
| dtype=torch.bfloat16, |
| device_map="auto", |
| ) |
|
|
| |
| 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 |
|
|
| |
| lora_config = LoraConfig( |
| task_type=TaskType.CAUSAL_LM, |
| r=args.lora_r, |
| lora_alpha=args.lora_alpha, |
| lora_dropout=args.lora_dropout, |
| bias="none", |
| target_modules=["q_proj", "v_proj"], |
| ) |
| model = get_peft_model(model, lora_config) |
| model.print_trainable_parameters() |
|
|
| |
| 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", |
| 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)], |
| ) |
|
|
| |
| 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() |
|
|
| |
| 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() |
|
|