SerendipLLM-news-classifier / scripts /train_classifier.py
Chamaka8's picture
Upload scripts/train_classifier.py with huggingface_hub
d7add55 verified
Raw
History Blame Contribute Delete
9.4 kB
import os, torch, random, json
os.environ['HF_HOME'] = '/workspace/hf_cache'
os.environ['HF_HUB_DISABLE_XET'] = '1'
from datasets import load_dataset
from transformers import (PreTrainedTokenizerFast, AutoModelForCausalLM, BitsAndBytesConfig)
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training, TaskType
from torch import nn
from torch.utils.data import Dataset, DataLoader
from sklearn.metrics import precision_recall_fscore_support, confusion_matrix
from torch.optim import AdamW
from transformers import get_cosine_schedule_with_warmup
TOKEN = 'YOUR_HF_TOKEN'
BASE = 'Chamaka8/Serendip-LLM-CPT-SFT-v2'
NEW_REPO = 'Chamaka8/SerendipLLM-news-classifier'
LABELS = ['ව්‍යාපාර', 'දේශපාලන', 'විනෝදාස්වාදය', 'ක්‍රීඩා', 'තාක්ෂණ']
SINLLAMA = {'P': 89.033, 'R': 86.787, 'F1': 86.402}
NUM_LABELS = len(LABELS)
DEVICE = 'cuda:0'
print('='*65)
print('SerendipLLM News Classifier (Classification Head)')
print('='*65)
# Dataset
print('\nLoading dataset...')
ds = load_dataset('NLPC-UOM/Sinhala-News-Category-classification')['train']
by_class = {i: [] for i in range(NUM_LABELS)}
for item in ds:
if isinstance(item['labels'], int) and item['labels'] < NUM_LABELS:
by_class[item['labels']].append(item)
print('Class sizes:', {LABELS[i]: len(by_class[i]) for i in range(NUM_LABELS)})
train_data, eval_data = [], []
for cls_idx in range(NUM_LABELS):
items = by_class[cls_idx].copy()
random.seed(42)
random.shuffle(items)
eval_data.extend(items[:40])
train_data.extend(items[40:])
# Oversample to balance
max_count = max(len([x for x in train_data if x['labels']==i]) for i in range(NUM_LABELS))
balanced_train = []
for cls_idx in range(NUM_LABELS):
cls_items = [x for x in train_data if x['labels']==cls_idx]
while len(cls_items) < max_count:
cls_items += cls_items
balanced_train.extend(cls_items[:max_count])
random.shuffle(balanced_train)
print(f'Train (balanced): {len(balanced_train)} | Eval: {len(eval_data)}')
# Model
print('\nLoading model...')
bnb_config = BitsAndBytesConfig(
load_in_4bit=True, bnb_4bit_quant_type='nf4',
bnb_4bit_compute_dtype=torch.float16, bnb_4bit_use_double_quant=True)
tokenizer = PreTrainedTokenizerFast.from_pretrained(BASE, token=TOKEN)
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = 'right'
base_model = AutoModelForCausalLM.from_pretrained(
BASE, token=TOKEN, quantization_config=bnb_config, device_map={'': DEVICE})
base_model = prepare_model_for_kbit_training(base_model)
lora_config = LoraConfig(
r=32, lora_alpha=64,
target_modules=['q_proj','k_proj','v_proj','o_proj','gate_proj','up_proj','down_proj'],
lora_dropout=0.05, bias='none', task_type=TaskType.CAUSAL_LM)
base_model = get_peft_model(base_model, lora_config)
base_model.print_trainable_parameters()
hidden_size = base_model.config.hidden_size
print(f'Hidden size: {hidden_size}')
# Classifier head — everything on DEVICE
classifier = nn.Linear(hidden_size, NUM_LABELS).to(DEVICE).float()
dropout = nn.Dropout(0.1).to(DEVICE)
def forward(input_ids, attention_mask, labels=None):
input_ids = input_ids.to(DEVICE)
attention_mask = attention_mask.to(DEVICE)
with torch.cuda.amp.autocast():
outputs = base_model(input_ids=input_ids, attention_mask=attention_mask,
output_hidden_states=True)
hidden = outputs.hidden_states[-1] # (B, seq, hidden)
# Last non-pad token
seq_len = attention_mask.sum(dim=1) - 1
last_hidden = hidden[torch.arange(hidden.size(0)), seq_len].float().to(DEVICE)
logits = classifier(dropout(last_hidden))
loss = None
if labels is not None:
loss = nn.CrossEntropyLoss()(logits, labels.to(DEVICE))
return loss, logits
# Dataset class
class NewsDataset(Dataset):
def __init__(self, data, tokenizer, max_len=256):
self.data = data
self.tokenizer = tokenizer
self.max_len = max_len
def __len__(self): return len(self.data)
def __getitem__(self, idx):
item = self.data[idx]
text = f'ප්‍රවෘත්ති: {str(item["comments"])[:400]}'
enc = self.tokenizer(text, truncation=True, max_length=self.max_len,
padding='max_length', return_tensors='pt')
return {
'input_ids': enc['input_ids'].squeeze(),
'attention_mask': enc['attention_mask'].squeeze(),
'labels': torch.tensor(item['labels'], dtype=torch.long)
}
train_dataset = NewsDataset(balanced_train, tokenizer)
eval_dataset = NewsDataset(eval_data, tokenizer)
train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, num_workers=2)
eval_loader = DataLoader(eval_dataset, batch_size=8, shuffle=False)
# Optimizer
lora_params = [p for n,p in base_model.named_parameters() if 'lora' in n and p.requires_grad]
optimizer = AdamW([
{'params': classifier.parameters(), 'lr': 2e-4},
{'params': lora_params, 'lr': 5e-5},
], weight_decay=0.01)
EPOCHS = 15
total_steps = len(train_loader) * EPOCHS
scheduler = get_cosine_schedule_with_warmup(
optimizer, num_warmup_steps=total_steps//10, num_training_steps=total_steps)
def evaluate():
base_model.eval()
classifier.eval()
true_l, pred_l = [], []
with torch.no_grad():
for batch in eval_loader:
_, logits = forward(batch['input_ids'], batch['attention_mask'])
preds = logits.argmax(dim=-1).cpu().tolist()
true_l.extend([LABELS[l] for l in batch['labels'].tolist()])
pred_l.extend([LABELS[p] for p in preds])
_, _, f1, _ = precision_recall_fscore_support(true_l, pred_l, average='macro', zero_division=0)
acc = sum(t==p for t,p in zip(true_l,pred_l))/len(true_l)
return round(f1*100,3), round(acc*100,1), true_l, pred_l
# Training
print('\nTraining...')
best_f1, best_classifier_state, best_lora_state = 0, None, None
for epoch in range(EPOCHS):
base_model.train()
classifier.train()
total_loss = 0
for batch in train_loader:
optimizer.zero_grad()
loss, _ = forward(batch['input_ids'], batch['attention_mask'], batch['labels'])
loss.backward()
torch.nn.utils.clip_grad_norm_(list(classifier.parameters()) + lora_params, 1.0)
optimizer.step()
scheduler.step()
total_loss += loss.item()
f1, acc, _, _ = evaluate()
avg_loss = round(total_loss/len(train_loader), 4)
print(f'Epoch {epoch+1:>2}/{EPOCHS} — loss: {avg_loss} | F1: {f1} | Acc: {acc}%',
'✓ BEST' if f1 > best_f1 else '')
if f1 > best_f1:
best_f1 = f1
best_classifier_state = {k: v.clone().cpu() for k, v in classifier.state_dict().items()}
best_lora_state = {k: v.clone().cpu() for k, v in base_model.state_dict().items()
if 'lora' in k}
# Load best
print(f'\nBest F1: {best_f1} — loading best checkpoint...')
classifier.load_state_dict({k: v.to(DEVICE) for k, v in best_classifier_state.items()})
f1, acc, true_l, pred_l = evaluate()
p_s, r_s, f1_s, _ = precision_recall_fscore_support(true_l, pred_l, average='macro', zero_division=0)
p_per, r_per, f1_per, sup = precision_recall_fscore_support(true_l, pred_l, labels=LABELS, zero_division=0)
our = {'P': round(p_s*100,3), 'R': round(r_s*100,3), 'F1': round(f1_s*100,3)}
print('\n' + '='*65)
print('PER-CLASS RESULTS')
print(f'{"Label":<22} {"P":>8} {"R":>8} {"F1":>8} {"Sup":>6}')
print('-'*65)
for i, label in enumerate(LABELS):
print(f'{label:<22} {round(p_per[i]*100,1):>8} {round(r_per[i]*100,1):>8} {round(f1_per[i]*100,1):>8} {sup[i]:>6}')
cm = confusion_matrix(true_l, pred_l, labels=LABELS)
print('\nConfusion Matrix:')
print(f'{"":>22}', end='')
for l in LABELS: print(f'{l[:4]:>8}', end='')
print()
for i, l in enumerate(LABELS):
print(f'{l:<22}', end='')
for j in range(len(LABELS)): print(f'{cm[i][j]:>8}', end='')
print()
print('\n' + '='*65)
print('SerendipLLM Classifier vs SinLlama')
print('='*65)
for metric in ['P', 'R', 'F1']:
diff = our[metric] - SINLLAMA[metric]
print(f'{metric:<15} {our[metric]:>15} {SINLLAMA[metric]:>15} {"▲" if diff>0 else "▼"}{abs(round(diff,3)):>9}')
diff_f1 = our['F1'] - SINLLAMA['F1']
print('='*65)
print(f'\n{"✅ BEATS" if diff_f1>0 else "❌ BELOW"} SinLlama by {abs(round(diff_f1,3))} F1 points!')
# Push LoRA + classifier head
print(f'\nPushing to {NEW_REPO}...')
base_model.push_to_hub(NEW_REPO, token=TOKEN)
tokenizer.push_to_hub(NEW_REPO, token=TOKEN)
torch.save(best_classifier_state, '/workspace/classifier_head.pt')
from huggingface_hub import HfApi
HfApi(token=TOKEN).upload_file(
path_or_fileobj='/workspace/classifier_head.pt',
path_in_repo='classifier_head.pt',
repo_id=NEW_REPO, repo_type='model')
print(f'Done! ✓ https://huggingface.co/{NEW_REPO}')
with open('/workspace/classifier_results.json','w',encoding='utf-8') as f:
json.dump({'model': NEW_REPO, 'serendipllm': our, 'sinllama': SINLLAMA,
'diff_f1': round(diff_f1,3), 'beats_sinllama': diff_f1>0,
'per_class': {LABELS[i]: {'P': round(p_per[i]*100,3),
'R': round(r_per[i]*100,3), 'F1': round(f1_per[i]*100,3)}
for i in range(NUM_LABELS)}}, f, indent=2, ensure_ascii=False)
print('Saved /workspace/classifier_results.json ✓')