整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
.venv/
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
@@ -0,0 +1,39 @@
|
||||
# Q2/Q3 algorithm selection built on Q1 alignment
|
||||
|
||||
## Decision about reusing Q1
|
||||
|
||||
The transferable part of Q1 is its explicit time correspondence and observation mask: features from different modalities share ordered positions, missing values are accompanied by masks, and a position can be traced to source time. That interface is useful for both Q2 local-gap handling and Q3 evidence localization.
|
||||
|
||||
The exact Q1 B1 extraction cannot be rerun over the 4,850 Attachment 2 training examples. Attachment 2 supplies precomputed aligned and unaligned feature tensors, but not the source audio/video or CTC word-time posteriors for the full training set. Its `aligned_50.pkl` also has 50 wordpiece positions and no Q1 `time_bounds_s`; those positions must not be described as the 50 equal-duration physical-time bins exported by `final/Q1`.
|
||||
|
||||
Accordingly, the Q2 experiment uses the official aligned feature set as its shared wordpiece axis, and compares it with a fixed equal-window pooling control made from the official unaligned audio/vision sequences. This is a downstream alignment-utility check, not a claim that B1 was recomputed on Attachment 2. The official train/validation split is retained; test labels are not used.
|
||||
|
||||
## Q2 candidates
|
||||
|
||||
All candidates use identical training examples, train-only median/MAD scaling, joint polarity/intensity objectives, and 15 validation corruptions (three contiguous missing rates by five modality patterns).
|
||||
|
||||
| Candidate | Fusion rule | What it tests |
|
||||
| --- | --- | --- |
|
||||
| `concat` | Project each modality, concatenate features and availability flags, then run a bidirectional GRU | Strong, simple early-fusion baseline |
|
||||
| `gate` | Learn per-slot modality weights, mask unavailable modalities, then run a bidirectional GRU | Whether explicit reliability-aware fusion handles local gaps |
|
||||
| `crossattn` | Apply masked cross-modal attention over the 50 shared slots, then temporal pooling | Whether contextual cross-modal exchange improves robustness |
|
||||
|
||||
The report keeps Macro-F1, MAE, and Pearson separate. The default selection is Macro-F1-first across local corruption conditions; MAE and Pearson remain explicit tradeoffs, not terms in a constructed total score. The selected architecture is also trained on fixed-window-resampled features as an alignment control. A separate validation control shifts audio and vision by 1–10 positions to measure sensitivity to cross-modal timing.
|
||||
|
||||
## Q3 explanation selection
|
||||
|
||||
The selected Q2 model is frozen. Integrated Gradients and five-slot grouped occlusion are compared on held-out Attachment 2 validation clips using deletion comprehensiveness, sufficiency, and local rank stability. Attachment 4 has original videos and transcripts, so B1's CTC hard word-time procedure can be applied to those 20 clips to map high-importance wordpiece positions back to seconds. The saved Attachment 4 pickle files do not include `time_bounds_s`; explanations therefore retain both the model slot and the CTC-derived word interval, with alignment quality recorded.
|
||||
|
||||
## Run
|
||||
|
||||
The project environment is managed by `uv` and installs the CUDA 13.0 PyTorch build:
|
||||
|
||||
```bash
|
||||
cd deep_learning/Q2
|
||||
uv sync
|
||||
uv run python -m q2.train_compare
|
||||
cd ../Q3
|
||||
uv run --project ../Q2 python -m q3.explain_selection
|
||||
```
|
||||
|
||||
The main outputs are written to `outputs/algorithm_selection/`; plots, CSV metrics, run metadata, and checkpoints stay under this directory. The source data, `math`, and `final/Q1` are read-only inputs.
|
||||
@@ -0,0 +1,54 @@
|
||||
# Q2 algorithm selection results
|
||||
|
||||
## Q1 alignment transfer decision
|
||||
|
||||
Q1 B1 aligns BERT word features and audio/vision observations with hard CTC word intervals, projects observed features onto a 0.1-second common grid, exports 50 equal-duration physical-time bins, and keeps observation masks. For Q2, the shared ordered axis and explicit masks transfer directly: a local gap stays a local gap after alignment and can be represented without inventing feature values.
|
||||
|
||||
The exact B1 extraction was not recomputed over Attachment 2. The official 4,850-row feature package contains precomputed aligned and unaligned tensors, but no full-set source audio/video or word-time posterior. Its `aligned_50.pkl` has 50 wordpiece positions and no per-slot `time_bounds_s`; those positions are not Q1's 50 equal-duration bins. This experiment therefore trains on the official aligned features and compares them with an equal-window audio/vision resampling control. The comparison tests the value of an aligned ordered representation for the downstream Q2 task; it does not claim to reproduce B1 on all 4,850 clips.
|
||||
|
||||
## Data and protocol
|
||||
|
||||
- Attachment 2 official split: 3,395 training clips and 728 validation clips. Their source-video ID sets do not overlap.
|
||||
- Each official aligned example has 50 positions with Text 768-D, Audio 74-D, Vision 35-D features and modality observation masks.
|
||||
- Attachment 2 test labels were not used.
|
||||
- The three candidates shared train-only median/MAD normalization, the joint polarity/intensity objective, and training-time contiguous block masking.
|
||||
- Validation corruption covered 10%, 20%, and 30% of 50 positions for Text, Audio, Vision, Audio+Vision, and all three modalities. This is a wordpiece-position proxy for a continuous time gap; full-set second-level timestamps are not supplied.
|
||||
- Each candidate was run with seeds 42, 3407, and 2026. Reported `±` values are seed standard deviations over the fixed official validation set and deterministic corruption draws; they are not confidence intervals over new videos.
|
||||
|
||||
## Fusion comparison
|
||||
|
||||
| Model | Clean Accuracy | Clean Macro-F1 | Corrupt Accuracy, mean | Corrupt Macro-F1, mean | Worst condition Macro-F1 | Corrupt MAE | Corrupt Pearson |
|
||||
| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
|
||||
| Early concatenation + BiGRU | 0.626 ± 0.013 | 0.580 ± 0.018 | 0.623 ± 0.011 | **0.575 ± 0.017** | **0.516** | **0.640 ± 0.004** | 0.607 ± 0.004 |
|
||||
| Reliability gate + BiGRU | 0.621 ± 0.011 | 0.570 ± 0.019 | 0.617 ± 0.011 | 0.565 ± 0.018 | 0.498 | 0.643 ± 0.008 | **0.610 ± 0.008** |
|
||||
| Masked cross-modal attention | 0.604 ± 0.012 | 0.540 ± 0.049 | 0.601 ± 0.009 | 0.539 ± 0.047 | 0.462 | 0.649 ± 0.024 | 0.592 ± 0.016 |
|
||||
|
||||
Early concatenation has the best mean corrupted Macro-F1 and MAE. The gate has slightly higher Pearson, so the metrics do not collapse to one score. Cross-modal attention is lower and more variable at this sample size. It is not selected for the next Q2 stage.
|
||||
|
||||
For the selected concatenation model, the hardest tested case is 30% Text masking: Macro-F1 0.547 and MAE 0.665, compared with clean Macro-F1 0.580 and MAE 0.636. Audio-only or Vision-only masking has a smaller effect in these runs. This is evidence about this feature set and these simulated spans; it does not establish a universal modality ranking.
|
||||
|
||||
## Alignment utility control
|
||||
|
||||
The same concatenation model was trained either on the supplied aligned wordpiece features or on equal-window-resampled audio/vision features from the official unaligned tensors.
|
||||
|
||||
| Representation | Clean Macro-F1 | Corrupt Macro-F1 | Corrupt MAE | Corrupt Pearson |
|
||||
| --- | ---: | ---: | ---: | ---: |
|
||||
| Supplied word-aligned 50 positions | 0.580 ± 0.018 | **0.575 ± 0.017** | **0.640 ± 0.004** | **0.607 ± 0.004** |
|
||||
| Equal-window resampled unaligned input | 0.501 ± 0.014 | 0.504 ± 0.016 | 0.665 ± 0.005 | 0.577 ± 0.010 |
|
||||
|
||||
On the selected model, shifting Audio and Vision by 1–10 positions changed aligned Macro-F1 from 0.580 to 0.557. That is a modest timing-sensitivity signal; it does not prove the model uses precise physical-time correspondence. Together with the fixed-window comparison, the result supports retaining the supplied aligned sequence for Q2.
|
||||
|
||||
## Selected Q2 direction
|
||||
|
||||
Continue with mask-aware early concatenation plus a bidirectional GRU, using local block masking during training. Keep the reliability gate as an ablation because its Pearson is slightly higher. Revisit cross-attention only if a later run has stronger evidence and enough data to control overfitting.
|
||||
|
||||
## Reproducible artifacts
|
||||
|
||||
- [Model and representation summary](outputs/algorithm_selection/summary.csv)
|
||||
- [Metrics by missing type and rate](outputs/algorithm_selection/validation_metrics_by_condition.csv)
|
||||
- [Aligned versus fixed-window and temporal-shift controls](outputs/algorithm_selection/alignment_transfer_ablation.csv)
|
||||
- [Training/data audit and run manifest](outputs/algorithm_selection/data_audit.json), [run manifest](outputs/algorithm_selection/run_manifest.json)
|
||||
- [Validation plot](outputs/algorithm_selection/missing_rate_comparison.png)
|
||||
- [Selected seed-42 checkpoint](outputs/algorithm_selection/models/aligned/concat/model_best.pt)
|
||||
|
||||
The fitted checkpoint is for algorithm selection, not the final Attachment 3 submission model. The final model should be trained on train+validation after the architecture and thresholds are frozen.
|
||||
Binary file not shown.
@@ -0,0 +1,7 @@
|
||||
method,representation,condition,n_valid,n_seeds,accuracy,accuracy_sd,macro_f1,macro_f1_sd,mae,mae_sd,pearson,pearson_sd,missing_rate
|
||||
concat,provided_word_aligned_50,clean,728,3,0.6259157509157509,0.012689016905266455,0.5803030257575642,0.01849482950112713,0.6362011035283407,0.0035474183737517766,0.6131094378711319,0.003049322677391937,
|
||||
concat,provided_word_aligned_50,audio_vision_shifted_1_to_10_slots,728,3,0.6144688644688645,0.017174907814570563,0.5566710058858901,0.02417454695632849,0.6264231006304423,0.003071737263575265,0.610966440919508,0.00045089311901344093,
|
||||
concat,provided_word_aligned_50,all_local_corruption_mean,728,3,0.6226800976800977,0.011084117339613314,0.5751264735191365,0.016846851540436875,0.6402280900213454,0.004374268275822229,0.6068152054284395,0.004485336212715804,0.20000000000000004
|
||||
concat,equal_window_resampled_unaligned,clean,728,3,0.586996336996337,0.010491244722884272,0.5005019574320758,0.013925639871727626,0.6615431904792786,0.005397772426619935,0.581966026863303,0.009642037378255953,
|
||||
concat,equal_window_resampled_unaligned,audio_vision_shifted_1_to_10_slots,728,3,0.5956959706959707,0.010309826235528995,0.5144238712048543,0.010932166832616294,0.6616438627243042,0.005509720037298075,0.5807893064362651,0.010274359593119802,
|
||||
concat,equal_window_resampled_unaligned,all_local_corruption_mean,728,3,0.5905677655677655,0.006480515023060099,0.5042920441199125,0.015580461355158729,0.6651763810051813,0.004962154081458186,0.5766306314815771,0.010310184996660417,0.20000000000000004
|
||||
|
@@ -0,0 +1,21 @@
|
||||
{
|
||||
"source": "/home/gloamxun/modeling_zhaocui/E题数据/附件2-数据集特征文件/aligned_50.pkl",
|
||||
"train_samples": 3395,
|
||||
"valid_samples": 728,
|
||||
"train_classes": [
|
||||
967,
|
||||
758,
|
||||
1670
|
||||
],
|
||||
"valid_classes": [
|
||||
206,
|
||||
184,
|
||||
338
|
||||
],
|
||||
"mean_observed_slots": {
|
||||
"text": 24.645655375552284,
|
||||
"audio": 22.626509572901327,
|
||||
"vision": 21.394108983799704
|
||||
},
|
||||
"train_valid_video_overlap": 0
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
|
After Width: | Height: | Size: 158 KiB |
Binary file not shown.
BIN
Binary file not shown.
+10
@@ -0,0 +1,10 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.05146148469713,0.9808105859127674
|
||||
2.0,0.8792424190927435,0.8854341854105939
|
||||
3.0,0.7658277087741427,0.8630200951963991
|
||||
4.0,0.7162007325225406,0.8734401222113725
|
||||
5.0,0.666606965440291,0.8838760400866414
|
||||
6.0,0.6495787705536242,0.9118019499621548
|
||||
7.0,0.5989178496378439,0.9559067448416909
|
||||
8.0,0.5539200300419772,0.9710462656649914
|
||||
9.0,0.5062780554095904,1.0261535592131563
|
||||
|
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0488198929362826,0.9925541471649002
|
||||
2.0,0.877586845446516,0.9104853272438049
|
||||
3.0,0.7854031710712998,0.8691042694416675
|
||||
4.0,0.7275559962899597,0.8561014971890293
|
||||
5.0,0.6785599307881461,0.8777891502275572
|
||||
6.0,0.6318395165381608,0.9196370329175677
|
||||
7.0,0.5972505140083807,0.9164495874237228
|
||||
8.0,0.5567877353341492,0.9559657193802216
|
||||
9.0,0.5066572507774388,1.0156079124618362
|
||||
10.0,0.47001609758094504,1.059160087134812
|
||||
|
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0536950241636347,0.9706045127176977
|
||||
2.0,0.8524193581607606,0.8777912927197886
|
||||
3.0,0.760877827251399,0.8636277507949661
|
||||
4.0,0.7150518761740791,0.8624753559028709
|
||||
5.0,0.6774439651657034,0.877832626248454
|
||||
6.0,0.6516634187212696,0.8901654671836685
|
||||
7.0,0.5939983526865641,0.918816069325248
|
||||
8.0,0.5618191918841114,0.9479289251369435
|
||||
9.0,0.513365975132695,1.013728333043528
|
||||
10.0,0.4933904481154901,1.030043561379988
|
||||
|
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0536950241636347,0.9706045127176977
|
||||
2.0,0.8524193581607606,0.8777912927197886
|
||||
3.0,0.760877827251399,0.8636277507949661
|
||||
4.0,0.7150518761740791,0.8624753559028709
|
||||
5.0,0.6774439651657034,0.877832626248454
|
||||
6.0,0.6516634187212696,0.8901654671836685
|
||||
7.0,0.5939983526865641,0.918816069325248
|
||||
8.0,0.5618191918841114,0.9479289251369435
|
||||
9.0,0.513365975132695,1.013728333043528
|
||||
10.0,0.4933904481154901,1.030043561379988
|
||||
|
Binary file not shown.
BIN
Binary file not shown.
+10
@@ -0,0 +1,10 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0184006403993677,0.9758902355864808
|
||||
2.0,0.8721449165432541,0.9131481346193251
|
||||
3.0,0.7584562389938919,0.9061737309445391
|
||||
4.0,0.70247816046079,0.908711409830785
|
||||
5.0,0.6474666976266437,0.9444795060943771
|
||||
6.0,0.6255715103061111,0.9916293214965652
|
||||
7.0,0.5631065650118722,1.0601372142414471
|
||||
8.0,0.5009032366452394,1.1017653536010574
|
||||
9.0,0.4493949540235378,1.177744794677902
|
||||
|
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0344635292335793,0.9769261263229034
|
||||
2.0,0.8467512428760529,0.8954217656628116
|
||||
3.0,0.7528229709024783,0.8857204809293642
|
||||
4.0,0.7020640406343672,0.8774335417118702
|
||||
5.0,0.6610019422239728,0.9038714190105815
|
||||
6.0,0.6052676102629414,0.9758975283130185
|
||||
7.0,0.5659312236088293,0.9750687735421317
|
||||
8.0,0.5310967906757638,1.049424912903335
|
||||
9.0,0.46808629096658144,1.1149894405197311
|
||||
10.0,0.4314451297676122,1.177863896548093
|
||||
|
BIN
Binary file not shown.
+9
@@ -0,0 +1,9 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.011995311136599,0.9519830608105921
|
||||
2.0,0.8360813723670112,0.8962717344472696
|
||||
3.0,0.7534678158936677,0.8991011263249995
|
||||
4.0,0.7069296527791906,0.9073571071519957
|
||||
5.0,0.6639300579274142,0.9432451947704776
|
||||
6.0,0.6353109103661997,0.9912105065125686
|
||||
7.0,0.5673890513954339,1.0443138106838687
|
||||
8.0,0.5254902547156369,1.1027433977022276
|
||||
|
+9
@@ -0,0 +1,9 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0119953784677718,0.9519833208440425
|
||||
2.0,0.8360813088991024,0.8962707964928596
|
||||
3.0,0.7534678666679947,0.8991013843934614
|
||||
4.0,0.7069298084135409,0.9073559026141743
|
||||
5.0,0.6639298437922089,0.9432453493495564
|
||||
6.0,0.6353108998801973,0.9912109388099922
|
||||
7.0,0.56738873405589,1.0443127351802783
|
||||
8.0,0.5254902160829968,1.1027441640476605
|
||||
|
Binary file not shown.
BIN
Binary file not shown.
+10
@@ -0,0 +1,10 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0707936783631642,1.0067818826371497
|
||||
2.0,0.9140718049473233,0.9000171115110208
|
||||
3.0,0.7863181178216581,0.8512143093151051
|
||||
4.0,0.7349886541013364,0.851207211122408
|
||||
5.0,0.6789810916891804,0.8575038864062383
|
||||
6.0,0.6610487986493994,0.8806245772393195
|
||||
7.0,0.6122590667671628,0.9165601101550427
|
||||
8.0,0.5692078007592095,0.9351011645662916
|
||||
9.0,0.5215919649711361,0.99570418714167
|
||||
|
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0577489623317011,0.9879156252840063
|
||||
2.0,0.8716377052995894,0.8907375866240197
|
||||
3.0,0.7750455615697084,0.8654810419449439
|
||||
4.0,0.7179531797214791,0.8613645519529071
|
||||
5.0,0.6809909796273267,0.8796235846949148
|
||||
6.0,0.6348593281926932,0.9029495820894347
|
||||
7.0,0.6029366790144531,0.9171899249265482
|
||||
8.0,0.5716203340777645,0.9463915248493572
|
||||
9.0,0.5201758698180869,0.9928344789442125
|
||||
10.0,0.5027363620422505,1.0327370245378096
|
||||
|
BIN
Binary file not shown.
+10
@@ -0,0 +1,10 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0575931425447818,0.9988141858970726
|
||||
2.0,0.8773076059641661,0.8911895647153749
|
||||
3.0,0.7700778461164899,0.8651458110128131
|
||||
4.0,0.7248913248380026,0.8694970201659988
|
||||
5.0,0.682613401501267,0.884684423823933
|
||||
6.0,0.659738369010113,0.9030193935383807
|
||||
7.0,0.6011747334290434,0.9310533351950593
|
||||
8.0,0.5665941779260282,0.9628150620303311
|
||||
9.0,0.5270494206084145,1.0122937671430818
|
||||
|
@@ -0,0 +1,10 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0575931425447818,0.9988141858970726
|
||||
2.0,0.8773076059641661,0.8911895647153749
|
||||
3.0,0.7700778461164899,0.8651458110128131
|
||||
4.0,0.7248913248380026,0.8694970201659988
|
||||
5.0,0.682613401501267,0.884684423823933
|
||||
6.0,0.659738369010113,0.9030193935383807
|
||||
7.0,0.6011747334290434,0.9310533351950593
|
||||
8.0,0.5665941779260282,0.9628150620303311
|
||||
9.0,0.5270494206084145,1.0122937671430818
|
||||
|
Binary file not shown.
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0648082616152588,1.0257792996836232
|
||||
2.0,0.9508535012050912,0.9455079901349414
|
||||
3.0,0.83619585191762,0.9239543119629661
|
||||
4.0,0.7835445602734884,0.9075145099189256
|
||||
5.0,0.7335039586932571,0.9146714800006741
|
||||
6.0,0.7189743309109299,0.9153845460860284
|
||||
7.0,0.669045564201143,0.9579015452783186
|
||||
8.0,0.6221646765867869,0.9763828352257445
|
||||
9.0,0.5898675388760037,1.0184288430999924
|
||||
10.0,0.5511867072847154,1.0475065275862976
|
||||
|
BIN
Binary file not shown.
+12
@@ -0,0 +1,12 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0638797470816859,1.028256350821191
|
||||
2.0,0.9404247038894229,0.9602446667440645
|
||||
3.0,0.8558459458527742,0.9207163734750433
|
||||
4.0,0.7991754125665735,0.9120825791096949
|
||||
5.0,0.7538611182460079,0.9113014386250422
|
||||
6.0,0.7000082863701714,0.9467641167588287
|
||||
7.0,0.6632694422646805,0.9526280204018394
|
||||
8.0,0.6257995438796503,0.9969304380836067
|
||||
9.0,0.5829168972041872,1.0277221163550576
|
||||
10.0,0.556718733575609,1.0591561126184987
|
||||
11.0,0.5129955758651098,1.1357925462198781
|
||||
|
BIN
Binary file not shown.
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0713691667274192,1.017243931581686
|
||||
2.0,0.9250873768771136,0.9292757629038213
|
||||
3.0,0.827078006333775,0.9095677916820233
|
||||
4.0,0.7792822652392917,0.9027277290166079
|
||||
5.0,0.7457061266457593,0.9076534837156862
|
||||
6.0,0.709590431716707,0.9116529927148924
|
||||
7.0,0.670481797169756,0.9217575978446793
|
||||
8.0,0.6348734536656627,0.9528291304032881
|
||||
9.0,0.5908853731773518,0.982715639439258
|
||||
10.0,0.5694722047558537,0.9955548689915583
|
||||
|
+11
@@ -0,0 +1,11 @@
|
||||
epoch,train_loss,valid_clean_loss
|
||||
1.0,1.0713691667274192,1.017243931581686
|
||||
2.0,0.9250873768771136,0.9292757629038213
|
||||
3.0,0.827078006333775,0.9095677916820233
|
||||
4.0,0.7792822652392917,0.9027277290166079
|
||||
5.0,0.7457061266457593,0.9076534837156862
|
||||
6.0,0.709590431716707,0.9116529927148924
|
||||
7.0,0.670481797169756,0.9217575978446793
|
||||
8.0,0.6348734536656627,0.9528291304032881
|
||||
9.0,0.5908853731773518,0.982715639439258
|
||||
10.0,0.5694722047558537,0.9955548689915583
|
||||
|
@@ -0,0 +1,54 @@
|
||||
{
|
||||
"source_feature": "/home/gloamxun/modeling_zhaocui/E题数据/附件2-数据集特征文件/aligned_50.pkl",
|
||||
"source_sha256": "66e867aa74bc70a844e806e5571e371c9abb4a35f9e2887ce9b4d97ff2cb8fcd",
|
||||
"device": "cuda",
|
||||
"cuda_name": "NVIDIA GeForce RTX 5070 Ti",
|
||||
"seeds": [
|
||||
42,
|
||||
3407,
|
||||
2026
|
||||
],
|
||||
"epochs_max": 32,
|
||||
"patience": 6,
|
||||
"batch_size": 64,
|
||||
"best_epochs": {
|
||||
"concat_seed_42": 4,
|
||||
"concat_seed_3407": 4,
|
||||
"concat_seed_2026": 3,
|
||||
"gate_seed_42": 3,
|
||||
"gate_seed_3407": 4,
|
||||
"gate_seed_2026": 3,
|
||||
"crossattn_seed_42": 2,
|
||||
"crossattn_seed_3407": 4,
|
||||
"crossattn_seed_2026": 3,
|
||||
"fixed_window_concat_seed_42": 4,
|
||||
"fixed_window_concat_seed_3407": 5,
|
||||
"fixed_window_concat_seed_2026": 4
|
||||
},
|
||||
"selected_macro_f1_first": "concat",
|
||||
"selection_policy": "report Macro-F1, MAE, and Pearson separately; selected model maximizes mean validation Macro-F1 across 15 contiguous corruption conditions, then uses MAE and lexical model name only as tie-breaks",
|
||||
"models": [
|
||||
"concat",
|
||||
"gate",
|
||||
"crossattn"
|
||||
],
|
||||
"corruption_rates": [
|
||||
0.1,
|
||||
0.2,
|
||||
0.3
|
||||
],
|
||||
"corruption_patterns": [
|
||||
"text",
|
||||
"audio",
|
||||
"vision",
|
||||
"audio_vision",
|
||||
"all_modalities"
|
||||
],
|
||||
"feature_scaling": "training split median/MAD; fallback to standard deviation for zero-MAD dimensions",
|
||||
"test_labels_used": false,
|
||||
"alignment_transfer_limit": "The official aligned_50 data use a 50-slot wordpiece sequence with no per-slot seconds or stored Q1 B1 time_bounds. The fixed-window comparison is a downstream alignment control, not a re-run of Q1 B1 on the full dataset.",
|
||||
"python": "3.14.7 (main, Aug 10 2026, 00:00:00) [GCC 16.1.1 20260515 (Red Hat 16.1.1-2)]",
|
||||
"torch": "2.14.0+cu130",
|
||||
"numpy": "2.5.3",
|
||||
"created_unix": 1790237371.5417986
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
Macro-F1-first validation selection: concat. See summary.csv for the full multi-metric tradeoff.
|
||||
@@ -0,0 +1,5 @@
|
||||
method,representation,n_seeds,clean_accuracy,clean_accuracy_sd,clean_macro_f1,clean_macro_f1_sd,clean_mae,clean_mae_sd,clean_pearson,clean_pearson_sd,corrupt_accuracy_mean,corrupt_accuracy_sd,corrupt_macro_f1_mean,corrupt_macro_f1_sd,corrupt_macro_f1_worst,corrupt_mae_mean,corrupt_mae_sd,corrupt_pearson_mean,corrupt_pearson_sd,f1_rate_10,accuracy_rate_10,mae_rate_10,f1_rate_20,accuracy_rate_20,mae_rate_20,f1_rate_30,accuracy_rate_30,mae_rate_30,pareto_nondominated
|
||||
concat,provided_word_aligned_50,3,0.6259157509157509,0.012689016905266455,0.5803030257575643,0.01849482950112713,0.6362011035283407,0.0035474183737517766,0.6131094378711319,0.003049322677391937,0.6226800976800976,0.011084117339613314,0.5751264735191365,0.016846851540436875,0.5160231153138954,0.6402280900213454,0.004374268275822229,0.6068152054284395,0.004485336212715804,0.5774116129988989,0.6244505494505495,0.6361218094825745,0.5732483562656322,0.6214285714285714,0.6398131450017294,0.5747194512928783,0.6221611721611722,0.6447493155797323,True
|
||||
gate,provided_word_aligned_50,3,0.6213369963369964,0.0106695789356511,0.570042406396691,0.0192417246427628,0.6392609675725301,0.010298306434114075,0.61843647657748,0.008153726338115605,0.6170940170940171,0.011429412792673272,0.5654484786987964,0.018209275544794713,0.4977788775985414,0.6432029167811076,0.00818900735475074,0.6097108251803484,0.008038674969140262,0.5690117947676202,0.6205128205128205,0.6386720657348633,0.5668714434894636,0.6183150183150183,0.6423242449760437,0.5604621978393048,0.6124542124542125,0.6486124396324157,True
|
||||
crossattn,provided_word_aligned_50,3,0.6039377289377289,0.011682555697960702,0.5402080959251137,0.04931996722297363,0.6486262281735738,0.0247913008377115,0.5962298100136721,0.016314693282960167,0.6013125763125764,0.00904432495782142,0.5394035484113701,0.047085193081398864,0.4615384615384615,0.6492320696512858,0.023553864046032207,0.5917838426035978,0.015944508668008214,0.5437388259395578,0.6055860805860807,0.6467597643534342,0.5379114650369052,0.5998168498168498,0.6491282820701599,0.5365603542576471,0.5985347985347985,0.6518081625302632,False
|
||||
concat,equal_window_resampled_unaligned,3,0.586996336996337,0.010491244722884272,0.5005019574320758,0.013925639871727626,0.6615431904792786,0.005397772426619935,0.581966026863303,0.009642037378255953,0.5905677655677655,0.006480515023060099,0.5042920441199125,0.015580461355158729,0.4641572706698656,0.6651763810051813,0.004962154081458186,0.5766306314815771,0.010310184996660417,0.5042141077437426,0.5906593406593407,0.6621770620346069,0.503384597978282,0.5899267399267399,0.6627050677935282,0.5052774266377128,0.5911172161172161,0.6706470131874084,True
|
||||
|
@@ -0,0 +1,199 @@
|
||||
method,representation,seed,condition,missing_rate,n_valid,accuracy,macro_f1,mae,pearson
|
||||
concat,provided_word_aligned_50,42,clean,0.0,728,0.6332417582417582,0.587941053090477,0.6350555419921875,0.6157338809096098
|
||||
concat,provided_word_aligned_50,42,text,0.1,728,0.6263736263736264,0.5758300795065502,0.6362224221229553,0.6099001459533547
|
||||
concat,provided_word_aligned_50,42,audio,0.1,728,0.6318681318681318,0.5859694579820366,0.6346138119697571,0.6186827904447096
|
||||
concat,provided_word_aligned_50,42,vision,0.1,728,0.6277472527472527,0.5825072763635762,0.6360324621200562,0.6155412177622099
|
||||
concat,provided_word_aligned_50,42,audio_vision,0.1,728,0.635989010989011,0.5919984665799197,0.6340617537498474,0.6171314707027704
|
||||
concat,provided_word_aligned_50,42,all_modalities,0.1,728,0.625,0.5767997019394343,0.6347602009773254,0.6127068856220025
|
||||
concat,provided_word_aligned_50,42,text,0.2,728,0.6112637362637363,0.5533514537785528,0.6472386121749878,0.5878231360637466
|
||||
concat,provided_word_aligned_50,42,audio,0.2,728,0.6346153846153846,0.5913737928222076,0.6337395906448364,0.6167897438460014
|
||||
concat,provided_word_aligned_50,42,vision,0.2,728,0.6277472527472527,0.5807552037464527,0.6369960308074951,0.6152903324501466
|
||||
concat,provided_word_aligned_50,42,audio_vision,0.2,728,0.6346153846153846,0.5889252158180606,0.6331195831298828,0.6187845560398619
|
||||
concat,provided_word_aligned_50,42,all_modalities,0.2,728,0.6263736263736264,0.5759393139946568,0.6418775916099548,0.6097362713984099
|
||||
concat,provided_word_aligned_50,42,text,0.3,728,0.6181318681318682,0.5669198053141907,0.6596490740776062,0.57596152605997
|
||||
concat,provided_word_aligned_50,42,audio,0.3,728,0.6414835164835165,0.5997485874560443,0.6367413997650146,0.6154119880721182
|
||||
concat,provided_word_aligned_50,42,vision,0.3,728,0.6332417582417582,0.5862230880240699,0.6370688080787659,0.61519025568542
|
||||
concat,provided_word_aligned_50,42,audio_vision,0.3,728,0.6428571428571429,0.5982133735032932,0.6347481608390808,0.6176032363345405
|
||||
concat,provided_word_aligned_50,42,all_modalities,0.3,728,0.625,0.5769720397379972,0.6341773867607117,0.6145213983185148
|
||||
concat,provided_word_aligned_50,3407,clean,0.0,728,0.6332417582417582,0.5937554950072336,0.6333680152893066,0.6138300943227009
|
||||
concat,provided_word_aligned_50,3407,text,0.1,728,0.6277472527472527,0.5848809977670061,0.6363487243652344,0.6056385975757032
|
||||
concat,provided_word_aligned_50,3407,audio,0.1,728,0.635989010989011,0.5968950441819239,0.6318458914756775,0.6158991632022357
|
||||
concat,provided_word_aligned_50,3407,vision,0.1,728,0.6332417582417582,0.5933513166880985,0.6328762769699097,0.6142965716303582
|
||||
concat,provided_word_aligned_50,3407,audio_vision,0.1,728,0.635989010989011,0.5972272747847399,0.6317789554595947,0.6131521212086648
|
||||
concat,provided_word_aligned_50,3407,all_modalities,0.1,728,0.6208791208791209,0.5757563714642701,0.6323474645614624,0.6174836749082849
|
||||
concat,provided_word_aligned_50,3407,text,0.2,728,0.6043956043956044,0.5490811521155289,0.6509508490562439,0.5895223408430469
|
||||
concat,provided_word_aligned_50,3407,audio,0.2,728,0.635989010989011,0.5977119556159886,0.6304426789283752,0.618383030295478
|
||||
concat,provided_word_aligned_50,3407,vision,0.2,728,0.6401098901098901,0.6034867151810619,0.628890335559845,0.6187934617340419
|
||||
concat,provided_word_aligned_50,3407,audio_vision,0.2,728,0.6346153846153846,0.600406263153788,0.633434534072876,0.6123211947292676
|
||||
concat,provided_word_aligned_50,3407,all_modalities,0.2,728,0.6153846153846154,0.5731546231546232,0.6435555815696716,0.6001160890533341
|
||||
concat,provided_word_aligned_50,3407,text,0.3,728,0.6071428571428571,0.5572172310238351,0.6613345146179199,0.5663984666907937
|
||||
concat,provided_word_aligned_50,3407,audio,0.3,728,0.6401098901098901,0.602821099359954,0.6328924298286438,0.6157652973695024
|
||||
concat,provided_word_aligned_50,3407,vision,0.3,728,0.6277472527472527,0.5882423141524775,0.6326207518577576,0.6153119550822198
|
||||
concat,provided_word_aligned_50,3407,audio_vision,0.3,728,0.6428571428571429,0.6062757576614914,0.6341086030006409,0.6170348507761535
|
||||
concat,provided_word_aligned_50,3407,all_modalities,0.3,728,0.6277472527472527,0.5839727159969967,0.646858811378479,0.5966126328673234
|
||||
concat,provided_word_aligned_50,2026,clean,0.0,728,0.6112637362637363,0.5592125291749822,0.6401797533035278,0.6097643383810852
|
||||
concat,provided_word_aligned_50,2026,text,0.1,728,0.6098901098901099,0.5538253146595995,0.6407675743103027,0.6035210214254456
|
||||
concat,provided_word_aligned_50,2026,audio,0.1,728,0.614010989010989,0.5630983962746592,0.6385665535926819,0.6110860092543575
|
||||
concat,provided_word_aligned_50,2026,vision,0.1,728,0.6181318681318682,0.5661140652625988,0.6405577659606934,0.6110400577216631
|
||||
concat,provided_word_aligned_50,2026,audio_vision,0.1,728,0.6098901098901099,0.5587567567850849,0.6381099224090576,0.6107705749941461
|
||||
concat,provided_word_aligned_50,2026,all_modalities,0.1,728,0.614010989010989,0.5581636747439855,0.6429373621940613,0.6088956569868328
|
||||
concat,provided_word_aligned_50,2026,text,0.2,728,0.6043956043956044,0.5429151983962918,0.6489536762237549,0.5886848356465891
|
||||
concat,provided_word_aligned_50,2026,audio,0.2,728,0.6112637362637363,0.5603026186889267,0.636631965637207,0.6105368566655438
|
||||
concat,provided_word_aligned_50,2026,vision,0.2,728,0.614010989010989,0.5607756511971072,0.6439995765686035,0.6079712044449676
|
||||
concat,provided_word_aligned_50,2026,audio_vision,0.2,728,0.6181318681318682,0.5696462960623787,0.6381081342697144,0.6114528702793328
|
||||
concat,provided_word_aligned_50,2026,all_modalities,0.2,728,0.6085164835164835,0.5508998902588577,0.6492584347724915,0.6012091713678648
|
||||
concat,provided_word_aligned_50,2026,text,0.3,728,0.5851648351648352,0.5160231153138954,0.6731752753257751,0.5521914823230472
|
||||
concat,provided_word_aligned_50,2026,audio,0.3,728,0.6181318681318682,0.5670636517410711,0.636723518371582,0.6097718270787369
|
||||
concat,provided_word_aligned_50,2026,vision,0.3,728,0.6085164835164835,0.5560149900702367,0.6435868144035339,0.6098501902389435
|
||||
concat,provided_word_aligned_50,2026,audio_vision,0.3,728,0.6098901098901099,0.5653848813301509,0.637891948223114,0.6075333619057974
|
||||
concat,provided_word_aligned_50,2026,all_modalities,0.3,728,0.6043956043956044,0.5496991187074715,0.6696622371673584,0.5843647212263229
|
||||
gate,provided_word_aligned_50,42,clean,0.0,728,0.6332417582417582,0.5807021489645882,0.649620771408081,0.6141327125880438
|
||||
gate,provided_word_aligned_50,42,text,0.1,728,0.6291208791208791,0.5724806267179149,0.6482935547828674,0.607812656143153
|
||||
gate,provided_word_aligned_50,42,audio,0.1,728,0.6318681318681318,0.5804741658290519,0.64830082654953,0.6160345444489322
|
||||
gate,provided_word_aligned_50,42,vision,0.1,728,0.635989010989011,0.5855122033916645,0.6496425867080688,0.6131362897637801
|
||||
gate,provided_word_aligned_50,42,audio_vision,0.1,728,0.6332417582417582,0.5831283881314556,0.6489386558532715,0.6140409647711148
|
||||
gate,provided_word_aligned_50,42,all_modalities,0.1,728,0.6332417582417582,0.583341398636846,0.6480620503425598,0.6083974001811914
|
||||
gate,provided_word_aligned_50,42,text,0.2,728,0.6071428571428571,0.5464225853270284,0.6594239473342896,0.5807870469081623
|
||||
gate,provided_word_aligned_50,42,audio,0.2,728,0.6373626373626373,0.5879790951691172,0.6486459970474243,0.6137318643257155
|
||||
gate,provided_word_aligned_50,42,vision,0.2,728,0.6401098901098901,0.5906020456660217,0.6478570103645325,0.6146850468047298
|
||||
gate,provided_word_aligned_50,42,audio_vision,0.2,728,0.6428571428571429,0.5956295807308858,0.6458060145378113,0.615765199266891
|
||||
gate,provided_word_aligned_50,42,all_modalities,0.2,728,0.6222527472527473,0.5699698676231126,0.6509362459182739,0.6090527347361592
|
||||
gate,provided_word_aligned_50,42,text,0.3,728,0.5947802197802198,0.5351473922902494,0.6707870364189148,0.5698602169256346
|
||||
gate,provided_word_aligned_50,42,audio,0.3,728,0.6414835164835165,0.5932626146881034,0.6488574147224426,0.6137698439998989
|
||||
gate,provided_word_aligned_50,42,vision,0.3,728,0.635989010989011,0.5859542607988605,0.6494671106338501,0.6129307301381368
|
||||
gate,provided_word_aligned_50,42,audio_vision,0.3,728,0.6373626373626373,0.5906348593126334,0.6477416753768921,0.6122889569353058
|
||||
gate,provided_word_aligned_50,42,all_modalities,0.3,728,0.6304945054945055,0.5730354311449867,0.6414216160774231,0.6210376760306293
|
||||
gate,provided_word_aligned_50,3407,clean,0.0,728,0.6181318681318682,0.5815951114421278,0.6290252804756165,0.6278403558722048
|
||||
gate,provided_word_aligned_50,3407,text,0.1,728,0.6126373626373627,0.5710061419053845,0.6293500661849976,0.6178298909468922
|
||||
gate,provided_word_aligned_50,3407,audio,0.1,728,0.6208791208791209,0.5860650310799707,0.6279832124710083,0.6279555276295478
|
||||
gate,provided_word_aligned_50,3407,vision,0.1,728,0.6181318681318682,0.5824584906426339,0.629248321056366,0.6268484455037675
|
||||
gate,provided_word_aligned_50,3407,audio_vision,0.1,728,0.6195054945054945,0.5832702882513066,0.6273886561393738,0.6257845884321586
|
||||
gate,provided_word_aligned_50,3407,all_modalities,0.1,728,0.6098901098901099,0.5711323268553031,0.6288275122642517,0.6273843507994947
|
||||
gate,provided_word_aligned_50,3407,text,0.2,728,0.5947802197802198,0.5362174090491446,0.6508536338806152,0.5974496153645171
|
||||
gate,provided_word_aligned_50,3407,audio,0.2,728,0.6236263736263736,0.5893879935672744,0.6285275220870972,0.6276866860061261
|
||||
gate,provided_word_aligned_50,3407,vision,0.2,728,0.6126373626373627,0.5806822211441783,0.6233463287353516,0.6299624942730432
|
||||
gate,provided_word_aligned_50,3407,audio_vision,0.2,728,0.6236263736263736,0.5941062537491897,0.6226418614387512,0.631152312923411
|
||||
gate,provided_word_aligned_50,3407,all_modalities,0.2,728,0.6085164835164835,0.5685926111566246,0.6388359069824219,0.6172122010198969
|
||||
gate,provided_word_aligned_50,3407,text,0.3,728,0.5892857142857143,0.5319135217224389,0.6806973814964294,0.5598124161849717
|
||||
gate,provided_word_aligned_50,3407,audio,0.3,728,0.614010989010989,0.5787631246271929,0.6292544603347778,0.6264481106992486
|
||||
gate,provided_word_aligned_50,3407,vision,0.3,728,0.6112637362637363,0.5785955746513064,0.6262590885162354,0.6269980353567415
|
||||
gate,provided_word_aligned_50,3407,audio_vision,0.3,728,0.614010989010989,0.5854704803100748,0.6231685280799866,0.627471635956635
|
||||
gate,provided_word_aligned_50,3407,all_modalities,0.3,728,0.603021978021978,0.5649390371505222,0.6471090912818909,0.6058508041636224
|
||||
gate,provided_word_aligned_50,2026,clean,0.0,728,0.6126373626373627,0.5478299587833569,0.6391368508338928,0.6133363612721914
|
||||
gate,provided_word_aligned_50,2026,text,0.1,728,0.6085164835164835,0.5374777669587797,0.6402928233146667,0.61032865906815
|
||||
gate,provided_word_aligned_50,2026,audio,0.1,728,0.614010989010989,0.5504415308771583,0.6376320719718933,0.6137133407936685
|
||||
gate,provided_word_aligned_50,2026,vision,0.1,728,0.614010989010989,0.549111309400672,0.6384209990501404,0.6116423863441457
|
||||
gate,provided_word_aligned_50,2026,audio_vision,0.1,728,0.6112637362637363,0.548206981836245,0.6367278695106506,0.6134016265067737
|
||||
gate,provided_word_aligned_50,2026,all_modalities,0.1,728,0.6153846153846154,0.5510702709999151,0.6409717798233032,0.6055737170119238
|
||||
gate,provided_word_aligned_50,2026,text,0.2,728,0.6043956043956044,0.5387500187126036,0.6619775295257568,0.5870594418667304
|
||||
gate,provided_word_aligned_50,2026,audio,0.2,728,0.614010989010989,0.5520324803225694,0.6362826824188232,0.6150266095836183
|
||||
gate,provided_word_aligned_50,2026,vision,0.2,728,0.614010989010989,0.549868355076148,0.6375797986984253,0.6131134463734892
|
||||
gate,provided_word_aligned_50,2026,audio_vision,0.2,728,0.6098901098901099,0.5462836267094535,0.6391955614089966,0.6117577661672767
|
||||
gate,provided_word_aligned_50,2026,all_modalities,0.2,728,0.6195054945054945,0.556547508338603,0.642953634262085,0.6004007468165689
|
||||
gate,provided_word_aligned_50,2026,text,0.3,728,0.5810439560439561,0.4977788775985414,0.6796172261238098,0.55000542409833
|
||||
gate,provided_word_aligned_50,2026,audio,0.3,728,0.614010989010989,0.5490560178594751,0.6381493806838989,0.6113943643034875
|
||||
gate,provided_word_aligned_50,2026,vision,0.3,728,0.614010989010989,0.5555904891120482,0.6359769105911255,0.6131701769471742
|
||||
gate,provided_word_aligned_50,2026,audio_vision,0.3,728,0.6126373626373627,0.5581596077386942,0.6396954655647278,0.6082282425073746
|
||||
gate,provided_word_aligned_50,2026,all_modalities,0.3,728,0.5934065934065934,0.5286316785844464,0.6709842085838318,0.5729928980874585
|
||||
crossattn,provided_word_aligned_50,42,clean,0.0,728,0.6043956043956044,0.48998863119202335,0.6759882569313049,0.5777291731774127
|
||||
crossattn,provided_word_aligned_50,42,text,0.1,728,0.603021978021978,0.49273388898137477,0.6812499165534973,0.5694193464166294
|
||||
crossattn,provided_word_aligned_50,42,audio,0.1,728,0.6112637362637363,0.504821770685027,0.6705242395401001,0.5815934451841169
|
||||
crossattn,provided_word_aligned_50,42,vision,0.1,728,0.6043956043956044,0.4937041426648234,0.6775157451629639,0.5754628320692556
|
||||
crossattn,provided_word_aligned_50,42,audio_vision,0.1,728,0.6112637362637363,0.5009137863541273,0.6698321104049683,0.5814767002889362
|
||||
crossattn,provided_word_aligned_50,42,all_modalities,0.1,728,0.6016483516483516,0.49390809329904145,0.6687620282173157,0.5750133906760472
|
||||
crossattn,provided_word_aligned_50,42,text,0.2,728,0.5837912087912088,0.4692083792640309,0.6945626139640808,0.5412206795833109
|
||||
crossattn,provided_word_aligned_50,42,audio,0.2,728,0.6057692307692307,0.4919398568662255,0.6648613810539246,0.5871766995932145
|
||||
crossattn,provided_word_aligned_50,42,vision,0.2,728,0.603021978021978,0.4986750074706885,0.6775456666946411,0.5770632994127847
|
||||
crossattn,provided_word_aligned_50,42,audio_vision,0.2,728,0.6112637362637363,0.506035010012657,0.6634278297424316,0.5873091839764368
|
||||
crossattn,provided_word_aligned_50,42,all_modalities,0.2,728,0.5906593406593407,0.48358818056113656,0.6767656803131104,0.5728545007104048
|
||||
crossattn,provided_word_aligned_50,42,text,0.3,728,0.5769230769230769,0.4615384615384615,0.7029426693916321,0.5265077727922097
|
||||
crossattn,provided_word_aligned_50,42,audio,0.3,728,0.6057692307692307,0.4999911736715954,0.6637417078018188,0.585929811841307
|
||||
crossattn,provided_word_aligned_50,42,vision,0.3,728,0.6002747252747253,0.488070946081782,0.6788135766983032,0.5753310885314236
|
||||
crossattn,provided_word_aligned_50,42,audio_vision,0.3,728,0.6071428571428571,0.5068731682791942,0.6635814309120178,0.5881054554275786
|
||||
crossattn,provided_word_aligned_50,42,all_modalities,0.3,728,0.592032967032967,0.48234759307110725,0.667457640171051,0.5803081998379233
|
||||
crossattn,provided_word_aligned_50,3407,clean,0.0,728,0.6153846153846154,0.5885764436484813,0.6276583671569824,0.6085564971364105
|
||||
crossattn,provided_word_aligned_50,3407,text,0.1,728,0.6112637362637363,0.5835380749854434,0.6310563683509827,0.599629174971605
|
||||
crossattn,provided_word_aligned_50,3407,audio,0.1,728,0.6222527472527473,0.5972489054492567,0.623246431350708,0.6125511992350159
|
||||
crossattn,provided_word_aligned_50,3407,vision,0.1,728,0.6098901098901099,0.583342725650418,0.6262677311897278,0.6099948694514816
|
||||
crossattn,provided_word_aligned_50,3407,audio_vision,0.1,728,0.6112637362637363,0.5867178901767943,0.623745322227478,0.6093859569553041
|
||||
crossattn,provided_word_aligned_50,3407,all_modalities,0.1,728,0.614010989010989,0.5888196778236806,0.6245514750480652,0.6117287980833325
|
||||
crossattn,provided_word_aligned_50,3407,text,0.2,728,0.6043956043956044,0.5717745964093709,0.6391856074333191,0.5848733951460722
|
||||
crossattn,provided_word_aligned_50,3407,audio,0.2,728,0.6112637362637363,0.5894389643701733,0.6257225871086121,0.6074896792266403
|
||||
crossattn,provided_word_aligned_50,3407,vision,0.2,728,0.6112637362637363,0.5851332896742695,0.6264434456825256,0.6127778365291627
|
||||
crossattn,provided_word_aligned_50,3407,audio_vision,0.2,728,0.614010989010989,0.5949110192264351,0.6187586188316345,0.6182033017472267
|
||||
crossattn,provided_word_aligned_50,3407,all_modalities,0.2,728,0.6016483516483516,0.5773507156134905,0.634911835193634,0.592528289751998
|
||||
crossattn,provided_word_aligned_50,3407,text,0.3,728,0.6002747252747253,0.5711675869572423,0.6530678868293762,0.5606507699299987
|
||||
crossattn,provided_word_aligned_50,3407,audio,0.3,728,0.614010989010989,0.5913067138879948,0.6232653856277466,0.608846448174426
|
||||
crossattn,provided_word_aligned_50,3407,vision,0.3,728,0.6071428571428571,0.5819589420485362,0.6271139979362488,0.6113427502800617
|
||||
crossattn,provided_word_aligned_50,3407,audio_vision,0.3,728,0.6098901098901099,0.5897044167286073,0.6164263486862183,0.6193569599948916
|
||||
crossattn,provided_word_aligned_50,3407,all_modalities,0.3,728,0.6181318681318682,0.5940082804540815,0.6317148804664612,0.5948538231621964
|
||||
crossattn,provided_word_aligned_50,2026,clean,0.0,728,0.592032967032967,0.5420592129348364,0.6422320604324341,0.6024037597271933
|
||||
crossattn,provided_word_aligned_50,2026,text,0.1,728,0.5975274725274725,0.5454707818440054,0.6418640613555908,0.6024112136660502
|
||||
crossattn,provided_word_aligned_50,2026,audio,0.1,728,0.5934065934065934,0.5436021768815864,0.6399821043014526,0.6049551097626092
|
||||
crossattn,provided_word_aligned_50,2026,vision,0.1,728,0.5892857142857143,0.5384039205789096,0.6415307521820068,0.6011185044495995
|
||||
crossattn,provided_word_aligned_50,2026,audio_vision,0.1,728,0.603021978021978,0.554150981347194,0.6382296681404114,0.6055650886270252
|
||||
crossattn,provided_word_aligned_50,2026,all_modalities,0.1,728,0.6002747252747253,0.5487055723716843,0.6430385112762451,0.5965045050053945
|
||||
crossattn,provided_word_aligned_50,2026,text,0.2,728,0.5906593406593407,0.5288690186841035,0.6494663953781128,0.5950657404432197
|
||||
crossattn,provided_word_aligned_50,2026,audio,0.2,728,0.5906593406593407,0.5424377950535354,0.6427614688873291,0.601656698677867
|
||||
crossattn,provided_word_aligned_50,2026,vision,0.2,728,0.5934065934065934,0.5455594737723953,0.6451690793037415,0.5999133532386576
|
||||
crossattn,provided_word_aligned_50,2026,audio_vision,0.2,728,0.5961538461538461,0.5478165327681003,0.6367671489715576,0.6098925146891432
|
||||
crossattn,provided_word_aligned_50,2026,all_modalities,0.2,728,0.5892857142857143,0.5359341358069671,0.6405748724937439,0.5950458235969649
|
||||
crossattn,provided_word_aligned_50,2026,text,0.3,728,0.5755494505494505,0.5101004599189818,0.6680996417999268,0.565804882523542
|
||||
crossattn,provided_word_aligned_50,2026,audio,0.3,728,0.603021978021978,0.5546585571147281,0.6443181037902832,0.6002440646348579
|
||||
crossattn,provided_word_aligned_50,2026,vision,0.3,728,0.5989010989010989,0.5523658108429751,0.6397479772567749,0.6069549742183601
|
||||
crossattn,provided_word_aligned_50,2026,audio_vision,0.3,728,0.5879120879120879,0.5420598732858583,0.6358648538589478,0.6059215495922518
|
||||
crossattn,provided_word_aligned_50,2026,all_modalities,0.3,728,0.5810439560439561,0.5222533299835609,0.6609663367271423,0.5802332350553648
|
||||
concat,provided_word_aligned_50,42,audio_vision_shifted_1_to_10_slots,0.0,728,0.6222527472527473,0.5650157195157686,0.6228974461555481,0.6112627581844078
|
||||
concat,provided_word_aligned_50,3407,audio_vision_shifted_1_to_10_slots,0.0,728,0.6263736263736264,0.5755677417356012,0.6278499364852905,0.6111890270289838
|
||||
concat,provided_word_aligned_50,2026,audio_vision_shifted_1_to_10_slots,0.0,728,0.5947802197802198,0.5294295564063006,0.6285219192504883,0.6104475375451327
|
||||
concat,equal_window_resampled_unaligned,42,clean,0.0,728,0.5892857142857143,0.4853955760688928,0.6553117632865906,0.587849225352828
|
||||
concat,equal_window_resampled_unaligned,42,text,0.1,728,0.5892857142857143,0.48234450535531,0.6542671918869019,0.5821794113270566
|
||||
concat,equal_window_resampled_unaligned,42,audio,0.1,728,0.5879120879120879,0.4840888292569874,0.6544017195701599,0.589859444837917
|
||||
concat,equal_window_resampled_unaligned,42,vision,0.1,728,0.5906593406593407,0.48684880290869176,0.6560104489326477,0.5873086725722289
|
||||
concat,equal_window_resampled_unaligned,42,audio_vision,0.1,728,0.5961538461538461,0.49502890457555865,0.6592482924461365,0.5895610005108267
|
||||
concat,equal_window_resampled_unaligned,42,all_modalities,0.1,728,0.5865384615384616,0.4839129480007753,0.661249577999115,0.5814457099031457
|
||||
concat,equal_window_resampled_unaligned,42,text,0.2,728,0.5741758241758241,0.4641572706698656,0.6567733883857727,0.5645263734507251
|
||||
concat,equal_window_resampled_unaligned,42,audio,0.2,728,0.5934065934065934,0.49411312080823294,0.6519993543624878,0.5919285095409706
|
||||
concat,equal_window_resampled_unaligned,42,vision,0.2,728,0.5851648351648352,0.4807564908587733,0.6555609107017517,0.5864882167667788
|
||||
concat,equal_window_resampled_unaligned,42,audio_vision,0.2,728,0.5947802197802198,0.486680557422301,0.6588707566261292,0.5950950548477731
|
||||
concat,equal_window_resampled_unaligned,42,all_modalities,0.2,728,0.5824175824175825,0.4826855628872202,0.6649248003959656,0.5778664904235934
|
||||
concat,equal_window_resampled_unaligned,42,text,0.3,728,0.5769230769230769,0.47568745986340694,0.6655409932136536,0.5549458556193816
|
||||
concat,equal_window_resampled_unaligned,42,audio,0.3,728,0.5865384615384616,0.49062463717636134,0.6570892930030823,0.5903076532449915
|
||||
concat,equal_window_resampled_unaligned,42,vision,0.3,728,0.5934065934065934,0.49023200117754157,0.6581456661224365,0.5844521656955831
|
||||
concat,equal_window_resampled_unaligned,42,audio_vision,0.3,728,0.6057692307692307,0.5036551721604455,0.6700985431671143,0.5948534685103303
|
||||
concat,equal_window_resampled_unaligned,42,all_modalities,0.3,728,0.5989010989010989,0.49385933420149347,0.6675514578819275,0.5765029580102057
|
||||
concat,equal_window_resampled_unaligned,42,audio_vision_shifted_1_to_10_slots,0.0,728,0.5961538461538461,0.5023963012985687,0.6553080081939697,0.5869747480356168
|
||||
concat,equal_window_resampled_unaligned,3407,clean,0.0,728,0.5961538461538461,0.512827084557042,0.664772629737854,0.5872103830577866
|
||||
concat,equal_window_resampled_unaligned,3407,text,0.1,728,0.6002747252747253,0.5143254275091239,0.6675248146057129,0.5790469117683982
|
||||
concat,equal_window_resampled_unaligned,3407,audio,0.1,728,0.5961538461538461,0.5141876873569328,0.6630795001983643,0.5887049023755868
|
||||
concat,equal_window_resampled_unaligned,3407,vision,0.1,728,0.592032967032967,0.5021484550452703,0.6637163758277893,0.588692970574274
|
||||
concat,equal_window_resampled_unaligned,3407,audio_vision,0.1,728,0.6002747252747253,0.5172277536075813,0.6633241176605225,0.592489640322385
|
||||
concat,equal_window_resampled_unaligned,3407,all_modalities,0.1,728,0.6043956043956044,0.5264327838402229,0.6659616231918335,0.5874171690472151
|
||||
concat,equal_window_resampled_unaligned,3407,text,0.2,728,0.5810439560439561,0.49649824913662327,0.6659781336784363,0.5747007869349695
|
||||
concat,equal_window_resampled_unaligned,3407,audio,0.2,728,0.6002747252747253,0.5210696576225557,0.662110447883606,0.5898815380715464
|
||||
concat,equal_window_resampled_unaligned,3407,vision,0.2,728,0.603021978021978,0.5134655842561412,0.6635054349899292,0.5874898942649338
|
||||
concat,equal_window_resampled_unaligned,3407,audio_vision,0.2,728,0.6043956043956044,0.5213539297376665,0.665846586227417,0.591038037009385
|
||||
concat,equal_window_resampled_unaligned,3407,all_modalities,0.2,728,0.592032967032967,0.5111232435941855,0.6659784317016602,0.5810186064947893
|
||||
concat,equal_window_resampled_unaligned,3407,text,0.3,728,0.5824175824175825,0.49292032420488613,0.6802199482917786,0.5488280644778166
|
||||
concat,equal_window_resampled_unaligned,3407,audio,0.3,728,0.603021978021978,0.5278870449050491,0.6631943583488464,0.587517131217198
|
||||
concat,equal_window_resampled_unaligned,3407,vision,0.3,728,0.6043956043956044,0.5144897492123951,0.6652408838272095,0.5849148208986014
|
||||
concat,equal_window_resampled_unaligned,3407,audio_vision,0.3,728,0.6057692307692307,0.5222120983952546,0.6709288954734802,0.5894788895188137
|
||||
concat,equal_window_resampled_unaligned,3407,all_modalities,0.3,728,0.5934065934065934,0.5118478541545217,0.6919083595275879,0.558672500491863
|
||||
concat,equal_window_resampled_unaligned,3407,audio_vision_shifted_1_to_10_slots,0.0,728,0.6057692307692307,0.5237566082920623,0.6653115153312683,0.5864640082461416
|
||||
concat,equal_window_resampled_unaligned,2026,clean,0.0,728,0.5755494505494505,0.5032832116702929,0.6645451784133911,0.5708384721792944
|
||||
concat,equal_window_resampled_unaligned,2026,text,0.1,728,0.5686813186813187,0.4953554567971044,0.6647433638572693,0.5624346756639816
|
||||
concat,equal_window_resampled_unaligned,2026,audio,0.1,728,0.592032967032967,0.524084368831769,0.6619700789451599,0.5727849021508449
|
||||
concat,equal_window_resampled_unaligned,2026,vision,0.1,728,0.5837912087912088,0.5121488634353761,0.6646548509597778,0.5731723699341335
|
||||
concat,equal_window_resampled_unaligned,2026,audio_vision,0.1,728,0.5879120879120879,0.5138457090126345,0.6634346842765808,0.5783368477304374
|
||||
concat,equal_window_resampled_unaligned,2026,all_modalities,0.1,728,0.5837912087912088,0.5112311206228025,0.6690692901611328,0.5634143872695312
|
||||
concat,equal_window_resampled_unaligned,2026,text,0.2,728,0.570054945054945,0.4938343318957674,0.6707639694213867,0.5459273537200591
|
||||
concat,equal_window_resampled_unaligned,2026,audio,0.2,728,0.5934065934065934,0.5230327607731714,0.6592981815338135,0.5755735913907287
|
||||
concat,equal_window_resampled_unaligned,2026,vision,0.2,728,0.5906593406593407,0.5211672989417647,0.6658343076705933,0.5720666083948681
|
||||
concat,equal_window_resampled_unaligned,2026,audio_vision,0.2,728,0.5989010989010989,0.5276783237584716,0.6637739539146423,0.583433626473117
|
||||
concat,equal_window_resampled_unaligned,2026,all_modalities,0.2,728,0.5851648351648352,0.5131525873114892,0.6693573594093323,0.5605391525623993
|
||||
concat,equal_window_resampled_unaligned,2026,text,0.3,728,0.5508241758241759,0.4775661528696233,0.6854801177978516,0.5184409269913169
|
||||
concat,equal_window_resampled_unaligned,2026,audio,0.3,728,0.5989010989010989,0.5333065606332371,0.6566126942634583,0.576381050540554
|
||||
concat,equal_window_resampled_unaligned,2026,vision,0.3,728,0.592032967032967,0.5197981617020956,0.6703022122383118,0.5691319241770694
|
||||
concat,equal_window_resampled_unaligned,2026,audio_vision,0.3,728,0.6016483516483516,0.5344217204682321,0.6686394214630127,0.5850381809653504
|
||||
concat,equal_window_resampled_unaligned,2026,all_modalities,0.3,728,0.5728021978021978,0.4906531284411469,0.6887523531913757,0.5344899699772931
|
||||
concat,equal_window_resampled_unaligned,2026,audio_vision_shifted_1_to_10_slots,0.0,728,0.5851648351648352,0.5171187040239319,0.6643120646476746,0.5689291630270369
|
||||
|
@@ -0,0 +1,19 @@
|
||||
[project]
|
||||
name = "deep-learning-q2-q3-selection"
|
||||
version = "0.1.0"
|
||||
requires-python = ">=3.14"
|
||||
dependencies = [
|
||||
"matplotlib>=3.11.2",
|
||||
"numpy>=2.5.3",
|
||||
"scikit-learn>=1.9.1",
|
||||
"torch>=2.14.0",
|
||||
"transformers>=5.17.0",
|
||||
]
|
||||
|
||||
[tool.uv.sources]
|
||||
torch = { index = "pytorch" }
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch"
|
||||
url = "https://download.pytorch.org/whl/cu130"
|
||||
explicit = true
|
||||
@@ -0,0 +1 @@
|
||||
"""Q2 robustness and Q3 explanation-selection experiments."""
|
||||
@@ -0,0 +1,213 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pickle
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[3]
|
||||
ATTACHMENT2 = ROOT / "E题数据" / "附件2-数据集特征文件"
|
||||
MODALITIES = ("text", "audio", "vision")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Split:
|
||||
x: tuple[np.ndarray, np.ndarray, np.ndarray]
|
||||
mask: np.ndarray # N x T x 3
|
||||
y_cls: np.ndarray
|
||||
y_reg: np.ndarray
|
||||
ids: list[str]
|
||||
|
||||
@property
|
||||
def n(self) -> int:
|
||||
return len(self.y_cls)
|
||||
|
||||
@property
|
||||
def steps(self) -> int:
|
||||
return int(self.x[0].shape[1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class RobustStats:
|
||||
center: tuple[np.ndarray, np.ndarray, np.ndarray]
|
||||
scale: tuple[np.ndarray, np.ndarray, np.ndarray]
|
||||
|
||||
def save(self, path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
np.savez_compressed(
|
||||
path,
|
||||
text_center=self.center[0], text_scale=self.scale[0],
|
||||
audio_center=self.center[1], audio_scale=self.scale[1],
|
||||
vision_center=self.center[2], vision_scale=self.scale[2],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def load(cls, path: Path) -> "RobustStats":
|
||||
with np.load(path) as data:
|
||||
return cls(
|
||||
tuple(data[f"{m}_center"].astype(np.float32) for m in MODALITIES),
|
||||
tuple(data[f"{m}_scale"].astype(np.float32) for m in MODALITIES),
|
||||
)
|
||||
|
||||
|
||||
def _unpickle(path: Path) -> dict[str, Any]:
|
||||
with path.open("rb") as stream:
|
||||
return pickle.load(stream, encoding="latin1")
|
||||
|
||||
|
||||
def _ids_and_targets(part: dict[str, Any]) -> tuple[list[str], np.ndarray, np.ndarray]:
|
||||
ids = [str(x) for x in part["id"]]
|
||||
y_cls = np.asarray(part["classification_labels"], dtype=np.int64).reshape(-1)
|
||||
y_reg = np.asarray(part["regression_labels"], dtype=np.float32).reshape(-1)
|
||||
return ids, y_cls, y_reg
|
||||
|
||||
|
||||
def _text_mask(part: dict[str, Any]) -> np.ndarray:
|
||||
tokens = np.asarray(part["text_bert"])
|
||||
if tokens.ndim != 3 or tokens.shape[1] < 2:
|
||||
raise ValueError(f"unexpected text_bert shape: {tokens.shape}")
|
||||
# MOSEI text_bert rows are input_ids, input_mask, segment_ids.
|
||||
return tokens[:, 1, :].astype(bool)
|
||||
|
||||
|
||||
def load_aligned(path: Path | None = None) -> dict[str, Split]:
|
||||
path = path or ATTACHMENT2 / "aligned_50.pkl"
|
||||
raw = _unpickle(path)
|
||||
result: dict[str, Split] = {}
|
||||
for name in ("train", "valid"):
|
||||
part = raw[name]
|
||||
xs = tuple(np.asarray(part[m], dtype=np.float32) for m in MODALITIES)
|
||||
masks = [
|
||||
_text_mask(part),
|
||||
np.any(np.isfinite(xs[1]) & (xs[1] != 0), axis=-1),
|
||||
np.any(np.isfinite(xs[2]) & (xs[2] != 0), axis=-1),
|
||||
]
|
||||
mask = np.stack(masks, axis=-1)
|
||||
ids, y_cls, y_reg = _ids_and_targets(part)
|
||||
if any(x.shape[1] != 50 for x in xs):
|
||||
raise ValueError(f"{name} aligned feature tensors must have 50 slots")
|
||||
result[name] = Split(xs, mask, y_cls, y_reg, ids)
|
||||
train_videos = {x.split("$_$", 1)[0] for x in result["train"].ids}
|
||||
valid_videos = {x.split("$_$", 1)[0] for x in result["valid"].ids}
|
||||
overlap = train_videos & valid_videos
|
||||
if overlap:
|
||||
raise ValueError(f"official train/valid split leaks {len(overlap)} source video ids")
|
||||
return result
|
||||
|
||||
|
||||
def _resample_rows_to_50(values: np.ndarray, lengths: list[int] | np.ndarray) -> tuple[np.ndarray, np.ndarray]:
|
||||
n, source_steps, dim = values.shape
|
||||
output = np.zeros((n, 50, dim), dtype=np.float32)
|
||||
mask = np.zeros((n, 50), dtype=bool)
|
||||
lengths_arr = np.asarray(lengths, dtype=np.int64).reshape(-1)
|
||||
for i in range(n):
|
||||
length = int(np.clip(lengths_arr[i], 0, source_steps))
|
||||
if length == 0:
|
||||
continue
|
||||
source = np.nan_to_num(values[i, :length], nan=0.0, posinf=0.0, neginf=0.0)
|
||||
observed = np.any(source != 0, axis=-1)
|
||||
for j in range(50):
|
||||
left = int(np.floor(j * length / 50))
|
||||
right = max(left + 1, int(np.ceil((j + 1) * length / 50)))
|
||||
right = min(right, length)
|
||||
use = observed[left:right]
|
||||
if use.any():
|
||||
output[i, j] = source[left:right][use].mean(axis=0)
|
||||
mask[i, j] = True
|
||||
return output, mask
|
||||
|
||||
|
||||
def load_fixed_window(path: Path | None = None) -> dict[str, Split]:
|
||||
"""Build a matched 50-slot equal-window control from the unaligned file."""
|
||||
path = path or ATTACHMENT2 / "unaligned_50.pkl"
|
||||
raw = _unpickle(path)
|
||||
result: dict[str, Split] = {}
|
||||
for name in ("train", "valid"):
|
||||
part = raw[name]
|
||||
text = np.asarray(part["text"], dtype=np.float32)
|
||||
audio, audio_mask = _resample_rows_to_50(part["audio"], part["audio_lengths"])
|
||||
vision, vision_mask = _resample_rows_to_50(part["vision"], part["vision_lengths"])
|
||||
text_mask = _text_mask(part)
|
||||
xs = (text, audio, vision)
|
||||
mask = np.stack((text_mask, audio_mask, vision_mask), axis=-1)
|
||||
ids, y_cls, y_reg = _ids_and_targets(part)
|
||||
result[name] = Split(xs, mask, y_cls, y_reg, ids)
|
||||
return result
|
||||
|
||||
|
||||
def fit_robust_stats(split: Split) -> RobustStats:
|
||||
centers: list[np.ndarray] = []
|
||||
scales: list[np.ndarray] = []
|
||||
for modality in range(3):
|
||||
observed = split.mask[:, :, modality].reshape(-1)
|
||||
values = split.x[modality].reshape(-1, split.x[modality].shape[-1])[observed]
|
||||
if not len(values):
|
||||
raise ValueError(f"no observed values for {MODALITIES[modality]}")
|
||||
values = np.nan_to_num(values, nan=0.0, posinf=0.0, neginf=0.0)
|
||||
center = np.median(values, axis=0)
|
||||
mad = np.median(np.abs(values - center), axis=0)
|
||||
scale = 1.4826 * mad
|
||||
std = np.std(values, axis=0)
|
||||
scale = np.where(scale > 1e-6, scale, std)
|
||||
scale = np.where(scale > 1e-6, scale, 1.0)
|
||||
centers.append(center.astype(np.float32))
|
||||
scales.append(scale.astype(np.float32))
|
||||
return RobustStats(tuple(centers), tuple(scales))
|
||||
|
||||
|
||||
def apply_robust_stats(split: Split, stats: RobustStats) -> Split:
|
||||
xs: list[np.ndarray] = []
|
||||
for modality in range(3):
|
||||
values = (split.x[modality] - stats.center[modality]) / stats.scale[modality]
|
||||
values = np.nan_to_num(values, nan=0.0, posinf=0.0, neginf=0.0)
|
||||
values *= split.mask[:, :, modality, None]
|
||||
xs.append(values.astype(np.float32, copy=False))
|
||||
return Split(tuple(xs), split.mask.copy(), split.y_cls, split.y_reg, split.ids)
|
||||
|
||||
|
||||
def corrupt_masks(
|
||||
base: np.ndarray,
|
||||
ratio: float,
|
||||
modalities: tuple[int, ...],
|
||||
seed: int,
|
||||
) -> np.ndarray:
|
||||
result = base.copy()
|
||||
rng = np.random.default_rng(seed)
|
||||
n, steps, _ = result.shape
|
||||
width = max(1, min(steps, int(round(ratio * steps))))
|
||||
starts = rng.integers(0, steps - width + 1, size=n)
|
||||
for row, start in enumerate(starts.tolist()):
|
||||
result[row, start:start + width, list(modalities)] = False
|
||||
return result
|
||||
|
||||
|
||||
def augment_masks(base: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
||||
result = base.copy()
|
||||
n, steps, _ = result.shape
|
||||
for row in range(n):
|
||||
if rng.random() >= 0.85:
|
||||
continue
|
||||
count = int(rng.integers(1, 4))
|
||||
modalities = rng.choice(3, size=count, replace=False)
|
||||
ratio = float(rng.choice((0.10, 0.20, 0.30)))
|
||||
width = max(1, int(round(ratio * steps)))
|
||||
start = int(rng.integers(0, steps - width + 1))
|
||||
result[row, start:start + width, modalities] = False
|
||||
return result
|
||||
|
||||
|
||||
def shift_audio_vision(split: Split, seed: int, max_shift: int = 10) -> Split:
|
||||
rng = np.random.default_rng(seed)
|
||||
xs = [x.copy() for x in split.x]
|
||||
masks = split.mask.copy()
|
||||
for row in range(split.n):
|
||||
for modality in (1, 2):
|
||||
shift = int(rng.integers(1, max_shift + 1))
|
||||
if rng.random() < 0.5:
|
||||
shift = -shift
|
||||
xs[modality][row] = np.roll(xs[modality][row], shift, axis=0)
|
||||
masks[row, :, modality] = np.roll(masks[row, :, modality], shift)
|
||||
return Split(tuple(xs), masks, split.y_cls, split.y_reg, split.ids)
|
||||
@@ -0,0 +1,29 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
from pathlib import Path
|
||||
|
||||
from .train_compare import _plot, _summary, _write_csv
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Rebuild Q2 summary tables from saved validation predictions")
|
||||
parser.add_argument("--output-dir", default=str(Path(__file__).resolve().parents[1] / "outputs" / "algorithm_selection"))
|
||||
args = parser.parse_args()
|
||||
output = Path(args.output_dir)
|
||||
with (output / "validation_metrics_by_condition.csv").open(encoding="utf-8-sig", newline="") as stream:
|
||||
rows = list(csv.DictReader(stream))
|
||||
for row in rows:
|
||||
for key in ("missing_rate", "accuracy", "macro_f1", "mae", "pearson", "n_valid"):
|
||||
row[key] = float(row[key])
|
||||
row["seed"] = int(row["seed"])
|
||||
summary = _summary(rows)
|
||||
_write_csv(output / "summary.csv", summary)
|
||||
aligned = [row for row in summary if row["representation"] == "provided_word_aligned_50"]
|
||||
_plot(aligned, rows, output / "missing_rate_comparison.png")
|
||||
print(f"rebuilt summary table and plot from {len(rows)} saved validation rows")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,109 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class AlignedFusionModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
kind: str,
|
||||
dims: tuple[int, int, int],
|
||||
steps: int = 50,
|
||||
hidden: int = 128,
|
||||
dropout: float = 0.15,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if kind not in {"concat", "gate", "crossattn"}:
|
||||
raise ValueError(f"unknown model kind: {kind}")
|
||||
self.kind = kind
|
||||
self.hidden = hidden
|
||||
self.projections = nn.ModuleList(
|
||||
nn.Sequential(nn.Linear(size, hidden), nn.GELU(), nn.LayerNorm(hidden))
|
||||
for size in dims
|
||||
)
|
||||
self.position = nn.Parameter(torch.randn(1, steps, hidden) * 0.02)
|
||||
self.modality = nn.Parameter(torch.randn(1, 1, 3, hidden) * 0.02)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
if kind == "concat":
|
||||
self.fusion = nn.Sequential(
|
||||
nn.Linear(hidden * 3 + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout)
|
||||
)
|
||||
elif kind == "gate":
|
||||
self.gate_score = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.Tanh(), nn.Linear(hidden // 2, 1))
|
||||
self.fusion = nn.Sequential(
|
||||
nn.Linear(hidden + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout)
|
||||
)
|
||||
else:
|
||||
layer = nn.TransformerEncoderLayer(
|
||||
d_model=hidden,
|
||||
nhead=4,
|
||||
dim_feedforward=hidden * 2,
|
||||
dropout=dropout,
|
||||
activation="gelu",
|
||||
batch_first=True,
|
||||
norm_first=True,
|
||||
)
|
||||
self.cross_encoder = nn.TransformerEncoder(layer, num_layers=2, enable_nested_tensor=False)
|
||||
self.fusion = nn.Sequential(
|
||||
nn.Linear(hidden + 3, hidden), nn.GELU(), nn.LayerNorm(hidden), nn.Dropout(dropout)
|
||||
)
|
||||
|
||||
self.temporal = nn.GRU(
|
||||
input_size=hidden,
|
||||
hidden_size=hidden // 2,
|
||||
num_layers=1,
|
||||
batch_first=True,
|
||||
bidirectional=True,
|
||||
)
|
||||
self.head = nn.Sequential(nn.Linear(hidden, hidden // 2), nn.GELU(), nn.Dropout(dropout))
|
||||
self.classifier = nn.Linear(hidden // 2, 3)
|
||||
self.regressor = nn.Linear(hidden // 2, 1)
|
||||
|
||||
def forward(self, xs: tuple[torch.Tensor, torch.Tensor, torch.Tensor], masks: torch.Tensor):
|
||||
masks = masks.bool()
|
||||
pos = self.position[:, :masks.shape[1]]
|
||||
encoded = []
|
||||
for modality, (projection, x) in enumerate(zip(self.projections, xs)):
|
||||
token = projection(x)
|
||||
token = token + pos + self.modality[:, :, modality, :]
|
||||
token = token * masks[:, :, modality, None]
|
||||
encoded.append(token)
|
||||
stack = torch.stack(encoded, dim=2) # B x T x M x D
|
||||
availability = masks.to(stack.dtype)
|
||||
gate_weights = None
|
||||
|
||||
if self.kind == "concat":
|
||||
fused = self.fusion(torch.cat((stack.flatten(2), availability), dim=-1))
|
||||
elif self.kind == "gate":
|
||||
scores = self.gate_score(stack).squeeze(-1)
|
||||
scores = scores.masked_fill(~masks, -1e4)
|
||||
gate_weights = torch.softmax(scores, dim=-1) * availability
|
||||
gate_weights = gate_weights / gate_weights.sum(dim=-1, keepdim=True).clamp_min(1e-8)
|
||||
weighted = (stack * gate_weights[..., None]).sum(dim=2)
|
||||
fused = self.fusion(torch.cat((weighted, availability), dim=-1))
|
||||
else:
|
||||
batch, steps, modalities, hidden = stack.shape
|
||||
flat = stack.reshape(batch, steps * modalities, hidden)
|
||||
valid = masks.reshape(batch, steps * modalities).clone()
|
||||
empty = ~valid.any(dim=1)
|
||||
if empty.any():
|
||||
valid[empty, 0] = True
|
||||
flat[empty, 0] = 0.0
|
||||
attended = self.cross_encoder(flat, src_key_padding_mask=~valid)
|
||||
attended = attended.reshape(batch, steps, modalities, hidden)
|
||||
observed_count = availability.sum(dim=2, keepdim=True)
|
||||
pooled = (attended * availability[..., None]).sum(dim=2) / observed_count.clamp_min(1.0)
|
||||
fused = self.fusion(torch.cat((pooled, availability), dim=-1))
|
||||
|
||||
temporal, _ = self.temporal(self.dropout(fused))
|
||||
time_weight = masks.any(dim=-1).to(temporal.dtype)
|
||||
empty_time = time_weight.sum(dim=1, keepdim=True) <= 0
|
||||
if empty_time.any():
|
||||
time_weight[empty_time.squeeze(1), 0] = 1.0
|
||||
pooled = (temporal * time_weight[..., None]).sum(dim=1) / time_weight.sum(dim=1, keepdim=True).clamp_min(1.0)
|
||||
hidden = self.head(pooled)
|
||||
logits = self.classifier(hidden)
|
||||
intensity = 3.0 * torch.tanh(self.regressor(hidden).squeeze(-1))
|
||||
return {"logits": logits, "intensity": intensity, "gate": gate_weights}
|
||||
@@ -0,0 +1,490 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import shutil
|
||||
import time
|
||||
from collections import Counter
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error
|
||||
from torch import nn
|
||||
|
||||
from .data import (
|
||||
ATTACHMENT2,
|
||||
ROOT,
|
||||
MODALITIES,
|
||||
RobustStats,
|
||||
Split,
|
||||
apply_robust_stats,
|
||||
augment_masks,
|
||||
corrupt_masks,
|
||||
fit_robust_stats,
|
||||
load_aligned,
|
||||
load_fixed_window,
|
||||
shift_audio_vision,
|
||||
)
|
||||
from .models import AlignedFusionModel
|
||||
|
||||
|
||||
PATTERNS = {
|
||||
"text": (0,),
|
||||
"audio": (1,),
|
||||
"vision": (2,),
|
||||
"audio_vision": (1, 2),
|
||||
"all_modalities": (0, 1, 2),
|
||||
}
|
||||
KINDS = ("concat", "gate", "crossattn")
|
||||
|
||||
|
||||
def seed_everything(seed: int) -> None:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
torch.backends.cudnn.benchmark = False
|
||||
|
||||
|
||||
def _tensor_split(split: Split, device: torch.device) -> tuple[tuple[torch.Tensor, ...], torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
xs = tuple(torch.as_tensor(x, dtype=torch.float32, device=device) for x in split.x)
|
||||
mask = torch.as_tensor(split.mask, dtype=torch.bool, device=device)
|
||||
y_cls = torch.as_tensor(split.y_cls, dtype=torch.long, device=device)
|
||||
y_reg = torch.as_tensor(split.y_reg, dtype=torch.float32, device=device)
|
||||
return xs, mask, y_cls, y_reg
|
||||
|
||||
|
||||
def _loss(output: dict[str, torch.Tensor], y_cls: torch.Tensor, y_reg: torch.Tensor) -> torch.Tensor:
|
||||
class_loss = F.cross_entropy(output["logits"], y_cls)
|
||||
intensity_loss = F.smooth_l1_loss(output["intensity"] / 3.0, y_reg / 3.0)
|
||||
return class_loss + 0.5 * intensity_loss
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def _score_arrays(
|
||||
model: AlignedFusionModel,
|
||||
split: Split,
|
||||
mask: np.ndarray,
|
||||
device: torch.device,
|
||||
batch_size: int = 128,
|
||||
) -> tuple[dict[str, float], dict[str, np.ndarray]]:
|
||||
model.eval()
|
||||
predictions: dict[str, list[np.ndarray]] = {"logits": [], "intensity": []}
|
||||
xs = split.x
|
||||
for start in range(0, split.n, batch_size):
|
||||
end = min(start + batch_size, split.n)
|
||||
xb = tuple(torch.as_tensor(x[start:end], dtype=torch.float32, device=device) for x in xs)
|
||||
mb = torch.as_tensor(mask[start:end], dtype=torch.bool, device=device)
|
||||
output = model(xb, mb)
|
||||
predictions["logits"].append(output["logits"].float().cpu().numpy())
|
||||
predictions["intensity"].append(output["intensity"].float().cpu().numpy())
|
||||
logits = np.concatenate(predictions["logits"], axis=0)
|
||||
intensity = np.clip(np.concatenate(predictions["intensity"], axis=0), -3.0, 3.0)
|
||||
pred_cls = logits.argmax(axis=-1)
|
||||
pearson = _pearson(split.y_reg, intensity)
|
||||
metrics = {
|
||||
"accuracy": float(accuracy_score(split.y_cls, pred_cls)),
|
||||
"macro_f1": float(f1_score(split.y_cls, pred_cls, labels=[0, 1, 2], average="macro", zero_division=0)),
|
||||
"mae": float(mean_absolute_error(split.y_reg, intensity)),
|
||||
"pearson": pearson,
|
||||
}
|
||||
return metrics, {"logits": logits, "intensity": intensity, "class": pred_cls}
|
||||
|
||||
|
||||
def _pearson(y: np.ndarray, pred: np.ndarray) -> float:
|
||||
a = np.asarray(y, dtype=np.float64)
|
||||
b = np.asarray(pred, dtype=np.float64)
|
||||
if a.std() < 1e-12 or b.std() < 1e-12:
|
||||
return 0.0
|
||||
return float(np.corrcoef(a, b)[0, 1])
|
||||
|
||||
|
||||
def _validation_loss(model: AlignedFusionModel, valid: Split, device: torch.device, batch_size: int) -> float:
|
||||
model.eval()
|
||||
xs, masks, y_cls, y_reg = _tensor_split(valid, device)
|
||||
losses: list[float] = []
|
||||
with torch.inference_mode():
|
||||
for start in range(0, valid.n, batch_size):
|
||||
idx = slice(start, min(start + batch_size, valid.n))
|
||||
output = model(tuple(x[idx] for x in xs), masks[idx])
|
||||
losses.append(float(_loss(output, y_cls[idx], y_reg[idx]).item()))
|
||||
return float(np.average(losses, weights=[min(batch_size, valid.n - i) for i in range(0, valid.n, batch_size)]))
|
||||
|
||||
|
||||
def _write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if not rows:
|
||||
return
|
||||
fields = list(dict.fromkeys(key for row in rows for key in row))
|
||||
with path.open("w", newline="", encoding="utf-8-sig") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=fields)
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def _train_one(
|
||||
kind: str,
|
||||
train: Split,
|
||||
valid: Split,
|
||||
output_dir: Path,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
epochs: int,
|
||||
patience: int,
|
||||
batch_size: int,
|
||||
) -> tuple[AlignedFusionModel, int, list[dict[str, float]]]:
|
||||
seed_everything(seed)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
model = AlignedFusionModel(kind, dims=dims).to(device)
|
||||
optimizer = torch.optim.AdamW(model.parameters(), lr=1.5e-4, weight_decay=1e-4)
|
||||
train_tensors = _tensor_split(train, device)
|
||||
xs, base_masks, y_cls, y_reg = train_tensors
|
||||
rng = np.random.default_rng(seed + 809)
|
||||
best_loss = math.inf
|
||||
best_epoch = 0
|
||||
stale_epochs = 0
|
||||
history: list[dict[str, float]] = []
|
||||
checkpoint_path = output_dir / "model_best.pt"
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
for epoch in range(1, epochs + 1):
|
||||
model.train()
|
||||
order = rng.permutation(train.n)
|
||||
batch_losses: list[float] = []
|
||||
for start in range(0, train.n, batch_size):
|
||||
ids_np = order[start:start + batch_size]
|
||||
ids = torch.as_tensor(ids_np, dtype=torch.long, device=device)
|
||||
masks_np = augment_masks(train.mask[ids_np], rng)
|
||||
masks = torch.as_tensor(masks_np, dtype=torch.bool, device=device)
|
||||
output = model(tuple(x.index_select(0, ids) for x in xs), masks)
|
||||
loss = _loss(output, y_cls.index_select(0, ids), y_reg.index_select(0, ids))
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
||||
optimizer.step()
|
||||
batch_losses.append(float(loss.detach().item()))
|
||||
valid_loss = _validation_loss(model, valid, device, batch_size)
|
||||
row = {"epoch": float(epoch), "train_loss": float(np.mean(batch_losses)), "valid_clean_loss": valid_loss}
|
||||
history.append(row)
|
||||
print(f"[{kind}] epoch={epoch:02d} train={row['train_loss']:.4f} valid={valid_loss:.4f}", flush=True)
|
||||
if valid_loss < best_loss - 1e-4:
|
||||
best_loss = valid_loss
|
||||
best_epoch = epoch
|
||||
stale_epochs = 0
|
||||
torch.save({"kind": kind, "dims": dims, "state_dict": model.state_dict(), "seed": seed, "best_epoch": epoch}, checkpoint_path)
|
||||
else:
|
||||
stale_epochs += 1
|
||||
if stale_epochs >= patience:
|
||||
break
|
||||
|
||||
saved = torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
model.load_state_dict(saved["state_dict"])
|
||||
model.eval()
|
||||
_write_csv(output_dir / "training_history.csv", history)
|
||||
return model, best_epoch, history
|
||||
|
||||
|
||||
def _conditions(valid: Split, seed: int) -> list[tuple[str, float, np.ndarray]]:
|
||||
result = [("clean", 0.0, valid.mask.copy())]
|
||||
for rate in (0.10, 0.20, 0.30):
|
||||
for pattern_id, (pattern, mods) in enumerate(PATTERNS.items()):
|
||||
result.append((pattern, rate, corrupt_masks(valid.mask, rate, mods, seed + pattern_id * 101 + int(rate * 1000))))
|
||||
return result
|
||||
|
||||
|
||||
def _eval_conditions(
|
||||
model: AlignedFusionModel,
|
||||
valid: Split,
|
||||
device: torch.device,
|
||||
seed: int,
|
||||
seed_run: int,
|
||||
method: str,
|
||||
representation: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
rows = []
|
||||
for condition, rate, masks in _conditions(valid, seed):
|
||||
metrics, _ = _score_arrays(model, valid, masks, device)
|
||||
rows.append({"method": method, "representation": representation, "seed": seed_run, "condition": condition,
|
||||
"missing_rate": rate, "n_valid": valid.n, **metrics})
|
||||
print(f"[{method}/{representation}] {condition:14s} rate={rate:.1f} "
|
||||
f"F1={metrics['macro_f1']:.3f} MAE={metrics['mae']:.3f} "
|
||||
f"P={metrics['pearson']:.3f}", flush=True)
|
||||
return rows
|
||||
|
||||
|
||||
def _summary(rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
groups = list(dict.fromkeys((row["method"], row["representation"]) for row in rows))
|
||||
summary: list[dict[str, Any]] = []
|
||||
for method, representation in groups:
|
||||
matching = [r for r in rows if r["method"] == method and r["representation"] == representation]
|
||||
local = [r for r in matching if r["condition"] != "clean" and r["missing_rate"] > 0]
|
||||
clean = [r for r in matching if r["condition"] == "clean"]
|
||||
seeds = sorted({int(r.get("seed", 0)) for r in matching})
|
||||
|
||||
def per_seed_mean(selected: list[dict[str, Any]], metric: str) -> list[float]:
|
||||
return [float(np.mean([r[metric] for r in selected if int(r.get("seed", 0)) == seed]))
|
||||
for seed in seeds if any(int(r.get("seed", 0)) == seed for r in selected)]
|
||||
|
||||
clean_f1 = per_seed_mean(clean, "macro_f1")
|
||||
clean_accuracy = per_seed_mean(clean, "accuracy")
|
||||
clean_mae = per_seed_mean(clean, "mae")
|
||||
clean_pearson = per_seed_mean(clean, "pearson")
|
||||
corrupt_f1 = per_seed_mean(local, "macro_f1")
|
||||
corrupt_accuracy = per_seed_mean(local, "accuracy")
|
||||
corrupt_mae = per_seed_mean(local, "mae")
|
||||
corrupt_pearson = per_seed_mean(local, "pearson")
|
||||
row: dict[str, Any] = {
|
||||
"method": method,
|
||||
"representation": representation,
|
||||
"n_seeds": len(seeds),
|
||||
"clean_accuracy": float(np.mean(clean_accuracy)),
|
||||
"clean_accuracy_sd": float(np.std(clean_accuracy, ddof=1)) if len(clean_accuracy) > 1 else 0.0,
|
||||
"clean_macro_f1": float(np.mean(clean_f1)),
|
||||
"clean_macro_f1_sd": float(np.std(clean_f1, ddof=1)) if len(clean_f1) > 1 else 0.0,
|
||||
"clean_mae": float(np.mean(clean_mae)),
|
||||
"clean_mae_sd": float(np.std(clean_mae, ddof=1)) if len(clean_mae) > 1 else 0.0,
|
||||
"clean_pearson": float(np.mean(clean_pearson)),
|
||||
"clean_pearson_sd": float(np.std(clean_pearson, ddof=1)) if len(clean_pearson) > 1 else 0.0,
|
||||
"corrupt_accuracy_mean": float(np.mean(corrupt_accuracy)),
|
||||
"corrupt_accuracy_sd": float(np.std(corrupt_accuracy, ddof=1)) if len(corrupt_accuracy) > 1 else 0.0,
|
||||
"corrupt_macro_f1_mean": float(np.mean(corrupt_f1)),
|
||||
"corrupt_macro_f1_sd": float(np.std(corrupt_f1, ddof=1)) if len(corrupt_f1) > 1 else 0.0,
|
||||
"corrupt_macro_f1_worst": float(np.min([r["macro_f1"] for r in local])),
|
||||
"corrupt_mae_mean": float(np.mean(corrupt_mae)),
|
||||
"corrupt_mae_sd": float(np.std(corrupt_mae, ddof=1)) if len(corrupt_mae) > 1 else 0.0,
|
||||
"corrupt_pearson_mean": float(np.mean(corrupt_pearson)),
|
||||
"corrupt_pearson_sd": float(np.std(corrupt_pearson, ddof=1)) if len(corrupt_pearson) > 1 else 0.0,
|
||||
}
|
||||
for rate in (0.10, 0.20, 0.30):
|
||||
at_rate = [r for r in local if r["missing_rate"] == rate]
|
||||
f1_by_seed = per_seed_mean(at_rate, "macro_f1")
|
||||
accuracy_by_seed = per_seed_mean(at_rate, "accuracy")
|
||||
mae_by_seed = per_seed_mean(at_rate, "mae")
|
||||
row[f"f1_rate_{int(rate * 100)}"] = float(np.mean(f1_by_seed))
|
||||
row[f"accuracy_rate_{int(rate * 100)}"] = float(np.mean(accuracy_by_seed))
|
||||
row[f"mae_rate_{int(rate * 100)}"] = float(np.mean(mae_by_seed))
|
||||
summary.append(row)
|
||||
for row in summary:
|
||||
row["pareto_nondominated"] = not any(
|
||||
other is not row and other["representation"] == row["representation"]
|
||||
and other["corrupt_macro_f1_mean"] >= row["corrupt_macro_f1_mean"]
|
||||
and other["corrupt_mae_mean"] <= row["corrupt_mae_mean"]
|
||||
and other["corrupt_pearson_mean"] >= row["corrupt_pearson_mean"]
|
||||
and (
|
||||
other["corrupt_macro_f1_mean"] > row["corrupt_macro_f1_mean"]
|
||||
or other["corrupt_mae_mean"] < row["corrupt_mae_mean"]
|
||||
or other["corrupt_pearson_mean"] > row["corrupt_pearson_mean"]
|
||||
)
|
||||
for other in summary
|
||||
)
|
||||
return summary
|
||||
|
||||
|
||||
def _plot(summary: list[dict[str, Any]], rows: list[dict[str, Any]], path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
colors = {"concat": "#4e79a7", "gate": "#f28e2b", "crossattn": "#59a14f"}
|
||||
fig, axes = plt.subplots(1, 2, figsize=(11, 4.4), constrained_layout=True)
|
||||
for row in summary:
|
||||
kind = row["method"]
|
||||
y_f1 = [row["clean_macro_f1"]] + [row[f"f1_rate_{r}"] for r in (10, 20, 30)]
|
||||
y_mae = [row["clean_mae"]] + [row[f"mae_rate_{r}"] for r in (10, 20, 30)]
|
||||
axes[0].plot([0, 10, 20, 30], y_f1, marker="o", label=kind, color=colors.get(kind))
|
||||
axes[1].plot([0, 10, 20, 30], y_mae, marker="o", label=kind, color=colors.get(kind))
|
||||
axes[0].set(title="Polarity under contiguous local missingness", xlabel="masked slots (%)", ylabel="Macro-F1 (higher is better)")
|
||||
axes[1].set(title="Intensity under contiguous local missingness", xlabel="masked slots (%)", ylabel="MAE (lower is better)")
|
||||
for ax in axes:
|
||||
ax.grid(alpha=0.25)
|
||||
ax.legend(frameon=False)
|
||||
fig.savefig(path, dpi=180)
|
||||
plt.close(fig)
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as stream:
|
||||
for block in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _run(args: argparse.Namespace) -> None:
|
||||
seed_everything(args.seeds[0])
|
||||
if args.device == "auto":
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
else:
|
||||
device = torch.device(args.device)
|
||||
torch.set_num_threads(args.threads)
|
||||
output = Path(args.output_dir)
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
aligned_raw = load_aligned()
|
||||
stats = fit_robust_stats(aligned_raw["train"])
|
||||
stats.save(output / "aligned_robust_stats.npz")
|
||||
aligned = {k: apply_robust_stats(v, stats) for k, v in aligned_raw.items()}
|
||||
audit = {
|
||||
"source": str(ATTACHMENT2 / "aligned_50.pkl"),
|
||||
"train_samples": aligned["train"].n,
|
||||
"valid_samples": aligned["valid"].n,
|
||||
"train_classes": np.bincount(aligned["train"].y_cls, minlength=3).tolist(),
|
||||
"valid_classes": np.bincount(aligned["valid"].y_cls, minlength=3).tolist(),
|
||||
"mean_observed_slots": {
|
||||
MODALITIES[m]: float(aligned["train"].mask[:, :, m].sum(axis=1).mean()) for m in range(3)
|
||||
},
|
||||
"train_valid_video_overlap": 0,
|
||||
}
|
||||
with (output / "data_audit.json").open("w", encoding="utf-8") as stream:
|
||||
json.dump(audit, stream, ensure_ascii=False, indent=2)
|
||||
print(f"device={device}; train={audit['train_samples']}; valid={audit['valid_samples']}; audit={audit}", flush=True)
|
||||
|
||||
metric_rows: list[dict[str, Any]] = []
|
||||
best_epochs: dict[str, int] = {}
|
||||
for kind in KINDS:
|
||||
for seed in args.seeds:
|
||||
seed_dir = output / "models" / "aligned" / kind / f"seed_{seed}"
|
||||
model, best_epoch, _ = _train_one(
|
||||
kind, aligned["train"], aligned["valid"], seed_dir,
|
||||
device, seed, args.epochs, args.patience, args.batch_size,
|
||||
)
|
||||
best_epochs[f"{kind}_seed_{seed}"] = best_epoch
|
||||
metric_rows.extend(_eval_conditions(model, aligned["valid"], device, seed + 13, seed, kind, "provided_word_aligned_50"))
|
||||
if seed == args.seeds[0]:
|
||||
shutil.copy2(seed_dir / "model_best.pt", output / "models" / "aligned" / kind / "model_best.pt")
|
||||
del model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
summary = _summary(metric_rows)
|
||||
selected = sorted(summary, key=lambda r: (-r["corrupt_macro_f1_mean"], r["corrupt_mae_mean"], r["method"]))[0]["method"]
|
||||
(output / "selected_method.txt").write_text(
|
||||
f"Macro-F1-first validation selection: {selected}. See summary.csv for the full multi-metric tradeoff.\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
# Matched audio/vision temporal-shift control for the selected architecture and every seed.
|
||||
for seed in args.seeds:
|
||||
aligned_payload = torch.load(output / "models" / "aligned" / selected / f"seed_{seed}" / "model_best.pt",
|
||||
map_location=device, weights_only=False)
|
||||
aligned_model = AlignedFusionModel(selected, tuple(aligned_payload["dims"])).to(device)
|
||||
aligned_model.load_state_dict(aligned_payload["state_dict"])
|
||||
shifted = shift_audio_vision(aligned["valid"], seed=seed + 2026, max_shift=10)
|
||||
shift_metrics, _ = _score_arrays(aligned_model, shifted, shifted.mask, device)
|
||||
metric_rows.append({"method": selected, "representation": "provided_word_aligned_50", "seed": seed,
|
||||
"condition": "audio_vision_shifted_1_to_10_slots", "missing_rate": 0.0,
|
||||
"n_valid": shifted.n, **shift_metrics})
|
||||
del aligned_model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Same selected fusion architecture, but equal-window audio/vision pooling of the unaligned source.
|
||||
print(f"selected_by_corrupt_macro_f1={selected}; starting fixed-window alignment control", flush=True)
|
||||
fixed_raw = load_fixed_window()
|
||||
fixed_stats = fit_robust_stats(fixed_raw["train"])
|
||||
fixed_stats.save(output / "fixed_window_robust_stats.npz")
|
||||
fixed = {k: apply_robust_stats(v, fixed_stats) for k, v in fixed_raw.items()}
|
||||
for seed in args.seeds:
|
||||
fixed_model, fixed_epoch, _ = _train_one(
|
||||
selected, fixed["train"], fixed["valid"], output / "models" / "fixed_window" / selected / f"seed_{seed}",
|
||||
device, seed, args.epochs, args.patience, args.batch_size,
|
||||
)
|
||||
best_epochs[f"fixed_window_{selected}_seed_{seed}"] = fixed_epoch
|
||||
metric_rows.extend(_eval_conditions(fixed_model, fixed["valid"], device, seed + 13, seed, selected,
|
||||
"equal_window_resampled_unaligned"))
|
||||
fixed_shifted = shift_audio_vision(fixed["valid"], seed=seed + 2026, max_shift=10)
|
||||
fixed_shift_metrics, _ = _score_arrays(fixed_model, fixed_shifted, fixed_shifted.mask, device)
|
||||
metric_rows.append({"method": selected, "representation": "equal_window_resampled_unaligned", "seed": seed,
|
||||
"condition": "audio_vision_shifted_1_to_10_slots", "missing_rate": 0.0,
|
||||
"n_valid": fixed_shifted.n, **fixed_shift_metrics})
|
||||
del fixed_model
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
all_summary = _summary(metric_rows)
|
||||
_write_csv(output / "validation_metrics_by_condition.csv", metric_rows)
|
||||
_write_csv(output / "summary.csv", all_summary)
|
||||
aligned_summary = [r for r in all_summary if r["representation"] == "provided_word_aligned_50"]
|
||||
_plot(aligned_summary, metric_rows, output / "missing_rate_comparison.png")
|
||||
alignment_rows = []
|
||||
for rep in ("provided_word_aligned_50", "equal_window_resampled_unaligned"):
|
||||
for condition in ("clean", "audio_vision_shifted_1_to_10_slots"):
|
||||
match = [r for r in metric_rows if r["method"] == selected and r["representation"] == rep
|
||||
and r["condition"] == condition]
|
||||
if match:
|
||||
row = {"method": selected, "representation": rep, "condition": condition,
|
||||
"n_valid": aligned["valid"].n, "n_seeds": len(match)}
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson"):
|
||||
values = [r[metric] for r in match]
|
||||
row[metric] = float(np.mean(values))
|
||||
row[f"{metric}_sd"] = float(np.std(values, ddof=1)) if len(values) > 1 else 0.0
|
||||
alignment_rows.append(row)
|
||||
corrupt = [r for r in metric_rows if r["method"] == selected and r["representation"] == rep
|
||||
and r["condition"] != "clean" and r["missing_rate"] > 0]
|
||||
if corrupt:
|
||||
per_seed = []
|
||||
for seed in args.seeds:
|
||||
local = [r for r in corrupt if int(r["seed"]) == seed]
|
||||
if local:
|
||||
per_seed.append({metric: float(np.mean([r[metric] for r in local])) for metric in
|
||||
("accuracy", "macro_f1", "mae", "pearson")})
|
||||
alignment_rows.append({
|
||||
"method": selected, "representation": rep, "condition": "all_local_corruption_mean",
|
||||
"missing_rate": float(np.mean([r["missing_rate"] for r in corrupt])),
|
||||
"n_valid": aligned["valid"].n, "n_seeds": len(per_seed),
|
||||
**{metric: float(np.mean([r[metric] for r in per_seed])) for metric in ("accuracy", "macro_f1", "mae", "pearson")},
|
||||
**{f"{metric}_sd": float(np.std([r[metric] for r in per_seed], ddof=1)) if len(per_seed) > 1 else 0.0
|
||||
for metric in ("accuracy", "macro_f1", "mae", "pearson")},
|
||||
})
|
||||
_write_csv(output / "alignment_transfer_ablation.csv", alignment_rows)
|
||||
|
||||
source_path = ATTACHMENT2 / "aligned_50.pkl"
|
||||
manifest = {
|
||||
"source_feature": str(source_path),
|
||||
"source_sha256": _sha256(source_path),
|
||||
"device": str(device),
|
||||
"cuda_name": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"seeds": args.seeds,
|
||||
"epochs_max": args.epochs,
|
||||
"patience": args.patience,
|
||||
"batch_size": args.batch_size,
|
||||
"best_epochs": best_epochs,
|
||||
"selected_macro_f1_first": selected,
|
||||
"selection_policy": "report Macro-F1, MAE, and Pearson separately; selected model maximizes mean validation Macro-F1 across 15 contiguous corruption conditions, then uses MAE and lexical model name only as tie-breaks",
|
||||
"models": list(KINDS),
|
||||
"corruption_rates": [0.10, 0.20, 0.30],
|
||||
"corruption_patterns": list(PATTERNS),
|
||||
"feature_scaling": "training split median/MAD; fallback to standard deviation for zero-MAD dimensions",
|
||||
"test_labels_used": False,
|
||||
"alignment_transfer_limit": "The official aligned_50 data use a 50-slot wordpiece sequence with no per-slot seconds or stored Q1 B1 time_bounds. The fixed-window comparison is a downstream alignment control, not a re-run of Q1 B1 on the full dataset.",
|
||||
"python": __import__("sys").version,
|
||||
"torch": torch.__version__,
|
||||
"numpy": np.__version__,
|
||||
"created_unix": time.time(),
|
||||
}
|
||||
with (output / "run_manifest.json").open("w", encoding="utf-8") as stream:
|
||||
json.dump(manifest, stream, ensure_ascii=False, indent=2)
|
||||
print(f"saved selection artifacts to {output}; selected={selected}; seeds={args.seeds}", flush=True)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Q2 local-missingness model and alignment transfer selection")
|
||||
parser.add_argument("--seeds", type=int, nargs="+", default=[42, 3407, 2026])
|
||||
parser.add_argument("--epochs", type=int, default=32)
|
||||
parser.add_argument("--patience", type=int, default=6)
|
||||
parser.add_argument("--batch-size", type=int, default=64)
|
||||
parser.add_argument("--threads", type=int, default=4)
|
||||
parser.add_argument("--device", default="auto")
|
||||
parser.add_argument("--output-dir", default=str(Path(__file__).resolve().parents[1] / "outputs" / "algorithm_selection"))
|
||||
args = parser.parse_args()
|
||||
_run(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Generated
+1285
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user