stanceeval2026 / code /src /llm_finetune.py
zaher-m's picture
Add files using upload-large-folder tool
7e9cfd1 verified
Raw
History Blame Contribute Delete
6.84 kB
"""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()