| """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) |
|
|
| |
| 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)) |
|
|
| |
| 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) |
|
|
| |
| 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() |
|
|