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 ✓')