A transformer-based model that predicts a discrete star rating (1–5) from customer review text.
Predicting exact 1–5 ratings is harder than binary sentiment because the “middle” classes (2★/3★/4★) are linguistically ambiguous. This project focuses on:
- strong NLP modeling
- reproducible data pipeline
- measurable evaluation
Model: microsoft/deberta-v3-large
Key ideas:
- Mean pooling: average token embeddings instead of relying only on
[CLS] - Dual heads:
- classification head → 5-class probabilities
- regression head → continuous rating signal (ordinal guidance)
- Hybrid ordinal loss:
- CrossEntropy (exact class)
- SmoothL1(regression vs true rating)
- SmoothL1(expected rating from probabilities vs true rating)
Dataset: McAuley-Lab/Amazon-Reviews-2023 (Hugging Face)
To reduce class imbalance and domain bias, training data is sampled with dual balancing:
- 4 domains: Electronics, Books, Clothing/Shoes/Jewelry, Home/Kitchen
- 5 rating classes
- equal samples per-class-per-domain (config-driven)
Validation/test sets are also generated as balanced splits so that macro-F1 is meaningful and classes are comparable.
Artifacts (data/checkpoints/results/logs) are stored outside git (recommended: on a larger disk for large runs).
Validation (25k balanced samples):
- Accuracy: 0.7078
- Macro-F1: 0.7066
Test (25k balanced samples):
- Accuracy: 0.6847
- Macro-F1: 0.6835
Additional quality signals (typical for the baseline):
- Off-by-one accuracy: ~0.97
- MAE: ~0.35 stars
These were tested with a clean protocol (tune on validation, test once):
- Ordinal post-processing / gating (argmax vs rounded expected rating)
- No meaningful improvement in exact 5-class test accuracy.
- Stacker models on transformer outputs (LogReg / HistGradientBoosting over probs + uncertainty + simple text features)
- Example test: Accuracy 0.6837, Macro-F1 0.6820 (worse than baseline).
- Ordinal-aware training loss (EMD / Wasserstein term)
- No validation gain over baseline (best val Macro-F1 ~0.7025 in the run shown).
- Leakage-safe metadata (verified_purchase only)
- Example test: Accuracy 0.6828, Macro-F1 0.6838 (no gain).
- Most errors are off-by-one (e.g., 4★ predicted as 5★), which indicates the model learns strong ordinal structure but exact boundaries between adjacent ratings remain difficult.
- The hardest classes are typically 2★ / 3★ / 4★ due to language ambiguity.
config.yaml— all hyperparameters and dataset settingssrc/— model, loss, trainer, preprocessingscripts/prepare_validation_set.py— builds balanced val/test splitsscripts/train.py— training entrypointscripts/evaluate.py— test evaluationscripts/inference.py— interactive and CLI inference
- Create env and install deps (example):
- PyTorch GPU build (CUDA-enabled)
transformers,datasets,accelerate, etc.
- Important GPU note:
If PyTorch reports no GPU but
nvidia-smiworks, check MIG mode:sudo nvidia-smi -i 0 -mig 0
From repo root:
-
Create balanced validation/test splits:
python scripts/prepare_validation_set.py -
Train (recommended inside tmux):
python scripts/train.py
Optional useful flags for experiments:
--run-name <name>to write artifacts undercheckpoints/<name>/andresults/<name>/--init-from <checkpoint.pt>to initialize weights from a prior run
-
Evaluate:
python scripts/evaluate.py -
Inference:
- Interactive:
python scripts/inference.py --interactive - One-off:
python scripts/inference.py --text "..." --category "Electronics"
Edit config.yaml to change:
- training steps, LR schedule, batch size, max_length
- dataset domains
- sample counts (train/val/test)
By default, the project writes artifacts to local folders:
checkpoints/,results/,data/,logs/
For large runs on cloud VMs, you can optionally move caches/artifacts to a larger disk by setting:
HF_HOME,HF_DATASETS_CACHE,TORCH_HOME,XDG_CACHE_HOMEor by symlinkingcheckpoints/ data/ results/ logs/to another mount.
None of this is required for small runs.
This repo does not redistribute the Amazon dataset. It downloads via Hugging Face at runtime.