"""LoRA finetune of a causal LM for stance classification. Each row becomes a short chat (system instruction, target+tweet, label as the reply), and we only backprop through the label tokens. Training across all targets together helps it generalize to targets it hasn't seen. python -m src.llm_finetune --train_csv data/track1/train.csv \\ --base_model ALLaM-AI/ALLaM-7B-Instruct-preview \\ --out_dir outputs/allam_t1 """ import argparse import json import os import random import numpy as np import torch from peft import LoraConfig, PeftModel, get_peft_model from torch.utils.data import DataLoader, Dataset from transformers import AutoModelForCausalLM, AutoTokenizer from src.data import load_split SYSTEM = ( "أنت مصنف موقف عربي دقيق. حدد موقف كاتب التغريدة تجاه الهدف المحدد. " "الموقف واحد من ثلاثة فقط: Favor أو Against أو None." ) def user_text(target, tweet): return f"الهدف: {target}\nالتغريدة: {tweet}\nالموقف:" class SFTDataset(Dataset): def __init__(self, df, tok, max_len): self.rows = df.to_dict("records") self.tok = tok self.max_len = max_len def __len__(self): return len(self.rows) def __getitem__(self, i): row = self.rows[i] msgs = [ {"role": "system", "content": SYSTEM}, {"role": "user", "content": user_text(row["target"], row["text"])}, ] prompt = self.tok.apply_chat_template( msgs, tokenize=False, add_generation_prompt=True ) full = prompt + " " + row["stance"] + self.tok.eos_token p_ids = self.tok(prompt, add_special_tokens=False)["input_ids"] f_ids = self.tok(full, add_special_tokens=False)["input_ids"] f_ids = f_ids[:self.max_len] labels = list(f_ids) for j in range(min(len(p_ids), len(labels))): labels[j] = -100 return {"input_ids": f_ids, "labels": labels} def collate(batch, pad_id): m = max(len(b["input_ids"]) for b in batch) ids, labs, att = [], [], [] for b in batch: n = m - len(b["input_ids"]) ids.append(b["input_ids"] + [pad_id] * n) labs.append(b["labels"] + [-100] * n) att.append([1] * len(b["input_ids"]) + [0] * n) return ( torch.tensor(ids), torch.tensor(labs), torch.tensor(att), ) def main(): ap = argparse.ArgumentParser() ap.add_argument("--train_csv", required=True) ap.add_argument("--base_model", required=True) ap.add_argument("--out_dir", required=True) ap.add_argument("--exclude_target", default=None, help="hold out a target (leave-one-out experiments)") ap.add_argument("--epochs", type=int, default=3) ap.add_argument("--lr", type=float, default=1e-4) ap.add_argument("--batch_size", type=int, default=8) ap.add_argument("--save_every", type=int, default=40) ap.add_argument("--max_len", type=int, default=192) ap.add_argument("--lora_r", type=int, default=16) ap.add_argument("--seed", type=int, default=42) args = ap.parse_args() random.seed(args.seed) np.random.seed(args.seed) torch.manual_seed(args.seed) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") df = load_split(args.train_csv, "preserve", has_labels=True) if args.exclude_target: df = df[df["target"] != args.exclude_target].reset_index(drop=True) print(f"[train] {len(df)} ex, targets={sorted(df['target'].unique())}") tok = AutoTokenizer.from_pretrained(args.base_model) if tok.pad_token is None: tok.pad_token = tok.eos_token model = AutoModelForCausalLM.from_pretrained( args.base_model, torch_dtype=torch.bfloat16 ).to(device) model.config.use_cache = False prog_path = os.path.join(args.out_dir, "progress.json") done_step = 0 resume = (os.path.exists(prog_path) and os.path.exists(os.path.join(args.out_dir, "adapter_config.json"))) if resume: done_step = json.load(open(prog_path)).get("global_step", 0) model = PeftModel.from_pretrained( model, args.out_dir, is_trainable=True ) print(f"[resume] loaded adapter at global_step={done_step}", flush=True) else: lora = LoraConfig( r=args.lora_r, lora_alpha=2 * args.lora_r, lora_dropout=0.05, bias="none", task_type="CAUSAL_LM", target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"], ) model = get_peft_model(model, lora) model.print_trainable_parameters() ds = SFTDataset(df, tok, args.max_len) loader = DataLoader( ds, batch_size=args.batch_size, shuffle=True, collate_fn=lambda b: collate(b, tok.pad_token_id), ) optim = torch.optim.AdamW( [p for p in model.parameters() if p.requires_grad], lr=args.lr ) optim_path = os.path.join(args.out_dir, "optim.pt") if resume and os.path.exists(optim_path): optim.load_state_dict(torch.load(optim_path, map_location=device)) print("[resume] restored optimizer state", flush=True) os.makedirs(args.out_dir, exist_ok=True) max_steps = args.epochs * len(loader) if done_step >= max_steps: print(f"[skip] already trained {done_step}/{max_steps} steps", flush=True) return def checkpoint(step): model.save_pretrained(args.out_dir) tok.save_pretrained(args.out_dir) torch.save(optim.state_dict(), optim_path) json.dump({"global_step": step, "max_steps": max_steps}, open(prog_path, "w")) model.train() gstep = done_step running, rn = 0.0, 0 while gstep < max_steps: for ids, labs, att in loader: if gstep >= max_steps: break ids, labs, att = ids.to(device), labs.to(device), att.to(device) out = model(input_ids=ids, attention_mask=att, labels=labs) out.loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optim.step() optim.zero_grad() gstep += 1 running += out.loss.item() rn += 1 if gstep % 25 == 0: print(f"step {gstep}/{max_steps} loss={running / rn:.4f}", flush=True) if gstep % args.save_every == 0: checkpoint(gstep) print(f"[ckpt] saved at step {gstep}", flush=True) checkpoint(max_steps) print(f"[saved] adapter -> {args.out_dir} ({max_steps} steps)", flush=True) if __name__ == "__main__": main()