Spaces:
Build error
Build error
| """Script 05: Train all baselines for comparison. | |
| Three baselines, all on the same train/val/test splits as the Proposed model: | |
| - Baseline 1: TF-IDF + LogReg (text only, overall 3-class) | |
| - Baseline 2: BERT-overall fine-tune (text only, overall 3-class) | |
| - Baseline 3: BERT-ACSA (no meta) (text only, per-aspect 3-class) | |
| ^- this is the key ablation for the paper: | |
| Baseline 3 vs Proposed isolates the value of | |
| metadata fusion. | |
| """ | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import pandas as pd | |
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) | |
| from src.utils import setup_logging | |
| from src import config as cfg | |
| from src.baselines import train_tfidf_baseline | |
| from src.trainer import train_bert_overall, train_acsa | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--skip_tfidf", action="store_true") | |
| parser.add_argument("--skip_bert_overall", action="store_true") | |
| parser.add_argument("--skip_acsa_no_meta", action="store_true") | |
| parser.add_argument("--epochs", type=int, default=cfg.DEFAULT_EPOCHS) | |
| parser.add_argument("--batch_size", type=int, default=cfg.DEFAULT_BATCH_SIZE) | |
| parser.add_argument("--bert_name", default=cfg.BERT_MODEL_NAME, | |
| help="HuggingFace model name (default: config.BERT_MODEL_NAME)") | |
| args = parser.parse_args() | |
| setup_logging() | |
| train_df = pd.read_parquet(cfg.TRAIN_PATH) | |
| val_df = pd.read_parquet(cfg.VAL_PATH) | |
| test_df = pd.read_parquet(cfg.TEST_PATH) | |
| # Baseline 1 | |
| if not args.skip_tfidf: | |
| print("=" * 60) | |
| print("Baseline 1: TF-IDF + Logistic Regression (overall 3-class)") | |
| print("=" * 60) | |
| _, m = train_tfidf_baseline(train_df, val_df, test_df) | |
| print(json.dumps({k: v for k, v in m.items() if not k.endswith("_report")}, | |
| indent=2)) | |
| # Baseline 2 | |
| if not args.skip_bert_overall: | |
| print("=" * 60) | |
| print("Baseline 2: BERT fine-tune (overall 3-class)") | |
| print("=" * 60) | |
| train_bert_overall(train_df=train_df, val_df=val_df, | |
| bert_name=args.bert_name, | |
| epochs=args.epochs, batch_size=args.batch_size) | |
| # Baseline 3 — KEY ablation for the paper | |
| if not args.skip_acsa_no_meta: | |
| print("=" * 60) | |
| print("Baseline 3: BERT-ACSA (per-aspect, NO metadata) — ablation") | |
| print("=" * 60) | |
| train_acsa(train_df=train_df, val_df=val_df, | |
| bert_name=args.bert_name, | |
| epochs=args.epochs, batch_size=args.batch_size) | |
| if __name__ == "__main__": | |
| main() | |