Spaces:
Running on Zero
Running on Zero
| """Training loops for ACSA and baseline models.""" | |
| import json | |
| import logging | |
| from pathlib import Path | |
| from typing import Optional | |
| import numpy as np | |
| import torch | |
| from torch.utils.data import DataLoader | |
| from transformers import AutoTokenizer, get_linear_schedule_with_warmup | |
| from tqdm import tqdm | |
| from . import config as cfg | |
| from .models import ( | |
| GatedAspectSemanticMetaFusionACSAModel, BertMetaFusionACSAModel, BertACSAModel, BertOverallModel, | |
| compute_class_weights, | |
| ) | |
| from .dataset import ACSADataset, MetaACSADataset, OverallSentimentDataset | |
| from .meta_encoder import MetaEncoder | |
| from .utils import set_seed, seeded_generator | |
| logger = logging.getLogger(__name__) | |
| def get_device(): | |
| if torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): | |
| return torch.device("mps") | |
| return torch.device("cpu") | |
| # --------------------------------------------------------------------------- | |
| # Generic per-aspect evaluation (used by both Proposed and Baseline 3) | |
| # --------------------------------------------------------------------------- | |
| def _evaluate_per_aspect(model, loader, device, with_meta: bool): | |
| from sklearn.metrics import f1_score, accuracy_score | |
| model.eval() | |
| all_preds = [[] for _ in range(cfg.NUM_ASPECTS)] | |
| all_labels = [[] for _ in range(cfg.NUM_ASPECTS)] | |
| with torch.no_grad(): | |
| for batch in loader: | |
| batch = {k: v.to(device) for k, v in batch.items()} | |
| if with_meta: | |
| out = model(batch["input_ids"], batch["attention_mask"], | |
| batch["meta_features"]) | |
| else: | |
| out = model(batch["input_ids"], batch["attention_mask"]) | |
| preds = out["logits"].argmax(dim=-1).cpu().numpy() | |
| labels = batch["labels"].cpu().numpy() | |
| for i in range(cfg.NUM_ASPECTS): | |
| all_preds[i].extend(preds[:, i].tolist()) | |
| all_labels[i].extend(labels[:, i].tolist()) | |
| metrics = {} | |
| f1s, accs = [], [] | |
| for i, aspect in enumerate(cfg.ASPECTS): | |
| f1 = f1_score(all_labels[i], all_preds[i], average="macro", zero_division=0) | |
| acc = accuracy_score(all_labels[i], all_preds[i]) | |
| metrics[f"f1_{aspect}"] = float(f1) | |
| metrics[f"acc_{aspect}"] = float(acc) | |
| f1s.append(f1); accs.append(acc) | |
| metrics["macro_f1_mean"] = float(np.mean(f1s)) if f1s else 0.0 | |
| metrics["accuracy_mean"] = float(np.mean(accs)) if accs else 0.0 | |
| return metrics | |
| def _evaluate_overall_head(model, loader, device): | |
| """Evaluate the overall sentiment auxiliary head on a dataloader. | |
| Returns metrics prefixed with 'overall_' so they don't collide with | |
| per-aspect metric keys. | |
| """ | |
| from sklearn.metrics import f1_score, accuracy_score | |
| model.eval() | |
| all_preds, all_labels = [], [] | |
| with torch.no_grad(): | |
| for batch in loader: | |
| batch = {k: v.to(device) for k, v in batch.items()} | |
| if "overall_labels" not in batch: | |
| return {} | |
| out = model(batch["input_ids"], batch["attention_mask"], | |
| batch["meta_features"]) | |
| preds = out["overall_logits"].argmax(dim=-1).cpu().numpy() | |
| all_preds.extend(preds.tolist()) | |
| all_labels.extend(batch["overall_labels"].cpu().numpy().tolist()) | |
| if not all_labels: | |
| return {} | |
| return { | |
| "overall_macro_f1": float(f1_score(all_labels, all_preds, average="macro", zero_division=0)), | |
| "overall_accuracy": float(accuracy_score(all_labels, all_preds)), | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Train: Proposed (BERT + Meta Cross-Attention + multi-head) | |
| # --------------------------------------------------------------------------- | |
| def train_meta_acsa( | |
| train_df, val_df, meta_encoder: MetaEncoder, | |
| bert_name: str = cfg.BERT_MODEL_NAME, | |
| epochs: int = cfg.DEFAULT_EPOCHS, | |
| batch_size: int = cfg.DEFAULT_BATCH_SIZE, | |
| lr_bert: float = cfg.DEFAULT_LR_BERT, | |
| lr_heads: float = cfg.DEFAULT_LR_HEADS, | |
| weight_decay: float = cfg.DEFAULT_WEIGHT_DECAY, | |
| use_class_weights: bool = True, | |
| output_dir: Optional[Path] = None, | |
| seed: int = cfg.RANDOM_SEED, | |
| ): | |
| if output_dir is None: | |
| output_dir = cfg.CHECKPOINT_DIR / "meta_acsa" | |
| output_dir = Path(output_dir); output_dir.mkdir(parents=True, exist_ok=True) | |
| set_seed(seed) | |
| shuffle_generator = seeded_generator(seed) | |
| device = get_device() | |
| logger.info("Device: %s", device) | |
| tokenizer = AutoTokenizer.from_pretrained(bert_name) | |
| aspect_cols = [f"aspect_{a}" for a in cfg.ASPECTS] | |
| class_weights = (compute_class_weights(train_df, aspect_cols, cfg.NUM_CLASSES).to(device) | |
| if use_class_weights else None) | |
| if class_weights is not None: | |
| logger.info("Per-aspect class weights:\n%s", class_weights.cpu().numpy()) | |
| model = GatedAspectSemanticMetaFusionACSAModel( | |
| bert_name=bert_name, | |
| meta_in_dim=meta_encoder.total_dim, | |
| class_weights=class_weights, | |
| ).to(device) | |
| train_ds = MetaACSADataset(train_df, tokenizer, meta_encoder) | |
| val_ds = MetaACSADataset(val_df, tokenizer, meta_encoder) | |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=0, | |
| generator=shuffle_generator) | |
| val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=0) | |
| bert_params = list(model.bert.named_parameters()) | |
| other_params = [(n, p) for n, p in model.named_parameters() | |
| if not n.startswith("bert.")] | |
| no_decay = ["bias", "LayerNorm.weight"] | |
| grouped = [ | |
| {"params": [p for n, p in bert_params if not any(nd in n for nd in no_decay)], | |
| "weight_decay": weight_decay, "lr": lr_bert}, | |
| {"params": [p for n, p in bert_params if any(nd in n for nd in no_decay)], | |
| "weight_decay": 0.0, "lr": lr_bert}, | |
| {"params": [p for n, p in other_params if not any(nd in n for nd in no_decay)], | |
| "weight_decay": weight_decay, "lr": lr_heads}, | |
| {"params": [p for n, p in other_params if any(nd in n for nd in no_decay)], | |
| "weight_decay": 0.0, "lr": lr_heads}, | |
| ] | |
| optimizer = torch.optim.AdamW(grouped) | |
| total_steps = max(len(train_loader) * epochs, 1) | |
| scheduler = get_linear_schedule_with_warmup( | |
| optimizer, num_warmup_steps=int(cfg.WARMUP_RATIO * total_steps), | |
| num_training_steps=total_steps, | |
| ) | |
| best_f1 = -1.0; history = [] | |
| for epoch in range(epochs): | |
| model.train(); running = 0.0 | |
| pbar = tqdm(train_loader, desc=f"[epoch {epoch+1}/{epochs}] meta_acsa") | |
| for batch in pbar: | |
| batch = {k: v.to(device) for k, v in batch.items()} | |
| optimizer.zero_grad() | |
| out = model(batch["input_ids"], batch["attention_mask"], | |
| batch["meta_features"], labels=batch["labels"], | |
| overall_labels=batch.get("overall_labels")) | |
| out["loss"].backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step(); scheduler.step() | |
| running += out["loss"].item() | |
| pbar.set_postfix({"loss": f"{out['loss'].item():.4f}"}) | |
| avg_loss = running / max(len(train_loader), 1) | |
| val = _evaluate_per_aspect(model, val_loader, device, with_meta=True) | |
| # Also evaluate overall head on val set | |
| val_overall = _evaluate_overall_head(model, val_loader, device) | |
| val.update(val_overall) | |
| logger.info("Epoch %d | loss=%.4f | val_macro_f1=%.4f | val_acc=%.4f | val_overall_f1=%.4f", | |
| epoch+1, avg_loss, val["macro_f1_mean"], val["accuracy_mean"], | |
| val.get("overall_macro_f1", 0.0)) | |
| history.append({"epoch": epoch+1, "train_loss": avg_loss, **val}) | |
| if val["macro_f1_mean"] > best_f1: | |
| best_f1 = val["macro_f1_mean"] | |
| torch.save({"model_state_dict": model.state_dict(), | |
| "config": {"bert_name": bert_name, | |
| "meta_in_dim": meta_encoder.total_dim, | |
| "architecture": "gated_aspect_semantic_meta_acsa", | |
| "meta_token_names": GatedAspectSemanticMetaFusionACSAModel.meta_token_names}}, | |
| output_dir / "best.pt") | |
| tokenizer.save_pretrained(output_dir / "tokenizer") | |
| logger.info("Saved new best meta_acsa (macro_f1=%.4f)", best_f1) | |
| with open(output_dir / "history.json", "w") as f: | |
| json.dump(history, f, indent=2) | |
| return model, history | |
| # --------------------------------------------------------------------------- | |
| # Train: Baseline 3 (BERT-ACSA, no meta 鈥?ablation) | |
| # --------------------------------------------------------------------------- | |
| def train_acsa( | |
| train_df, val_df, | |
| bert_name: str = cfg.BERT_MODEL_NAME, | |
| epochs: int = cfg.DEFAULT_EPOCHS, | |
| batch_size: int = cfg.DEFAULT_BATCH_SIZE, | |
| lr_bert: float = cfg.DEFAULT_LR_BERT, | |
| lr_heads: float = cfg.DEFAULT_LR_HEADS, | |
| weight_decay: float = cfg.DEFAULT_WEIGHT_DECAY, | |
| use_class_weights: bool = True, | |
| output_dir: Optional[Path] = None, | |
| seed: int = cfg.RANDOM_SEED, | |
| ): | |
| if output_dir is None: | |
| output_dir = cfg.CHECKPOINT_DIR / "acsa" | |
| output_dir = Path(output_dir); output_dir.mkdir(parents=True, exist_ok=True) | |
| set_seed(seed) | |
| shuffle_generator = seeded_generator(seed) | |
| device = get_device() | |
| tokenizer = AutoTokenizer.from_pretrained(bert_name) | |
| aspect_cols = [f"aspect_{a}" for a in cfg.ASPECTS] | |
| class_weights = (compute_class_weights(train_df, aspect_cols, cfg.NUM_CLASSES).to(device) | |
| if use_class_weights else None) | |
| model = BertACSAModel(bert_name=bert_name, class_weights=class_weights).to(device) | |
| train_ds = ACSADataset(train_df, tokenizer) | |
| val_ds = ACSADataset(val_df, tokenizer) | |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=0, | |
| generator=shuffle_generator) | |
| val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=0) | |
| bert_params = list(model.bert.named_parameters()) | |
| head_params = list(model.heads.named_parameters()) | |
| no_decay = ["bias", "LayerNorm.weight"] | |
| grouped = [ | |
| {"params": [p for n, p in bert_params if not any(nd in n for nd in no_decay)], | |
| "weight_decay": weight_decay, "lr": lr_bert}, | |
| {"params": [p for n, p in bert_params if any(nd in n for nd in no_decay)], | |
| "weight_decay": 0.0, "lr": lr_bert}, | |
| {"params": [p for _, p in head_params], "weight_decay": weight_decay, "lr": lr_heads}, | |
| ] | |
| optimizer = torch.optim.AdamW(grouped) | |
| total_steps = max(len(train_loader) * epochs, 1) | |
| scheduler = get_linear_schedule_with_warmup( | |
| optimizer, int(cfg.WARMUP_RATIO * total_steps), total_steps, | |
| ) | |
| best_f1 = -1.0; history = [] | |
| for epoch in range(epochs): | |
| model.train(); running = 0.0 | |
| pbar = tqdm(train_loader, desc=f"[epoch {epoch+1}/{epochs}] acsa-no-meta") | |
| for batch in pbar: | |
| batch = {k: v.to(device) for k, v in batch.items()} | |
| optimizer.zero_grad() | |
| out = model(batch["input_ids"], batch["attention_mask"], labels=batch["labels"]) | |
| out["loss"].backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step(); scheduler.step() | |
| running += out["loss"].item() | |
| pbar.set_postfix({"loss": f"{out['loss'].item():.4f}"}) | |
| val = _evaluate_per_aspect(model, val_loader, device, with_meta=False) | |
| logger.info("Epoch %d | val_macro_f1=%.4f | val_acc=%.4f", | |
| epoch+1, val["macro_f1_mean"], val["accuracy_mean"]) | |
| history.append({"epoch": epoch+1, "train_loss": running/max(len(train_loader),1), | |
| **val}) | |
| if val["macro_f1_mean"] > best_f1: | |
| best_f1 = val["macro_f1_mean"] | |
| torch.save({"model_state_dict": model.state_dict(), | |
| "config": {"bert_name": bert_name}}, | |
| output_dir / "best.pt") | |
| tokenizer.save_pretrained(output_dir / "tokenizer") | |
| with open(output_dir / "history.json", "w") as f: | |
| json.dump(history, f, indent=2) | |
| return model, history | |
| # --------------------------------------------------------------------------- | |
| # Train: Baseline 2 (BERT-overall) | |
| # --------------------------------------------------------------------------- | |
| def train_bert_overall( | |
| train_df, val_df, | |
| bert_name: str = cfg.BERT_MODEL_NAME, | |
| epochs: int = cfg.DEFAULT_EPOCHS, | |
| batch_size: int = cfg.DEFAULT_BATCH_SIZE, | |
| lr: float = cfg.DEFAULT_LR_BERT, | |
| weight_decay: float = cfg.DEFAULT_WEIGHT_DECAY, | |
| output_dir: Optional[Path] = None, | |
| seed: int = cfg.RANDOM_SEED, | |
| ): | |
| from sklearn.metrics import f1_score, accuracy_score | |
| if output_dir is None: | |
| output_dir = cfg.CHECKPOINT_DIR / "bert_overall" | |
| output_dir = Path(output_dir); output_dir.mkdir(parents=True, exist_ok=True) | |
| set_seed(seed) | |
| shuffle_generator = seeded_generator(seed) | |
| device = get_device() | |
| tokenizer = AutoTokenizer.from_pretrained(bert_name) | |
| model = BertOverallModel(bert_name=bert_name).to(device) | |
| train_ds = OverallSentimentDataset(train_df, tokenizer) | |
| val_ds = OverallSentimentDataset(val_df, tokenizer) | |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=0, | |
| generator=shuffle_generator) | |
| val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=0) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay) | |
| total_steps = max(len(train_loader) * epochs, 1) | |
| scheduler = get_linear_schedule_with_warmup( | |
| optimizer, int(cfg.WARMUP_RATIO * total_steps), total_steps, | |
| ) | |
| best_f1 = -1.0; history = [] | |
| for epoch in range(epochs): | |
| model.train(); running = 0.0 | |
| pbar = tqdm(train_loader, desc=f"[epoch {epoch+1}/{epochs}] bert-overall") | |
| for batch in pbar: | |
| batch = {k: v.to(device) for k, v in batch.items()} | |
| optimizer.zero_grad() | |
| out = model(**batch) | |
| out["loss"].backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step(); scheduler.step() | |
| running += out["loss"].item() | |
| pbar.set_postfix({"loss": f"{out['loss'].item():.4f}"}) | |
| model.eval(); all_p, all_l = [], [] | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| batch = {k: v.to(device) for k, v in batch.items()} | |
| out = model(**batch) | |
| all_p.extend(out["logits"].argmax(dim=-1).cpu().numpy().tolist()) | |
| all_l.extend(batch["labels"].cpu().numpy().tolist()) | |
| metrics = { | |
| "macro_f1": float(f1_score(all_l, all_p, average="macro", zero_division=0)), | |
| "accuracy": float(accuracy_score(all_l, all_p)), | |
| "weighted_f1": float(f1_score(all_l, all_p, average="weighted", zero_division=0)), | |
| } | |
| logger.info("Epoch %d | val_macro_f1=%.4f | val_acc=%.4f", | |
| epoch+1, metrics["macro_f1"], metrics["accuracy"]) | |
| history.append({"epoch": epoch+1, "train_loss": running/max(len(train_loader),1), | |
| **metrics}) | |
| if metrics["macro_f1"] > best_f1: | |
| best_f1 = metrics["macro_f1"] | |
| torch.save({"model_state_dict": model.state_dict(), | |
| "config": {"bert_name": bert_name}}, | |
| output_dir / "best.pt") | |
| tokenizer.save_pretrained(output_dir / "tokenizer") | |
| with open(output_dir / "history.json", "w") as f: | |
| json.dump(history, f, indent=2) | |
| return model, history | |