JinghuiLuAstronaut commited on
Commit
470b581
·
verified ·
1 Parent(s): 4979f07

Add files using upload-large-folder tool

Browse files
Files changed (20) hide show
  1. LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/infer_lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523_step_0010000_t1p45.log +36 -0
  2. LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/infer_lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523_step_0012000_t1p45.log +36 -0
  3. LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/processed_lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523.txt +5 -0
  4. LTA_openwebtext_dualt/logs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717.log.nohup +205 -0
  5. LTA_openwebtext_dualt/logs/lta_owt_bert_absrope_adaln_dirichlet_len1024_Cv_to_2v_mask1_sameT_gbs512_b4x4_1m_save1k_watch_20260525.log +0 -0
  6. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE +21 -0
  7. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/pygments/modeline.py +43 -0
  8. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/nt.py +163 -0
  9. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/_core.py +3 -0
  10. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/proc.py +83 -0
  11. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/image_processing_aria.py +226 -0
  12. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/processing_aria.py +177 -0
  13. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/__init__.py +28 -0
  14. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/configuration_eomt_dinov3.py +107 -0
  15. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modeling_eomt_dinov3.py +1374 -0
  16. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modular_eomt_dinov3.py +364 -0
  17. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/__init__.py +27 -0
  18. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/configuration_patchtsmixer.py +166 -0
  19. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/modeling_patchtsmixer.py +2121 -0
  20. LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/timesformer/configuration_timesformer.py +66 -0
LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/infer_lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523_step_0010000_t1p45.log ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [watch-classic-1k] 2026-05-23_20:26:46 infer runs/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0010000.pt -> docs/lta_samples/metrics_20260523/lm1b_classic_dirichlet_len512_every1k_normal_steps_state_t1p45_c1024_n256/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0010000
2
+ [ckpt] runs/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0010000.pt step=10000
3
+ [decode] steps128_c1024_t1p45 generated 8/256
4
+ [decode] steps128_c1024_t1p45 generated 16/256
5
+ [decode] steps128_c1024_t1p45 generated 24/256
6
+ [decode] steps128_c1024_t1p45 generated 32/256
7
+ [decode] steps128_c1024_t1p45 generated 40/256
8
+ [decode] steps128_c1024_t1p45 generated 48/256
9
+ [decode] steps128_c1024_t1p45 generated 56/256
10
+ [decode] steps128_c1024_t1p45 generated 64/256
11
+ [decode] steps128_c1024_t1p45 generated 72/256
12
+ [decode] steps128_c1024_t1p45 generated 80/256
13
+ [decode] steps128_c1024_t1p45 generated 88/256
14
+ [decode] steps128_c1024_t1p45 generated 96/256
15
+ [decode] steps128_c1024_t1p45 generated 104/256
16
+ [decode] steps128_c1024_t1p45 generated 112/256
17
+ [decode] steps128_c1024_t1p45 generated 120/256
18
+ [decode] steps128_c1024_t1p45 generated 128/256
19
+ [decode] steps128_c1024_t1p45 generated 136/256
20
+ [decode] steps128_c1024_t1p45 generated 144/256
21
+ [decode] steps128_c1024_t1p45 generated 152/256
22
+ [decode] steps128_c1024_t1p45 generated 160/256
23
+ [decode] steps128_c1024_t1p45 generated 168/256
24
+ [decode] steps128_c1024_t1p45 generated 176/256
25
+ [decode] steps128_c1024_t1p45 generated 184/256
26
+ [decode] steps128_c1024_t1p45 generated 192/256
27
+ [decode] steps128_c1024_t1p45 generated 200/256
28
+ [decode] steps128_c1024_t1p45 generated 208/256
29
+ [decode] steps128_c1024_t1p45 generated 216/256
30
+ [decode] steps128_c1024_t1p45 generated 224/256
31
+ [decode] steps128_c1024_t1p45 generated 232/256
32
+ [decode] steps128_c1024_t1p45 generated 240/256
33
+ [decode] steps128_c1024_t1p45 generated 248/256
34
+ [decode] steps128_c1024_t1p45 generated 256/256
35
+ [summary] {"name": "steps128_c1024_t1p45", "step": 10000, "decode_steps": 128, "concentration_max": 1024.0, "raw_genppl": 32.44292253963206, "stripped_genppl": 36.139052745033965, "sample_entropy": 4.137907218789172, "distinct_1": 0.02729034423828125, "distinct_2": 0.19622217465753425, "top_token_mass": 0.13262939453125, "raw_kept": 256, "stripped_kept": 256}
36
+ [watch-classic-1k] 2026-05-23_20:33:13 done step_0010000
LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/infer_lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523_step_0012000_t1p45.log ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [watch-classic-1k] 2026-05-23_21:10:45 infer runs/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0012000.pt -> docs/lta_samples/metrics_20260523/lm1b_classic_dirichlet_len512_every1k_normal_steps_state_t1p45_c1024_n256/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0012000
2
+ [ckpt] runs/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0012000.pt step=12000
3
+ [decode] steps128_c1024_t1p45 generated 8/256
4
+ [decode] steps128_c1024_t1p45 generated 16/256
5
+ [decode] steps128_c1024_t1p45 generated 24/256
6
+ [decode] steps128_c1024_t1p45 generated 32/256
7
+ [decode] steps128_c1024_t1p45 generated 40/256
8
+ [decode] steps128_c1024_t1p45 generated 48/256
9
+ [decode] steps128_c1024_t1p45 generated 56/256
10
+ [decode] steps128_c1024_t1p45 generated 64/256
11
+ [decode] steps128_c1024_t1p45 generated 72/256
12
+ [decode] steps128_c1024_t1p45 generated 80/256
13
+ [decode] steps128_c1024_t1p45 generated 88/256
14
+ [decode] steps128_c1024_t1p45 generated 96/256
15
+ [decode] steps128_c1024_t1p45 generated 104/256
16
+ [decode] steps128_c1024_t1p45 generated 112/256
17
+ [decode] steps128_c1024_t1p45 generated 120/256
18
+ [decode] steps128_c1024_t1p45 generated 128/256
19
+ [decode] steps128_c1024_t1p45 generated 136/256
20
+ [decode] steps128_c1024_t1p45 generated 144/256
21
+ [decode] steps128_c1024_t1p45 generated 152/256
22
+ [decode] steps128_c1024_t1p45 generated 160/256
23
+ [decode] steps128_c1024_t1p45 generated 168/256
24
+ [decode] steps128_c1024_t1p45 generated 176/256
25
+ [decode] steps128_c1024_t1p45 generated 184/256
26
+ [decode] steps128_c1024_t1p45 generated 192/256
27
+ [decode] steps128_c1024_t1p45 generated 200/256
28
+ [decode] steps128_c1024_t1p45 generated 208/256
29
+ [decode] steps128_c1024_t1p45 generated 216/256
30
+ [decode] steps128_c1024_t1p45 generated 224/256
31
+ [decode] steps128_c1024_t1p45 generated 232/256
32
+ [decode] steps128_c1024_t1p45 generated 240/256
33
+ [decode] steps128_c1024_t1p45 generated 248/256
34
+ [decode] steps128_c1024_t1p45 generated 256/256
35
+ [summary] {"name": "steps128_c1024_t1p45", "step": 12000, "decode_steps": 128, "concentration_max": 1024.0, "raw_genppl": 28.203606139465418, "stripped_genppl": 31.641789039636308, "sample_entropy": 4.094898267469448, "distinct_1": 0.03934478759765625, "distinct_2": 0.23760854941291584, "top_token_mass": 0.15522003173828125, "raw_kept": 256, "stripped_kept": 256}
36
+ [watch-classic-1k] 2026-05-23_21:17:11 done step_0012000
LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/processed_lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0001000.pt
2
+ runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0002000.pt
3
+ runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0003000.pt
4
+ runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0004000.pt
5
+ runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0005000.pt
LTA_openwebtext_dualt/logs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717.log.nohup ADDED
@@ -0,0 +1,205 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [launch] method=categorical_fullvocab_c1024_fullycoupled host=di-20260411014000-djqhq time=2026-05-13T18:47:17+00:00
2
+ [launch] cwd=/e2e-data/evad-tech-vla/wanghan58/workspace/LTA_openwebtext_dualt
3
+ [launch] run_name=lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717
4
+ [launch] save_dir=runs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717
5
+ [launch] log_file=logs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717.log
6
+ NCCL version 2.25.1+cuda12.8
7
+ {
8
+ "device": "cuda:0",
9
+ "rank": 0,
10
+ "world_size": 4,
11
+ "samples": "wrapped_stream",
12
+ "vocab_size": 30522,
13
+ "tokenizer_vocab_size": 30522,
14
+ "save_dir": "runs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717",
15
+ "batch_size": 64,
16
+ "grad_accum": 2,
17
+ "effective_batch_size": 512,
18
+ "global_batch_size": 512,
19
+ "lr_schedule": "constant_warmup",
20
+ "optimizer": "adamw",
21
+ "warmup_steps": 2500,
22
+ "min_lr": 6e-05,
23
+ "weight_decay": 0.0,
24
+ "adamw_param_groups": "nanogpt",
25
+ "adam_beta1": 0.9,
26
+ "adam_beta2": 0.999,
27
+ "adam_eps": 1e-08,
28
+ "muon_momentum": 0.95,
29
+ "muon_ns_steps": 5,
30
+ "muon_update_scale": 1.0,
31
+ "ema_decay": 0.0,
32
+ "ema_start_step": 0,
33
+ "model_type": "ddit",
34
+ "dual_t": true,
35
+ "corrupt_t_mode": "same",
36
+ "corrupt_min_t": 0.0,
37
+ "corrupt_max_t": 1.0,
38
+ "prefix_block_prob": 0.0,
39
+ "prefix_block_len": 128,
40
+ "mask_ratio_floor_schedule": "none",
41
+ "dirichlet_endpoint_mode": "categorical_dual_t",
42
+ "dirichlet_semantic_t_mode": "same",
43
+ "dirichlet_semantic_t_value": 0.0,
44
+ "endpoint_sequence_random_prob_alpha": 0.0,
45
+ "categorical_wrong_from_full_vocab": true,
46
+ "categorical_wrong_from_batch_valid_tokens": false,
47
+ "mask_mixture_original_prob": 0.0,
48
+ "mask_mixture_lowk_prob": 0.0,
49
+ "mask_mixture_lowcorrupt_prob": 0.0,
50
+ "mask_mixture_block_prob": 0.0,
51
+ "mask_mixture_all_prob": 0.0,
52
+ "mask_mixture_lowk_clean_tokens": "1,2,4,8,16,32,64",
53
+ "mask_mixture_lowcorrupt_tokens": "1,2,4,8,16,32,64",
54
+ "mask_mixture_block_tokens": "64,128",
55
+ "simplex_bridge_sampler": "dirichlet",
56
+ "logistic_normal_sigma_min": 0.18,
57
+ "logistic_normal_sigma_max": 2.2,
58
+ "logistic_normal_tau_min": 0.65,
59
+ "logistic_normal_tau_max": 1.15,
60
+ "torch_compile": false,
61
+ "compile_mode": "max-autotune",
62
+ "state_format": "prob",
63
+ "target_loss": "hard_ce",
64
+ "meanflow_weight": 0.0,
65
+ "rollout_train_prob": 0.0,
66
+ "rollout_train_steps": 1,
67
+ "rollout_train_infer_steps": 64,
68
+ "rollout_train_temp": 1.45,
69
+ "rollout_train_max_gamma": 1.0,
70
+ "rollout_train_corrupt_only": true,
71
+ "rollout_train_samplewise": false,
72
+ "rollout_train_compute_always": false,
73
+ "bridge_noise_init": "logistic_normal",
74
+ "noise_sigma": -1.0,
75
+ "allow_tf32": true,
76
+ "activation_checkpointing": false,
77
+ "activation_checkpoint_interval": 1,
78
+ "activation_checkpoint_scope": "block",
79
+ "ddp_static_graph": false,
80
+ "ddp_gradient_as_bucket_view": true,
81
+ "blocking_data_transfer": false,
82
+ "dataloader_prefetch_factor": 2,
83
+ "full_train_stats": false,
84
+ "record_pad_truncate": false,
85
+ "record_add_eos": false,
86
+ "record_add_special_tokens": false,
87
+ "record_pad_token": "pad",
88
+ "record_shuffle_buffer": 10000,
89
+ "wrap": true,
90
+ "wrap_mode": "stream",
91
+ "wrap_record_buffer_size": 200,
92
+ "owt_cached_chunks": false,
93
+ "owt_chunk_cache_dir": "",
94
+ "owt_chunk_cache_rebuild": false,
95
+ "owt_chunk_cache_write_batch": 4096,
96
+ "owt_exact_repeat_per_chunk": 0,
97
+ "online_chunk_shuffle": false,
98
+ "online_chunk_shuffle_buffer": 10000,
99
+ "openwebtext_split": "all",
100
+ "detokenizer": "auto",
101
+ "resolved_detokenizer": "lm1b",
102
+ "num_workers": 0,
103
+ "latest_every": 1000,
104
+ "resume_path": ""
105
+ }
106
+ step=100 micro_steps=200 elapsed=25.1s lr=1.212000e-05 loss=10.1834 loss_recon=10.1834 loss_meanflow=0.0000 mean_model_t=0.5044 mean_corrupt_t=0.5044 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.1761 acc_corrupt=0.1090 corrupt_frac=0.4869 loss_all=9.7762 loss_corrupt=9.7867 acc_corrupt_t_0p0_0p2=0.0437 corrupt_frac_t_0p0_0p2=0.1780 acc_corrupt_t_0p2_0p4=0.0658 corrupt_frac_t_0p2_0p4=0.1599 acc_corrupt_t_0p4_0p6=0.1004 corrupt_frac_t_0p4_0p6=0.3071 acc_corrupt_t_0p6_0p8=0.1235 corrupt_frac_t_0p6_0p8=0.1828 acc_corrupt_t_0p8_1p0=0.2169 corrupt_frac_t_0p8_1p0=0.1722 wrong_frac=0.4986 init_acc_corrupt=0.4703 init_gold_top10=0.4971 init_gold_top100=0.5009
107
+ step=200 micro_steps=400 elapsed=25.9s lr=2.412000e-05 loss=8.9765 loss_recon=8.9765 loss_meanflow=0.0000 mean_model_t=0.4986 mean_corrupt_t=0.4986 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.1156 acc_corrupt=0.0861 corrupt_frac=0.5730 loss_all=8.0280 loss_corrupt=8.0534 acc_corrupt_t_0p0_0p2=0.0476 corrupt_frac_t_0p0_0p2=0.2060 acc_corrupt_t_0p2_0p4=0.0737 corrupt_frac_t_0p2_0p4=0.2399 acc_corrupt_t_0p4_0p6=0.0817 corrupt_frac_t_0p4_0p6=0.2190 acc_corrupt_t_0p6_0p8=0.1055 corrupt_frac_t_0p6_0p8=0.2060 acc_corrupt_t_0p8_1p0=0.1469 corrupt_frac_t_0p8_1p0=0.1291 wrong_frac=0.5279 init_acc_corrupt=0.4325 init_gold_top10=0.4657 init_gold_top100=0.4721
108
+ step=300 micro_steps=600 elapsed=29.4s lr=3.612000e-05 loss=7.1205 loss_recon=7.1205 loss_meanflow=0.0000 mean_model_t=0.5034 mean_corrupt_t=0.5034 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.2524 acc_corrupt=0.1873 corrupt_frac=0.5977 loss_all=5.8884 loss_corrupt=6.2545 acc_corrupt_t_0p0_0p2=0.0590 corrupt_frac_t_0p0_0p2=0.2042 acc_corrupt_t_0p2_0p4=0.1470 corrupt_frac_t_0p2_0p4=0.2210 acc_corrupt_t_0p4_0p6=0.1735 corrupt_frac_t_0p4_0p6=0.2177 acc_corrupt_t_0p6_0p8=0.2652 corrupt_frac_t_0p6_0p8=0.1748 acc_corrupt_t_0p8_1p0=0.3217 corrupt_frac_t_0p8_1p0=0.1822 wrong_frac=0.5037 init_acc_corrupt=0.4618 init_gold_top10=0.4908 init_gold_top100=0.4955
109
+ step=400 micro_steps=800 elapsed=30.6s lr=4.812000e-05 loss=4.8976 loss_recon=4.8976 loss_meanflow=0.0000 mean_model_t=0.5032 mean_corrupt_t=0.5032 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.6721 acc_corrupt=0.4187 corrupt_frac=0.4825 loss_all=2.6266 loss_corrupt=4.5334 acc_corrupt_t_0p0_0p2=0.0964 corrupt_frac_t_0p0_0p2=0.3122 acc_corrupt_t_0p2_0p4=0.2916 corrupt_frac_t_0p2_0p4=0.2386 acc_corrupt_t_0p4_0p6=0.5181 corrupt_frac_t_0p4_0p6=0.0908 acc_corrupt_t_0p6_0p8=0.6515 corrupt_frac_t_0p6_0p8=0.1553 acc_corrupt_t_0p8_1p0=0.8406 corrupt_frac_t_0p8_1p0=0.2031 wrong_frac=0.5742 init_acc_corrupt=0.3881 init_gold_top10=0.4156 init_gold_top100=0.4237
110
+ step=500 micro_steps=1000 elapsed=31.1s lr=6.012000e-05 loss=3.8670 loss_recon=3.8670 loss_meanflow=0.0000 mean_model_t=0.5017 mean_corrupt_t=0.5017 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.6992 acc_corrupt=0.4903 corrupt_frac=0.5557 loss_all=2.4022 loss_corrupt=4.0040 acc_corrupt_t_0p0_0p2=0.1230 corrupt_frac_t_0p0_0p2=0.1821 acc_corrupt_t_0p2_0p4=0.3115 corrupt_frac_t_0p2_0p4=0.3181 acc_corrupt_t_0p4_0p6=0.5806 corrupt_frac_t_0p4_0p6=0.1744 acc_corrupt_t_0p6_0p8=0.7542 corrupt_frac_t_0p6_0p8=0.1430 acc_corrupt_t_0p8_1p0=0.8759 corrupt_frac_t_0p8_1p0=0.1823 wrong_frac=0.5330 init_acc_corrupt=0.4225 init_gold_top10=0.4616 init_gold_top100=0.4679
111
+ step=600 micro_steps=1200 elapsed=31.7s lr=7.212000e-05 loss=3.5477 loss_recon=3.5477 loss_meanflow=0.0000 mean_model_t=0.5057 mean_corrupt_t=0.5057 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.6798 acc_corrupt=0.4914 corrupt_frac=0.5928 loss_all=2.4158 loss_corrupt=3.8192 acc_corrupt_t_0p0_0p2=0.1457 corrupt_frac_t_0p0_0p2=0.2261 acc_corrupt_t_0p2_0p4=0.3474 corrupt_frac_t_0p2_0p4=0.2644 acc_corrupt_t_0p4_0p6=0.6021 corrupt_frac_t_0p4_0p6=0.1936 acc_corrupt_t_0p6_0p8=0.7500 corrupt_frac_t_0p6_0p8=0.2282 acc_corrupt_t_0p8_1p0=0.8991 corrupt_frac_t_0p8_1p0=0.0877 wrong_frac=0.5523 init_acc_corrupt=0.4094 init_gold_top10=0.4401 init_gold_top100=0.4465
112
+ step=700 micro_steps=1400 elapsed=32.5s lr=8.412000e-05 loss=3.4190 loss_recon=3.4190 loss_meanflow=0.0000 mean_model_t=0.4963 mean_corrupt_t=0.4963 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7461 acc_corrupt=0.5368 corrupt_frac=0.5289 loss_all=1.8703 loss_corrupt=3.3598 acc_corrupt_t_0p0_0p2=0.1706 corrupt_frac_t_0p0_0p2=0.2340 acc_corrupt_t_0p2_0p4=0.3301 corrupt_frac_t_0p2_0p4=0.1881 acc_corrupt_t_0p4_0p6=0.5655 corrupt_frac_t_0p4_0p6=0.1673 acc_corrupt_t_0p6_0p8=0.7674 corrupt_frac_t_0p6_0p8=0.2382 acc_corrupt_t_0p8_1p0=0.9130 corrupt_frac_t_0p8_1p0=0.1724 wrong_frac=0.5093 init_acc_corrupt=0.4510 init_gold_top10=0.4842 init_gold_top100=0.4909
113
+ step=800 micro_steps=1600 elapsed=32.4s lr=9.612000e-05 loss=3.3271 loss_recon=3.3271 loss_meanflow=0.0000 mean_model_t=0.4974 mean_corrupt_t=0.4974 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7971 acc_corrupt=0.6090 corrupt_frac=0.4939 loss_all=1.4957 loss_corrupt=2.8453 acc_corrupt_t_0p0_0p2=0.1655 corrupt_frac_t_0p0_0p2=0.1060 acc_corrupt_t_0p2_0p4=0.3159 corrupt_frac_t_0p2_0p4=0.1510 acc_corrupt_t_0p4_0p6=0.5866 corrupt_frac_t_0p4_0p6=0.2954 acc_corrupt_t_0p6_0p8=0.7184 corrupt_frac_t_0p6_0p8=0.2071 acc_corrupt_t_0p8_1p0=0.9219 corrupt_frac_t_0p8_1p0=0.2405 wrong_frac=0.4474 init_acc_corrupt=0.5220 init_gold_top10=0.5494 init_gold_top100=0.5534
114
+ step=900 micro_steps=1800 elapsed=32.4s lr=1.081200e-04 loss=3.2167 loss_recon=3.2167 loss_meanflow=0.0000 mean_model_t=0.5013 mean_corrupt_t=0.5013 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7612 acc_corrupt=0.6084 corrupt_frac=0.5883 loss_all=1.6892 loss_corrupt=2.7423 acc_corrupt_t_0p0_0p2=0.1997 corrupt_frac_t_0p0_0p2=0.1237 acc_corrupt_t_0p2_0p4=0.3986 corrupt_frac_t_0p2_0p4=0.2395 acc_corrupt_t_0p4_0p6=0.5352 corrupt_frac_t_0p4_0p6=0.1888 acc_corrupt_t_0p6_0p8=0.7550 corrupt_frac_t_0p6_0p8=0.1558 acc_corrupt_t_0p8_1p0=0.9226 corrupt_frac_t_0p8_1p0=0.2922 wrong_frac=0.4598 init_acc_corrupt=0.5186 init_gold_top10=0.5360 init_gold_top100=0.5389
115
+ step=1000 micro_steps=2000 elapsed=32.6s lr=1.201200e-04 loss=3.1430 loss_recon=3.1430 loss_meanflow=0.0000 mean_model_t=0.5023 mean_corrupt_t=0.5023 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7438 acc_corrupt=0.5793 corrupt_frac=0.5696 loss_all=1.8455 loss_corrupt=3.0114 acc_corrupt_t_0p0_0p2=0.1932 corrupt_frac_t_0p0_0p2=0.2085 acc_corrupt_t_0p2_0p4=0.3678 corrupt_frac_t_0p2_0p4=0.1929 acc_corrupt_t_0p4_0p6=0.6213 corrupt_frac_t_0p4_0p6=0.2032 acc_corrupt_t_0p6_0p8=0.7635 corrupt_frac_t_0p6_0p8=0.1468 acc_corrupt_t_0p8_1p0=0.9241 corrupt_frac_t_0p8_1p0=0.2486 wrong_frac=0.4846 init_acc_corrupt=0.4764 init_gold_top10=0.5103 init_gold_top100=0.5167
116
+ step=1100 micro_steps=2200 elapsed=35.6s lr=1.321200e-04 loss=3.0752 loss_recon=3.0752 loss_meanflow=0.0000 mean_model_t=0.5000 mean_corrupt_t=0.5000 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7717 acc_corrupt=0.5831 corrupt_frac=0.5247 loss_all=1.6289 loss_corrupt=2.9267 acc_corrupt_t_0p0_0p2=0.1876 corrupt_frac_t_0p0_0p2=0.1575 acc_corrupt_t_0p2_0p4=0.3874 corrupt_frac_t_0p2_0p4=0.2324 acc_corrupt_t_0p4_0p6=0.5764 corrupt_frac_t_0p4_0p6=0.1752 acc_corrupt_t_0p6_0p8=0.7801 corrupt_frac_t_0p6_0p8=0.2645 acc_corrupt_t_0p8_1p0=0.9167 corrupt_frac_t_0p8_1p0=0.1703 wrong_frac=0.4786 init_acc_corrupt=0.4828 init_gold_top10=0.5149 init_gold_top100=0.5214
117
+ step=1200 micro_steps=2400 elapsed=32.8s lr=1.441200e-04 loss=3.0372 loss_recon=3.0372 loss_meanflow=0.0000 mean_model_t=0.5020 mean_corrupt_t=0.5020 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7493 acc_corrupt=0.6050 corrupt_frac=0.5918 loss_all=1.7732 loss_corrupt=2.7718 acc_corrupt_t_0p0_0p2=0.1487 corrupt_frac_t_0p0_0p2=0.1970 acc_corrupt_t_0p2_0p4=0.3654 corrupt_frac_t_0p2_0p4=0.1801 acc_corrupt_t_0p4_0p6=0.6269 corrupt_frac_t_0p4_0p6=0.1592 acc_corrupt_t_0p6_0p8=0.8058 corrupt_frac_t_0p6_0p8=0.1838 acc_corrupt_t_0p8_1p0=0.9359 corrupt_frac_t_0p8_1p0=0.2799 wrong_frac=0.4528 init_acc_corrupt=0.5210 init_gold_top10=0.5429 init_gold_top100=0.5470
118
+ step=1300 micro_steps=2600 elapsed=32.9s lr=1.561200e-04 loss=3.0159 loss_recon=3.0159 loss_meanflow=0.0000 mean_model_t=0.4970 mean_corrupt_t=0.4970 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7704 acc_corrupt=0.5906 corrupt_frac=0.5293 loss_all=1.6082 loss_corrupt=2.8270 acc_corrupt_t_0p0_0p2=0.2367 corrupt_frac_t_0p0_0p2=0.1900 acc_corrupt_t_0p2_0p4=0.4088 corrupt_frac_t_0p2_0p4=0.1946 acc_corrupt_t_0p4_0p6=0.5746 corrupt_frac_t_0p4_0p6=0.1963 acc_corrupt_t_0p6_0p8=0.7861 corrupt_frac_t_0p6_0p8=0.2447 acc_corrupt_t_0p8_1p0=0.9233 corrupt_frac_t_0p8_1p0=0.1744 wrong_frac=0.4892 init_acc_corrupt=0.4709 init_gold_top10=0.5060 init_gold_top100=0.5106
119
+ step=1400 micro_steps=2800 elapsed=32.5s lr=1.681200e-04 loss=2.9175 loss_recon=2.9175 loss_meanflow=0.0000 mean_model_t=0.5085 mean_corrupt_t=0.5085 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7028 acc_corrupt=0.5007 corrupt_frac=0.5623 loss_all=2.0643 loss_corrupt=3.4455 acc_corrupt_t_0p0_0p2=0.1902 corrupt_frac_t_0p0_0p2=0.2579 acc_corrupt_t_0p2_0p4=0.3040 corrupt_frac_t_0p2_0p4=0.2199 acc_corrupt_t_0p4_0p6=0.5713 corrupt_frac_t_0p4_0p6=0.2269 acc_corrupt_t_0p6_0p8=0.8011 corrupt_frac_t_0p6_0p8=0.1201 acc_corrupt_t_0p8_1p0=0.9071 corrupt_frac_t_0p8_1p0=0.1752 wrong_frac=0.5536 init_acc_corrupt=0.3906 init_gold_top10=0.4394 init_gold_top100=0.4475
120
+ step=1500 micro_steps=3000 elapsed=32.8s lr=1.801200e-04 loss=2.9483 loss_recon=2.9483 loss_meanflow=0.0000 mean_model_t=0.5002 mean_corrupt_t=0.5002 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7744 acc_corrupt=0.5755 corrupt_frac=0.5063 loss_all=1.5783 loss_corrupt=2.9346 acc_corrupt_t_0p0_0p2=0.2041 corrupt_frac_t_0p0_0p2=0.1984 acc_corrupt_t_0p2_0p4=0.3830 corrupt_frac_t_0p2_0p4=0.2411 acc_corrupt_t_0p4_0p6=0.6497 corrupt_frac_t_0p4_0p6=0.1748 acc_corrupt_t_0p6_0p8=0.7817 corrupt_frac_t_0p6_0p8=0.2131 acc_corrupt_t_0p8_1p0=0.9413 corrupt_frac_t_0p8_1p0=0.1726 wrong_frac=0.5007 init_acc_corrupt=0.4535 init_gold_top10=0.4906 init_gold_top100=0.4993
121
+ step=1600 micro_steps=3200 elapsed=32.7s lr=1.921200e-04 loss=2.9063 loss_recon=2.9063 loss_meanflow=0.0000 mean_model_t=0.5013 mean_corrupt_t=0.5013 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7737 acc_corrupt=0.5454 corrupt_frac=0.4747 loss_all=1.6032 loss_corrupt=3.1656 acc_corrupt_t_0p0_0p2=0.2086 corrupt_frac_t_0p0_0p2=0.2947 acc_corrupt_t_0p2_0p4=0.4480 corrupt_frac_t_0p2_0p4=0.1905 acc_corrupt_t_0p4_0p6=0.6407 corrupt_frac_t_0p4_0p6=0.1882 acc_corrupt_t_0p6_0p8=0.7808 corrupt_frac_t_0p6_0p8=0.1502 acc_corrupt_t_0p8_1p0=0.9111 corrupt_frac_t_0p8_1p0=0.1764 wrong_frac=0.5513 init_acc_corrupt=0.4076 init_gold_top10=0.4407 init_gold_top100=0.4490
122
+ step=1700 micro_steps=3400 elapsed=32.7s lr=2.041200e-04 loss=2.9063 loss_recon=2.9063 loss_meanflow=0.0000 mean_model_t=0.4944 mean_corrupt_t=0.4944 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7395 acc_corrupt=0.5679 corrupt_frac=0.5514 loss_all=1.7862 loss_corrupt=2.9407 acc_corrupt_t_0p0_0p2=0.2490 corrupt_frac_t_0p0_0p2=0.2302 acc_corrupt_t_0p2_0p4=0.4100 corrupt_frac_t_0p2_0p4=0.1722 acc_corrupt_t_0p4_0p6=0.6026 corrupt_frac_t_0p4_0p6=0.2741 acc_corrupt_t_0p6_0p8=0.7707 corrupt_frac_t_0p6_0p8=0.1603 acc_corrupt_t_0p8_1p0=0.9267 corrupt_frac_t_0p8_1p0=0.1632 wrong_frac=0.5349 init_acc_corrupt=0.4341 init_gold_top10=0.4600 init_gold_top100=0.4660
123
+ step=1800 micro_steps=3600 elapsed=33.1s lr=2.161200e-04 loss=2.8104 loss_recon=2.8104 loss_meanflow=0.0000 mean_model_t=0.5022 mean_corrupt_t=0.5022 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7971 acc_corrupt=0.6304 corrupt_frac=0.5232 loss_all=1.4026 loss_corrupt=2.5385 acc_corrupt_t_0p0_0p2=0.2392 corrupt_frac_t_0p0_0p2=0.1843 acc_corrupt_t_0p2_0p4=0.4688 corrupt_frac_t_0p2_0p4=0.1085 acc_corrupt_t_0p4_0p6=0.6210 corrupt_frac_t_0p4_0p6=0.2825 acc_corrupt_t_0p6_0p8=0.8017 corrupt_frac_t_0p6_0p8=0.2954 acc_corrupt_t_0p8_1p0=0.9531 corrupt_frac_t_0p8_1p0=0.1293 wrong_frac=0.4722 init_acc_corrupt=0.5047 init_gold_top10=0.5215 init_gold_top100=0.5287
124
+ step=1900 micro_steps=3800 elapsed=32.7s lr=2.281200e-04 loss=2.8400 loss_recon=2.8400 loss_meanflow=0.0000 mean_model_t=0.4991 mean_corrupt_t=0.4991 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7673 acc_corrupt=0.5870 corrupt_frac=0.5311 loss_all=1.5840 loss_corrupt=2.7945 acc_corrupt_t_0p0_0p2=0.2045 corrupt_frac_t_0p0_0p2=0.2046 acc_corrupt_t_0p2_0p4=0.4286 corrupt_frac_t_0p2_0p4=0.1834 acc_corrupt_t_0p4_0p6=0.6160 corrupt_frac_t_0p4_0p6=0.2071 acc_corrupt_t_0p6_0p8=0.7692 corrupt_frac_t_0p6_0p8=0.2310 acc_corrupt_t_0p8_1p0=0.9273 corrupt_frac_t_0p8_1p0=0.1740 wrong_frac=0.5059 init_acc_corrupt=0.4610 init_gold_top10=0.4898 init_gold_top100=0.4939
125
+ step=2000 micro_steps=4000 elapsed=32.7s lr=2.401200e-04 loss=2.7629 loss_recon=2.7629 loss_meanflow=0.0000 mean_model_t=0.5030 mean_corrupt_t=0.5030 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7139 acc_corrupt=0.5357 corrupt_frac=0.5756 loss_all=2.0052 loss_corrupt=3.2353 acc_corrupt_t_0p0_0p2=0.2236 corrupt_frac_t_0p0_0p2=0.2893 acc_corrupt_t_0p2_0p4=0.3738 corrupt_frac_t_0p2_0p4=0.2269 acc_corrupt_t_0p4_0p6=0.6094 corrupt_frac_t_0p4_0p6=0.0766 acc_corrupt_t_0p6_0p8=0.7480 corrupt_frac_t_0p6_0p8=0.2121 acc_corrupt_t_0p8_1p0=0.9272 corrupt_frac_t_0p8_1p0=0.1951 wrong_frac=0.5574 init_acc_corrupt=0.3994 init_gold_top10=0.4348 init_gold_top100=0.4409
126
+ step=2100 micro_steps=4200 elapsed=35.4s lr=2.521200e-04 loss=2.8096 loss_recon=2.8096 loss_meanflow=0.0000 mean_model_t=0.4988 mean_corrupt_t=0.4988 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7480 acc_corrupt=0.5846 corrupt_frac=0.5698 loss_all=1.7203 loss_corrupt=2.8122 acc_corrupt_t_0p0_0p2=0.2335 corrupt_frac_t_0p0_0p2=0.1982 acc_corrupt_t_0p2_0p4=0.3762 corrupt_frac_t_0p2_0p4=0.1765 acc_corrupt_t_0p4_0p6=0.6254 corrupt_frac_t_0p4_0p6=0.2264 acc_corrupt_t_0p6_0p8=0.7874 corrupt_frac_t_0p6_0p8=0.2761 acc_corrupt_t_0p8_1p0=0.9197 corrupt_frac_t_0p8_1p0=0.1228 wrong_frac=0.5154 init_acc_corrupt=0.4522 init_gold_top10=0.4779 init_gold_top100=0.4844
127
+ step=2200 micro_steps=4400 elapsed=33.0s lr=2.641200e-04 loss=2.8370 loss_recon=2.8370 loss_meanflow=0.0000 mean_model_t=0.4943 mean_corrupt_t=0.4943 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7657 acc_corrupt=0.6355 corrupt_frac=0.6089 loss_all=1.6353 loss_corrupt=2.5344 acc_corrupt_t_0p0_0p2=0.2129 corrupt_frac_t_0p0_0p2=0.1864 acc_corrupt_t_0p2_0p4=0.4363 corrupt_frac_t_0p2_0p4=0.1512 acc_corrupt_t_0p4_0p6=0.6530 corrupt_frac_t_0p4_0p6=0.1520 acc_corrupt_t_0p6_0p8=0.7694 corrupt_frac_t_0p6_0p8=0.3077 acc_corrupt_t_0p8_1p0=0.9565 corrupt_frac_t_0p8_1p0=0.2027 wrong_frac=0.4479 init_acc_corrupt=0.5253 init_gold_top10=0.5465 init_gold_top100=0.5521
128
+ step=2300 micro_steps=4600 elapsed=32.7s lr=2.761200e-04 loss=2.7322 loss_recon=2.7322 loss_meanflow=0.0000 mean_model_t=0.5004 mean_corrupt_t=0.5004 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7596 acc_corrupt=0.6017 corrupt_frac=0.5759 loss_all=1.6318 loss_corrupt=2.6851 acc_corrupt_t_0p0_0p2=0.1849 corrupt_frac_t_0p0_0p2=0.1937 acc_corrupt_t_0p2_0p4=0.4594 corrupt_frac_t_0p2_0p4=0.2298 acc_corrupt_t_0p4_0p6=0.6619 corrupt_frac_t_0p4_0p6=0.1774 acc_corrupt_t_0p6_0p8=0.7679 corrupt_frac_t_0p6_0p8=0.1927 acc_corrupt_t_0p8_1p0=0.9446 corrupt_frac_t_0p8_1p0=0.2064 wrong_frac=0.4939 init_acc_corrupt=0.4758 init_gold_top10=0.4992 init_gold_top100=0.5070
129
+ step=2400 micro_steps=4800 elapsed=33.2s lr=2.881200e-04 loss=2.7271 loss_recon=2.7271 loss_meanflow=0.0000 mean_model_t=0.5002 mean_corrupt_t=0.5002 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7828 acc_corrupt=0.5899 corrupt_frac=0.4863 loss_all=1.4524 loss_corrupt=2.7114 acc_corrupt_t_0p0_0p2=0.2472 corrupt_frac_t_0p0_0p2=0.2427 acc_corrupt_t_0p2_0p4=0.4311 corrupt_frac_t_0p2_0p4=0.2096 acc_corrupt_t_0p4_0p6=0.6620 corrupt_frac_t_0p4_0p6=0.1968 acc_corrupt_t_0p6_0p8=0.8051 corrupt_frac_t_0p6_0p8=0.1880 acc_corrupt_t_0p8_1p0=0.9692 corrupt_frac_t_0p8_1p0=0.1629 wrong_frac=0.5203 init_acc_corrupt=0.4488 init_gold_top10=0.4734 init_gold_top100=0.4794
130
+ step=2500 micro_steps=5000 elapsed=35.7s lr=3.000000e-04 loss=2.7222 loss_recon=2.7222 loss_meanflow=0.0000 mean_model_t=0.4978 mean_corrupt_t=0.4978 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7560 acc_corrupt=0.6051 corrupt_frac=0.5814 loss_all=1.6523 loss_corrupt=2.6605 acc_corrupt_t_0p0_0p2=0.2413 corrupt_frac_t_0p0_0p2=0.1757 acc_corrupt_t_0p2_0p4=0.4189 corrupt_frac_t_0p2_0p4=0.2461 acc_corrupt_t_0p4_0p6=0.6331 corrupt_frac_t_0p4_0p6=0.1482 acc_corrupt_t_0p6_0p8=0.7900 corrupt_frac_t_0p6_0p8=0.2400 acc_corrupt_t_0p8_1p0=0.9271 corrupt_frac_t_0p8_1p0=0.1900 wrong_frac=0.4961 init_acc_corrupt=0.4692 init_gold_top10=0.4988 init_gold_top100=0.5043
131
+ step=2600 micro_steps=5200 elapsed=33.9s lr=3.000000e-04 loss=2.6895 loss_recon=2.6895 loss_meanflow=0.0000 mean_model_t=0.4940 mean_corrupt_t=0.4940 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7802 acc_corrupt=0.6134 corrupt_frac=0.5367 loss_all=1.5002 loss_corrupt=2.6191 acc_corrupt_t_0p0_0p2=0.2584 corrupt_frac_t_0p0_0p2=0.2033 acc_corrupt_t_0p2_0p4=0.3884 corrupt_frac_t_0p2_0p4=0.1926 acc_corrupt_t_0p4_0p6=0.6302 corrupt_frac_t_0p4_0p6=0.1808 acc_corrupt_t_0p6_0p8=0.8215 corrupt_frac_t_0p6_0p8=0.2051 acc_corrupt_t_0p8_1p0=0.9333 corrupt_frac_t_0p8_1p0=0.2181 wrong_frac=0.4783 init_acc_corrupt=0.4865 init_gold_top10=0.5138 init_gold_top100=0.5224
132
+ step=2700 micro_steps=5400 elapsed=32.8s lr=3.000000e-04 loss=2.6913 loss_recon=2.6913 loss_meanflow=0.0000 mean_model_t=0.4975 mean_corrupt_t=0.4975 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7703 acc_corrupt=0.6079 corrupt_frac=0.5516 loss_all=1.5171 loss_corrupt=2.5527 acc_corrupt_t_0p0_0p2=0.2437 corrupt_frac_t_0p0_0p2=0.2806 acc_corrupt_t_0p2_0p4=0.4662 corrupt_frac_t_0p2_0p4=0.0983 acc_corrupt_t_0p4_0p6=0.6303 corrupt_frac_t_0p4_0p6=0.1580 acc_corrupt_t_0p6_0p8=0.7961 corrupt_frac_t_0p6_0p8=0.2952 acc_corrupt_t_0p8_1p0=0.9473 corrupt_frac_t_0p8_1p0=0.1680 wrong_frac=0.4939 init_acc_corrupt=0.4678 init_gold_top10=0.4981 init_gold_top100=0.5063
133
+ step=2800 micro_steps=5600 elapsed=32.6s lr=3.000000e-04 loss=2.6582 loss_recon=2.6582 loss_meanflow=0.0000 mean_model_t=0.4969 mean_corrupt_t=0.4969 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7960 acc_corrupt=0.6110 corrupt_frac=0.4845 loss_all=1.3712 loss_corrupt=2.5989 acc_corrupt_t_0p0_0p2=0.3040 corrupt_frac_t_0p0_0p2=0.2063 acc_corrupt_t_0p2_0p4=0.4395 corrupt_frac_t_0p2_0p4=0.1978 acc_corrupt_t_0p4_0p6=0.6050 corrupt_frac_t_0p4_0p6=0.2028 acc_corrupt_t_0p6_0p8=0.7486 corrupt_frac_t_0p6_0p8=0.1754 acc_corrupt_t_0p8_1p0=0.9525 corrupt_frac_t_0p8_1p0=0.2177 wrong_frac=0.5117 init_acc_corrupt=0.4578 init_gold_top10=0.4840 init_gold_top100=0.4885
134
+ step=2900 micro_steps=5800 elapsed=32.7s lr=3.000000e-04 loss=2.6140 loss_recon=2.6140 loss_meanflow=0.0000 mean_model_t=0.5025 mean_corrupt_t=0.5025 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7776 acc_corrupt=0.6555 corrupt_frac=0.6084 loss_all=1.4709 loss_corrupt=2.2665 acc_corrupt_t_0p0_0p2=0.2029 corrupt_frac_t_0p0_0p2=0.1928 acc_corrupt_t_0p2_0p4=0.4750 corrupt_frac_t_0p2_0p4=0.1162 acc_corrupt_t_0p4_0p6=0.6511 corrupt_frac_t_0p4_0p6=0.1806 acc_corrupt_t_0p6_0p8=0.7998 corrupt_frac_t_0p6_0p8=0.2616 acc_corrupt_t_0p8_1p0=0.9419 corrupt_frac_t_0p8_1p0=0.2488 wrong_frac=0.4424 init_acc_corrupt=0.5377 init_gold_top10=0.5540 init_gold_top100=0.5574
135
+ step=3000 micro_steps=6000 elapsed=33.1s lr=3.000000e-04 loss=2.6014 loss_recon=2.6014 loss_meanflow=0.0000 mean_model_t=0.5019 mean_corrupt_t=0.5019 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7902 acc_corrupt=0.5905 corrupt_frac=0.4707 loss_all=1.4004 loss_corrupt=2.7095 acc_corrupt_t_0p0_0p2=0.3031 corrupt_frac_t_0p0_0p2=0.2601 acc_corrupt_t_0p2_0p4=0.4446 corrupt_frac_t_0p2_0p4=0.2573 acc_corrupt_t_0p4_0p6=0.6436 corrupt_frac_t_0p4_0p6=0.1499 acc_corrupt_t_0p6_0p8=0.8176 corrupt_frac_t_0p6_0p8=0.1294 acc_corrupt_t_0p8_1p0=0.9592 corrupt_frac_t_0p8_1p0=0.2033 wrong_frac=0.5329 init_acc_corrupt=0.4199 init_gold_top10=0.4590 init_gold_top100=0.4668
136
+ step=3100 micro_steps=6200 elapsed=36.3s lr=3.000000e-04 loss=2.5873 loss_recon=2.5873 loss_meanflow=0.0000 mean_model_t=0.5023 mean_corrupt_t=0.5023 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7916 acc_corrupt=0.6029 corrupt_frac=0.4943 loss_all=1.3728 loss_corrupt=2.5977 acc_corrupt_t_0p0_0p2=0.2815 corrupt_frac_t_0p0_0p2=0.1413 acc_corrupt_t_0p2_0p4=0.4271 corrupt_frac_t_0p2_0p4=0.3267 acc_corrupt_t_0p4_0p6=0.6184 corrupt_frac_t_0p4_0p6=0.1637 acc_corrupt_t_0p6_0p8=0.8156 corrupt_frac_t_0p6_0p8=0.1741 acc_corrupt_t_0p8_1p0=0.9288 corrupt_frac_t_0p8_1p0=0.1941 wrong_frac=0.5009 init_acc_corrupt=0.4576 init_gold_top10=0.4932 init_gold_top100=0.4986
137
+ step=3200 micro_steps=6400 elapsed=33.9s lr=3.000000e-04 loss=2.6312 loss_recon=2.6312 loss_meanflow=0.0000 mean_model_t=0.4954 mean_corrupt_t=0.4954 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7373 acc_corrupt=0.5457 corrupt_frac=0.5344 loss_all=1.7428 loss_corrupt=3.0053 acc_corrupt_t_0p0_0p2=0.2741 corrupt_frac_t_0p0_0p2=0.3559 acc_corrupt_t_0p2_0p4=0.4393 corrupt_frac_t_0p2_0p4=0.1674 acc_corrupt_t_0p4_0p6=0.6071 corrupt_frac_t_0p4_0p6=0.1035 acc_corrupt_t_0p6_0p8=0.7998 corrupt_frac_t_0p6_0p8=0.2716 acc_corrupt_t_0p8_1p0=0.9303 corrupt_frac_t_0p8_1p0=0.1016 wrong_frac=0.5893 init_acc_corrupt=0.3646 init_gold_top10=0.4025 init_gold_top100=0.4107
138
+ step=3300 micro_steps=6600 elapsed=32.7s lr=3.000000e-04 loss=2.6239 loss_recon=2.6239 loss_meanflow=0.0000 mean_model_t=0.4934 mean_corrupt_t=0.4934 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7620 acc_corrupt=0.6141 corrupt_frac=0.5906 loss_all=1.6176 loss_corrupt=2.6085 acc_corrupt_t_0p0_0p2=0.1868 corrupt_frac_t_0p0_0p2=0.1782 acc_corrupt_t_0p2_0p4=0.4540 corrupt_frac_t_0p2_0p4=0.2090 acc_corrupt_t_0p4_0p6=0.6652 corrupt_frac_t_0p4_0p6=0.2427 acc_corrupt_t_0p6_0p8=0.7991 corrupt_frac_t_0p6_0p8=0.1420 acc_corrupt_t_0p8_1p0=0.9248 corrupt_frac_t_0p8_1p0=0.2282 wrong_frac=0.4957 init_acc_corrupt=0.4756 init_gold_top10=0.4992 init_gold_top100=0.5037
139
+ step=3400 micro_steps=6800 elapsed=32.8s lr=3.000000e-04 loss=2.5839 loss_recon=2.5839 loss_meanflow=0.0000 mean_model_t=0.4976 mean_corrupt_t=0.4976 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7504 acc_corrupt=0.5574 corrupt_frac=0.5359 loss_all=1.6673 loss_corrupt=2.9311 acc_corrupt_t_0p0_0p2=0.1792 corrupt_frac_t_0p0_0p2=0.2428 acc_corrupt_t_0p2_0p4=0.4278 corrupt_frac_t_0p2_0p4=0.2460 acc_corrupt_t_0p4_0p6=0.6306 corrupt_frac_t_0p4_0p6=0.1622 acc_corrupt_t_0p6_0p8=0.7854 corrupt_frac_t_0p6_0p8=0.1592 acc_corrupt_t_0p8_1p0=0.9556 corrupt_frac_t_0p8_1p0=0.1897 wrong_frac=0.5510 init_acc_corrupt=0.4125 init_gold_top10=0.4428 init_gold_top100=0.4487
140
+ step=3500 micro_steps=7000 elapsed=32.9s lr=3.000000e-04 loss=2.5431 loss_recon=2.5431 loss_meanflow=0.0000 mean_model_t=0.5002 mean_corrupt_t=0.5002 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7958 acc_corrupt=0.6559 corrupt_frac=0.5570 loss_all=1.2813 loss_corrupt=2.1481 acc_corrupt_t_0p0_0p2=0.2704 corrupt_frac_t_0p0_0p2=0.1370 acc_corrupt_t_0p2_0p4=0.4290 corrupt_frac_t_0p2_0p4=0.2115 acc_corrupt_t_0p4_0p6=0.6747 corrupt_frac_t_0p4_0p6=0.2176 acc_corrupt_t_0p6_0p8=0.7941 corrupt_frac_t_0p6_0p8=0.2172 acc_corrupt_t_0p8_1p0=0.9636 corrupt_frac_t_0p8_1p0=0.2167 wrong_frac=0.4563 init_acc_corrupt=0.5146 init_gold_top10=0.5389 init_gold_top100=0.5435
141
+ step=3600 micro_steps=7200 elapsed=32.7s lr=3.000000e-04 loss=2.5584 loss_recon=2.5584 loss_meanflow=0.0000 mean_model_t=0.5020 mean_corrupt_t=0.5020 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8075 acc_corrupt=0.6596 corrupt_frac=0.5189 loss_all=1.2272 loss_corrupt=2.1615 acc_corrupt_t_0p0_0p2=0.2956 corrupt_frac_t_0p0_0p2=0.1807 acc_corrupt_t_0p2_0p4=0.4819 corrupt_frac_t_0p2_0p4=0.1299 acc_corrupt_t_0p4_0p6=0.6043 corrupt_frac_t_0p4_0p6=0.2842 acc_corrupt_t_0p6_0p8=0.8492 corrupt_frac_t_0p6_0p8=0.1404 acc_corrupt_t_0p8_1p0=0.9538 corrupt_frac_t_0p8_1p0=0.2649 wrong_frac=0.4587 init_acc_corrupt=0.5171 init_gold_top10=0.5361 init_gold_top100=0.5422
142
+ step=3700 micro_steps=7400 elapsed=32.9s lr=3.000000e-04 loss=2.5212 loss_recon=2.5212 loss_meanflow=0.0000 mean_model_t=0.5023 mean_corrupt_t=0.5023 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8555 acc_corrupt=0.7266 corrupt_frac=0.4983 loss_all=0.9338 loss_corrupt=1.7476 acc_corrupt_t_0p0_0p2=0.2394 corrupt_frac_t_0p0_0p2=0.1330 acc_corrupt_t_0p2_0p4=0.5758 corrupt_frac_t_0p2_0p4=0.1115 acc_corrupt_t_0p4_0p6=0.7033 corrupt_frac_t_0p4_0p6=0.1247 acc_corrupt_t_0p6_0p8=0.8062 corrupt_frac_t_0p6_0p8=0.3844 acc_corrupt_t_0p8_1p0=0.9453 corrupt_frac_t_0p8_1p0=0.2464 wrong_frac=0.3768 init_acc_corrupt=0.6024 init_gold_top10=0.6213 init_gold_top100=0.6235
143
+ step=3800 micro_steps=7600 elapsed=32.5s lr=3.000000e-04 loss=2.5015 loss_recon=2.5015 loss_meanflow=0.0000 mean_model_t=0.5044 mean_corrupt_t=0.5044 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8256 acc_corrupt=0.6825 corrupt_frac=0.5114 loss_all=1.1426 loss_corrupt=2.0673 acc_corrupt_t_0p0_0p2=0.2514 corrupt_frac_t_0p0_0p2=0.1767 acc_corrupt_t_0p2_0p4=0.4989 corrupt_frac_t_0p2_0p4=0.1105 acc_corrupt_t_0p4_0p6=0.6610 corrupt_frac_t_0p4_0p6=0.1697 acc_corrupt_t_0p6_0p8=0.8127 corrupt_frac_t_0p6_0p8=0.2600 acc_corrupt_t_0p8_1p0=0.9165 corrupt_frac_t_0p8_1p0=0.2831 wrong_frac=0.4388 init_acc_corrupt=0.5397 init_gold_top10=0.5565 init_gold_top100=0.5608
144
+ step=3900 micro_steps=7800 elapsed=32.6s lr=3.000000e-04 loss=2.5639 loss_recon=2.5639 loss_meanflow=0.0000 mean_model_t=0.4941 mean_corrupt_t=0.4941 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7728 acc_corrupt=0.6005 corrupt_frac=0.5417 loss_all=1.5214 loss_corrupt=2.6664 acc_corrupt_t_0p0_0p2=0.2213 corrupt_frac_t_0p0_0p2=0.1904 acc_corrupt_t_0p2_0p4=0.4838 corrupt_frac_t_0p2_0p4=0.2156 acc_corrupt_t_0p4_0p6=0.6420 corrupt_frac_t_0p4_0p6=0.2341 acc_corrupt_t_0p6_0p8=0.8175 corrupt_frac_t_0p6_0p8=0.2605 acc_corrupt_t_0p8_1p0=0.9138 corrupt_frac_t_0p8_1p0=0.0994 wrong_frac=0.5273 init_acc_corrupt=0.4369 init_gold_top10=0.4662 init_gold_top100=0.4725
145
+ step=4000 micro_steps=8000 elapsed=32.6s lr=3.000000e-04 loss=2.5095 loss_recon=2.5095 loss_meanflow=0.0000 mean_model_t=0.4992 mean_corrupt_t=0.4992 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7638 acc_corrupt=0.5968 corrupt_frac=0.5637 loss_all=1.5411 loss_corrupt=2.6093 acc_corrupt_t_0p0_0p2=0.1990 corrupt_frac_t_0p0_0p2=0.2100 acc_corrupt_t_0p2_0p4=0.4960 corrupt_frac_t_0p2_0p4=0.2689 acc_corrupt_t_0p4_0p6=0.6839 corrupt_frac_t_0p4_0p6=0.1966 acc_corrupt_t_0p6_0p8=0.8283 corrupt_frac_t_0p6_0p8=0.1589 acc_corrupt_t_0p8_1p0=0.9398 corrupt_frac_t_0p8_1p0=0.1654 wrong_frac=0.5223 init_acc_corrupt=0.4355 init_gold_top10=0.4738 init_gold_top100=0.4779
146
+ step=4100 micro_steps=8200 elapsed=36.5s lr=3.000000e-04 loss=2.5317 loss_recon=2.5317 loss_meanflow=0.0000 mean_model_t=0.4962 mean_corrupt_t=0.4962 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7899 acc_corrupt=0.6008 corrupt_frac=0.5033 loss_all=1.3959 loss_corrupt=2.6191 acc_corrupt_t_0p0_0p2=0.2629 corrupt_frac_t_0p0_0p2=0.2399 acc_corrupt_t_0p2_0p4=0.4746 corrupt_frac_t_0p2_0p4=0.1906 acc_corrupt_t_0p4_0p6=0.6150 corrupt_frac_t_0p4_0p6=0.1613 acc_corrupt_t_0p6_0p8=0.8189 corrupt_frac_t_0p6_0p8=0.2624 acc_corrupt_t_0p8_1p0=0.9135 corrupt_frac_t_0p8_1p0=0.1458 wrong_frac=0.5224 init_acc_corrupt=0.4400 init_gold_top10=0.4710 init_gold_top100=0.4776
147
+ step=4200 micro_steps=8400 elapsed=33.5s lr=3.000000e-04 loss=2.5253 loss_recon=2.5253 loss_meanflow=0.0000 mean_model_t=0.4938 mean_corrupt_t=0.4938 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7943 acc_corrupt=0.6361 corrupt_frac=0.5314 loss_all=1.3446 loss_corrupt=2.3587 acc_corrupt_t_0p0_0p2=0.2855 corrupt_frac_t_0p0_0p2=0.1762 acc_corrupt_t_0p2_0p4=0.3905 corrupt_frac_t_0p2_0p4=0.2171 acc_corrupt_t_0p4_0p6=0.6643 corrupt_frac_t_0p4_0p6=0.1649 acc_corrupt_t_0p6_0p8=0.8380 corrupt_frac_t_0p6_0p8=0.2538 acc_corrupt_t_0p8_1p0=0.9511 corrupt_frac_t_0p8_1p0=0.1879 wrong_frac=0.4852 init_acc_corrupt=0.4776 init_gold_top10=0.5098 init_gold_top100=0.5155
148
+ step=4300 micro_steps=8600 elapsed=32.8s lr=3.000000e-04 loss=2.5230 loss_recon=2.5230 loss_meanflow=0.0000 mean_model_t=0.4960 mean_corrupt_t=0.4960 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7672 acc_corrupt=0.5929 corrupt_frac=0.5479 loss_all=1.5292 loss_corrupt=2.6535 acc_corrupt_t_0p0_0p2=0.3198 corrupt_frac_t_0p0_0p2=0.1589 acc_corrupt_t_0p2_0p4=0.4367 corrupt_frac_t_0p2_0p4=0.2796 acc_corrupt_t_0p4_0p6=0.6412 corrupt_frac_t_0p4_0p6=0.3030 acc_corrupt_t_0p6_0p8=0.8322 corrupt_frac_t_0p6_0p8=0.1647 acc_corrupt_t_0p8_1p0=0.9454 corrupt_frac_t_0p8_1p0=0.0938 wrong_frac=0.5428 init_acc_corrupt=0.4129 init_gold_top10=0.4525 init_gold_top100=0.4579
149
+ step=4400 micro_steps=8800 elapsed=32.5s lr=3.000000e-04 loss=2.5007 loss_recon=2.5007 loss_meanflow=0.0000 mean_model_t=0.5046 mean_corrupt_t=0.5046 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8005 acc_corrupt=0.6388 corrupt_frac=0.5192 loss_all=1.3267 loss_corrupt=2.3917 acc_corrupt_t_0p0_0p2=0.2153 corrupt_frac_t_0p0_0p2=0.2086 acc_corrupt_t_0p2_0p4=0.3725 corrupt_frac_t_0p2_0p4=0.1439 acc_corrupt_t_0p4_0p6=0.7050 corrupt_frac_t_0p4_0p6=0.1801 acc_corrupt_t_0p6_0p8=0.8262 corrupt_frac_t_0p6_0p8=0.1867 acc_corrupt_t_0p8_1p0=0.9229 corrupt_frac_t_0p8_1p0=0.2807 wrong_frac=0.4632 init_acc_corrupt=0.5041 init_gold_top10=0.5288 init_gold_top100=0.5356
150
+ step=4500 micro_steps=9000 elapsed=32.4s lr=3.000000e-04 loss=2.4506 loss_recon=2.4506 loss_meanflow=0.0000 mean_model_t=0.5026 mean_corrupt_t=0.5026 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7817 acc_corrupt=0.6405 corrupt_frac=0.5844 loss_all=1.4468 loss_corrupt=2.3733 acc_corrupt_t_0p0_0p2=0.1813 corrupt_frac_t_0p0_0p2=0.1717 acc_corrupt_t_0p2_0p4=0.4369 corrupt_frac_t_0p2_0p4=0.2037 acc_corrupt_t_0p4_0p6=0.6953 corrupt_frac_t_0p4_0p6=0.1604 acc_corrupt_t_0p6_0p8=0.8024 corrupt_frac_t_0p6_0p8=0.1744 acc_corrupt_t_0p8_1p0=0.9279 corrupt_frac_t_0p8_1p0=0.2897 wrong_frac=0.4567 init_acc_corrupt=0.5072 init_gold_top10=0.5369 init_gold_top100=0.5429
151
+ step=4600 micro_steps=9200 elapsed=32.3s lr=3.000000e-04 loss=2.4457 loss_recon=2.4457 loss_meanflow=0.0000 mean_model_t=0.5063 mean_corrupt_t=0.5063 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7791 acc_corrupt=0.6502 corrupt_frac=0.6085 loss_all=1.4816 loss_corrupt=2.3300 acc_corrupt_t_0p0_0p2=0.2290 corrupt_frac_t_0p0_0p2=0.1840 acc_corrupt_t_0p2_0p4=0.4478 corrupt_frac_t_0p2_0p4=0.1384 acc_corrupt_t_0p4_0p6=0.6513 corrupt_frac_t_0p4_0p6=0.1916 acc_corrupt_t_0p6_0p8=0.7922 corrupt_frac_t_0p6_0p8=0.2355 acc_corrupt_t_0p8_1p0=0.9367 corrupt_frac_t_0p8_1p0=0.2506 wrong_frac=0.4586 init_acc_corrupt=0.5121 init_gold_top10=0.5378 init_gold_top100=0.5414
152
+ step=4700 micro_steps=9400 elapsed=32.5s lr=3.000000e-04 loss=2.4545 loss_recon=2.4545 loss_meanflow=0.0000 mean_model_t=0.5058 mean_corrupt_t=0.5058 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7086 acc_corrupt=0.5420 corrupt_frac=0.6108 loss_all=1.9604 loss_corrupt=3.0636 acc_corrupt_t_0p0_0p2=0.2045 corrupt_frac_t_0p0_0p2=0.3000 acc_corrupt_t_0p2_0p4=0.4304 corrupt_frac_t_0p2_0p4=0.1996 acc_corrupt_t_0p4_0p6=0.6316 corrupt_frac_t_0p4_0p6=0.1709 acc_corrupt_t_0p6_0p8=0.7888 corrupt_frac_t_0p6_0p8=0.1505 acc_corrupt_t_0p8_1p0=0.9386 corrupt_frac_t_0p8_1p0=0.1791 wrong_frac=0.5675 init_acc_corrupt=0.3859 init_gold_top10=0.4249 init_gold_top100=0.4315
153
+ step=4800 micro_steps=9600 elapsed=32.6s lr=3.000000e-04 loss=2.4391 loss_recon=2.4391 loss_meanflow=0.0000 mean_model_t=0.5036 mean_corrupt_t=0.5036 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7823 acc_corrupt=0.6230 corrupt_frac=0.5531 loss_all=1.4630 loss_corrupt=2.5318 acc_corrupt_t_0p0_0p2=0.2464 corrupt_frac_t_0p0_0p2=0.0923 acc_corrupt_t_0p2_0p4=0.4334 corrupt_frac_t_0p2_0p4=0.3315 acc_corrupt_t_0p4_0p6=0.6403 corrupt_frac_t_0p4_0p6=0.1896 acc_corrupt_t_0p6_0p8=0.8045 corrupt_frac_t_0p6_0p8=0.1953 acc_corrupt_t_0p8_1p0=0.9308 corrupt_frac_t_0p8_1p0=0.1913 wrong_frac=0.4981 init_acc_corrupt=0.4646 init_gold_top10=0.4988 init_gold_top100=0.5021
154
+ step=4900 micro_steps=9800 elapsed=32.7s lr=3.000000e-04 loss=2.4564 loss_recon=2.4564 loss_meanflow=0.0000 mean_model_t=0.5003 mean_corrupt_t=0.5003 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7985 acc_corrupt=0.6209 corrupt_frac=0.5010 loss_all=1.3425 loss_corrupt=2.5106 acc_corrupt_t_0p0_0p2=0.2153 corrupt_frac_t_0p0_0p2=0.3090 acc_corrupt_t_0p2_0p4=0.5072 corrupt_frac_t_0p2_0p4=0.0846 acc_corrupt_t_0p4_0p6=0.6897 corrupt_frac_t_0p4_0p6=0.1908 acc_corrupt_t_0p6_0p8=0.8615 corrupt_frac_t_0p6_0p8=0.1830 acc_corrupt_t_0p8_1p0=0.9550 corrupt_frac_t_0p8_1p0=0.2327 wrong_frac=0.4963 init_acc_corrupt=0.4730 init_gold_top10=0.4961 init_gold_top100=0.5032
155
+ step=5000 micro_steps=10000 elapsed=32.5s lr=3.000000e-04 loss=2.4344 loss_recon=2.4344 loss_meanflow=0.0000 mean_model_t=0.5035 mean_corrupt_t=0.5035 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7706 acc_corrupt=0.5883 corrupt_frac=0.5289 loss_all=1.4750 loss_corrupt=2.6390 acc_corrupt_t_0p0_0p2=0.2517 corrupt_frac_t_0p0_0p2=0.2677 acc_corrupt_t_0p2_0p4=0.4609 corrupt_frac_t_0p2_0p4=0.1918 acc_corrupt_t_0p4_0p6=0.6821 corrupt_frac_t_0p4_0p6=0.2389 acc_corrupt_t_0p6_0p8=0.8268 corrupt_frac_t_0p6_0p8=0.1479 acc_corrupt_t_0p8_1p0=0.9580 corrupt_frac_t_0p8_1p0=0.1537 wrong_frac=0.5430 init_acc_corrupt=0.4145 init_gold_top10=0.4503 init_gold_top100=0.4567
156
+ step=5100 micro_steps=10200 elapsed=37.6s lr=3.000000e-04 loss=2.4830 loss_recon=2.4830 loss_meanflow=0.0000 mean_model_t=0.4969 mean_corrupt_t=0.4969 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7603 acc_corrupt=0.5806 corrupt_frac=0.5391 loss_all=1.5845 loss_corrupt=2.7692 acc_corrupt_t_0p0_0p2=0.2298 corrupt_frac_t_0p0_0p2=0.3193 acc_corrupt_t_0p2_0p4=0.4959 corrupt_frac_t_0p2_0p4=0.1365 acc_corrupt_t_0p4_0p6=0.6750 corrupt_frac_t_0p4_0p6=0.1540 acc_corrupt_t_0p6_0p8=0.7960 corrupt_frac_t_0p6_0p8=0.2131 acc_corrupt_t_0p8_1p0=0.9373 corrupt_frac_t_0p8_1p0=0.1771 wrong_frac=0.5403 init_acc_corrupt=0.4173 init_gold_top10=0.4497 init_gold_top100=0.4601
157
+ step=5200 micro_steps=10400 elapsed=32.5s lr=3.000000e-04 loss=2.4681 loss_recon=2.4681 loss_meanflow=0.0000 mean_model_t=0.4963 mean_corrupt_t=0.4963 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7826 acc_corrupt=0.6126 corrupt_frac=0.5482 loss_all=1.4255 loss_corrupt=2.5267 acc_corrupt_t_0p0_0p2=0.2335 corrupt_frac_t_0p0_0p2=0.1792 acc_corrupt_t_0p2_0p4=0.4310 corrupt_frac_t_0p2_0p4=0.2129 acc_corrupt_t_0p4_0p6=0.6397 corrupt_frac_t_0p4_0p6=0.1953 acc_corrupt_t_0p6_0p8=0.8191 corrupt_frac_t_0p6_0p8=0.2843 acc_corrupt_t_0p8_1p0=0.9444 corrupt_frac_t_0p8_1p0=0.1283 wrong_frac=0.5006 init_acc_corrupt=0.4634 init_gold_top10=0.4930 init_gold_top100=0.4990
158
+ step=5300 micro_steps=10600 elapsed=32.3s lr=3.000000e-04 loss=2.4746 loss_recon=2.4746 loss_meanflow=0.0000 mean_model_t=0.4938 mean_corrupt_t=0.4938 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7745 acc_corrupt=0.6295 corrupt_frac=0.5829 loss_all=1.4651 loss_corrupt=2.3953 acc_corrupt_t_0p0_0p2=0.2724 corrupt_frac_t_0p0_0p2=0.1546 acc_corrupt_t_0p2_0p4=0.3986 corrupt_frac_t_0p2_0p4=0.2312 acc_corrupt_t_0p4_0p6=0.6560 corrupt_frac_t_0p4_0p6=0.1960 acc_corrupt_t_0p6_0p8=0.8219 corrupt_frac_t_0p6_0p8=0.2069 acc_corrupt_t_0p8_1p0=0.9306 corrupt_frac_t_0p8_1p0=0.2113 wrong_frac=0.4936 init_acc_corrupt=0.4643 init_gold_top10=0.4999 init_gold_top100=0.5066
159
+ step=5400 micro_steps=10800 elapsed=32.8s lr=3.000000e-04 loss=2.4106 loss_recon=2.4106 loss_meanflow=0.0000 mean_model_t=0.5054 mean_corrupt_t=0.5054 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8022 acc_corrupt=0.6526 corrupt_frac=0.5460 loss_all=1.2926 loss_corrupt=2.2511 acc_corrupt_t_0p0_0p2=0.2210 corrupt_frac_t_0p0_0p2=0.1983 acc_corrupt_t_0p2_0p4=0.4158 corrupt_frac_t_0p2_0p4=0.1355 acc_corrupt_t_0p4_0p6=0.6766 corrupt_frac_t_0p4_0p6=0.1645 acc_corrupt_t_0p6_0p8=0.8182 corrupt_frac_t_0p6_0p8=0.2238 acc_corrupt_t_0p8_1p0=0.9284 corrupt_frac_t_0p8_1p0=0.2779 wrong_frac=0.4664 init_acc_corrupt=0.5001 init_gold_top10=0.5265 init_gold_top100=0.5321
160
+ step=5500 micro_steps=11000 elapsed=32.5s lr=3.000000e-04 loss=2.4468 loss_recon=2.4468 loss_meanflow=0.0000 mean_model_t=0.5043 mean_corrupt_t=0.5043 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7991 acc_corrupt=0.6389 corrupt_frac=0.5321 loss_all=1.3142 loss_corrupt=2.3258 acc_corrupt_t_0p0_0p2=0.2688 corrupt_frac_t_0p0_0p2=0.2381 acc_corrupt_t_0p2_0p4=0.5144 corrupt_frac_t_0p2_0p4=0.1757 acc_corrupt_t_0p4_0p6=0.6767 corrupt_frac_t_0p4_0p6=0.1902 acc_corrupt_t_0p6_0p8=0.8341 corrupt_frac_t_0p6_0p8=0.1562 acc_corrupt_t_0p8_1p0=0.9407 corrupt_frac_t_0p8_1p0=0.2397 wrong_frac=0.4861 init_acc_corrupt=0.4786 init_gold_top10=0.5047 init_gold_top100=0.5136
161
+ step=5600 micro_steps=11200 elapsed=32.6s lr=3.000000e-04 loss=2.4220 loss_recon=2.4220 loss_meanflow=0.0000 mean_model_t=0.5013 mean_corrupt_t=0.5013 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7941 acc_corrupt=0.6276 corrupt_frac=0.5294 loss_all=1.3502 loss_corrupt=2.4375 acc_corrupt_t_0p0_0p2=0.2587 corrupt_frac_t_0p0_0p2=0.2460 acc_corrupt_t_0p2_0p4=0.4496 corrupt_frac_t_0p2_0p4=0.1282 acc_corrupt_t_0p4_0p6=0.6560 corrupt_frac_t_0p4_0p6=0.1877 acc_corrupt_t_0p6_0p8=0.8123 corrupt_frac_t_0p6_0p8=0.2174 acc_corrupt_t_0p8_1p0=0.9363 corrupt_frac_t_0p8_1p0=0.2207 wrong_frac=0.4807 init_acc_corrupt=0.4844 init_gold_top10=0.5112 init_gold_top100=0.5188
162
+ step=5700 micro_steps=11400 elapsed=32.5s lr=3.000000e-04 loss=2.4515 loss_recon=2.4515 loss_meanflow=0.0000 mean_model_t=0.4947 mean_corrupt_t=0.4947 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7612 acc_corrupt=0.5910 corrupt_frac=0.5656 loss_all=1.5814 loss_corrupt=2.6812 acc_corrupt_t_0p0_0p2=0.2500 corrupt_frac_t_0p0_0p2=0.2297 acc_corrupt_t_0p2_0p4=0.4344 corrupt_frac_t_0p2_0p4=0.2420 acc_corrupt_t_0p4_0p6=0.6774 corrupt_frac_t_0p4_0p6=0.1545 acc_corrupt_t_0p6_0p8=0.7939 corrupt_frac_t_0p6_0p8=0.1854 acc_corrupt_t_0p8_1p0=0.9370 corrupt_frac_t_0p8_1p0=0.1884 wrong_frac=0.5325 init_acc_corrupt=0.4220 init_gold_top10=0.4619 init_gold_top100=0.4682
163
+ step=5800 micro_steps=11600 elapsed=32.4s lr=3.000000e-04 loss=2.4125 loss_recon=2.4125 loss_meanflow=0.0000 mean_model_t=0.5063 mean_corrupt_t=0.5063 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8220 acc_corrupt=0.6771 corrupt_frac=0.5248 loss_all=1.1475 loss_corrupt=2.0689 acc_corrupt_t_0p0_0p2=0.2766 corrupt_frac_t_0p0_0p2=0.1707 acc_corrupt_t_0p2_0p4=0.4457 corrupt_frac_t_0p2_0p4=0.1649 acc_corrupt_t_0p4_0p6=0.6473 corrupt_frac_t_0p4_0p6=0.1682 acc_corrupt_t_0p6_0p8=0.8585 corrupt_frac_t_0p6_0p8=0.1891 acc_corrupt_t_0p8_1p0=0.9288 corrupt_frac_t_0p8_1p0=0.3070 wrong_frac=0.4359 init_acc_corrupt=0.5327 init_gold_top10=0.5576 init_gold_top100=0.5641
164
+ step=5900 micro_steps=11800 elapsed=33.1s lr=3.000000e-04 loss=2.4244 loss_recon=2.4244 loss_meanflow=0.0000 mean_model_t=0.4958 mean_corrupt_t=0.4958 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7633 acc_corrupt=0.6075 corrupt_frac=0.5831 loss_all=1.5579 loss_corrupt=2.5681 acc_corrupt_t_0p0_0p2=0.2986 corrupt_frac_t_0p0_0p2=0.2152 acc_corrupt_t_0p2_0p4=0.3992 corrupt_frac_t_0p2_0p4=0.2596 acc_corrupt_t_0p4_0p6=0.6643 corrupt_frac_t_0p4_0p6=0.1179 acc_corrupt_t_0p6_0p8=0.7985 corrupt_frac_t_0p6_0p8=0.1662 acc_corrupt_t_0p8_1p0=0.9479 corrupt_frac_t_0p8_1p0=0.2412 wrong_frac=0.5106 init_acc_corrupt=0.4398 init_gold_top10=0.4815 init_gold_top100=0.4901
165
+ step=6000 micro_steps=12000 elapsed=32.5s lr=3.000000e-04 loss=2.4245 loss_recon=2.4245 loss_meanflow=0.0000 mean_model_t=0.4974 mean_corrupt_t=0.4974 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7959 acc_corrupt=0.6458 corrupt_frac=0.5508 loss_all=1.2778 loss_corrupt=2.1981 acc_corrupt_t_0p0_0p2=0.2380 corrupt_frac_t_0p0_0p2=0.2682 acc_corrupt_t_0p2_0p4=0.5064 corrupt_frac_t_0p2_0p4=0.1033 acc_corrupt_t_0p4_0p6=0.7004 corrupt_frac_t_0p4_0p6=0.1095 acc_corrupt_t_0p6_0p8=0.8166 corrupt_frac_t_0p6_0p8=0.3105 acc_corrupt_t_0p8_1p0=0.9564 corrupt_frac_t_0p8_1p0=0.2086 wrong_frac=0.4774 init_acc_corrupt=0.4869 init_gold_top10=0.5166 init_gold_top100=0.5224
166
+ step=6100 micro_steps=12200 elapsed=37.6s lr=3.000000e-04 loss=2.4481 loss_recon=2.4481 loss_meanflow=0.0000 mean_model_t=0.4961 mean_corrupt_t=0.4961 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7778 acc_corrupt=0.6372 corrupt_frac=0.5912 loss_all=1.4559 loss_corrupt=2.3530 acc_corrupt_t_0p0_0p2=0.2159 corrupt_frac_t_0p0_0p2=0.1893 acc_corrupt_t_0p2_0p4=0.4456 corrupt_frac_t_0p2_0p4=0.2507 acc_corrupt_t_0p4_0p6=0.6592 corrupt_frac_t_0p4_0p6=0.0648 acc_corrupt_t_0p6_0p8=0.8173 corrupt_frac_t_0p6_0p8=0.2057 acc_corrupt_t_0p8_1p0=0.9458 corrupt_frac_t_0p8_1p0=0.2895 wrong_frac=0.4586 init_acc_corrupt=0.4989 init_gold_top10=0.5356 init_gold_top100=0.5400
167
+ step=6200 micro_steps=12400 elapsed=33.1s lr=3.000000e-04 loss=2.3905 loss_recon=2.3905 loss_meanflow=0.0000 mean_model_t=0.5014 mean_corrupt_t=0.5014 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7460 acc_corrupt=0.5732 corrupt_frac=0.5775 loss_all=1.7144 loss_corrupt=2.8578 acc_corrupt_t_0p0_0p2=0.2201 corrupt_frac_t_0p0_0p2=0.3149 acc_corrupt_t_0p2_0p4=0.4610 corrupt_frac_t_0p2_0p4=0.1706 acc_corrupt_t_0p4_0p6=0.6982 corrupt_frac_t_0p4_0p6=0.1296 acc_corrupt_t_0p6_0p8=0.8184 corrupt_frac_t_0p6_0p8=0.2257 acc_corrupt_t_0p8_1p0=0.9429 corrupt_frac_t_0p8_1p0=0.1592 wrong_frac=0.5462 init_acc_corrupt=0.4105 init_gold_top10=0.4437 init_gold_top100=0.4538
168
+ step=6300 micro_steps=12600 elapsed=32.5s lr=3.000000e-04 loss=2.3994 loss_recon=2.3994 loss_meanflow=0.0000 mean_model_t=0.5005 mean_corrupt_t=0.5005 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7892 acc_corrupt=0.6288 corrupt_frac=0.5511 loss_all=1.4068 loss_corrupt=2.4614 acc_corrupt_t_0p0_0p2=0.2055 corrupt_frac_t_0p0_0p2=0.1681 acc_corrupt_t_0p2_0p4=0.4378 corrupt_frac_t_0p2_0p4=0.2651 acc_corrupt_t_0p4_0p6=0.7108 corrupt_frac_t_0p4_0p6=0.1455 acc_corrupt_t_0p6_0p8=0.8432 corrupt_frac_t_0p6_0p8=0.2062 acc_corrupt_t_0p8_1p0=0.9341 corrupt_frac_t_0p8_1p0=0.2151 wrong_frac=0.4859 init_acc_corrupt=0.4724 init_gold_top10=0.5068 init_gold_top100=0.5147
169
+ step=6400 micro_steps=12800 elapsed=32.4s lr=3.000000e-04 loss=2.3480 loss_recon=2.3480 loss_meanflow=0.0000 mean_model_t=0.5049 mean_corrupt_t=0.5049 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8137 acc_corrupt=0.6455 corrupt_frac=0.5031 loss_all=1.1801 loss_corrupt=2.2205 acc_corrupt_t_0p0_0p2=0.2568 corrupt_frac_t_0p0_0p2=0.1437 acc_corrupt_t_0p2_0p4=0.4931 corrupt_frac_t_0p2_0p4=0.2451 acc_corrupt_t_0p4_0p6=0.6363 corrupt_frac_t_0p4_0p6=0.1995 acc_corrupt_t_0p6_0p8=0.8122 corrupt_frac_t_0p6_0p8=0.1667 acc_corrupt_t_0p8_1p0=0.9198 corrupt_frac_t_0p8_1p0=0.2451 wrong_frac=0.4858 init_acc_corrupt=0.4771 init_gold_top10=0.5096 init_gold_top100=0.5130
170
+ step=6500 micro_steps=13000 elapsed=32.4s lr=3.000000e-04 loss=2.4436 loss_recon=2.4436 loss_meanflow=0.0000 mean_model_t=0.4955 mean_corrupt_t=0.4955 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7321 acc_corrupt=0.5374 corrupt_frac=0.5608 loss_all=1.7895 loss_corrupt=3.0628 acc_corrupt_t_0p0_0p2=0.2173 corrupt_frac_t_0p0_0p2=0.3045 acc_corrupt_t_0p2_0p4=0.4596 corrupt_frac_t_0p2_0p4=0.1833 acc_corrupt_t_0p4_0p6=0.5965 corrupt_frac_t_0p4_0p6=0.2007 acc_corrupt_t_0p6_0p8=0.8041 corrupt_frac_t_0p6_0p8=0.1922 acc_corrupt_t_0p8_1p0=0.9453 corrupt_frac_t_0p8_1p0=0.1193 wrong_frac=0.5869 init_acc_corrupt=0.3755 init_gold_top10=0.4066 init_gold_top100=0.4140
171
+ step=6600 micro_steps=13200 elapsed=42.3s lr=3.000000e-04 loss=2.3843 loss_recon=2.3843 loss_meanflow=0.0000 mean_model_t=0.5030 mean_corrupt_t=0.5030 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7758 acc_corrupt=0.6467 corrupt_frac=0.6144 loss_all=1.5029 loss_corrupt=2.3503 acc_corrupt_t_0p0_0p2=0.2272 corrupt_frac_t_0p0_0p2=0.2291 acc_corrupt_t_0p2_0p4=0.5172 corrupt_frac_t_0p2_0p4=0.1560 acc_corrupt_t_0p4_0p6=0.6663 corrupt_frac_t_0p4_0p6=0.2043 acc_corrupt_t_0p6_0p8=0.8206 corrupt_frac_t_0p6_0p8=0.1218 acc_corrupt_t_0p8_1p0=0.9622 corrupt_frac_t_0p8_1p0=0.2889 wrong_frac=0.4568 init_acc_corrupt=0.5172 init_gold_top10=0.5394 init_gold_top100=0.5430
172
+ step=6700 micro_steps=13400 elapsed=33.4s lr=3.000000e-04 loss=2.4084 loss_recon=2.4084 loss_meanflow=0.0000 mean_model_t=0.4973 mean_corrupt_t=0.4973 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8079 acc_corrupt=0.6287 corrupt_frac=0.4918 loss_all=1.2711 loss_corrupt=2.4412 acc_corrupt_t_0p0_0p2=0.2835 corrupt_frac_t_0p0_0p2=0.2355 acc_corrupt_t_0p2_0p4=0.4284 corrupt_frac_t_0p2_0p4=0.2184 acc_corrupt_t_0p4_0p6=0.6678 corrupt_frac_t_0p4_0p6=0.0740 acc_corrupt_t_0p6_0p8=0.8468 corrupt_frac_t_0p6_0p8=0.2755 acc_corrupt_t_0p8_1p0=0.9444 corrupt_frac_t_0p8_1p0=0.1966 wrong_frac=0.5034 init_acc_corrupt=0.4552 init_gold_top10=0.4892 init_gold_top100=0.4966
173
+ step=6800 micro_steps=13600 elapsed=32.4s lr=3.000000e-04 loss=2.4036 loss_recon=2.4036 loss_meanflow=0.0000 mean_model_t=0.4942 mean_corrupt_t=0.4942 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7295 acc_corrupt=0.5820 corrupt_frac=0.6150 loss_all=1.7332 loss_corrupt=2.6699 acc_corrupt_t_0p0_0p2=0.2282 corrupt_frac_t_0p0_0p2=0.2279 acc_corrupt_t_0p2_0p4=0.4094 corrupt_frac_t_0p2_0p4=0.2070 acc_corrupt_t_0p4_0p6=0.6713 corrupt_frac_t_0p4_0p6=0.2295 acc_corrupt_t_0p6_0p8=0.8174 corrupt_frac_t_0p6_0p8=0.1967 acc_corrupt_t_0p8_1p0=0.9386 corrupt_frac_t_0p8_1p0=0.1389 wrong_frac=0.5345 init_acc_corrupt=0.4242 init_gold_top10=0.4595 init_gold_top100=0.4649
174
+ step=6900 micro_steps=13800 elapsed=32.4s lr=3.000000e-04 loss=2.3715 loss_recon=2.3715 loss_meanflow=0.0000 mean_model_t=0.5004 mean_corrupt_t=0.5004 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7765 acc_corrupt=0.6192 corrupt_frac=0.5658 loss_all=1.4830 loss_corrupt=2.5086 acc_corrupt_t_0p0_0p2=0.2265 corrupt_frac_t_0p0_0p2=0.2248 acc_corrupt_t_0p2_0p4=0.4439 corrupt_frac_t_0p2_0p4=0.2289 acc_corrupt_t_0p4_0p6=0.6961 corrupt_frac_t_0p4_0p6=0.1172 acc_corrupt_t_0p6_0p8=0.8488 corrupt_frac_t_0p6_0p8=0.1584 acc_corrupt_t_0p8_1p0=0.9259 corrupt_frac_t_0p8_1p0=0.2708 wrong_frac=0.4889 init_acc_corrupt=0.4647 init_gold_top10=0.5066 init_gold_top100=0.5122
175
+ step=7000 micro_steps=14000 elapsed=32.4s lr=3.000000e-04 loss=2.3797 loss_recon=2.3797 loss_meanflow=0.0000 mean_model_t=0.5042 mean_corrupt_t=0.5042 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8234 acc_corrupt=0.6913 corrupt_frac=0.5504 loss_all=1.1000 loss_corrupt=1.9145 acc_corrupt_t_0p0_0p2=0.3303 corrupt_frac_t_0p0_0p2=0.0960 acc_corrupt_t_0p2_0p4=0.4384 corrupt_frac_t_0p2_0p4=0.1457 acc_corrupt_t_0p4_0p6=0.6367 corrupt_frac_t_0p4_0p6=0.2686 acc_corrupt_t_0p6_0p8=0.8262 corrupt_frac_t_0p6_0p8=0.3253 acc_corrupt_t_0p8_1p0=0.9487 corrupt_frac_t_0p8_1p0=0.1643 wrong_frac=0.4351 init_acc_corrupt=0.5387 init_gold_top10=0.5613 init_gold_top100=0.5669
176
+ step=7100 micro_steps=14200 elapsed=36.3s lr=3.000000e-04 loss=2.3705 loss_recon=2.3705 loss_meanflow=0.0000 mean_model_t=0.4987 mean_corrupt_t=0.4987 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7817 acc_corrupt=0.5911 corrupt_frac=0.5114 loss_all=1.4174 loss_corrupt=2.6212 acc_corrupt_t_0p0_0p2=0.2580 corrupt_frac_t_0p0_0p2=0.3053 acc_corrupt_t_0p2_0p4=0.4535 corrupt_frac_t_0p2_0p4=0.1258 acc_corrupt_t_0p4_0p6=0.6581 corrupt_frac_t_0p4_0p6=0.1955 acc_corrupt_t_0p6_0p8=0.8037 corrupt_frac_t_0p6_0p8=0.1812 acc_corrupt_t_0p8_1p0=0.9416 corrupt_frac_t_0p8_1p0=0.1922 wrong_frac=0.5335 init_acc_corrupt=0.4249 init_gold_top10=0.4581 init_gold_top100=0.4648
177
+ step=7200 micro_steps=14400 elapsed=33.6s lr=3.000000e-04 loss=2.3550 loss_recon=2.3550 loss_meanflow=0.0000 mean_model_t=0.5025 mean_corrupt_t=0.5025 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8049 acc_corrupt=0.6077 corrupt_frac=0.4841 loss_all=1.2914 loss_corrupt=2.5661 acc_corrupt_t_0p0_0p2=0.2202 corrupt_frac_t_0p0_0p2=0.2267 acc_corrupt_t_0p2_0p4=0.4738 corrupt_frac_t_0p2_0p4=0.2410 acc_corrupt_t_0p4_0p6=0.6773 corrupt_frac_t_0p4_0p6=0.1899 acc_corrupt_t_0p6_0p8=0.8246 corrupt_frac_t_0p6_0p8=0.0819 acc_corrupt_t_0p8_1p0=0.9497 corrupt_frac_t_0p8_1p0=0.2605 wrong_frac=0.5108 init_acc_corrupt=0.4423 init_gold_top10=0.4801 init_gold_top100=0.4881
178
+ step=7300 micro_steps=14600 elapsed=32.9s lr=3.000000e-04 loss=2.4168 loss_recon=2.4168 loss_meanflow=0.0000 mean_model_t=0.4967 mean_corrupt_t=0.4967 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7893 acc_corrupt=0.6238 corrupt_frac=0.5403 loss_all=1.3718 loss_corrupt=2.4203 acc_corrupt_t_0p0_0p2=0.2840 corrupt_frac_t_0p0_0p2=0.2307 acc_corrupt_t_0p2_0p4=0.4683 corrupt_frac_t_0p2_0p4=0.1959 acc_corrupt_t_0p4_0p6=0.6545 corrupt_frac_t_0p4_0p6=0.1726 acc_corrupt_t_0p6_0p8=0.8180 corrupt_frac_t_0p6_0p8=0.2185 acc_corrupt_t_0p8_1p0=0.9591 corrupt_frac_t_0p8_1p0=0.1823 wrong_frac=0.5099 init_acc_corrupt=0.4607 init_gold_top10=0.4858 init_gold_top100=0.4896
179
+ step=7400 micro_steps=14800 elapsed=32.3s lr=3.000000e-04 loss=2.3490 loss_recon=2.3490 loss_meanflow=0.0000 mean_model_t=0.5032 mean_corrupt_t=0.5032 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8151 acc_corrupt=0.6857 corrupt_frac=0.5667 loss_all=1.1789 loss_corrupt=1.9786 acc_corrupt_t_0p0_0p2=0.3457 corrupt_frac_t_0p0_0p2=0.1159 acc_corrupt_t_0p2_0p4=0.4903 corrupt_frac_t_0p2_0p4=0.2342 acc_corrupt_t_0p4_0p6=0.6839 corrupt_frac_t_0p4_0p6=0.1963 acc_corrupt_t_0p6_0p8=0.8146 corrupt_frac_t_0p6_0p8=0.2359 acc_corrupt_t_0p8_1p0=0.9387 corrupt_frac_t_0p8_1p0=0.2178 wrong_frac=0.4505 init_acc_corrupt=0.5265 init_gold_top10=0.5472 init_gold_top100=0.5504
180
+ step=7500 micro_steps=15000 elapsed=32.5s lr=3.000000e-04 loss=2.3399 loss_recon=2.3399 loss_meanflow=0.0000 mean_model_t=0.5055 mean_corrupt_t=0.5055 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7861 acc_corrupt=0.6251 corrupt_frac=0.5564 loss_all=1.4150 loss_corrupt=2.4611 acc_corrupt_t_0p0_0p2=0.2451 corrupt_frac_t_0p0_0p2=0.2122 acc_corrupt_t_0p2_0p4=0.4603 corrupt_frac_t_0p2_0p4=0.2102 acc_corrupt_t_0p4_0p6=0.6397 corrupt_frac_t_0p4_0p6=0.1272 acc_corrupt_t_0p6_0p8=0.8209 corrupt_frac_t_0p6_0p8=0.2475 acc_corrupt_t_0p8_1p0=0.9449 corrupt_frac_t_0p8_1p0=0.2029 wrong_frac=0.4901 init_acc_corrupt=0.4752 init_gold_top10=0.5018 init_gold_top100=0.5101
181
+ step=7600 micro_steps=15200 elapsed=32.5s lr=3.000000e-04 loss=2.3615 loss_recon=2.3615 loss_meanflow=0.0000 mean_model_t=0.5041 mean_corrupt_t=0.5041 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7592 acc_corrupt=0.5780 corrupt_frac=0.5476 loss_all=1.5827 loss_corrupt=2.7427 acc_corrupt_t_0p0_0p2=0.2921 corrupt_frac_t_0p0_0p2=0.2579 acc_corrupt_t_0p2_0p4=0.4650 corrupt_frac_t_0p2_0p4=0.2454 acc_corrupt_t_0p4_0p6=0.6878 corrupt_frac_t_0p4_0p6=0.2499 acc_corrupt_t_0p6_0p8=0.8128 corrupt_frac_t_0p6_0p8=0.1322 acc_corrupt_t_0p8_1p0=0.9533 corrupt_frac_t_0p8_1p0=0.1146 wrong_frac=0.5742 init_acc_corrupt=0.3839 init_gold_top10=0.4195 init_gold_top100=0.4249
182
+ step=7700 micro_steps=15400 elapsed=32.8s lr=3.000000e-04 loss=2.3903 loss_recon=2.3903 loss_meanflow=0.0000 mean_model_t=0.5021 mean_corrupt_t=0.5021 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7935 acc_corrupt=0.6444 corrupt_frac=0.5568 loss_all=1.3285 loss_corrupt=2.2563 acc_corrupt_t_0p0_0p2=0.3152 corrupt_frac_t_0p0_0p2=0.1934 acc_corrupt_t_0p2_0p4=0.4771 corrupt_frac_t_0p2_0p4=0.1815 acc_corrupt_t_0p4_0p6=0.6544 corrupt_frac_t_0p4_0p6=0.1783 acc_corrupt_t_0p6_0p8=0.8102 corrupt_frac_t_0p6_0p8=0.2934 acc_corrupt_t_0p8_1p0=0.9286 corrupt_frac_t_0p8_1p0=0.1535 wrong_frac=0.4973 init_acc_corrupt=0.4692 init_gold_top10=0.4979 init_gold_top100=0.5030
183
+ step=7800 micro_steps=15600 elapsed=32.5s lr=3.000000e-04 loss=2.3804 loss_recon=2.3804 loss_meanflow=0.0000 mean_model_t=0.5005 mean_corrupt_t=0.5005 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8365 acc_corrupt=0.6929 corrupt_frac=0.5143 loss_all=1.0226 loss_corrupt=1.9068 acc_corrupt_t_0p0_0p2=0.2763 corrupt_frac_t_0p0_0p2=0.1014 acc_corrupt_t_0p2_0p4=0.4734 corrupt_frac_t_0p2_0p4=0.1870 acc_corrupt_t_0p4_0p6=0.7243 corrupt_frac_t_0p4_0p6=0.3211 acc_corrupt_t_0p6_0p8=0.7952 corrupt_frac_t_0p6_0p8=0.2086 acc_corrupt_t_0p8_1p0=0.9778 corrupt_frac_t_0p8_1p0=0.1818 wrong_frac=0.4515 init_acc_corrupt=0.5191 init_gold_top10=0.5431 init_gold_top100=0.5490
184
+ step=7900 micro_steps=15800 elapsed=32.3s lr=3.000000e-04 loss=2.3912 loss_recon=2.3912 loss_meanflow=0.0000 mean_model_t=0.5022 mean_corrupt_t=0.5022 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7556 acc_corrupt=0.5700 corrupt_frac=0.5474 loss_all=1.5953 loss_corrupt=2.7852 acc_corrupt_t_0p0_0p2=0.2391 corrupt_frac_t_0p0_0p2=0.2565 acc_corrupt_t_0p2_0p4=0.4363 corrupt_frac_t_0p2_0p4=0.2346 acc_corrupt_t_0p4_0p6=0.6509 corrupt_frac_t_0p4_0p6=0.1552 acc_corrupt_t_0p6_0p8=0.8308 corrupt_frac_t_0p6_0p8=0.2043 acc_corrupt_t_0p8_1p0=0.9075 corrupt_frac_t_0p8_1p0=0.1494 wrong_frac=0.5500 init_acc_corrupt=0.4072 init_gold_top10=0.4413 init_gold_top100=0.4503
185
+ step=8000 micro_steps=16000 elapsed=32.5s lr=3.000000e-04 loss=2.3689 loss_recon=2.3689 loss_meanflow=0.0000 mean_model_t=0.4999 mean_corrupt_t=0.4999 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7850 acc_corrupt=0.6320 corrupt_frac=0.5605 loss_all=1.3855 loss_corrupt=2.3546 acc_corrupt_t_0p0_0p2=0.2671 corrupt_frac_t_0p0_0p2=0.2267 acc_corrupt_t_0p2_0p4=0.4596 corrupt_frac_t_0p2_0p4=0.1455 acc_corrupt_t_0p4_0p6=0.6648 corrupt_frac_t_0p4_0p6=0.2306 acc_corrupt_t_0p6_0p8=0.8443 corrupt_frac_t_0p6_0p8=0.2154 acc_corrupt_t_0p8_1p0=0.9317 corrupt_frac_t_0p8_1p0=0.1818 wrong_frac=0.5087 init_acc_corrupt=0.4556 init_gold_top10=0.4858 init_gold_top100=0.4904
186
+ step=8100 micro_steps=16200 elapsed=43.4s lr=3.000000e-04 loss=2.4130 loss_recon=2.4130 loss_meanflow=0.0000 mean_model_t=0.4973 mean_corrupt_t=0.4973 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7635 acc_corrupt=0.5840 corrupt_frac=0.5569 loss_all=1.5969 loss_corrupt=2.7884 acc_corrupt_t_0p0_0p2=0.2493 corrupt_frac_t_0p0_0p2=0.2207 acc_corrupt_t_0p2_0p4=0.4245 corrupt_frac_t_0p2_0p4=0.2742 acc_corrupt_t_0p4_0p6=0.6915 corrupt_frac_t_0p4_0p6=0.1968 acc_corrupt_t_0p6_0p8=0.8569 corrupt_frac_t_0p6_0p8=0.1471 acc_corrupt_t_0p8_1p0=0.9333 corrupt_frac_t_0p8_1p0=0.1611 wrong_frac=0.5285 init_acc_corrupt=0.4158 init_gold_top10=0.4645 init_gold_top100=0.4722
187
+ step=8200 micro_steps=16400 elapsed=47.5s lr=3.000000e-04 loss=2.3786 loss_recon=2.3786 loss_meanflow=0.0000 mean_model_t=0.5008 mean_corrupt_t=0.5008 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8087 acc_corrupt=0.6504 corrupt_frac=0.5304 loss_all=1.2304 loss_corrupt=2.2301 acc_corrupt_t_0p0_0p2=0.2541 corrupt_frac_t_0p0_0p2=0.1975 acc_corrupt_t_0p2_0p4=0.4862 corrupt_frac_t_0p2_0p4=0.1751 acc_corrupt_t_0p4_0p6=0.7006 corrupt_frac_t_0p4_0p6=0.1853 acc_corrupt_t_0p6_0p8=0.8128 corrupt_frac_t_0p6_0p8=0.2226 acc_corrupt_t_0p8_1p0=0.9308 corrupt_frac_t_0p8_1p0=0.2196 wrong_frac=0.4764 init_acc_corrupt=0.4891 init_gold_top10=0.5169 init_gold_top100=0.5234
188
+ step=8300 micro_steps=16600 elapsed=52.3s lr=3.000000e-04 loss=2.3408 loss_recon=2.3408 loss_meanflow=0.0000 mean_model_t=0.5017 mean_corrupt_t=0.5017 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7971 acc_corrupt=0.6638 corrupt_frac=0.5896 loss_all=1.2689 loss_corrupt=2.0859 acc_corrupt_t_0p0_0p2=0.3463 corrupt_frac_t_0p0_0p2=0.1381 acc_corrupt_t_0p2_0p4=0.4667 corrupt_frac_t_0p2_0p4=0.2178 acc_corrupt_t_0p4_0p6=0.6592 corrupt_frac_t_0p4_0p6=0.2418 acc_corrupt_t_0p6_0p8=0.8142 corrupt_frac_t_0p6_0p8=0.1716 acc_corrupt_t_0p8_1p0=0.9327 corrupt_frac_t_0p8_1p0=0.2306 wrong_frac=0.4718 init_acc_corrupt=0.4973 init_gold_top10=0.5234 init_gold_top100=0.5282
189
+ step=8400 micro_steps=16800 elapsed=32.2s lr=3.000000e-04 loss=2.3167 loss_recon=2.3167 loss_meanflow=0.0000 mean_model_t=0.5048 mean_corrupt_t=0.5048 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7776 acc_corrupt=0.6027 corrupt_frac=0.5371 loss_all=1.4438 loss_corrupt=2.5679 acc_corrupt_t_0p0_0p2=0.3136 corrupt_frac_t_0p0_0p2=0.2145 acc_corrupt_t_0p2_0p4=0.4466 corrupt_frac_t_0p2_0p4=0.2636 acc_corrupt_t_0p4_0p6=0.6600 corrupt_frac_t_0p4_0p6=0.1711 acc_corrupt_t_0p6_0p8=0.8161 corrupt_frac_t_0p6_0p8=0.1952 acc_corrupt_t_0p8_1p0=0.9357 corrupt_frac_t_0p8_1p0=0.1555 wrong_frac=0.5407 init_acc_corrupt=0.4198 init_gold_top10=0.4539 init_gold_top100=0.4595
190
+ step=8500 micro_steps=17000 elapsed=32.5s lr=3.000000e-04 loss=2.3919 loss_recon=2.3919 loss_meanflow=0.0000 mean_model_t=0.4953 mean_corrupt_t=0.4953 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7719 acc_corrupt=0.6052 corrupt_frac=0.5599 loss_all=1.4995 loss_corrupt=2.5646 acc_corrupt_t_0p0_0p2=0.2153 corrupt_frac_t_0p0_0p2=0.2370 acc_corrupt_t_0p2_0p4=0.4455 corrupt_frac_t_0p2_0p4=0.2261 acc_corrupt_t_0p4_0p6=0.7067 corrupt_frac_t_0p4_0p6=0.1145 acc_corrupt_t_0p6_0p8=0.8179 corrupt_frac_t_0p6_0p8=0.2359 acc_corrupt_t_0p8_1p0=0.9626 corrupt_frac_t_0p8_1p0=0.1866 wrong_frac=0.5278 init_acc_corrupt=0.4384 init_gold_top10=0.4661 init_gold_top100=0.4726
191
+ step=8600 micro_steps=17200 elapsed=32.4s lr=3.000000e-04 loss=2.3623 loss_recon=2.3623 loss_meanflow=0.0000 mean_model_t=0.4970 mean_corrupt_t=0.4970 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8030 acc_corrupt=0.6424 corrupt_frac=0.5281 loss_all=1.2839 loss_corrupt=2.3189 acc_corrupt_t_0p0_0p2=0.2156 corrupt_frac_t_0p0_0p2=0.2702 acc_corrupt_t_0p2_0p4=0.5357 corrupt_frac_t_0p2_0p4=0.0453 acc_corrupt_t_0p4_0p6=0.6749 corrupt_frac_t_0p4_0p6=0.2233 acc_corrupt_t_0p6_0p8=0.8419 corrupt_frac_t_0p6_0p8=0.2647 acc_corrupt_t_0p8_1p0=0.9482 corrupt_frac_t_0p8_1p0=0.1965 wrong_frac=0.4820 init_acc_corrupt=0.4988 init_gold_top10=0.5116 init_gold_top100=0.5183
192
+ step=8700 micro_steps=17400 elapsed=32.3s lr=3.000000e-04 loss=2.2614 loss_recon=2.2614 loss_meanflow=0.0000 mean_model_t=0.5106 mean_corrupt_t=0.5106 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8051 acc_corrupt=0.6550 corrupt_frac=0.5427 loss_all=1.2625 loss_corrupt=2.2045 acc_corrupt_t_0p0_0p2=0.2907 corrupt_frac_t_0p0_0p2=0.1725 acc_corrupt_t_0p2_0p4=0.4543 corrupt_frac_t_0p2_0p4=0.1698 acc_corrupt_t_0p4_0p6=0.6771 corrupt_frac_t_0p4_0p6=0.2222 acc_corrupt_t_0p6_0p8=0.7910 corrupt_frac_t_0p6_0p8=0.1894 acc_corrupt_t_0p8_1p0=0.9241 corrupt_frac_t_0p8_1p0=0.2461 wrong_frac=0.4802 init_acc_corrupt=0.4939 init_gold_top10=0.5153 init_gold_top100=0.5196
193
+ step=8800 micro_steps=17600 elapsed=32.6s lr=3.000000e-04 loss=2.3606 loss_recon=2.3606 loss_meanflow=0.0000 mean_model_t=0.5033 mean_corrupt_t=0.5033 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7980 acc_corrupt=0.6842 corrupt_frac=0.6193 loss_all=1.3130 loss_corrupt=2.0470 acc_corrupt_t_0p0_0p2=0.2682 corrupt_frac_t_0p0_0p2=0.1543 acc_corrupt_t_0p2_0p4=0.4499 corrupt_frac_t_0p2_0p4=0.1713 acc_corrupt_t_0p4_0p6=0.6616 corrupt_frac_t_0p4_0p6=0.1287 acc_corrupt_t_0p6_0p8=0.8318 corrupt_frac_t_0p6_0p8=0.3059 acc_corrupt_t_0p8_1p0=0.9433 corrupt_frac_t_0p8_1p0=0.2397 wrong_frac=0.4384 init_acc_corrupt=0.5293 init_gold_top10=0.5567 init_gold_top100=0.5626
194
+ step=8900 micro_steps=17800 elapsed=32.5s lr=3.000000e-04 loss=2.3782 loss_recon=2.3782 loss_meanflow=0.0000 mean_model_t=0.4971 mean_corrupt_t=0.4971 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8019 acc_corrupt=0.6508 corrupt_frac=0.5505 loss_all=1.3063 loss_corrupt=2.2857 acc_corrupt_t_0p0_0p2=0.2642 corrupt_frac_t_0p0_0p2=0.2191 acc_corrupt_t_0p2_0p4=0.4369 corrupt_frac_t_0p2_0p4=0.1177 acc_corrupt_t_0p4_0p6=0.6514 corrupt_frac_t_0p4_0p6=0.2124 acc_corrupt_t_0p6_0p8=0.8333 corrupt_frac_t_0p6_0p8=0.1836 acc_corrupt_t_0p8_1p0=0.9361 corrupt_frac_t_0p8_1p0=0.2672 wrong_frac=0.4916 init_acc_corrupt=0.4858 init_gold_top10=0.5047 init_gold_top100=0.5089
195
+ step=9000 micro_steps=18000 elapsed=32.6s lr=3.000000e-04 loss=2.3931 loss_recon=2.3931 loss_meanflow=0.0000 mean_model_t=0.4952 mean_corrupt_t=0.4952 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7948 acc_corrupt=0.6385 corrupt_frac=0.5464 loss_all=1.3362 loss_corrupt=2.3406 acc_corrupt_t_0p0_0p2=0.2481 corrupt_frac_t_0p0_0p2=0.1747 acc_corrupt_t_0p2_0p4=0.5011 corrupt_frac_t_0p2_0p4=0.2109 acc_corrupt_t_0p4_0p6=0.6327 corrupt_frac_t_0p4_0p6=0.2475 acc_corrupt_t_0p6_0p8=0.8234 corrupt_frac_t_0p6_0p8=0.1126 acc_corrupt_t_0p8_1p0=0.9446 corrupt_frac_t_0p8_1p0=0.2542 wrong_frac=0.4899 init_acc_corrupt=0.4846 init_gold_top10=0.5058 init_gold_top100=0.5098
196
+ step=9100 micro_steps=18200 elapsed=37.7s lr=3.000000e-04 loss=2.3922 loss_recon=2.3922 loss_meanflow=0.0000 mean_model_t=0.4939 mean_corrupt_t=0.4939 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7892 acc_corrupt=0.6223 corrupt_frac=0.5463 loss_all=1.3571 loss_corrupt=2.4083 acc_corrupt_t_0p0_0p2=0.1860 corrupt_frac_t_0p0_0p2=0.1178 acc_corrupt_t_0p2_0p4=0.4215 corrupt_frac_t_0p2_0p4=0.2306 acc_corrupt_t_0p4_0p6=0.6636 corrupt_frac_t_0p4_0p6=0.2644 acc_corrupt_t_0p6_0p8=0.7981 corrupt_frac_t_0p6_0p8=0.2766 acc_corrupt_t_0p8_1p0=0.9677 corrupt_frac_t_0p8_1p0=0.1106 wrong_frac=0.5169 init_acc_corrupt=0.4547 init_gold_top10=0.4809 init_gold_top100=0.4822
197
+ step=9200 micro_steps=18400 elapsed=32.6s lr=3.000000e-04 loss=2.3107 loss_recon=2.3107 loss_meanflow=0.0000 mean_model_t=0.5026 mean_corrupt_t=0.5026 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7659 acc_corrupt=0.5941 corrupt_frac=0.5543 loss_all=1.5224 loss_corrupt=2.6231 acc_corrupt_t_0p0_0p2=0.2399 corrupt_frac_t_0p0_0p2=0.3112 acc_corrupt_t_0p2_0p4=0.4830 corrupt_frac_t_0p2_0p4=0.1163 acc_corrupt_t_0p4_0p6=0.6919 corrupt_frac_t_0p4_0p6=0.2066 acc_corrupt_t_0p6_0p8=0.7984 corrupt_frac_t_0p6_0p8=0.1911 acc_corrupt_t_0p8_1p0=0.9597 corrupt_frac_t_0p8_1p0=0.1749 wrong_frac=0.5426 init_acc_corrupt=0.4272 init_gold_top10=0.4510 init_gold_top100=0.4583
198
+ step=9300 micro_steps=18600 elapsed=32.2s lr=3.000000e-04 loss=2.3304 loss_recon=2.3304 loss_meanflow=0.0000 mean_model_t=0.5023 mean_corrupt_t=0.5023 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7822 acc_corrupt=0.6199 corrupt_frac=0.5540 loss_all=1.4160 loss_corrupt=2.4604 acc_corrupt_t_0p0_0p2=0.2491 corrupt_frac_t_0p0_0p2=0.1911 acc_corrupt_t_0p2_0p4=0.4799 corrupt_frac_t_0p2_0p4=0.2406 acc_corrupt_t_0p4_0p6=0.6622 corrupt_frac_t_0p4_0p6=0.1814 acc_corrupt_t_0p6_0p8=0.8290 corrupt_frac_t_0p6_0p8=0.2552 acc_corrupt_t_0p8_1p0=0.9498 corrupt_frac_t_0p8_1p0=0.1318 wrong_frac=0.5126 init_acc_corrupt=0.4425 init_gold_top10=0.4819 init_gold_top100=0.4877
199
+ step=9400 micro_steps=18800 elapsed=32.5s lr=3.000000e-04 loss=2.3230 loss_recon=2.3230 loss_meanflow=0.0000 mean_model_t=0.5000 mean_corrupt_t=0.5000 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7828 acc_corrupt=0.6168 corrupt_frac=0.5422 loss_all=1.3839 loss_corrupt=2.4266 acc_corrupt_t_0p0_0p2=0.2129 corrupt_frac_t_0p0_0p2=0.2337 acc_corrupt_t_0p2_0p4=0.5255 corrupt_frac_t_0p2_0p4=0.1499 acc_corrupt_t_0p4_0p6=0.6771 corrupt_frac_t_0p4_0p6=0.2593 acc_corrupt_t_0p6_0p8=0.8404 corrupt_frac_t_0p6_0p8=0.2384 acc_corrupt_t_0p8_1p0=0.9469 corrupt_frac_t_0p8_1p0=0.1186 wrong_frac=0.5331 init_acc_corrupt=0.4397 init_gold_top10=0.4629 init_gold_top100=0.4680
200
+ step=9500 micro_steps=19000 elapsed=32.5s lr=3.000000e-04 loss=2.3165 loss_recon=2.3165 loss_meanflow=0.0000 mean_model_t=0.4985 mean_corrupt_t=0.4985 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.8507 acc_corrupt=0.7211 corrupt_frac=0.5214 loss_all=0.9183 loss_corrupt=1.7041 acc_corrupt_t_0p0_0p2=0.3224 corrupt_frac_t_0p0_0p2=0.0995 acc_corrupt_t_0p2_0p4=0.4780 corrupt_frac_t_0p2_0p4=0.1651 acc_corrupt_t_0p4_0p6=0.6488 corrupt_frac_t_0p4_0p6=0.1880 acc_corrupt_t_0p6_0p8=0.8014 corrupt_frac_t_0p6_0p8=0.2334 acc_corrupt_t_0p8_1p0=0.9590 corrupt_frac_t_0p8_1p0=0.3140 wrong_frac=0.3987 init_acc_corrupt=0.5793 init_gold_top10=0.5973 init_gold_top100=0.6017
201
+ step=9600 micro_steps=19200 elapsed=32.9s lr=3.000000e-04 loss=2.3114 loss_recon=2.3114 loss_meanflow=0.0000 mean_model_t=0.5027 mean_corrupt_t=0.5027 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7990 acc_corrupt=0.6431 corrupt_frac=0.5452 loss_all=1.2865 loss_corrupt=2.2516 acc_corrupt_t_0p0_0p2=0.2576 corrupt_frac_t_0p0_0p2=0.1834 acc_corrupt_t_0p2_0p4=0.4611 corrupt_frac_t_0p2_0p4=0.2015 acc_corrupt_t_0p4_0p6=0.7011 corrupt_frac_t_0p4_0p6=0.2143 acc_corrupt_t_0p6_0p8=0.8283 corrupt_frac_t_0p6_0p8=0.2203 acc_corrupt_t_0p8_1p0=0.9429 corrupt_frac_t_0p8_1p0=0.1805 wrong_frac=0.4966 init_acc_corrupt=0.4698 init_gold_top10=0.4978 init_gold_top100=0.5047
202
+ step=9700 micro_steps=19400 elapsed=32.7s lr=3.000000e-04 loss=2.3457 loss_recon=2.3457 loss_meanflow=0.0000 mean_model_t=0.5034 mean_corrupt_t=0.5034 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7977 acc_corrupt=0.6516 corrupt_frac=0.5718 loss_all=1.3128 loss_corrupt=2.2397 acc_corrupt_t_0p0_0p2=0.2351 corrupt_frac_t_0p0_0p2=0.1652 acc_corrupt_t_0p2_0p4=0.4175 corrupt_frac_t_0p2_0p4=0.1657 acc_corrupt_t_0p4_0p6=0.6818 corrupt_frac_t_0p4_0p6=0.2355 acc_corrupt_t_0p6_0p8=0.8325 corrupt_frac_t_0p6_0p8=0.2472 acc_corrupt_t_0p8_1p0=0.9507 corrupt_frac_t_0p8_1p0=0.1864 wrong_frac=0.4763 init_acc_corrupt=0.4904 init_gold_top10=0.5194 init_gold_top100=0.5258
203
+ step=9800 micro_steps=19600 elapsed=32.4s lr=3.000000e-04 loss=2.3153 loss_recon=2.3153 loss_meanflow=0.0000 mean_model_t=0.5041 mean_corrupt_t=0.5041 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7815 acc_corrupt=0.6160 corrupt_frac=0.5544 loss_all=1.3938 loss_corrupt=2.4232 acc_corrupt_t_0p0_0p2=0.2207 corrupt_frac_t_0p0_0p2=0.1975 acc_corrupt_t_0p2_0p4=0.4402 corrupt_frac_t_0p2_0p4=0.1546 acc_corrupt_t_0p4_0p6=0.6526 corrupt_frac_t_0p4_0p6=0.2871 acc_corrupt_t_0p6_0p8=0.8371 corrupt_frac_t_0p6_0p8=0.2501 acc_corrupt_t_0p8_1p0=0.9722 corrupt_frac_t_0p8_1p0=0.1107 wrong_frac=0.5112 init_acc_corrupt=0.4615 init_gold_top10=0.4819 init_gold_top100=0.4894
204
+ step=9900 micro_steps=19800 elapsed=32.5s lr=3.000000e-04 loss=2.3570 loss_recon=2.3570 loss_meanflow=0.0000 mean_model_t=0.4989 mean_corrupt_t=0.4989 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7633 acc_corrupt=0.6095 corrupt_frac=0.5927 loss_all=1.5402 loss_corrupt=2.5243 acc_corrupt_t_0p0_0p2=0.2407 corrupt_frac_t_0p0_0p2=0.1823 acc_corrupt_t_0p2_0p4=0.4419 corrupt_frac_t_0p2_0p4=0.2358 acc_corrupt_t_0p4_0p6=0.6891 corrupt_frac_t_0p4_0p6=0.2445 acc_corrupt_t_0p6_0p8=0.8089 corrupt_frac_t_0p6_0p8=0.1907 acc_corrupt_t_0p8_1p0=0.9452 corrupt_frac_t_0p8_1p0=0.1467 wrong_frac=0.5267 init_acc_corrupt=0.4377 init_gold_top10=0.4674 init_gold_top100=0.4739
205
+ step=10000 micro_steps=20000 elapsed=33.3s lr=3.000000e-04 loss=2.3284 loss_recon=2.3284 loss_meanflow=0.0000 mean_model_t=0.5016 mean_corrupt_t=0.5016 mean_loss_t_weight=1.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 acc_all=0.7976 acc_corrupt=0.6730 corrupt_frac=0.6011 loss_all=1.3364 loss_corrupt=2.1521 acc_corrupt_t_0p0_0p2=0.2234 corrupt_frac_t_0p0_0p2=0.1755 acc_corrupt_t_0p2_0p4=0.5024 corrupt_frac_t_0p2_0p4=0.1290 acc_corrupt_t_0p4_0p6=0.6469 corrupt_frac_t_0p4_0p6=0.1783 acc_corrupt_t_0p6_0p8=0.8134 corrupt_frac_t_0p6_0p8=0.2797 acc_corrupt_t_0p8_1p0=0.9521 corrupt_frac_t_0p8_1p0=0.2376 wrong_frac=0.4559 init_acc_corrupt=0.5148 init_gold_top10=0.5370 init_gold_top100=0.5441
LTA_openwebtext_dualt/logs/lta_owt_bert_absrope_adaln_dirichlet_len1024_Cv_to_2v_mask1_sameT_gbs512_b4x4_1m_save1k_watch_20260525.log ADDED
The diff for this file is too large to render. See raw diff
 
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ The MIT License (MIT)
2
+
3
+ Copyright (c) 2025 Sebastián Ramírez
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in
13
+ all copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
21
+ THE SOFTWARE.
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/pygments/modeline.py ADDED
@@ -0,0 +1,43 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ pygments.modeline
3
+ ~~~~~~~~~~~~~~~~~
4
+
5
+ A simple modeline parser (based on pymodeline).
6
+
7
+ :copyright: Copyright 2006-present by the Pygments team, see AUTHORS.
8
+ :license: BSD, see LICENSE for details.
9
+ """
10
+
11
+ import re
12
+
13
+ __all__ = ['get_filetype_from_buffer']
14
+
15
+
16
+ modeline_re = re.compile(r'''
17
+ (?: vi | vim | ex ) (?: [<=>]? \d* )? :
18
+ .* (?: ft | filetype | syn | syntax ) = ( [^:\s]+ )
19
+ ''', re.VERBOSE)
20
+
21
+
22
+ def get_filetype_from_line(l): # noqa: E741
23
+ m = modeline_re.search(l)
24
+ if m:
25
+ return m.group(1)
26
+
27
+
28
+ def get_filetype_from_buffer(buf, max_lines=5):
29
+ """
30
+ Scan the buffer for modelines and return filetype if one is found.
31
+ """
32
+ lines = buf.splitlines()
33
+ for line in lines[-1:-max_lines-1:-1]:
34
+ ret = get_filetype_from_line(line)
35
+ if ret:
36
+ return ret
37
+ for i in range(max_lines, -1, -1):
38
+ if i < len(lines):
39
+ ret = get_filetype_from_line(lines[i])
40
+ if ret:
41
+ return ret
42
+
43
+ return None
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/nt.py ADDED
@@ -0,0 +1,163 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import contextlib
2
+ import ctypes
3
+ import os
4
+
5
+ from ctypes.wintypes import (
6
+ BOOL,
7
+ CHAR,
8
+ DWORD,
9
+ HANDLE,
10
+ LONG,
11
+ LPWSTR,
12
+ MAX_PATH,
13
+ PDWORD,
14
+ ULONG,
15
+ )
16
+
17
+ from shellingham._core import SHELL_NAMES
18
+
19
+
20
+ INVALID_HANDLE_VALUE = HANDLE(-1).value
21
+ ERROR_NO_MORE_FILES = 18
22
+ ERROR_INSUFFICIENT_BUFFER = 122
23
+ TH32CS_SNAPPROCESS = 2
24
+ PROCESS_QUERY_LIMITED_INFORMATION = 0x1000
25
+
26
+
27
+ kernel32 = ctypes.windll.kernel32
28
+
29
+
30
+ def _check_handle(error_val=0):
31
+ def check(ret, func, args):
32
+ if ret == error_val:
33
+ raise ctypes.WinError()
34
+ return ret
35
+
36
+ return check
37
+
38
+
39
+ def _check_expected(expected):
40
+ def check(ret, func, args):
41
+ if ret:
42
+ return True
43
+ code = ctypes.GetLastError()
44
+ if code == expected:
45
+ return False
46
+ raise ctypes.WinError(code)
47
+
48
+ return check
49
+
50
+
51
+ class ProcessEntry32(ctypes.Structure):
52
+ _fields_ = (
53
+ ("dwSize", DWORD),
54
+ ("cntUsage", DWORD),
55
+ ("th32ProcessID", DWORD),
56
+ ("th32DefaultHeapID", ctypes.POINTER(ULONG)),
57
+ ("th32ModuleID", DWORD),
58
+ ("cntThreads", DWORD),
59
+ ("th32ParentProcessID", DWORD),
60
+ ("pcPriClassBase", LONG),
61
+ ("dwFlags", DWORD),
62
+ ("szExeFile", CHAR * MAX_PATH),
63
+ )
64
+
65
+
66
+ kernel32.CloseHandle.argtypes = [HANDLE]
67
+ kernel32.CloseHandle.restype = BOOL
68
+
69
+ kernel32.CreateToolhelp32Snapshot.argtypes = [DWORD, DWORD]
70
+ kernel32.CreateToolhelp32Snapshot.restype = HANDLE
71
+ kernel32.CreateToolhelp32Snapshot.errcheck = _check_handle( # type: ignore
72
+ INVALID_HANDLE_VALUE,
73
+ )
74
+
75
+ kernel32.Process32First.argtypes = [HANDLE, ctypes.POINTER(ProcessEntry32)]
76
+ kernel32.Process32First.restype = BOOL
77
+ kernel32.Process32First.errcheck = _check_expected( # type: ignore
78
+ ERROR_NO_MORE_FILES,
79
+ )
80
+
81
+ kernel32.Process32Next.argtypes = [HANDLE, ctypes.POINTER(ProcessEntry32)]
82
+ kernel32.Process32Next.restype = BOOL
83
+ kernel32.Process32Next.errcheck = _check_expected( # type: ignore
84
+ ERROR_NO_MORE_FILES,
85
+ )
86
+
87
+ kernel32.GetCurrentProcessId.argtypes = []
88
+ kernel32.GetCurrentProcessId.restype = DWORD
89
+
90
+ kernel32.OpenProcess.argtypes = [DWORD, BOOL, DWORD]
91
+ kernel32.OpenProcess.restype = HANDLE
92
+ kernel32.OpenProcess.errcheck = _check_handle( # type: ignore
93
+ INVALID_HANDLE_VALUE,
94
+ )
95
+
96
+ kernel32.QueryFullProcessImageNameW.argtypes = [HANDLE, DWORD, LPWSTR, PDWORD]
97
+ kernel32.QueryFullProcessImageNameW.restype = BOOL
98
+ kernel32.QueryFullProcessImageNameW.errcheck = _check_expected( # type: ignore
99
+ ERROR_INSUFFICIENT_BUFFER,
100
+ )
101
+
102
+
103
+ @contextlib.contextmanager
104
+ def _handle(f, *args, **kwargs):
105
+ handle = f(*args, **kwargs)
106
+ try:
107
+ yield handle
108
+ finally:
109
+ kernel32.CloseHandle(handle)
110
+
111
+
112
+ def _iter_processes():
113
+ f = kernel32.CreateToolhelp32Snapshot
114
+ with _handle(f, TH32CS_SNAPPROCESS, 0) as snap:
115
+ entry = ProcessEntry32()
116
+ entry.dwSize = ctypes.sizeof(entry)
117
+ ret = kernel32.Process32First(snap, entry)
118
+ while ret:
119
+ yield entry
120
+ ret = kernel32.Process32Next(snap, entry)
121
+
122
+
123
+ def _get_full_path(proch):
124
+ size = DWORD(MAX_PATH)
125
+ while True:
126
+ path_buff = ctypes.create_unicode_buffer("", size.value)
127
+ if kernel32.QueryFullProcessImageNameW(proch, 0, path_buff, size):
128
+ return path_buff.value
129
+ size.value *= 2
130
+
131
+
132
+ def get_shell(pid=None, max_depth=10):
133
+ proc_map = {
134
+ proc.th32ProcessID: (proc.th32ParentProcessID, proc.szExeFile)
135
+ for proc in _iter_processes()
136
+ }
137
+ pid = pid or os.getpid()
138
+
139
+ for _ in range(0, max_depth + 1):
140
+ try:
141
+ ppid, executable = proc_map[pid]
142
+ except KeyError: # No such process? Give up.
143
+ break
144
+
145
+ # The executable name would be encoded with the current code page if
146
+ # we're in ANSI mode (usually). Try to decode it into str/unicode,
147
+ # replacing invalid characters to be safe (not thoeratically necessary,
148
+ # I think). Note that we need to use 'mbcs' instead of encoding
149
+ # settings from sys because this is from the Windows API, not Python
150
+ # internals (which those settings reflect). (pypa/pipenv#3382)
151
+ if isinstance(executable, bytes):
152
+ executable = executable.decode("mbcs", "replace")
153
+
154
+ name = executable.rpartition(".")[0].lower()
155
+ if name not in SHELL_NAMES:
156
+ pid = ppid
157
+ continue
158
+
159
+ key = PROCESS_QUERY_LIMITED_INFORMATION
160
+ with _handle(kernel32.OpenProcess, key, 0, pid) as proch:
161
+ return (name, _get_full_path(proch))
162
+
163
+ return None
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/_core.py ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ import collections
2
+
3
+ Process = collections.namedtuple("Process", "args pid ppid")
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/proc.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import io
2
+ import os
3
+ import re
4
+ import sys
5
+
6
+ from ._core import Process
7
+
8
+ # FreeBSD: https://www.freebsd.org/cgi/man.cgi?query=procfs
9
+ # NetBSD: https://man.netbsd.org/NetBSD-9.3-STABLE/mount_procfs.8
10
+ # DragonFlyBSD: https://www.dragonflybsd.org/cgi/web-man?command=procfs
11
+ BSD_STAT_PPID = 2
12
+
13
+ # See https://docs.kernel.org/filesystems/proc.html
14
+ LINUX_STAT_PPID = 3
15
+
16
+ STAT_PATTERN = re.compile(r"\(.+\)|\S+")
17
+
18
+
19
+ def detect_proc():
20
+ """Detect /proc filesystem style.
21
+
22
+ This checks the /proc/{pid} directory for possible formats. Returns one of
23
+ the following as str:
24
+
25
+ * `stat`: Linux-style, i.e. ``/proc/{pid}/stat``.
26
+ * `status`: BSD-style, i.e. ``/proc/{pid}/status``.
27
+ """
28
+ pid = os.getpid()
29
+ for name in ("stat", "status"):
30
+ if os.path.exists(os.path.join("/proc", str(pid), name)):
31
+ return name
32
+ raise ProcFormatError("unsupported proc format")
33
+
34
+
35
+ def _use_bsd_stat_format():
36
+ try:
37
+ return os.uname().sysname.lower() in ("freebsd", "netbsd", "dragonfly")
38
+ except Exception:
39
+ return False
40
+
41
+
42
+ def _get_ppid(pid, name):
43
+ path = os.path.join("/proc", str(pid), name)
44
+ with io.open(path, encoding="ascii", errors="replace") as f:
45
+ parts = STAT_PATTERN.findall(f.read())
46
+ # We only care about TTY and PPID -- both are numbers.
47
+ if _use_bsd_stat_format():
48
+ return parts[BSD_STAT_PPID]
49
+ return parts[LINUX_STAT_PPID]
50
+
51
+
52
+ def _get_cmdline(pid):
53
+ path = os.path.join("/proc", str(pid), "cmdline")
54
+ encoding = sys.getfilesystemencoding() or "utf-8"
55
+ with io.open(path, encoding=encoding, errors="replace") as f:
56
+ # XXX: Command line arguments can be arbitrary byte sequences, not
57
+ # necessarily decodable. For Shellingham's purpose, however, we don't
58
+ # care. (pypa/pipenv#2820)
59
+ # cmdline appends an extra NULL at the end, hence the [:-1].
60
+ return tuple(f.read().split("\0")[:-1])
61
+
62
+
63
+ class ProcFormatError(EnvironmentError):
64
+ pass
65
+
66
+
67
+ def iter_process_parents(pid, max_depth=10):
68
+ """Try to look up the process tree via the /proc interface."""
69
+ stat_name = detect_proc()
70
+
71
+ # Inner generator function so we correctly throw an error eagerly if proc
72
+ # is not supported, rather than on the first call to the iterator. This
73
+ # allows the call site detects the correct implementation.
74
+ def _iter_process_parents(pid, max_depth):
75
+ for _ in range(max_depth):
76
+ ppid = _get_ppid(pid, stat_name)
77
+ args = _get_cmdline(pid)
78
+ yield Process(args=args, pid=pid, ppid=ppid)
79
+ if ppid == "0":
80
+ break
81
+ pid = ppid
82
+
83
+ return _iter_process_parents(pid, max_depth)
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/image_processing_aria.py ADDED
@@ -0,0 +1,226 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
2
+ # This file was automatically generated from src/transformers/models/aria/modular_aria.py.
3
+ # Do NOT edit this file manually as any edits will be overwritten by the generation of
4
+ # the file from the modular. If any change should be done, please apply the change to the
5
+ # modular_aria.py file directly. One of our CI enforces this.
6
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
7
+ # Copyright 2024 The Rhymes-AI Teams Authors and The HuggingFace Inc. team. All rights reserved.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+ import torch
21
+ from torchvision.transforms.v2 import functional as tvF
22
+
23
+ from ...image_processing_backends import TorchvisionBackend
24
+ from ...image_processing_utils import BatchFeature, get_patch_output_size, select_best_resolution
25
+ from ...image_transforms import divide_to_patches
26
+ from ...image_utils import ChannelDimension, PILImageResampling, SizeDict, get_image_size
27
+ from ...processing_utils import ImagesKwargs, Unpack
28
+ from ...utils import TensorType, auto_docstring
29
+
30
+
31
+ class AriaImageProcessorKwargs(ImagesKwargs, total=False):
32
+ r"""
33
+ max_image_size (`int`, *optional*, defaults to `self.max_image_size`):
34
+ Maximum image size. Must be either 490 or 980.
35
+ min_image_size (`int`, *optional*, defaults to `self.min_image_size`):
36
+ Minimum image size. Images smaller than this in any dimension will be scaled up.
37
+ split_resolutions (`list[list[int]]`, *optional*, defaults to `self.split_resolutions`):
38
+ A list of possible resolutions as (height, width) pairs for splitting high-resolution images into patches.
39
+ split_image (`bool`, *optional*, defaults to `self.split_image`):
40
+ Whether to split the image into patches using the best matching resolution from `split_resolutions`.
41
+ """
42
+
43
+ max_image_size: int
44
+ min_image_size: int
45
+ split_resolutions: list[list[int]]
46
+ split_image: bool
47
+
48
+
49
+ @auto_docstring
50
+ class AriaImageProcessor(TorchvisionBackend):
51
+ model_input_names = ["pixel_values", "pixel_mask", "num_crops"]
52
+ valid_kwargs = AriaImageProcessorKwargs
53
+
54
+ resample = PILImageResampling.BICUBIC
55
+ image_mean = [0.5, 0.5, 0.5]
56
+ image_std = [0.5, 0.5, 0.5]
57
+ max_image_size = 980
58
+ min_image_size = 336
59
+ split_image = False
60
+ split_resolutions = None
61
+ do_convert_rgb = True
62
+ do_rescale = True
63
+ do_normalize = True
64
+
65
+ def __init__(self, **kwargs: Unpack[AriaImageProcessorKwargs]):
66
+ if kwargs.get("split_resolutions") is None:
67
+ default_resolutions = [(1, 2), (1, 3), (1, 4), (1, 5), (1, 6), (1, 7), (1, 8), (2, 4), (2, 3), (2, 2), (2, 1), (3, 1), (3, 2), (4, 1), (4, 2), (5, 1), (6, 1), (7, 1), (8, 1)] # fmt: skip
68
+ kwargs["split_resolutions"] = [[el[0] * 490, el[1] * 490] for el in default_resolutions]
69
+ super().__init__(**kwargs)
70
+
71
+ def _get_padding_size(self, original_resolution: tuple, target_resolution: tuple) -> list[int]:
72
+ """Get padding size for patching, returns [left, top, right, bottom] for tvF.pad."""
73
+ original_height, original_width = original_resolution
74
+ target_height, target_width = target_resolution
75
+ paste_x, r_x = divmod(target_width - original_width, 2)
76
+ paste_y, r_y = divmod(target_height - original_height, 2)
77
+ return [paste_x, paste_y, paste_x + r_x, paste_y + r_y]
78
+
79
+ def _resize_for_patching(
80
+ self,
81
+ image: "torch.Tensor",
82
+ target_resolution: tuple,
83
+ resample: "PILImageResampling | tvF.InterpolationMode | int | None",
84
+ ) -> "torch.Tensor":
85
+ """Resize an image to a target resolution while maintaining aspect ratio."""
86
+ new_height, new_width = get_patch_output_size(
87
+ image, target_resolution, input_data_format=ChannelDimension.FIRST
88
+ )
89
+ return self.resize(image, SizeDict(height=new_height, width=new_width), resample)
90
+
91
+ def _pad_for_patching(
92
+ self,
93
+ image: "torch.Tensor",
94
+ target_resolution: tuple,
95
+ ) -> "torch.Tensor":
96
+ """Pad an image to a target resolution while maintaining aspect ratio."""
97
+ new_resolution = get_patch_output_size(image, target_resolution, input_data_format=ChannelDimension.FIRST)
98
+ padding = self._get_padding_size(new_resolution, target_resolution)
99
+ return tvF.pad(image, padding=padding)
100
+
101
+ def get_image_patches(
102
+ self,
103
+ image: "torch.Tensor",
104
+ grid_pinpoints: list[list[int]],
105
+ patch_size: int,
106
+ resample: "PILImageResampling | tvF.InterpolationMode | int | None",
107
+ ) -> list["torch.Tensor"]:
108
+ """
109
+ Process an image with variable resolutions by dividing it into patches.
110
+
111
+ Args:
112
+ image (`torch.Tensor`):
113
+ The input image to be processed (channels-first format).
114
+ grid_pinpoints (`list[list[int]]`):
115
+ A list of possible resolutions as (height, width) pairs.
116
+ patch_size (`int`):
117
+ Size of each square patch to divide the image into.
118
+ resample (`PILImageResampling | tvF.InterpolationMode | int | None`):
119
+ Resampling filter to use when resizing.
120
+
121
+ Returns:
122
+ `list[torch.Tensor]`: A list of image patches in channels-first format.
123
+ """
124
+ if not isinstance(grid_pinpoints, list):
125
+ raise TypeError("grid_pinpoints must be a list of possible resolutions.")
126
+
127
+ image_size = get_image_size(image, channel_dim=ChannelDimension.FIRST)
128
+ best_resolution = select_best_resolution(image_size, grid_pinpoints)
129
+ resized_image = self._resize_for_patching(image, best_resolution, resample)
130
+ padded_image = self._pad_for_patching(resized_image, best_resolution)
131
+ patches = divide_to_patches(padded_image, patch_size=patch_size)
132
+ return patches
133
+
134
+ def _preprocess(
135
+ self,
136
+ images: list["torch.Tensor"],
137
+ do_rescale: bool,
138
+ rescale_factor: float,
139
+ do_normalize: bool,
140
+ image_mean: float | list[float] | None,
141
+ image_std: float | list[float] | None,
142
+ disable_grouping: bool | None,
143
+ return_tensors: str | TensorType | None,
144
+ max_image_size: int = 980,
145
+ min_image_size: int = 336,
146
+ split_resolutions: list[list[int]] | None = None,
147
+ split_image: bool = False,
148
+ resample: "PILImageResampling | tvF.InterpolationMode | int | None" = None,
149
+ **kwargs,
150
+ ) -> BatchFeature:
151
+ if max_image_size not in [490, 980]:
152
+ raise ValueError("max_image_size must be either 490 or 980")
153
+
154
+ pixel_masks = []
155
+ processed_crops = []
156
+ num_crops = None
157
+
158
+ for image in images:
159
+ if split_image:
160
+ crop_images = self.get_image_patches(image, split_resolutions, max_image_size, resample)
161
+ else:
162
+ crop_images = [image]
163
+
164
+ if num_crops is None or len(crop_images) > num_crops:
165
+ num_crops = len(crop_images)
166
+
167
+ for crop_image in crop_images:
168
+ h, w = crop_image.shape[-2], crop_image.shape[-1]
169
+ scale = max_image_size / max(h, w)
170
+ if w >= h:
171
+ new_h = max(int(h * scale), min_image_size)
172
+ new_w = max_image_size
173
+ else:
174
+ new_h = max_image_size
175
+ new_w = max(int(w * scale), min_image_size)
176
+
177
+ crop_image = self.resize(crop_image, SizeDict(height=new_h, width=new_w), resample)
178
+
179
+ padding_bottom = max_image_size - new_h
180
+ padding_right = max_image_size - new_w
181
+ crop_image = tvF.pad(crop_image, [0, 0, padding_right, padding_bottom])
182
+
183
+ pixel_mask = torch.zeros((max_image_size, max_image_size), dtype=torch.bool)
184
+ pixel_mask[:new_h, :new_w] = True
185
+ pixel_masks.append(pixel_mask)
186
+ processed_crops.append(crop_image)
187
+
188
+ stacked_images = torch.stack(processed_crops, dim=0)
189
+ stacked_images = self.rescale_and_normalize(
190
+ stacked_images, do_rescale, rescale_factor, do_normalize, image_mean, image_std
191
+ )
192
+ stacked_masks = torch.stack(pixel_masks, dim=0)
193
+
194
+ return BatchFeature(
195
+ data={
196
+ "pixel_values": stacked_images,
197
+ "pixel_mask": stacked_masks,
198
+ "num_crops": num_crops,
199
+ },
200
+ tensor_type=return_tensors,
201
+ )
202
+
203
+ def get_number_of_image_patches(self, height: int, width: int, images_kwargs=None):
204
+ """
205
+ A utility that returns number of image patches for a given image size.
206
+
207
+ Args:
208
+ height (`int`):
209
+ Height of the input image.
210
+ width (`int`):
211
+ Width of the input image.
212
+ images_kwargs (`dict`, *optional*):
213
+ Any kwargs to override defaults of the image processor.
214
+
215
+ Returns:
216
+ `int`: Number of patches per image.
217
+ """
218
+ split_image = images_kwargs.get("split_image", self.split_image)
219
+ max_image_size = images_kwargs.get("max_image_size", self.max_image_size)
220
+
221
+ resized_height, resized_width = select_best_resolution((height, width), self.split_resolutions)
222
+ num_patches = 1 if not split_image else resized_height // max_image_size * resized_width // max_image_size
223
+ return num_patches
224
+
225
+
226
+ __all__ = ["AriaImageProcessor"]
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/processing_aria.py ADDED
@@ -0,0 +1,177 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
2
+ # This file was automatically generated from src/transformers/models/aria/modular_aria.py.
3
+ # Do NOT edit this file manually as any edits will be overwritten by the generation of
4
+ # the file from the modular. If any change should be done, please apply the change to the
5
+ # modular_aria.py file directly. One of our CI enforces this.
6
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
7
+ # Copyright 2024 The Rhymes-AI Teams Authors and The HuggingFace Inc. team. All rights reserved.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+ from ...image_processing_utils import BatchFeature
21
+ from ...image_utils import ImageInput
22
+ from ...processing_utils import ImagesKwargs, MultiModalData, ProcessingKwargs, ProcessorMixin, Unpack
23
+ from ...tokenization_python import PreTokenizedInput, TextInput
24
+ from ...utils import TensorType, auto_docstring
25
+ from ..auto import AutoTokenizer
26
+
27
+
28
+ class AriaImagesKwargs(ImagesKwargs, total=False):
29
+ """
30
+ split_image (`bool`, *optional*, defaults to `False`):
31
+ Whether to split large images into multiple crops. When enabled, images exceeding the maximum size are
32
+ divided into overlapping crops that are processed separately and then combined. This allows processing
33
+ of very high-resolution images that exceed the model's input size limits.
34
+ max_image_size (`int`, *optional*, defaults to `980`):
35
+ Maximum image size (in pixels) for a single image crop. Images larger than this will be split into
36
+ multiple crops when `split_image=True`, or resized if splitting is disabled. This parameter controls
37
+ the maximum resolution of individual image patches processed by the model.
38
+ min_image_size (`int`, *optional*):
39
+ Minimum image size (in pixels) for a single image crop. Images smaller than this will be upscaled to
40
+ meet the minimum requirement. If not specified, images are processed at their original size (subject
41
+ to the maximum size constraint).
42
+ """
43
+
44
+ split_image: bool
45
+ max_image_size: int
46
+ min_image_size: int
47
+
48
+
49
+ class AriaProcessorKwargs(ProcessingKwargs, total=False):
50
+ images_kwargs: AriaImagesKwargs
51
+
52
+ _defaults = {
53
+ "text_kwargs": {
54
+ "padding": False,
55
+ "return_mm_token_type_ids": False,
56
+ },
57
+ "images_kwargs": {
58
+ "max_image_size": 980,
59
+ "split_image": False,
60
+ },
61
+ "return_tensors": TensorType.PYTORCH,
62
+ }
63
+
64
+
65
+ @auto_docstring
66
+ class AriaProcessor(ProcessorMixin):
67
+ def __init__(
68
+ self,
69
+ image_processor=None,
70
+ tokenizer: AutoTokenizer | str = None,
71
+ chat_template: str | None = None,
72
+ size_conversion: dict[float | int, int] | None = None,
73
+ ):
74
+ r"""
75
+ size_conversion (`Dict`, *optional*):
76
+ A dictionary indicating size conversions for images.
77
+ """
78
+ if size_conversion is None:
79
+ size_conversion = {490: 128, 980: 256}
80
+ self.size_conversion = {int(k): v for k, v in size_conversion.items()}
81
+
82
+ self.image_token = tokenizer.image_token
83
+ self.image_token_id = tokenizer.image_token_id
84
+ if tokenizer is not None and tokenizer.pad_token is None:
85
+ tokenizer.pad_token = tokenizer.unk_token
86
+
87
+ super().__init__(image_processor, tokenizer, chat_template=chat_template)
88
+
89
+ @auto_docstring
90
+ def __call__(
91
+ self,
92
+ text: TextInput | PreTokenizedInput | list[TextInput] | list[PreTokenizedInput],
93
+ images: ImageInput | None = None,
94
+ **kwargs: Unpack[AriaProcessorKwargs],
95
+ ) -> BatchFeature:
96
+ r"""
97
+ Returns:
98
+ [`BatchFeature`]: A [`BatchFeature`] with the following fields:
99
+ - **input_ids** -- List of token ids to be fed to a model. Returned when `text` is not `None`.
100
+ - **attention_mask** -- List of indices specifying which tokens should be attended to by the model (when
101
+ `return_attention_mask=True` or if *"attention_mask"* is in `self.model_input_names` and if `text` is not
102
+ `None`).
103
+ - **pixel_values** -- Pixel values to be fed to a model. Returned when `images` is not `None`.
104
+ - **pixel_mask** -- Pixel mask to be fed to a model. Returned when `images` is not `None`.
105
+ """
106
+ output_kwargs = self._merge_kwargs(
107
+ AriaProcessorKwargs,
108
+ tokenizer_init_kwargs=self.tokenizer.init_kwargs,
109
+ **kwargs,
110
+ )
111
+
112
+ if isinstance(text, str):
113
+ text = [text]
114
+ elif not isinstance(text, list) and not isinstance(text[0], str):
115
+ raise TypeError("Invalid input text. Please provide a string, or a list of strings")
116
+
117
+ if images is not None:
118
+ image_inputs = self.image_processor(images, **output_kwargs["images_kwargs"])
119
+ # expand the image_token according to the num_crops and tokens per image
120
+ tokens_per_image = self.size_conversion[image_inputs.pixel_values.shape[2]]
121
+ prompt_strings = []
122
+ num_crops = image_inputs.pop("num_crops") * tokens_per_image
123
+ for sample in text:
124
+ sample = sample.replace(self.tokenizer.image_token, self.tokenizer.image_token * num_crops)
125
+ prompt_strings.append(sample)
126
+
127
+ else:
128
+ image_inputs = {}
129
+ prompt_strings = text
130
+
131
+ return_tensors = output_kwargs["text_kwargs"].pop("return_tensors", None)
132
+ return_mm_token_type_ids = output_kwargs["text_kwargs"].pop("return_mm_token_type_ids", False)
133
+ text_inputs = self.tokenizer(prompt_strings, **output_kwargs["text_kwargs"], return_tensors=None)
134
+ self._check_special_mm_tokens(prompt_strings, text_inputs, modalities=["image"])
135
+
136
+ if return_mm_token_type_ids:
137
+ text_inputs["mm_token_type_ids"] = self.create_mm_token_type_ids(text_inputs["input_ids"])
138
+ return BatchFeature(data={**text_inputs, **image_inputs}, tensor_type=return_tensors)
139
+
140
+ def _get_num_multimodal_tokens(self, image_sizes=None, **kwargs):
141
+ """
142
+ Computes the number of placeholder tokens needed for multimodal inputs with the given sizes.
143
+ Args:
144
+ image_sizes (`list[list[int]]`, *optional*):
145
+ The input sizes formatted as (height, width) per each image.
146
+ Returns:
147
+ `MultiModalData`: A `MultiModalData` object holding number of tokens per each of the provided
148
+ input modalities, along with other useful data.
149
+ """
150
+
151
+ vision_data = {}
152
+ if image_sizes is not None:
153
+ images_kwargs = AriaProcessorKwargs._defaults.get("images_kwargs", {})
154
+ images_kwargs.update(kwargs)
155
+
156
+ max_size = images_kwargs.get("max_image_size", None) or self.image_processor.max_image_size
157
+ num_image_patches = [
158
+ self.image_processor.get_number_of_image_patches(*image_size, images_kwargs)
159
+ for image_size in image_sizes
160
+ ]
161
+ num_image_tokens = [self.size_conversion[max_size] * num_patches for num_patches in num_image_patches]
162
+ vision_data.update({"num_image_tokens": num_image_tokens, "num_image_patches": num_image_patches})
163
+
164
+ return MultiModalData(**vision_data)
165
+
166
+ @property
167
+ def model_input_names(self):
168
+ tokenizer_input_names = self.tokenizer.model_input_names
169
+ image_processor_input_names = self.image_processor.model_input_names
170
+
171
+ # Remove `num_crops`, it is popped and used only when processing. Make a copy of list when removing
172
+ # otherwise `self.image_processor.model_input_names` is also modified
173
+ image_processor_input_names = [name for name in image_processor_input_names if name != "num_crops"]
174
+ return list(dict.fromkeys(tokenizer_input_names + image_processor_input_names))
175
+
176
+
177
+ __all__ = ["AriaProcessor"]
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/__init__.py ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 the HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ from typing import TYPE_CHECKING
16
+
17
+ from ...utils import _LazyModule
18
+ from ...utils.import_utils import define_import_structure
19
+
20
+
21
+ if TYPE_CHECKING:
22
+ from .configuration_eomt_dinov3 import *
23
+ from .modeling_eomt_dinov3 import *
24
+ else:
25
+ import sys
26
+
27
+ _file = globals()["__file__"]
28
+ sys.modules[__name__] = _LazyModule(__name__, _file, define_import_structure(_file), module_spec=__spec__)
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/configuration_eomt_dinov3.py ADDED
@@ -0,0 +1,107 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
2
+ # This file was automatically generated from src/transformers/models/eomt_dinov3/modular_eomt_dinov3.py.
3
+ # Do NOT edit this file manually as any edits will be overwritten by the generation of
4
+ # the file from the modular. If any change should be done, please apply the change to the
5
+ # modular_eomt_dinov3.py file directly. One of our CI enforces this.
6
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
7
+ # Copyright 2026 the HuggingFace Team. All rights reserved.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+ from huggingface_hub.dataclasses import strict
21
+
22
+ from ...configuration_utils import PreTrainedConfig
23
+ from ...modeling_rope_utils import RopeParameters
24
+ from ...utils import auto_docstring
25
+
26
+
27
+ @auto_docstring(checkpoint="tue-mps/coco_panoptic_eomt_large_640_dinov3")
28
+ @strict
29
+ class EomtDinov3Config(PreTrainedConfig):
30
+ r"""
31
+ layerscale_value (`float`, *optional*, defaults to 1.0):
32
+ Initial value for the LayerScale parameter.
33
+ num_upscale_blocks (`int`, *optional*, defaults to 2):
34
+ Number of upsampling blocks used in the decoder or segmentation head.
35
+ num_blocks (`int`, *optional*, defaults to 4):
36
+ Number of feature blocks or stages in the architecture.
37
+ no_object_weight (`float`, *optional*, defaults to 0.1):
38
+ Loss weight for the "no object" class in panoptic/instance segmentation.
39
+ train_num_points (`int`, *optional*, defaults to 12544):
40
+ Number of points to sample for mask loss computation during training.
41
+ oversample_ratio (`float`, *optional*, defaults to 3.0):
42
+ Oversampling ratio used in point sampling for mask training.
43
+ importance_sample_ratio (`float`, *optional*, defaults to 0.75):
44
+ Ratio of points to sample based on importance during training.
45
+ num_queries (`int`, *optional*, defaults to 200):
46
+ Number of object queries in the Transformer.
47
+ num_register_tokens (`int`, *optional*, defaults to 4):
48
+ Number of learnable register tokens added to the transformer input.
49
+ query_bias (`bool`, *optional*, defaults to `True`):
50
+ Whether to use bias in query projection.
51
+ key_bias (`bool`, *optional*, defaults to `False`):
52
+ Whether to use bias in key projection.
53
+ value_bias (`bool`, *optional*, defaults to `True`):
54
+ Whether to use bias in value projection.
55
+ proj_bias (`bool`, *optional*, defaults to `True`):
56
+ Whether to use bias in output projection.
57
+ use_gated_mlp (`bool`, *optional*, defaults to `False`):
58
+ Whether to use gated MLP layers.
59
+ pos_embed_shift (`float`, *optional*):
60
+ Shift value for position embeddings.
61
+ pos_embed_jitter (`float`, *optional*):
62
+ Jitter value for position embeddings.
63
+ pos_embed_rescale (`float`, *optional*, defaults to 2.0):
64
+ Rescale value for position embeddings.
65
+ """
66
+
67
+ model_type = "eomt_dinov3"
68
+
69
+ hidden_size: int = 1024
70
+ num_hidden_layers: int = 24
71
+ num_attention_heads: int = 16
72
+ hidden_act: str = "gelu"
73
+ hidden_dropout_prob: float | int = 0.0
74
+ initializer_range: float = 0.02
75
+ layer_norm_eps: float = 1e-6
76
+ image_size: int | list[int] | tuple[int, int] = 640
77
+ patch_size: int | list[int] | tuple[int, int] = 16
78
+ num_channels: int = 3
79
+ layerscale_value: float = 1.0
80
+ drop_path_rate: float | int = 0.0
81
+ num_upscale_blocks: int = 2
82
+ attention_dropout: float | int = 0.0
83
+ num_blocks: int = 4
84
+ no_object_weight: float = 0.1
85
+ class_weight: float = 2.0
86
+ mask_weight: float = 5.0
87
+ dice_weight: float = 5.0
88
+ train_num_points: int = 12544
89
+ oversample_ratio: float = 3.0
90
+ importance_sample_ratio: float = 0.75
91
+ num_queries: int = 200
92
+ num_register_tokens: int = 4
93
+ default_theta = 100.0
94
+ intermediate_size: int = 4096
95
+ rope_parameters: RopeParameters | dict | None = None
96
+ query_bias: bool = True
97
+ key_bias: bool = False
98
+ value_bias: bool = True
99
+ proj_bias: bool = True
100
+ mlp_bias: bool = True
101
+ use_gated_mlp: bool = False
102
+ pos_embed_shift: float | None = None
103
+ pos_embed_jitter: float | None = None
104
+ pos_embed_rescale: float | None = 2.0
105
+
106
+
107
+ __all__ = ["EomtDinov3Config"]
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modeling_eomt_dinov3.py ADDED
@@ -0,0 +1,1374 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
2
+ # This file was automatically generated from src/transformers/models/eomt_dinov3/modular_eomt_dinov3.py.
3
+ # Do NOT edit this file manually as any edits will be overwritten by the generation of
4
+ # the file from the modular. If any change should be done, please apply the change to the
5
+ # modular_eomt_dinov3.py file directly. One of our CI enforces this.
6
+ # 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
7
+ # Copyright 2026 the HuggingFace Team. All rights reserved.
8
+ #
9
+ # Licensed under the Apache License, Version 2.0 (the "License");
10
+ # you may not use this file except in compliance with the License.
11
+ # You may obtain a copy of the License at
12
+ #
13
+ # http://www.apache.org/licenses/LICENSE-2.0
14
+ #
15
+ # Unless required by applicable law or agreed to in writing, software
16
+ # distributed under the License is distributed on an "AS IS" BASIS,
17
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
18
+ # See the License for the specific language governing permissions and
19
+ # limitations under the License.
20
+
21
+ import math
22
+ from collections.abc import Callable
23
+ from dataclasses import dataclass
24
+ from typing import Optional
25
+
26
+ import numpy as np
27
+ import torch
28
+ import torch.nn.functional as F
29
+ from torch import Tensor, nn
30
+
31
+ from ... import initialization as init
32
+ from ...activations import ACT2FN
33
+ from ...file_utils import ModelOutput, is_scipy_available, requires_backends
34
+ from ...modeling_layers import GradientCheckpointingLayer
35
+ from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
36
+ from ...processing_utils import Unpack
37
+ from ...pytorch_utils import compile_compatible_method_lru_cache
38
+ from ...utils import TransformersKwargs, auto_docstring, is_accelerate_available
39
+ from ...utils.generic import maybe_autocast, merge_with_config_defaults
40
+ from ...utils.output_capturing import capture_outputs
41
+ from .configuration_eomt_dinov3 import EomtDinov3Config
42
+
43
+
44
+ if is_scipy_available():
45
+ from scipy.optimize import linear_sum_assignment
46
+
47
+ if is_accelerate_available():
48
+ from accelerate import PartialState
49
+ from accelerate.utils import reduce
50
+
51
+
52
+ def rotate_half(x):
53
+ """Rotates half the hidden dims of the input."""
54
+ x1 = x[..., : x.shape[-1] // 2]
55
+ x2 = x[..., x.shape[-1] // 2 :]
56
+ return torch.cat((-x2, x1), dim=-1)
57
+
58
+
59
+ def eager_attention_forward(
60
+ module: nn.Module,
61
+ query: torch.Tensor,
62
+ key: torch.Tensor,
63
+ value: torch.Tensor,
64
+ attention_mask: torch.Tensor | None,
65
+ scaling: float | None = None,
66
+ dropout: float = 0.0,
67
+ **kwargs: Unpack[TransformersKwargs],
68
+ ):
69
+ if scaling is None:
70
+ scaling = query.size(-1) ** -0.5
71
+
72
+ # Take the dot product between "query" and "key" to get the raw attention scores.
73
+ attn_weights = torch.matmul(query, key.transpose(2, 3)) * scaling
74
+
75
+ if attention_mask is not None:
76
+ attn_weights = attn_weights + attention_mask
77
+
78
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1)
79
+ attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)
80
+
81
+ attn_output = torch.matmul(attn_weights, value)
82
+ attn_output = attn_output.transpose(1, 2).contiguous()
83
+
84
+ return attn_output, attn_weights
85
+
86
+
87
+ def apply_rotary_pos_emb(
88
+ q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, **kwargs
89
+ ) -> tuple[torch.Tensor, torch.Tensor]:
90
+ """Applies Rotary Position Embedding to the query and key tensors, but only to the patch tokens,
91
+ ignoring the prefix tokens (cls token and register tokens).
92
+
93
+ Args:
94
+ q (`torch.Tensor`): The query tensor.
95
+ k (`torch.Tensor`): The key tensor.
96
+ cos (`torch.Tensor`): The cosine part of the rotary embedding.
97
+ sin (`torch.Tensor`): The sine part of the rotary embedding.
98
+
99
+ Returns:
100
+ `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
101
+ """
102
+
103
+ num_tokens = q.shape[-2]
104
+ num_patches = sin.shape[-2]
105
+ num_prefix_tokens = num_tokens - num_patches # cls token + register tokens
106
+
107
+ q_prefix_tokens, q_patches = q.split((num_prefix_tokens, num_patches), dim=-2)
108
+ k_prefix_tokens, k_patches = k.split((num_prefix_tokens, num_patches), dim=-2)
109
+
110
+ # apply rope only to patch tokens
111
+ q_patches = (q_patches * cos) + (rotate_half(q_patches) * sin)
112
+ k_patches = (k_patches * cos) + (rotate_half(k_patches) * sin)
113
+
114
+ q = torch.cat((q_prefix_tokens, q_patches), dim=-2)
115
+ k = torch.cat((k_prefix_tokens, k_patches), dim=-2)
116
+
117
+ return q, k
118
+
119
+
120
+ class EomtDinov3Attention(nn.Module):
121
+ """
122
+ Multi-headed attention compatible with ALL_ATTENTION_FUNCTIONS.
123
+ """
124
+
125
+ def __init__(self, config: EomtDinov3Config):
126
+ super().__init__()
127
+ self.config = config
128
+ self.embed_dim = config.hidden_size
129
+ self.num_heads = config.num_attention_heads
130
+ self.head_dim = self.embed_dim // self.num_heads
131
+ self.is_causal = False
132
+
133
+ self.scaling = self.head_dim**-0.5
134
+ self.is_causal = False
135
+
136
+ self.dropout = config.attention_dropout
137
+ self.k_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.key_bias)
138
+ self.v_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.value_bias)
139
+
140
+ self.q_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.query_bias)
141
+ self.o_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.proj_bias)
142
+
143
+ def forward(
144
+ self,
145
+ hidden_states: torch.Tensor,
146
+ attention_mask: torch.Tensor | None = None,
147
+ position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None,
148
+ **kwargs: Unpack[TransformersKwargs],
149
+ ) -> tuple[torch.Tensor, torch.Tensor | None]:
150
+ """Input shape: Batch x Time x Channel"""
151
+
152
+ batch_size, patches, _ = hidden_states.size()
153
+
154
+ query_states = self.q_proj(hidden_states)
155
+ key_states = self.k_proj(hidden_states)
156
+ value_states = self.v_proj(hidden_states)
157
+
158
+ query_states = query_states.view(batch_size, patches, self.num_heads, self.head_dim).transpose(1, 2)
159
+ key_states = key_states.view(batch_size, patches, self.num_heads, self.head_dim).transpose(1, 2)
160
+ value_states = value_states.view(batch_size, patches, self.num_heads, self.head_dim).transpose(1, 2)
161
+
162
+ cos, sin = position_embeddings
163
+ query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)
164
+
165
+ attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
166
+ self.config._attn_implementation, eager_attention_forward
167
+ )
168
+
169
+ attn_output, attn_weights = attention_interface(
170
+ self,
171
+ query_states,
172
+ key_states,
173
+ value_states,
174
+ attention_mask,
175
+ dropout=0.0 if not self.training else self.dropout,
176
+ scaling=self.scaling,
177
+ **kwargs,
178
+ )
179
+
180
+ attn_output = attn_output.reshape(batch_size, patches, -1).contiguous()
181
+ attn_output = self.o_proj(attn_output)
182
+
183
+ return attn_output, attn_weights
184
+
185
+
186
+ class EomtDinov3Embeddings(nn.Module):
187
+ """
188
+ Construct the CLS token, mask token, position and patch embeddings.
189
+ """
190
+
191
+ def __init__(self, config: EomtDinov3Config):
192
+ super().__init__()
193
+ self.config = config
194
+ self.cls_token = nn.Parameter(torch.randn(1, 1, config.hidden_size))
195
+ self.mask_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size))
196
+ self.register_tokens = nn.Parameter(torch.empty(1, config.num_register_tokens, config.hidden_size))
197
+ self.patch_embeddings = nn.Conv2d(
198
+ config.num_channels, config.hidden_size, kernel_size=config.patch_size, stride=config.patch_size
199
+ )
200
+ self.num_prefix_tokens = 1 + config.num_register_tokens
201
+
202
+ def forward(self, pixel_values: torch.Tensor, bool_masked_pos: torch.Tensor | None = None) -> torch.Tensor:
203
+ batch_size = pixel_values.shape[0]
204
+ target_dtype = self.patch_embeddings.weight.dtype
205
+
206
+ # (batch_size, num_channels, height, width) -> (batch_size, num_patches, hidden_size)
207
+ patch_embeddings = self.patch_embeddings(pixel_values.to(dtype=target_dtype))
208
+ patch_embeddings = patch_embeddings.flatten(2).transpose(1, 2)
209
+
210
+ if bool_masked_pos is not None:
211
+ mask_token = self.mask_token.to(patch_embeddings.dtype)
212
+ patch_embeddings = torch.where(bool_masked_pos.unsqueeze(-1), mask_token, patch_embeddings)
213
+
214
+ # Add CLS and register tokens
215
+ cls_token = self.cls_token.expand(batch_size, -1, -1)
216
+ register_tokens = self.register_tokens.expand(batch_size, -1, -1)
217
+ embeddings = torch.cat([cls_token, register_tokens, patch_embeddings], dim=1)
218
+
219
+ return embeddings
220
+
221
+
222
+ class EomtDinov3MLP(nn.Module):
223
+ def __init__(self, config):
224
+ super().__init__()
225
+ self.config = config
226
+ self.hidden_size = config.hidden_size
227
+ self.intermediate_size = config.intermediate_size
228
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias)
229
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.mlp_bias)
230
+ self.act_fn = ACT2FN[config.hidden_act]
231
+
232
+ def forward(self, x):
233
+ return self.down_proj(self.act_fn(self.up_proj(x)))
234
+
235
+
236
+ class EomtDinov3GatedMLP(nn.Module):
237
+ def __init__(self, config):
238
+ super().__init__()
239
+ self.config = config
240
+ self.hidden_size = config.hidden_size
241
+ self.intermediate_size = config.intermediate_size
242
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias)
243
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias)
244
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.mlp_bias)
245
+ self.act_fn = ACT2FN[config.hidden_act]
246
+
247
+ def forward(self, x):
248
+ down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
249
+ return down_proj
250
+
251
+
252
+ class EomtDinov3DropPath(nn.Module):
253
+ """Stochastic depth (DropPath) per sample, for residual blocks.
254
+
255
+ Identity when ``drop_prob`` is 0 or outside training. See `Deep Networks with Stochastic Depth
256
+ <https://arxiv.org/abs/1603.09382>`_.
257
+ """
258
+
259
+ def __init__(self, drop_prob: float = 0.0) -> None:
260
+ super().__init__()
261
+ self.drop_prob = drop_prob
262
+
263
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
264
+ if self.drop_prob == 0.0 or not self.training:
265
+ return hidden_states
266
+ keep_prob = 1 - self.drop_prob
267
+ shape = (hidden_states.shape[0],) + (1,) * (hidden_states.ndim - 1)
268
+ random_tensor = torch.rand(shape, dtype=hidden_states.dtype, device=hidden_states.device)
269
+ random_tensor = torch.floor(random_tensor + keep_prob)
270
+ return hidden_states.div(keep_prob) * random_tensor
271
+
272
+ def extra_repr(self) -> str:
273
+ return f"p={self.drop_prob}"
274
+
275
+
276
+ class EomtDinov3Layer(GradientCheckpointingLayer):
277
+ """This corresponds to the Block class in the original implementation."""
278
+
279
+ def __init__(self, config: EomtDinov3Config):
280
+ super().__init__()
281
+
282
+ self.norm1 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
283
+ self.attention = EomtDinov3Attention(config)
284
+ self.layer_scale1 = EomtDinov3LayerScale(config)
285
+ self.drop_path = EomtDinov3DropPath(config.drop_path_rate) if config.drop_path_rate > 0.0 else nn.Identity()
286
+
287
+ self.norm2 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
288
+
289
+ if config.use_gated_mlp:
290
+ self.mlp = EomtDinov3GatedMLP(config)
291
+ else:
292
+ self.mlp = EomtDinov3MLP(config)
293
+ self.layer_scale2 = EomtDinov3LayerScale(config)
294
+
295
+ def forward(
296
+ self,
297
+ hidden_states: torch.Tensor,
298
+ attention_mask: torch.Tensor | None = None,
299
+ position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None,
300
+ **kwargs: Unpack[TransformersKwargs],
301
+ ) -> torch.Tensor:
302
+ # Attention with residual connection
303
+ residual = hidden_states
304
+ hidden_states = self.norm1(hidden_states)
305
+ hidden_states, _ = self.attention(
306
+ hidden_states,
307
+ attention_mask=attention_mask,
308
+ position_embeddings=position_embeddings,
309
+ **kwargs,
310
+ )
311
+ hidden_states = self.layer_scale1(hidden_states)
312
+ hidden_states = self.drop_path(hidden_states) + residual
313
+
314
+ # MLP with residual connection
315
+ residual = hidden_states
316
+ hidden_states = self.norm2(hidden_states)
317
+ hidden_states = self.mlp(hidden_states)
318
+ hidden_states = self.layer_scale2(hidden_states)
319
+ hidden_states = self.drop_path(hidden_states) + residual
320
+
321
+ return hidden_states
322
+
323
+
324
+ class EomtDinov3LayerScale(nn.Module):
325
+ def __init__(self, config) -> None:
326
+ super().__init__()
327
+ self.lambda1 = nn.Parameter(config.layerscale_value * torch.ones(config.hidden_size))
328
+
329
+ def forward(self, hidden_state: torch.Tensor) -> torch.Tensor:
330
+ return hidden_state * self.lambda1
331
+
332
+
333
+ @compile_compatible_method_lru_cache(maxsize=32)
334
+ def get_patches_center_coordinates(
335
+ num_patches_h: int, num_patches_w: int, dtype: torch.dtype, device: torch.device
336
+ ) -> torch.Tensor:
337
+ """
338
+ Computes the 2D coordinates of the centers of image patches, normalized to the range [-1, +1].
339
+ The center of each patch is exactly halfway between its top-left and bottom-right corners.
340
+
341
+ Args:
342
+ num_patches_h (int): Number of patches along the vertical (height) axis.
343
+ num_patches_w (int): Number of patches along the horizontal (width) axis.
344
+ dtype (torch.dtype): The desired data type of the returned tensor.
345
+
346
+ Returns:
347
+ torch.Tensor: A tensor of shape (height * width, 2), where each row contains the (y, x)
348
+ coordinates of a patch center, normalized to [-1, +1].
349
+ """
350
+ coords_h = torch.arange(0.5, num_patches_h, dtype=dtype, device=device)
351
+ coords_w = torch.arange(0.5, num_patches_w, dtype=dtype, device=device)
352
+ coords_h = coords_h / num_patches_h
353
+ coords_w = coords_w / num_patches_w
354
+ # (height, width, 2) -> (height * width, 2)
355
+ coords = torch.stack(torch.meshgrid(coords_h, coords_w, indexing="ij"), dim=-1)
356
+ coords = coords.flatten(0, 1)
357
+ # Shift range [0, 1] to [-1, +1]
358
+ coords = 2.0 * coords - 1.0
359
+ return coords
360
+
361
+
362
+ def augment_patches_center_coordinates(
363
+ coords: torch.Tensor,
364
+ shift: float | None = None,
365
+ jitter: float | None = None,
366
+ rescale: float | None = None,
367
+ ) -> torch.Tensor:
368
+ # Shift coords by adding a uniform value in [-shift, shift]
369
+ if shift is not None:
370
+ shift_hw = torch.empty((1, 2), device=coords.device, dtype=coords.dtype)
371
+ shift_hw = shift_hw.uniform_(-shift, shift)
372
+ coords = coords + shift_hw
373
+
374
+ # Jitter coords by multiplying the range [-1, 1] by a log-uniform value in [1/jitter, jitter]
375
+ if jitter is not None:
376
+ jitter_range = np.log(jitter)
377
+ jitter_hw = torch.empty((1, 2), device=coords.device, dtype=coords.dtype)
378
+ jitter_hw = jitter_hw.uniform_(-jitter_range, jitter_range).exp()
379
+ coords = coords * jitter_hw
380
+
381
+ # Rescale coords by multiplying the range [-1, 1] by a log-uniform value in [1/rescale, rescale]
382
+ if rescale is not None:
383
+ rescale_range = np.log(rescale)
384
+ rescale_hw = torch.empty(1, device=coords.device, dtype=coords.dtype)
385
+ rescale_hw = rescale_hw.uniform_(-rescale_range, rescale_range).exp()
386
+ coords = coords * rescale_hw
387
+
388
+ return coords
389
+
390
+
391
+ class EomtDinov3RotaryEmbedding(nn.Module):
392
+ inv_freq: Tensor
393
+
394
+ def __init__(self, config: EomtDinov3Config, device=None):
395
+ super().__init__()
396
+ self.config = config
397
+
398
+ self.rope_type = self.config.rope_parameters["rope_type"]
399
+ rope_init_fn: Callable = self.compute_default_rope_parameters
400
+ if self.rope_type != "default":
401
+ raise ValueError("`EomtDinov3` only supports `default` RoPE! Please check your `rope_type`")
402
+ inv_freq, self.attention_scaling = rope_init_fn(self.config, device)
403
+
404
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
405
+ self.register_buffer("original_inv_freq", inv_freq.clone(), persistent=False)
406
+
407
+ def forward(self, pixel_values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
408
+ _, _, height, width = pixel_values.shape
409
+ num_patches_h = height // self.config.patch_size
410
+ num_patches_w = width // self.config.patch_size
411
+
412
+ device = pixel_values.device
413
+ device_type = device.type if isinstance(device.type, str) and device.type != "mps" else "cpu"
414
+
415
+ with maybe_autocast(device_type=device_type, enabled=False): # Force float32
416
+ # Although we could precompute static patch_coords from image_size and patch_size in the config,
417
+ # the model was trained with random_scale, so it can process images of varying sizes.
418
+ # Therefore, it's better to compute patch_coords dynamically (with lru_cache).
419
+ patch_coords = get_patches_center_coordinates(
420
+ num_patches_h, num_patches_w, dtype=torch.float32, device=device
421
+ )
422
+ if self.training:
423
+ patch_coords = augment_patches_center_coordinates(
424
+ patch_coords,
425
+ shift=self.config.pos_embed_shift,
426
+ jitter=self.config.pos_embed_jitter,
427
+ rescale=self.config.pos_embed_rescale,
428
+ )
429
+
430
+ # (height * width, 2, head_dim / 4) -> (height * width, head_dim / 2) -> (height * width, head_dim)
431
+ angles = 2 * math.pi * patch_coords[:, :, None] * self.inv_freq[None, None, :]
432
+ angles = angles.flatten(1, 2)
433
+ angles = angles.tile(2)
434
+
435
+ cos = torch.cos(angles)
436
+ sin = torch.sin(angles)
437
+
438
+ dtype = pixel_values.dtype
439
+ return cos.to(dtype=dtype), sin.to(dtype=dtype)
440
+
441
+ @staticmethod
442
+ def compute_default_rope_parameters(
443
+ config: EomtDinov3Config | None = None,
444
+ device: Optional["torch.device"] = None,
445
+ seq_len: int | None = None,
446
+ ) -> torch.Tensor:
447
+ """
448
+ Computes the inverse frequencies according to the original RoPE implementation
449
+ Args:
450
+ config ([`~transformers.PreTrainedConfig`]):
451
+ The model configuration.
452
+ device (`torch.device`):
453
+ The device to use for initialization of the inverse frequencies.
454
+ seq_len (`int`, *optional*):
455
+ The current sequence length. Unused for this type of RoPE.
456
+ Returns:
457
+ Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
458
+ post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
459
+ """
460
+ base = config.rope_parameters["rope_theta"]
461
+ head_dim = config.hidden_size // config.num_attention_heads
462
+
463
+ attention_factor = 1.0 # Unused in this type of RoPE
464
+
465
+ # Compute the inverse frequencies
466
+ inv_freq = 1 / base ** torch.arange(0, 1, 4 / head_dim, dtype=torch.float32, device=device)
467
+ return inv_freq, attention_factor
468
+
469
+
470
+ # Adapted from https://github.com/facebookresearch/detectron2/blob/main/projects/PointRend/point_rend/point_features.py
471
+ def sample_point(
472
+ input_features: torch.Tensor, point_coordinates: torch.Tensor, add_dim=False, **kwargs
473
+ ) -> torch.Tensor:
474
+ """
475
+ A wrapper around `torch.nn.functional.grid_sample` to support 3D point_coordinates tensors.
476
+
477
+ Args:
478
+ input_features (`torch.Tensor` of shape (batch_size, channels, height, width)):
479
+ A tensor that contains features map on a height * width grid
480
+ point_coordinates (`torch.Tensor` of shape (batch_size, num_points, 2) or (batch_size, grid_height, grid_width,:
481
+ 2)):
482
+ A tensor that contains [0, 1] * [0, 1] normalized point coordinates
483
+ add_dim (`bool`):
484
+ boolean value to keep track of added dimension
485
+
486
+ Returns:
487
+ point_features (`torch.Tensor` of shape (batch_size, channels, num_points) or (batch_size, channels,
488
+ height_grid, width_grid):
489
+ A tensor that contains features for points in `point_coordinates`.
490
+ """
491
+ if point_coordinates.dim() == 3:
492
+ add_dim = True
493
+ point_coordinates = point_coordinates.unsqueeze(2)
494
+
495
+ # use nn.function.grid_sample to get features for points in `point_coordinates` via bilinear interpolation
496
+ point_features = torch.nn.functional.grid_sample(input_features, 2.0 * point_coordinates - 1.0, **kwargs)
497
+ if add_dim:
498
+ point_features = point_features.squeeze(3)
499
+
500
+ return point_features
501
+
502
+
503
+ def pair_wise_dice_loss(inputs: Tensor, labels: Tensor) -> Tensor:
504
+ """
505
+ A pair wise version of the dice loss, see `dice_loss` for usage.
506
+
507
+ Args:
508
+ inputs (`torch.Tensor`):
509
+ A tensor representing a mask
510
+ labels (`torch.Tensor`):
511
+ A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
512
+ (0 for the negative class and 1 for the positive class).
513
+
514
+ Returns:
515
+ `torch.Tensor`: The computed loss between each pairs.
516
+ """
517
+ inputs = inputs.sigmoid().flatten(1)
518
+ numerator = 2 * torch.matmul(inputs, labels.T)
519
+ # using broadcasting to get a [num_queries, NUM_CLASSES] matrix
520
+ denominator = inputs.sum(-1)[:, None] + labels.sum(-1)[None, :]
521
+ loss = 1 - (numerator + 1) / (denominator + 1)
522
+ return loss
523
+
524
+
525
+ def pair_wise_sigmoid_cross_entropy_loss(inputs: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
526
+ r"""
527
+ A pair wise version of the cross entropy loss, see `sigmoid_cross_entropy_loss` for usage.
528
+
529
+ Args:
530
+ inputs (`torch.Tensor`):
531
+ A tensor representing a mask.
532
+ labels (`torch.Tensor`):
533
+ A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
534
+ (0 for the negative class and 1 for the positive class).
535
+
536
+ Returns:
537
+ loss (`torch.Tensor`): The computed loss between each pairs.
538
+ """
539
+
540
+ height_and_width = inputs.shape[1]
541
+
542
+ criterion = nn.BCEWithLogitsLoss(reduction="none")
543
+ cross_entropy_loss_pos = criterion(inputs, torch.ones_like(inputs))
544
+ cross_entropy_loss_neg = criterion(inputs, torch.zeros_like(inputs))
545
+
546
+ loss_pos = torch.matmul(cross_entropy_loss_pos / height_and_width, labels.T)
547
+ loss_neg = torch.matmul(cross_entropy_loss_neg / height_and_width, (1 - labels).T)
548
+ loss = loss_pos + loss_neg
549
+ return loss
550
+
551
+
552
+ # Adapted from https://github.com/facebookresearch/EomtDinov3/blob/main/eomt_dinov3/modeling/matcher.py
553
+ class EomtDinov3HungarianMatcher(nn.Module):
554
+ """This class computes an assignment between the labels and the predictions of the network.
555
+
556
+ For efficiency reasons, the labels don't include the no_object. Because of this, in general, there are more
557
+ predictions than labels. In this case, we do a 1-to-1 matching of the best predictions, while the others are
558
+ un-matched (and thus treated as non-objects).
559
+ """
560
+
561
+ def __init__(
562
+ self, cost_class: float = 1.0, cost_mask: float = 1.0, cost_dice: float = 1.0, num_points: int = 12544
563
+ ):
564
+ """Creates the matcher
565
+
566
+ Params:
567
+ cost_class (`float`, *optional*, defaults to 1.0):
568
+ Relative weight of the classification error in the matching cost.
569
+ cost_mask (`float`, *optional*, defaults to 1.0):
570
+ This is the relative weight of the focal loss of the binary mask in the matching cost.
571
+ cost_dice (`float`, *optional*, defaults to 1.0):
572
+ This is the relative weight of the dice loss of the binary mask in the matching cost.
573
+ num_points (`int`, *optional*, defaults to 12544):
574
+ No. of points to sample on which the mask loss will be calculated. The same set of K points are
575
+ uniformly sampled for all prediction and ground truth masks to construct the cost matrix for bipartite
576
+ matching.
577
+ """
578
+ super().__init__()
579
+ if cost_class == 0 and cost_mask == 0 and cost_dice == 0:
580
+ raise ValueError("All costs can't be 0")
581
+
582
+ self.num_points = num_points
583
+ self.cost_class = cost_class
584
+ self.cost_mask = cost_mask
585
+ self.cost_dice = cost_dice
586
+
587
+ @torch.no_grad()
588
+ def forward(
589
+ self,
590
+ masks_queries_logits: torch.Tensor,
591
+ class_queries_logits: torch.Tensor,
592
+ mask_labels: torch.Tensor,
593
+ class_labels: torch.Tensor,
594
+ ) -> list[tuple[Tensor]]:
595
+ """
596
+ Params:
597
+ masks_queries_logits (`torch.Tensor`):
598
+ A tensor of dim `batch_size, num_queries, num_labels` with the classification logits.
599
+ class_queries_logits (`torch.Tensor`):
600
+ A tensor of dim `batch_size, num_queries, height, width` with the predicted masks.
601
+ class_labels (`torch.Tensor`):
602
+ A tensor of dim `num_target_boxes` (where num_target_boxes is the number of ground-truth objects in the
603
+ target) containing the class labels.
604
+ mask_labels (`torch.Tensor`):
605
+ A tensor of dim `num_target_boxes, height, width` containing the target masks.
606
+
607
+ Returns:
608
+ matched_indices (`list[tuple[Tensor]]`): A list of size batch_size, containing tuples of (index_i, index_j)
609
+ where:
610
+ - index_i is the indices of the selected predictions (in order)
611
+ - index_j is the indices of the corresponding selected labels (in order)
612
+ For each batch element, it holds:
613
+ len(index_i) = len(index_j) = min(num_queries, num_target_boxes).
614
+ """
615
+ indices: list[tuple[np.array]] = []
616
+
617
+ # iterate through batch size
618
+ batch_size = masks_queries_logits.shape[0]
619
+ for i in range(batch_size):
620
+ pred_probs = class_queries_logits[i].softmax(-1)
621
+ pred_mask = masks_queries_logits[i]
622
+
623
+ # Compute the classification cost. Contrary to the loss, we don't use the NLL, but approximate it in 1 - proba[target class]. The 1 is a constant that doesn't change the matching, it can be omitted.
624
+ cost_class = -pred_probs[:, class_labels[i]]
625
+ target_mask = mask_labels[i].to(pred_mask)
626
+ target_mask = target_mask[:, None]
627
+ pred_mask = pred_mask[:, None]
628
+
629
+ # Sample ground truth and predicted masks
630
+ point_coordinates = torch.rand(1, self.num_points, 2, device=pred_mask.device)
631
+
632
+ target_coordinates = point_coordinates.repeat(target_mask.shape[0], 1, 1)
633
+ target_mask = sample_point(target_mask, target_coordinates, align_corners=False).squeeze(1)
634
+
635
+ pred_coordinates = point_coordinates.repeat(pred_mask.shape[0], 1, 1)
636
+ pred_mask = sample_point(pred_mask, pred_coordinates, align_corners=False).squeeze(1)
637
+
638
+ # compute the cross entropy loss between each mask pairs -> shape (num_queries, num_labels)
639
+ cost_mask = pair_wise_sigmoid_cross_entropy_loss(pred_mask, target_mask)
640
+ # Compute the dice loss between each mask pairs -> shape (num_queries, num_labels)
641
+ cost_dice = pair_wise_dice_loss(pred_mask, target_mask)
642
+ # final cost matrix
643
+ cost_matrix = self.cost_mask * cost_mask + self.cost_class * cost_class + self.cost_dice * cost_dice
644
+ # eliminate infinite values in cost_matrix to avoid the error ``ValueError: cost matrix is infeasible``
645
+ cost_matrix = torch.minimum(cost_matrix, torch.tensor(1e10))
646
+ cost_matrix = torch.maximum(cost_matrix, torch.tensor(-1e10))
647
+ cost_matrix = torch.nan_to_num(cost_matrix, 0)
648
+ # do the assignment using the hungarian algorithm in scipy
649
+ assigned_indices: tuple[np.array] = linear_sum_assignment(cost_matrix.cpu())
650
+ indices.append(assigned_indices)
651
+
652
+ # It could be stacked in one tensor
653
+ matched_indices = [
654
+ (torch.as_tensor(i, dtype=torch.int64), torch.as_tensor(j, dtype=torch.int64)) for i, j in indices
655
+ ]
656
+ return matched_indices
657
+
658
+
659
+ def dice_loss(inputs: Tensor, labels: Tensor, num_masks: int) -> Tensor:
660
+ r"""
661
+ Compute the DICE loss, similar to generalized IOU for masks as follows:
662
+
663
+ $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x \cap y }{x \cup y + 1}} $$
664
+
665
+ In practice, since `labels` is a binary mask, (only 0s and 1s), dice can be computed as follow
666
+
667
+ $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x * y }{x + y + 1}} $$
668
+
669
+ Args:
670
+ inputs (`torch.Tensor`):
671
+ A tensor representing a mask.
672
+ labels (`torch.Tensor`):
673
+ A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
674
+ (0 for the negative class and 1 for the positive class).
675
+ num_masks (`int`):
676
+ The number of masks present in the current batch, used for normalization.
677
+
678
+ Returns:
679
+ `torch.Tensor`: The computed loss.
680
+ """
681
+ probs = inputs.sigmoid().flatten(1)
682
+ numerator = 2 * (probs * labels).sum(-1)
683
+ denominator = probs.sum(-1) + labels.sum(-1)
684
+ loss = 1 - (numerator + 1) / (denominator + 1)
685
+ loss = loss.sum() / num_masks
686
+ return loss
687
+
688
+
689
+ def sigmoid_cross_entropy_loss(inputs: torch.Tensor, labels: torch.Tensor, num_masks: int) -> torch.Tensor:
690
+ r"""
691
+ Args:
692
+ inputs (`torch.Tensor`):
693
+ A float tensor of arbitrary shape.
694
+ labels (`torch.Tensor`):
695
+ A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs
696
+ (0 for the negative class and 1 for the positive class).
697
+
698
+ Returns:
699
+ loss (`torch.Tensor`): The computed loss.
700
+ """
701
+ criterion = nn.BCEWithLogitsLoss(reduction="none")
702
+ cross_entropy_loss = criterion(inputs, labels)
703
+
704
+ loss = cross_entropy_loss.mean(1).sum() / num_masks
705
+ return loss
706
+
707
+
708
+ # Adapted from https://github.com/facebookresearch/EomtDinov3/blob/main/eomt_dinov3/modeling/criterion.py
709
+ class EomtDinov3Loss(nn.Module):
710
+ def __init__(self, config: EomtDinov3Config, weight_dict: dict[str, float]):
711
+ """
712
+ The EomtDinov3 Loss. The loss is computed very similar to DETR. The process happens in two steps: 1) we
713
+ compute hungarian assignment between ground truth masks and the outputs of the model 2) we supervise each pair
714
+ of matched ground-truth / prediction (supervise class and mask)
715
+
716
+ Args:
717
+ config (`EomtDinov3Config`):
718
+ The configuration for EomtDinov3 model also containing loss calculation specific parameters.
719
+ weight_dict (`dict[str, float]`):
720
+ A dictionary of weights to be applied to the different losses.
721
+ """
722
+ super().__init__()
723
+ requires_backends(self, ["scipy"])
724
+ self.num_labels = config.num_labels
725
+ self.weight_dict = weight_dict
726
+
727
+ # Weight to apply to the null class
728
+ self.eos_coef = config.no_object_weight
729
+ empty_weight = torch.ones(self.num_labels + 1)
730
+ empty_weight[-1] = self.eos_coef
731
+ self.register_buffer("empty_weight", empty_weight)
732
+
733
+ # pointwise mask loss parameters
734
+ self.num_points = config.train_num_points
735
+ self.oversample_ratio = config.oversample_ratio
736
+ self.importance_sample_ratio = config.importance_sample_ratio
737
+
738
+ self.matcher = EomtDinov3HungarianMatcher(
739
+ cost_class=config.class_weight,
740
+ cost_dice=config.dice_weight,
741
+ cost_mask=config.mask_weight,
742
+ num_points=self.num_points,
743
+ )
744
+
745
+ def _max_by_axis(self, sizes: list[list[int]]) -> list[int]:
746
+ maxes = sizes[0]
747
+ for sublist in sizes[1:]:
748
+ for index, item in enumerate(sublist):
749
+ maxes[index] = max(maxes[index], item)
750
+ return maxes
751
+
752
+ # Adapted from nested_tensor_from_tensor_list() in original implementation
753
+ def _pad_images_to_max_in_batch(self, tensors: list[Tensor]) -> tuple[Tensor, Tensor]:
754
+ # get the maximum size in the batch
755
+ max_size = self._max_by_axis([list(tensor.shape) for tensor in tensors])
756
+ # compute final size
757
+ batch_shape = [len(tensors)] + max_size
758
+ batch_size, _, height, width = batch_shape
759
+ dtype = tensors[0].dtype
760
+ device = tensors[0].device
761
+ padded_tensors = torch.zeros(batch_shape, dtype=dtype, device=device)
762
+ padding_masks = torch.ones((batch_size, height, width), dtype=torch.bool, device=device)
763
+ # pad the tensors to the size of the biggest one
764
+ for tensor, padded_tensor, padding_mask in zip(tensors, padded_tensors, padding_masks):
765
+ padded_tensor[: tensor.shape[0], : tensor.shape[1], : tensor.shape[2]].copy_(tensor)
766
+ padding_mask[: tensor.shape[1], : tensor.shape[2]] = False
767
+
768
+ return padded_tensors, padding_masks
769
+
770
+ def loss_labels(
771
+ self, class_queries_logits: Tensor, class_labels: list[Tensor], indices: tuple[np.array]
772
+ ) -> dict[str, Tensor]:
773
+ """Compute the losses related to the labels using cross entropy.
774
+
775
+ Args:
776
+ class_queries_logits (`torch.Tensor`):
777
+ A tensor of shape `batch_size, num_queries, num_labels`
778
+ class_labels (`list[torch.Tensor]`):
779
+ List of class labels of shape `(labels)`.
780
+ indices (`tuple[np.array])`:
781
+ The indices computed by the Hungarian matcher.
782
+
783
+ Returns:
784
+ `dict[str, Tensor]`: A dict of `torch.Tensor` containing the following key:
785
+ - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels.
786
+ """
787
+ pred_logits = class_queries_logits
788
+ batch_size, num_queries, _ = pred_logits.shape
789
+ criterion = nn.CrossEntropyLoss(weight=self.empty_weight)
790
+ idx = self._get_predictions_permutation_indices(indices) # shape of (batch_size, num_queries)
791
+ target_classes_o = torch.cat(
792
+ [target[j] for target, (_, j) in zip(class_labels, indices)]
793
+ ) # shape of (batch_size, num_queries)
794
+ target_classes = torch.full(
795
+ (batch_size, num_queries), fill_value=self.num_labels, dtype=torch.int64, device=pred_logits.device
796
+ )
797
+ target_classes[idx] = target_classes_o
798
+ # Permute target_classes (batch_size, num_queries, num_labels) -> (batch_size, num_labels, num_queries)
799
+ pred_logits_transposed = pred_logits.transpose(1, 2)
800
+ loss_ce = criterion(pred_logits_transposed, target_classes)
801
+ losses = {"loss_cross_entropy": loss_ce}
802
+ return losses
803
+
804
+ def loss_masks(
805
+ self,
806
+ masks_queries_logits: torch.Tensor,
807
+ mask_labels: list[torch.Tensor],
808
+ indices: tuple[np.array],
809
+ num_masks: int,
810
+ ) -> dict[str, torch.Tensor]:
811
+ """Compute the losses related to the masks using sigmoid_cross_entropy_loss and dice loss.
812
+
813
+ Args:
814
+ masks_queries_logits (`torch.Tensor`):
815
+ A tensor of shape `(batch_size, num_queries, height, width)`.
816
+ mask_labels (`torch.Tensor`):
817
+ List of mask labels of shape `(labels, height, width)`.
818
+ indices (`tuple[np.array])`:
819
+ The indices computed by the Hungarian matcher.
820
+ num_masks (`int)`:
821
+ The number of masks, used for normalization.
822
+
823
+ Returns:
824
+ losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing two keys:
825
+ - **loss_mask** -- The loss computed using sigmoid cross entropy loss on the predicted and ground truth.
826
+ masks.
827
+ - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth,
828
+ masks.
829
+ """
830
+ src_idx = self._get_predictions_permutation_indices(indices)
831
+ tgt_idx = self._get_targets_permutation_indices(indices)
832
+ # shape (batch_size * num_queries, height, width)
833
+ pred_masks = masks_queries_logits[src_idx]
834
+ # shape (batch_size, num_queries, height, width)
835
+ # pad all and stack the targets to the num_labels dimension
836
+ target_masks, _ = self._pad_images_to_max_in_batch(mask_labels)
837
+ target_masks = target_masks[tgt_idx]
838
+
839
+ # No need to upsample predictions as we are using normalized coordinates
840
+ pred_masks = pred_masks[:, None]
841
+ target_masks = target_masks[:, None]
842
+
843
+ # Sample point coordinates
844
+ with torch.no_grad():
845
+ point_coordinates = self.sample_points_using_uncertainty(
846
+ pred_masks,
847
+ lambda logits: self.calculate_uncertainty(logits),
848
+ self.num_points,
849
+ self.oversample_ratio,
850
+ self.importance_sample_ratio,
851
+ )
852
+
853
+ point_labels = sample_point(target_masks, point_coordinates, align_corners=False).squeeze(1)
854
+
855
+ point_logits = sample_point(pred_masks, point_coordinates, align_corners=False).squeeze(1)
856
+
857
+ losses = {
858
+ "loss_mask": sigmoid_cross_entropy_loss(point_logits, point_labels, num_masks),
859
+ "loss_dice": dice_loss(point_logits, point_labels, num_masks),
860
+ }
861
+
862
+ del pred_masks
863
+ del target_masks
864
+ return losses
865
+
866
+ def _get_predictions_permutation_indices(self, indices):
867
+ # Permute predictions following indices
868
+ batch_indices = torch.cat([torch.full_like(src, i) for i, (src, _) in enumerate(indices)])
869
+ predictions_indices = torch.cat([src for (src, _) in indices])
870
+ return batch_indices, predictions_indices
871
+
872
+ def _get_targets_permutation_indices(self, indices):
873
+ # Permute labels following indices
874
+ batch_indices = torch.cat([torch.full_like(tgt, i) for i, (_, tgt) in enumerate(indices)])
875
+ target_indices = torch.cat([tgt for (_, tgt) in indices])
876
+ return batch_indices, target_indices
877
+
878
+ def calculate_uncertainty(self, logits: torch.Tensor) -> torch.Tensor:
879
+ """
880
+ In EomtDinov3 paper, uncertainty is estimated as L1 distance between 0.0 and the logit prediction in 'logits'
881
+ for the foreground class in `classes`.
882
+
883
+ Args:
884
+ logits (`torch.Tensor`):
885
+ A tensor of shape (R, 1, ...) for class-specific or class-agnostic, where R is the total number of predicted masks in all images and C is:
886
+ the number of foreground classes. The values are logits.
887
+
888
+ Returns:
889
+ scores (`torch.Tensor`): A tensor of shape (R, 1, ...) that contains uncertainty scores with the most
890
+ uncertain locations having the highest uncertainty score.
891
+ """
892
+ uncertainty_scores = -(torch.abs(logits))
893
+ return uncertainty_scores
894
+
895
+ def sample_points_using_uncertainty(
896
+ self,
897
+ logits: torch.Tensor,
898
+ uncertainty_function,
899
+ num_points: int,
900
+ oversample_ratio: int,
901
+ importance_sample_ratio: float,
902
+ ) -> torch.Tensor:
903
+ """
904
+ This function is meant for sampling points in [0, 1] * [0, 1] coordinate space based on their uncertainty. The
905
+ uncertainty is calculated for each point using the passed `uncertainty function` that takes points logit
906
+ prediction as input.
907
+
908
+ Args:
909
+ logits (`float`):
910
+ Logit predictions for P points.
911
+ uncertainty_function:
912
+ A function that takes logit predictions for P points and returns their uncertainties.
913
+ num_points (`int`):
914
+ The number of points P to sample.
915
+ oversample_ratio (`int`):
916
+ Oversampling parameter.
917
+ importance_sample_ratio (`float`):
918
+ Ratio of points that are sampled via importance sampling.
919
+
920
+ Returns:
921
+ point_coordinates (`torch.Tensor`):
922
+ Coordinates for P sampled points.
923
+ """
924
+
925
+ num_boxes = logits.shape[0]
926
+ num_points_sampled = int(num_points * oversample_ratio)
927
+
928
+ # Get random point coordinates
929
+ point_coordinates = torch.rand(num_boxes, num_points_sampled, 2, device=logits.device)
930
+ # Get sampled prediction value for the point coordinates
931
+ point_logits = sample_point(logits, point_coordinates, align_corners=False)
932
+ # Calculate the uncertainties based on the sampled prediction values of the points
933
+ point_uncertainties = uncertainty_function(point_logits)
934
+
935
+ num_uncertain_points = int(importance_sample_ratio * num_points)
936
+ num_random_points = num_points - num_uncertain_points
937
+
938
+ idx = torch.topk(point_uncertainties[:, 0, :], k=num_uncertain_points, dim=1)[1]
939
+ shift = num_points_sampled * torch.arange(num_boxes, dtype=torch.long, device=logits.device)
940
+ idx += shift[:, None]
941
+ point_coordinates = point_coordinates.view(-1, 2)[idx.view(-1), :].view(num_boxes, num_uncertain_points, 2)
942
+
943
+ if num_random_points > 0:
944
+ point_coordinates = torch.cat(
945
+ [point_coordinates, torch.rand(num_boxes, num_random_points, 2, device=logits.device)],
946
+ dim=1,
947
+ )
948
+ return point_coordinates
949
+
950
+ def forward(
951
+ self,
952
+ masks_queries_logits: torch.Tensor,
953
+ class_queries_logits: torch.Tensor,
954
+ mask_labels: list[torch.Tensor],
955
+ class_labels: list[torch.Tensor],
956
+ auxiliary_predictions: dict[str, torch.Tensor] | None = None,
957
+ ) -> dict[str, torch.Tensor]:
958
+ """
959
+ This performs the loss computation.
960
+
961
+ Args:
962
+ masks_queries_logits (`torch.Tensor`):
963
+ A tensor of shape `(batch_size, num_queries, height, width)`.
964
+ class_queries_logits (`torch.Tensor`):
965
+ A tensor of shape `(batch_size, num_queries, num_labels)`.
966
+ mask_labels (`torch.Tensor`):
967
+ List of mask labels of shape `(labels, height, width)`.
968
+ class_labels (`list[torch.Tensor]`):
969
+ List of class labels of shape `(labels)`.
970
+ auxiliary_predictions (`dict[str, torch.Tensor]`, *optional*):
971
+ if `use_auxiliary_loss` was set to `true` in [`EomtDinov3Config`], then it contains the logits from
972
+ the inner layers of the EomtDinov3MaskedAttentionDecoder.
973
+
974
+ Returns:
975
+ losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing three keys:
976
+ - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels.
977
+ - **loss_mask** -- The loss computed using sigmoid cross_entropy loss on the predicted and ground truth
978
+ masks.
979
+ - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth
980
+ masks.
981
+ if `use_auxiliary_loss` was set to `true` in [`EomtDinov3Config`], the dictionary contains additional
982
+ losses for each auxiliary predictions.
983
+ """
984
+
985
+ # retrieve the matching between the outputs of the last layer and the labels
986
+ indices = self.matcher(masks_queries_logits, class_queries_logits, mask_labels, class_labels)
987
+ # compute the average number of target masks for normalization purposes
988
+ num_masks = self.get_num_masks(class_labels, device=class_labels[0].device)
989
+ # get all the losses
990
+ losses: dict[str, Tensor] = {
991
+ **self.loss_masks(masks_queries_logits, mask_labels, indices, num_masks),
992
+ **self.loss_labels(class_queries_logits, class_labels, indices),
993
+ }
994
+ # in case of auxiliary losses, we repeat this process with the output of each intermediate layer.
995
+ if auxiliary_predictions is not None:
996
+ for idx, aux_outputs in enumerate(auxiliary_predictions):
997
+ masks_queries_logits = aux_outputs["masks_queries_logits"]
998
+ class_queries_logits = aux_outputs["class_queries_logits"]
999
+ loss_dict = self.forward(masks_queries_logits, class_queries_logits, mask_labels, class_labels)
1000
+ loss_dict = {f"{key}_{idx}": value for key, value in loss_dict.items()}
1001
+ losses.update(loss_dict)
1002
+
1003
+ return losses
1004
+
1005
+ def get_num_masks(self, class_labels: torch.Tensor, device: torch.device) -> torch.Tensor:
1006
+ """
1007
+ Computes the average number of target masks across the batch, for normalization purposes.
1008
+ """
1009
+ num_masks = sum(len(classes) for classes in class_labels)
1010
+ num_masks = torch.as_tensor(num_masks, dtype=torch.float, device=device)
1011
+ world_size = 1
1012
+ if is_accelerate_available():
1013
+ if PartialState._shared_state != {}:
1014
+ num_masks = reduce(num_masks)
1015
+ world_size = PartialState().num_processes
1016
+
1017
+ num_masks = torch.clamp(num_masks / world_size, min=1)
1018
+ return num_masks
1019
+
1020
+
1021
+ @auto_docstring(
1022
+ custom_intro="""
1023
+ Class for outputs of [`EomtDinov3ForUniversalSegmentationOutput`].
1024
+
1025
+ This output can be directly passed to [`~EomtDinov3ImageProcessor.post_process_semantic_segmentation`] or
1026
+ [`~EomtDinov3ImageProcessor.post_process_instance_segmentation`] or
1027
+ [`~EomtDinov3ImageProcessor.post_process_panoptic_segmentation`] to compute final segmentation maps. Please, see
1028
+ [`~EomtDinov3ImageProcessor] for details regarding usage.
1029
+ """
1030
+ )
1031
+ @dataclass
1032
+ class EomtDinov3ForUniversalSegmentationOutput(ModelOutput):
1033
+ r"""
1034
+ loss (`torch.Tensor`, *optional*):
1035
+ The computed loss, returned when labels are present.
1036
+ class_queries_logits (`torch.FloatTensor`):
1037
+ A tensor of shape `(batch_size, num_queries, num_labels + 1)` representing the proposed classes for each
1038
+ query. Note the `+ 1` is needed because we incorporate the null class.
1039
+ masks_queries_logits (`torch.FloatTensor`):
1040
+ A tensor of shape `(batch_size, num_queries, height, width)` representing the proposed masks for each
1041
+ query.
1042
+ last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
1043
+ Last hidden states (final feature map) of the last layer.
1044
+ hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
1045
+ Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each stage) of
1046
+ shape `(batch_size, sequence_length, hidden_size)`. Hidden-states all layers of the model.
1047
+ attentions (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
1048
+ Tuple of `tuple(torch.FloatTensor)` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
1049
+ sequence_length)`. Self and Cross Attentions weights from transformer decoder.
1050
+ patch_offsets (`list[torch.Tensor]`, *optional*):
1051
+ list of tuples indicating the image index and start and end positions of patches for semantic segmentation.
1052
+ """
1053
+
1054
+ loss: torch.FloatTensor | None = None
1055
+ class_queries_logits: torch.FloatTensor | None = None
1056
+ masks_queries_logits: torch.FloatTensor | None = None
1057
+ last_hidden_state: torch.FloatTensor | None = None
1058
+ hidden_states: tuple[torch.FloatTensor] | None = None
1059
+ attentions: tuple[torch.FloatTensor] | None = None
1060
+ patch_offsets: list[torch.Tensor] | None = None
1061
+
1062
+
1063
+ @auto_docstring
1064
+ class EomtDinov3PreTrainedModel(PreTrainedModel):
1065
+ """
1066
+ An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
1067
+ models.
1068
+ """
1069
+
1070
+ config: EomtDinov3Config
1071
+ base_model_prefix = "eomt_dinov3"
1072
+ main_input_name = "pixel_values"
1073
+ input_modalities = ("image",)
1074
+ supports_gradient_checkpointing = False
1075
+ _no_split_modules = ["EomtDinov3Layer"]
1076
+ _supports_sdpa = True
1077
+ _can_record_outputs = {
1078
+ "hidden_states": EomtDinov3Layer,
1079
+ "attentions": EomtDinov3Attention,
1080
+ }
1081
+ config_class = EomtDinov3Config
1082
+
1083
+ @torch.no_grad()
1084
+ def _init_weights(self, module: nn.Module) -> None:
1085
+ super()._init_weights(module)
1086
+ std = self.config.initializer_range
1087
+ if isinstance(module, EomtDinov3LayerScale):
1088
+ if hasattr(module, "lambda1"):
1089
+ init.constant_(module.lambda1, self.config.layerscale_value)
1090
+ elif isinstance(module, EomtDinov3Embeddings):
1091
+ init.trunc_normal_(module.cls_token, mean=0.0, std=std)
1092
+ init.zeros_(module.register_tokens)
1093
+ elif isinstance(module, EomtDinov3Loss):
1094
+ empty_weight = torch.ones(module.num_labels + 1)
1095
+ empty_weight[-1] = module.eos_coef
1096
+ init.copy_(module.empty_weight, empty_weight)
1097
+ elif isinstance(module, EomtDinov3ForUniversalSegmentation):
1098
+ init.ones_(module.attn_mask_probs)
1099
+
1100
+
1101
+ class EomtDinov3LayerNorm2d(nn.LayerNorm):
1102
+ def __init__(self, num_channels, eps=1e-6, affine=True):
1103
+ super().__init__(num_channels, eps=eps, elementwise_affine=affine)
1104
+
1105
+ def forward(self, hidden_state: torch.Tensor) -> torch.Tensor:
1106
+ hidden_state = hidden_state.permute(0, 2, 3, 1)
1107
+ hidden_state = F.layer_norm(hidden_state, self.normalized_shape, self.weight, self.bias, self.eps)
1108
+ hidden_state = hidden_state.permute(0, 3, 1, 2)
1109
+ return hidden_state
1110
+
1111
+
1112
+ class EomtDinov3ScaleLayer(nn.Module):
1113
+ def __init__(self, config: EomtDinov3Config):
1114
+ super().__init__()
1115
+ hidden_size = config.hidden_size
1116
+ self.conv1 = nn.ConvTranspose2d(hidden_size, hidden_size, kernel_size=2, stride=2)
1117
+ self.activation = ACT2FN[config.hidden_act]
1118
+ self.conv2 = nn.Conv2d(
1119
+ hidden_size,
1120
+ hidden_size,
1121
+ kernel_size=3,
1122
+ padding=1,
1123
+ groups=hidden_size,
1124
+ bias=False,
1125
+ )
1126
+
1127
+ self.layernorm2d = EomtDinov3LayerNorm2d(hidden_size)
1128
+
1129
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
1130
+ hidden_states = self.conv1(hidden_states)
1131
+ hidden_states = self.activation(hidden_states)
1132
+ hidden_states = self.conv2(hidden_states)
1133
+ hidden_states = self.layernorm2d(hidden_states)
1134
+ return hidden_states
1135
+
1136
+
1137
+ class EomtDinov3ScaleBlock(nn.Module):
1138
+ def __init__(self, config: EomtDinov3Config):
1139
+ super().__init__()
1140
+ self.num_blocks = config.num_upscale_blocks
1141
+ self.block = nn.ModuleList([EomtDinov3ScaleLayer(config) for _ in range(self.num_blocks)])
1142
+
1143
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
1144
+ for block in self.block:
1145
+ hidden_states = block(hidden_states)
1146
+ return hidden_states
1147
+
1148
+
1149
+ class EomtDinov3MaskHead(nn.Module):
1150
+ def __init__(self, config: EomtDinov3Config):
1151
+ super().__init__()
1152
+
1153
+ hidden_size = config.hidden_size
1154
+ self.fc1 = nn.Linear(hidden_size, hidden_size)
1155
+ self.fc2 = nn.Linear(hidden_size, hidden_size)
1156
+ self.fc3 = nn.Linear(hidden_size, hidden_size)
1157
+ self.activation = ACT2FN[config.hidden_act]
1158
+
1159
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
1160
+ hidden_states = self.activation(self.fc1(hidden_states))
1161
+ hidden_states = self.activation(self.fc2(hidden_states))
1162
+ hidden_states = self.fc3(hidden_states)
1163
+ return hidden_states
1164
+
1165
+
1166
+ @auto_docstring(
1167
+ custom_intro="""
1168
+ The EoMT-DINOv3 model with head on top for instance/semantic/panoptic segmentation.
1169
+ """,
1170
+ )
1171
+ class EomtDinov3ForUniversalSegmentation(EomtDinov3PreTrainedModel):
1172
+ main_input_name = "pixel_values"
1173
+
1174
+ def __init__(self, config: EomtDinov3Config):
1175
+ super().__init__(config)
1176
+ self.config = config
1177
+ self.num_hidden_layers = config.num_hidden_layers
1178
+ self.embeddings = EomtDinov3Embeddings(config)
1179
+ self.layernorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
1180
+
1181
+ self.query = nn.Embedding(config.num_queries, config.hidden_size)
1182
+ self.layers = nn.ModuleList([EomtDinov3Layer(config) for _ in range(config.num_hidden_layers)])
1183
+
1184
+ self.upscale_block = EomtDinov3ScaleBlock(config)
1185
+ self.mask_head = EomtDinov3MaskHead(config)
1186
+
1187
+ self.class_predictor = nn.Linear(config.hidden_size, config.num_labels + 1)
1188
+
1189
+ self.grid_size = (config.image_size // config.patch_size, config.image_size // config.patch_size)
1190
+ self.weight_dict: dict[str, float] = {
1191
+ "loss_cross_entropy": config.class_weight,
1192
+ "loss_mask": config.mask_weight,
1193
+ "loss_dice": config.dice_weight,
1194
+ }
1195
+
1196
+ self.criterion = EomtDinov3Loss(config=config, weight_dict=self.weight_dict)
1197
+
1198
+ self.register_buffer("attn_mask_probs", torch.ones(config.num_blocks))
1199
+
1200
+ self.num_prefix_tokens = 1 + config.num_register_tokens
1201
+ self.dropout = nn.Dropout(config.hidden_dropout_prob)
1202
+ self.embeddings.register_parameter("mask_token", None)
1203
+
1204
+ self.rope_embeddings = EomtDinov3RotaryEmbedding(config)
1205
+
1206
+ self.post_init()
1207
+
1208
+ def get_loss_dict(
1209
+ self,
1210
+ masks_queries_logits: Tensor,
1211
+ class_queries_logits: Tensor,
1212
+ mask_labels: Tensor,
1213
+ class_labels: Tensor,
1214
+ auxiliary_predictions: dict[str, Tensor],
1215
+ ) -> dict[str, Tensor]:
1216
+ loss_dict: dict[str, Tensor] = self.criterion(
1217
+ masks_queries_logits=masks_queries_logits,
1218
+ class_queries_logits=class_queries_logits,
1219
+ mask_labels=mask_labels,
1220
+ class_labels=class_labels,
1221
+ auxiliary_predictions=auxiliary_predictions,
1222
+ )
1223
+
1224
+ # weight each loss by `self.weight_dict[<LOSS_NAME>]` including auxiliary losses
1225
+ for key, weight in self.weight_dict.items():
1226
+ for loss_key, loss in loss_dict.items():
1227
+ if key in loss_key:
1228
+ loss *= weight
1229
+
1230
+ return loss_dict
1231
+
1232
+ def get_loss(self, loss_dict: dict[str, Tensor]) -> Tensor:
1233
+ return sum(loss_dict.values())
1234
+
1235
+ @merge_with_config_defaults
1236
+ @capture_outputs
1237
+ @auto_docstring
1238
+ def forward(
1239
+ self,
1240
+ pixel_values: Tensor,
1241
+ mask_labels: list[Tensor] | None = None,
1242
+ class_labels: list[Tensor] | None = None,
1243
+ patch_offsets: list[Tensor] | None = None,
1244
+ **kwargs: Unpack[TransformersKwargs],
1245
+ ) -> EomtDinov3ForUniversalSegmentationOutput:
1246
+ r"""
1247
+ mask_labels (`list[torch.Tensor]`, *optional*):
1248
+ list of mask labels of shape `(num_labels, height, width)` to be fed to a model
1249
+ class_labels (`list[torch.LongTensor]`, *optional*):
1250
+ list of target class labels of shape `(num_labels, height, width)` to be fed to a model. They identify the
1251
+ labels of `mask_labels`, e.g. the label of `mask_labels[i][j]` if `class_labels[i][j]`.
1252
+ patch_offsets (`list[torch.Tensor]`, *optional*):
1253
+ list of tuples indicating the image index and start and end positions of patches for semantic segmentation.
1254
+ """
1255
+ masks_queries_logits_per_layer, class_queries_logits_per_layer = (), ()
1256
+
1257
+ hidden_states = self.dropout(self.embeddings(pixel_values))
1258
+ position_embeddings = self.rope_embeddings(pixel_values.to(hidden_states.dtype))
1259
+ attention_mask = None
1260
+
1261
+ for idx, layer_module in enumerate(self.layers):
1262
+ if idx == self.num_hidden_layers - self.config.num_blocks:
1263
+ query = self.query.weight[None, :, :].expand(hidden_states.shape[0], -1, -1).to(hidden_states.device)
1264
+ hidden_states = torch.cat((query, hidden_states), dim=1)
1265
+
1266
+ if idx >= self.num_hidden_layers - self.config.num_blocks and (
1267
+ self.training or self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks] > 0
1268
+ ):
1269
+ norm_hidden_states = self.layernorm(hidden_states)
1270
+ masks_queries_logits, class_queries_logits = self.predict(norm_hidden_states)
1271
+
1272
+ masks_queries_logits_per_layer += (masks_queries_logits,)
1273
+ class_queries_logits_per_layer += (class_queries_logits,)
1274
+
1275
+ attention_mask = torch.ones(
1276
+ hidden_states.shape[0],
1277
+ hidden_states.shape[1],
1278
+ hidden_states.shape[1],
1279
+ device=hidden_states.device,
1280
+ dtype=torch.bool,
1281
+ )
1282
+
1283
+ interpolated_logits = F.interpolate(masks_queries_logits, size=self.grid_size, mode="bilinear")
1284
+ interpolated_logits = interpolated_logits.view(
1285
+ interpolated_logits.size(0), interpolated_logits.size(1), -1
1286
+ )
1287
+
1288
+ num_query_tokens = self.config.num_queries
1289
+ encoder_start_tokens = num_query_tokens + self.num_prefix_tokens
1290
+
1291
+ # Set attention mask for queries to focus on encoder tokens based on interpolated logits
1292
+ attention_mask[:, :num_query_tokens, encoder_start_tokens:] = interpolated_logits > 0
1293
+
1294
+ # Disable attention mask for random query tokens.
1295
+ attention_mask = self._disable_attention_mask(
1296
+ attention_mask,
1297
+ prob=self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks],
1298
+ num_query_tokens=num_query_tokens,
1299
+ encoder_start_tokens=encoder_start_tokens,
1300
+ device=attention_mask.device,
1301
+ )
1302
+
1303
+ # Expand attention mask to 4d mask.
1304
+ attention_mask = attention_mask[:, None, ...].expand(-1, self.config.num_attention_heads, -1, -1)
1305
+ dtype_min = torch.finfo(hidden_states.dtype).min
1306
+ attention_mask = attention_mask.to(hidden_states.dtype).masked_fill(~attention_mask, dtype_min)
1307
+
1308
+ hidden_states = layer_module(
1309
+ hidden_states,
1310
+ attention_mask=attention_mask,
1311
+ position_embeddings=position_embeddings,
1312
+ )
1313
+
1314
+ sequence_output = self.layernorm(hidden_states)
1315
+
1316
+ masks_queries_logits, class_queries_logits = self.predict(sequence_output)
1317
+ masks_queries_logits_per_layer += (masks_queries_logits,)
1318
+ class_queries_logits_per_layer += (class_queries_logits,)
1319
+
1320
+ loss = None
1321
+ if mask_labels is not None and class_labels is not None:
1322
+ loss = 0.0
1323
+ for masks_queries_logits, class_queries_logits in zip(
1324
+ masks_queries_logits_per_layer, class_queries_logits_per_layer
1325
+ ):
1326
+ loss_dict = self.get_loss_dict(
1327
+ masks_queries_logits=masks_queries_logits,
1328
+ class_queries_logits=class_queries_logits,
1329
+ mask_labels=mask_labels,
1330
+ class_labels=class_labels,
1331
+ auxiliary_predictions=None,
1332
+ )
1333
+ loss += self.get_loss(loss_dict)
1334
+
1335
+ return EomtDinov3ForUniversalSegmentationOutput(
1336
+ loss=loss,
1337
+ masks_queries_logits=masks_queries_logits,
1338
+ class_queries_logits=class_queries_logits,
1339
+ last_hidden_state=sequence_output,
1340
+ patch_offsets=patch_offsets,
1341
+ )
1342
+
1343
+ def get_input_embeddings(self):
1344
+ return self.embeddings.patch_embeddings
1345
+
1346
+ def predict(self, logits: torch.Tensor):
1347
+ query_tokens = logits[:, : self.config.num_queries, :]
1348
+ class_logits = self.class_predictor(query_tokens)
1349
+
1350
+ prefix_tokens = logits[:, self.config.num_queries + self.embeddings.num_prefix_tokens :, :]
1351
+ prefix_tokens = prefix_tokens.transpose(1, 2)
1352
+
1353
+ prefix_tokens = prefix_tokens.reshape(prefix_tokens.shape[0], -1, *self.grid_size)
1354
+
1355
+ query_tokens = self.mask_head(query_tokens)
1356
+ prefix_tokens = self.upscale_block(prefix_tokens)
1357
+
1358
+ mask_logits = torch.einsum("bqc, bchw -> bqhw", query_tokens, prefix_tokens)
1359
+
1360
+ return mask_logits, class_logits
1361
+
1362
+ @staticmethod
1363
+ def _disable_attention_mask(attn_mask, prob, num_query_tokens, encoder_start_tokens, device):
1364
+ if prob < 1:
1365
+ # Generate random queries to disable based on the probs
1366
+ random_queries = torch.rand(attn_mask.shape[0], num_query_tokens, device=device) > prob
1367
+
1368
+ # Disable attention to the query tokens, considering the prefix tokens
1369
+ attn_mask[:, :num_query_tokens, encoder_start_tokens:][random_queries] = 1
1370
+
1371
+ return attn_mask
1372
+
1373
+
1374
+ __all__ = ["EomtDinov3PreTrainedModel", "EomtDinov3ForUniversalSegmentation"]
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modular_eomt_dinov3.py ADDED
@@ -0,0 +1,364 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 the HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ """PyTorch EoMT model backed by DINOv3."""
15
+
16
+ from collections.abc import Callable
17
+ from typing import Optional
18
+
19
+ import torch
20
+ import torch.nn.functional as F
21
+ from huggingface_hub.dataclasses import strict
22
+ from torch import Tensor, nn
23
+
24
+ from ... import initialization as init
25
+ from ...modeling_rope_utils import RopeParameters
26
+ from ...modeling_utils import PreTrainedModel
27
+ from ...processing_utils import Unpack
28
+ from ...utils import (
29
+ TransformersKwargs,
30
+ auto_docstring,
31
+ )
32
+ from ...utils.generic import merge_with_config_defaults
33
+ from ...utils.output_capturing import capture_outputs
34
+ from ..dinov3_vit.modeling_dinov3_vit import (
35
+ DINOv3ViTAttention,
36
+ DINOv3ViTEmbeddings,
37
+ DINOv3ViTLayer,
38
+ DINOv3ViTLayerScale,
39
+ DINOv3ViTRopePositionEmbedding,
40
+ )
41
+ from ..eomt.configuration_eomt import EomtConfig
42
+ from ..eomt.modeling_eomt import (
43
+ EomtForUniversalSegmentation,
44
+ EomtForUniversalSegmentationOutput,
45
+ EomtLoss,
46
+ EomtPreTrainedModel,
47
+ )
48
+
49
+
50
+ @auto_docstring(checkpoint="tue-mps/coco_panoptic_eomt_large_640_dinov3")
51
+ @strict
52
+ class EomtDinov3Config(EomtConfig):
53
+ r"""
54
+ layerscale_value (`float`, *optional*, defaults to 1.0):
55
+ Initial value for the LayerScale parameter.
56
+ num_upscale_blocks (`int`, *optional*, defaults to 2):
57
+ Number of upsampling blocks used in the decoder or segmentation head.
58
+ num_blocks (`int`, *optional*, defaults to 4):
59
+ Number of feature blocks or stages in the architecture.
60
+ no_object_weight (`float`, *optional*, defaults to 0.1):
61
+ Loss weight for the "no object" class in panoptic/instance segmentation.
62
+ train_num_points (`int`, *optional*, defaults to 12544):
63
+ Number of points to sample for mask loss computation during training.
64
+ oversample_ratio (`float`, *optional*, defaults to 3.0):
65
+ Oversampling ratio used in point sampling for mask training.
66
+ importance_sample_ratio (`float`, *optional*, defaults to 0.75):
67
+ Ratio of points to sample based on importance during training.
68
+ num_queries (`int`, *optional*, defaults to 200):
69
+ Number of object queries in the Transformer.
70
+ num_register_tokens (`int`, *optional*, defaults to 4):
71
+ Number of learnable register tokens added to the transformer input.
72
+ query_bias (`bool`, *optional*, defaults to `True`):
73
+ Whether to use bias in query projection.
74
+ key_bias (`bool`, *optional*, defaults to `False`):
75
+ Whether to use bias in key projection.
76
+ value_bias (`bool`, *optional*, defaults to `True`):
77
+ Whether to use bias in value projection.
78
+ proj_bias (`bool`, *optional*, defaults to `True`):
79
+ Whether to use bias in output projection.
80
+ use_gated_mlp (`bool`, *optional*, defaults to `False`):
81
+ Whether to use gated MLP layers.
82
+ pos_embed_shift (`float`, *optional*):
83
+ Shift value for position embeddings.
84
+ pos_embed_jitter (`float`, *optional*):
85
+ Jitter value for position embeddings.
86
+ pos_embed_rescale (`float`, *optional*, defaults to 2.0):
87
+ Rescale value for position embeddings.
88
+ """
89
+
90
+ model_type = "eomt_dinov3"
91
+ default_theta = 100.0
92
+
93
+ hidden_size: int = 1024
94
+ num_hidden_layers: int = 24
95
+ num_attention_heads: int = 16
96
+ intermediate_size: int = 4096
97
+ hidden_act: str = "gelu"
98
+ hidden_dropout_prob: float | int = 0.0
99
+ initializer_range: float = 0.02
100
+ layer_norm_eps: float = 1e-6
101
+ image_size: int | list[int] | tuple[int, int] = 640
102
+ patch_size: int | list[int] | tuple[int, int] = 16
103
+ num_channels: int = 3
104
+ layerscale_value: float = 1.0
105
+ drop_path_rate: float | int = 0.0
106
+ num_upscale_blocks: int = 2
107
+ attention_dropout: float | int = 0.0
108
+ num_blocks: int = 4
109
+ no_object_weight: float = 0.1
110
+ class_weight: float = 2.0
111
+ mask_weight: float = 5.0
112
+ dice_weight: float = 5.0
113
+ train_num_points: int = 12544
114
+ oversample_ratio: float = 3.0
115
+ importance_sample_ratio: float = 0.75
116
+ num_queries: int = 200
117
+ num_register_tokens: int = 4
118
+ rope_parameters: RopeParameters | dict | None = None
119
+ query_bias: bool = True
120
+ key_bias: bool = False
121
+ value_bias: bool = True
122
+ proj_bias: bool = True
123
+ mlp_bias: bool = True
124
+ use_gated_mlp: bool = False
125
+ pos_embed_shift: float | None = None
126
+ pos_embed_jitter: float | None = None
127
+ pos_embed_rescale: float | None = 2.0
128
+
129
+ mlp_ratio = AttributeError()
130
+ use_swiglu_ffn = AttributeError()
131
+
132
+
133
+ class EomtDinov3Attention(DINOv3ViTAttention):
134
+ pass
135
+
136
+
137
+ class EomtDinov3Embeddings(DINOv3ViTEmbeddings):
138
+ def __init__(self, config: EomtDinov3Config):
139
+ super().__init__(config)
140
+ self.num_prefix_tokens = 1 + config.num_register_tokens
141
+
142
+
143
+ class EomtDinov3Layer(DINOv3ViTLayer):
144
+ pass
145
+
146
+
147
+ class EomtDinov3LayerScale(DINOv3ViTLayerScale):
148
+ pass
149
+
150
+
151
+ class EomtDinov3RotaryEmbedding(DINOv3ViTRopePositionEmbedding):
152
+ inv_freq: Tensor
153
+
154
+ def __init__(self, config: EomtDinov3Config, device=None):
155
+ nn.Module.__init__(self)
156
+ self.config = config
157
+
158
+ self.rope_type = self.config.rope_parameters["rope_type"]
159
+ rope_init_fn: Callable = self.compute_default_rope_parameters
160
+ if self.rope_type != "default":
161
+ raise ValueError("`EomtDinov3` only supports `default` RoPE! Please check your `rope_type`")
162
+ inv_freq, self.attention_scaling = rope_init_fn(self.config, device)
163
+
164
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
165
+ self.register_buffer("original_inv_freq", inv_freq.clone(), persistent=False)
166
+
167
+ @staticmethod
168
+ def compute_default_rope_parameters(
169
+ config: EomtDinov3Config | None = None,
170
+ device: Optional["torch.device"] = None,
171
+ seq_len: int | None = None,
172
+ ) -> torch.Tensor:
173
+ """
174
+ Computes the inverse frequencies according to the original RoPE implementation
175
+ Args:
176
+ config ([`~transformers.PreTrainedConfig`]):
177
+ The model configuration.
178
+ device (`torch.device`):
179
+ The device to use for initialization of the inverse frequencies.
180
+ seq_len (`int`, *optional*):
181
+ The current sequence length. Unused for this type of RoPE.
182
+ Returns:
183
+ Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
184
+ post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
185
+ """
186
+ base = config.rope_parameters["rope_theta"]
187
+ head_dim = config.hidden_size // config.num_attention_heads
188
+
189
+ attention_factor = 1.0 # Unused in this type of RoPE
190
+
191
+ # Compute the inverse frequencies
192
+ inv_freq = 1 / base ** torch.arange(0, 1, 4 / head_dim, dtype=torch.float32, device=device)
193
+ return inv_freq, attention_factor
194
+
195
+
196
+ class EomtDinov3Loss(EomtLoss):
197
+ pass
198
+
199
+
200
+ class EomtDinov3ForUniversalSegmentationOutput(EomtForUniversalSegmentationOutput):
201
+ pass
202
+
203
+
204
+ class EomtDinov3PreTrainedModel(EomtPreTrainedModel):
205
+ config_class = EomtDinov3Config
206
+ base_model_prefix = "eomt_dinov3"
207
+ _no_split_modules = ["EomtDinov3Layer"]
208
+ _can_record_outputs = {
209
+ "hidden_states": EomtDinov3Layer,
210
+ "attentions": EomtDinov3Attention,
211
+ }
212
+
213
+ def _init_weights(self, module: nn.Module) -> None:
214
+ PreTrainedModel._init_weights(module)
215
+ std = self.config.initializer_range
216
+ if isinstance(module, EomtDinov3LayerScale):
217
+ if hasattr(module, "lambda1"):
218
+ init.constant_(module.lambda1, self.config.layerscale_value)
219
+ elif isinstance(module, EomtDinov3Embeddings):
220
+ init.trunc_normal_(module.cls_token, mean=0.0, std=std)
221
+ init.zeros_(module.register_tokens)
222
+ elif isinstance(module, EomtDinov3Loss):
223
+ empty_weight = torch.ones(module.num_labels + 1)
224
+ empty_weight[-1] = module.eos_coef
225
+ init.copy_(module.empty_weight, empty_weight)
226
+ elif isinstance(module, EomtDinov3ForUniversalSegmentation):
227
+ init.ones_(module.attn_mask_probs)
228
+
229
+
230
+ @auto_docstring(
231
+ custom_intro="""
232
+ The EoMT-DINOv3 model with head on top for instance/semantic/panoptic segmentation.
233
+ """,
234
+ )
235
+ class EomtDinov3ForUniversalSegmentation(EomtDinov3PreTrainedModel, EomtForUniversalSegmentation):
236
+ def __init__(self, config: EomtDinov3Config):
237
+ super().__init__(config)
238
+
239
+ self.num_prefix_tokens = 1 + config.num_register_tokens
240
+ self.dropout = nn.Dropout(config.hidden_dropout_prob)
241
+ self.embeddings = EomtDinov3Embeddings(config)
242
+ self.embeddings.register_parameter("mask_token", None)
243
+
244
+ self.rope_embeddings = EomtDinov3RotaryEmbedding(config)
245
+ self.layers = nn.ModuleList([EomtDinov3Layer(config) for _ in range(config.num_hidden_layers)])
246
+
247
+ self.post_init()
248
+
249
+ # We redefine forward here because EoMT-DINOv3 uses DINOv3 backbone components (RoPE embeddings, layers)
250
+ # which require different integration than the base EoMT model that uses a separate encoder.
251
+ @merge_with_config_defaults
252
+ @capture_outputs
253
+ @auto_docstring
254
+ def forward(
255
+ self,
256
+ pixel_values: Tensor,
257
+ mask_labels: list[Tensor] | None = None,
258
+ class_labels: list[Tensor] | None = None,
259
+ patch_offsets: list[Tensor] | None = None,
260
+ **kwargs: Unpack[TransformersKwargs],
261
+ ) -> EomtDinov3ForUniversalSegmentationOutput:
262
+ r"""
263
+ mask_labels (`list[torch.Tensor]`, *optional*):
264
+ list of mask labels of shape `(num_labels, height, width)` to be fed to a model
265
+ class_labels (`list[torch.LongTensor]`, *optional*):
266
+ list of target class labels of shape `(num_labels, height, width)` to be fed to a model. They identify the
267
+ labels of `mask_labels`, e.g. the label of `mask_labels[i][j]` if `class_labels[i][j]`.
268
+ patch_offsets (`list[torch.Tensor]`, *optional*):
269
+ list of tuples indicating the image index and start and end positions of patches for semantic segmentation.
270
+ """
271
+ masks_queries_logits_per_layer, class_queries_logits_per_layer = (), ()
272
+
273
+ hidden_states = self.dropout(self.embeddings(pixel_values))
274
+ position_embeddings = self.rope_embeddings(pixel_values.to(hidden_states.dtype))
275
+ attention_mask = None
276
+
277
+ for idx, layer_module in enumerate(self.layers):
278
+ if idx == self.num_hidden_layers - self.config.num_blocks:
279
+ query = self.query.weight[None, :, :].expand(hidden_states.shape[0], -1, -1).to(hidden_states.device)
280
+ hidden_states = torch.cat((query, hidden_states), dim=1)
281
+
282
+ if idx >= self.num_hidden_layers - self.config.num_blocks and (
283
+ self.training or self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks] > 0
284
+ ):
285
+ norm_hidden_states = self.layernorm(hidden_states)
286
+ masks_queries_logits, class_queries_logits = self.predict(norm_hidden_states)
287
+
288
+ masks_queries_logits_per_layer += (masks_queries_logits,)
289
+ class_queries_logits_per_layer += (class_queries_logits,)
290
+
291
+ attention_mask = torch.ones(
292
+ hidden_states.shape[0],
293
+ hidden_states.shape[1],
294
+ hidden_states.shape[1],
295
+ device=hidden_states.device,
296
+ dtype=torch.bool,
297
+ )
298
+
299
+ interpolated_logits = F.interpolate(masks_queries_logits, size=self.grid_size, mode="bilinear")
300
+ interpolated_logits = interpolated_logits.view(
301
+ interpolated_logits.size(0), interpolated_logits.size(1), -1
302
+ )
303
+
304
+ num_query_tokens = self.config.num_queries
305
+ encoder_start_tokens = num_query_tokens + self.num_prefix_tokens
306
+
307
+ # Set attention mask for queries to focus on encoder tokens based on interpolated logits
308
+ attention_mask[:, :num_query_tokens, encoder_start_tokens:] = interpolated_logits > 0
309
+
310
+ # Disable attention mask for random query tokens.
311
+ attention_mask = self._disable_attention_mask(
312
+ attention_mask,
313
+ prob=self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks],
314
+ num_query_tokens=num_query_tokens,
315
+ encoder_start_tokens=encoder_start_tokens,
316
+ device=attention_mask.device,
317
+ )
318
+
319
+ # Expand attention mask to 4d mask.
320
+ attention_mask = attention_mask[:, None, ...].expand(-1, self.config.num_attention_heads, -1, -1)
321
+ dtype_min = torch.finfo(hidden_states.dtype).min
322
+ attention_mask = attention_mask.to(hidden_states.dtype).masked_fill(~attention_mask, dtype_min)
323
+
324
+ hidden_states = layer_module(
325
+ hidden_states,
326
+ attention_mask=attention_mask,
327
+ position_embeddings=position_embeddings,
328
+ )
329
+
330
+ sequence_output = self.layernorm(hidden_states)
331
+
332
+ masks_queries_logits, class_queries_logits = self.predict(sequence_output)
333
+ masks_queries_logits_per_layer += (masks_queries_logits,)
334
+ class_queries_logits_per_layer += (class_queries_logits,)
335
+
336
+ loss = None
337
+ if mask_labels is not None and class_labels is not None:
338
+ loss = 0.0
339
+ for masks_queries_logits, class_queries_logits in zip(
340
+ masks_queries_logits_per_layer, class_queries_logits_per_layer
341
+ ):
342
+ loss_dict = self.get_loss_dict(
343
+ masks_queries_logits=masks_queries_logits,
344
+ class_queries_logits=class_queries_logits,
345
+ mask_labels=mask_labels,
346
+ class_labels=class_labels,
347
+ auxiliary_predictions=None,
348
+ )
349
+ loss += self.get_loss(loss_dict)
350
+
351
+ return EomtDinov3ForUniversalSegmentationOutput(
352
+ loss=loss,
353
+ masks_queries_logits=masks_queries_logits,
354
+ class_queries_logits=class_queries_logits,
355
+ last_hidden_state=sequence_output,
356
+ patch_offsets=patch_offsets,
357
+ )
358
+
359
+
360
+ __all__ = [
361
+ "EomtDinov3Config",
362
+ "EomtDinov3PreTrainedModel",
363
+ "EomtDinov3ForUniversalSegmentation",
364
+ ]
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/__init__.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2024 The HuggingFace Team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ from typing import TYPE_CHECKING
15
+
16
+ from ...utils import _LazyModule
17
+ from ...utils.import_utils import define_import_structure
18
+
19
+
20
+ if TYPE_CHECKING:
21
+ from .configuration_patchtsmixer import *
22
+ from .modeling_patchtsmixer import *
23
+ else:
24
+ import sys
25
+
26
+ _file = globals()["__file__"]
27
+ sys.modules[__name__] = _LazyModule(__name__, _file, define_import_structure(_file), module_spec=__spec__)
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/configuration_patchtsmixer.py ADDED
@@ -0,0 +1,166 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2023 IBM and HuggingFace Inc. team. All Rights Reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ """PatchTSMixer model configuration"""
15
+
16
+ from huggingface_hub.dataclasses import strict
17
+
18
+ from ...configuration_utils import PreTrainedConfig
19
+ from ...utils import auto_docstring
20
+
21
+
22
+ @auto_docstring(checkpoint="ibm/patchtsmixer-etth1-pretrain")
23
+ @strict
24
+ class PatchTSMixerConfig(PreTrainedConfig):
25
+ r"""
26
+ context_length (`int`, *optional*, defaults to 32):
27
+ The context/history length for the input sequence.
28
+ patch_length (`int`, *optional*, defaults to 8):
29
+ The patch length for the input sequence.
30
+ patch_stride (`int`, *optional*, defaults to 8):
31
+ Determines the overlap between two consecutive patches. Set it to patch_length (or greater), if we want
32
+ non-overlapping patches.
33
+ num_parallel_samples (`int`, *optional*, defaults to 100):
34
+ The number of samples to generate in parallel for probabilistic forecast.
35
+ expansion_factor (`int`, *optional*, defaults to 2):
36
+ Expansion factor to use inside MLP. Recommended range is 2-5. Larger value indicates more complex model.
37
+ mode (`str`, *optional*, defaults to `"common_channel"`):
38
+ Mixer Mode. Determines how to process the channels. Allowed values: "common_channel", "mix_channel". In
39
+ "common_channel" mode, we follow Channel-independent modelling with no explicit channel-mixing. Channel
40
+ mixing happens in an implicit manner via shared weights across channels. (preferred first approach) In
41
+ "mix_channel" mode, we follow explicit channel-mixing in addition to patch and feature mixer. (preferred
42
+ approach when channel correlations are very important to model)
43
+ gated_attn (`bool`, *optional*, defaults to `True`):
44
+ Enable Gated Attention.
45
+ norm_mlp (`str`, *optional*, defaults to `"LayerNorm"`):
46
+ Normalization layer (BatchNorm or LayerNorm).
47
+ self_attn (`bool`, *optional*, defaults to `False`):
48
+ Enable Tiny self attention across patches. This can be enabled when the output of Vanilla PatchTSMixer with
49
+ gated attention is not satisfactory. Enabling this leads to explicit pair-wise attention and modelling
50
+ across patches.
51
+ self_attn_heads (`int`, *optional*, defaults to 1):
52
+ Number of self-attention heads. Works only when `self_attn` is set to `True`.
53
+ use_positional_encoding (`bool`, *optional*, defaults to `False`):
54
+ Enable the use of positional embedding for the tiny self-attention layers. Works only when `self_attn` is
55
+ set to `True`.
56
+ positional_encoding_type (`str`, *optional*, defaults to `"sincos"`):
57
+ Positional encodings. Options `"random"` and `"sincos"` are supported. Works only when
58
+ `use_positional_encoding` is set to `True`
59
+ scaling (`string` or `bool`, *optional*, defaults to `"std"`):
60
+ Whether to scale the input targets via "mean" scaler, "std" scaler or no scaler if `None`. If `True`, the
61
+ scaler is set to "mean".
62
+ loss (`string`, *optional*, defaults to `"mse"`):
63
+ The loss function for the model corresponding to the `distribution_output` head. For parametric
64
+ distributions it is the negative log likelihood ("nll") and for point estimates it is the mean squared
65
+ error "mse".
66
+ norm_eps (`float`, *optional*, defaults to 1e-05):
67
+ A value added to the denominator for numerical stability of normalization.
68
+ mask_type (`str`, *optional*, defaults to `"random"`):
69
+ Type of masking to use for Masked Pretraining mode. Allowed values are "random", "forecast". In Random
70
+ masking, points are masked randomly. In Forecast masking, points are masked towards the end.
71
+ random_mask_ratio (`float`, *optional*, defaults to 0.5):
72
+ Masking ratio to use when `mask_type` is `random`. Higher value indicates more masking.
73
+ num_forecast_mask_patches (`int` or `list`, *optional*, defaults to `[2]`):
74
+ Number of patches to be masked at the end of each batch sample. If it is an integer, all the samples in the
75
+ batch will have the same number of masked patches. If it is a list, samples in the batch will be randomly
76
+ masked by numbers defined in the list. This argument is only used for forecast pretraining.
77
+ mask_value (`float`, *optional*, defaults to `0.0`):
78
+ Mask value to use.
79
+ masked_loss (`bool`, *optional*, defaults to `True`):
80
+ Whether to compute pretraining loss only at the masked portions, or on the entire output.
81
+ channel_consistent_masking (`bool`, *optional*, defaults to `True`):
82
+ When true, masking will be same across all channels of a timeseries. Otherwise, masking positions will vary
83
+ across channels.
84
+ unmasked_channel_indices (`list`, *optional*):
85
+ Channels that are not masked during pretraining.
86
+ head_dropout (`float`, *optional*, defaults to 0.2):
87
+ The dropout probability the `PatchTSMixer` head.
88
+ distribution_output (`string`, *optional*, defaults to `"student_t"`):
89
+ The distribution emission head for the model when loss is "nll". Could be either "student_t", "normal" or
90
+ "negative_binomial".
91
+ prediction_length (`int`, *optional*, defaults to 16):
92
+ Number of time steps to forecast for a forecasting task. Also known as the Forecast Horizon.
93
+ prediction_channel_indices (`list`, *optional*):
94
+ List of channel indices to forecast. If None, forecast all channels. Target data is expected to have all
95
+ channels and we explicitly filter the channels in prediction and target before loss computation.
96
+ num_targets (`int`, *optional*, defaults to 3):
97
+ Number of targets (dimensionality of the regressed variable) for a regression task.
98
+ output_range (`list`, *optional*):
99
+ Output range to restrict for the regression task. Defaults to None.
100
+ head_aggregation (`str`, *optional*, defaults to `"max_pool"`):
101
+ Aggregation mode to enable for classification or regression task. Allowed values are `None`, "use_last",
102
+ "max_pool", "avg_pool".
103
+
104
+ Example:
105
+
106
+ ```python
107
+ >>> from transformers import PatchTSMixerConfig, PatchTSMixerModel
108
+
109
+ >>> # Initializing a default PatchTSMixer configuration
110
+ >>> configuration = PatchTSMixerConfig()
111
+
112
+ >>> # Randomly initializing a model (with random weights) from the configuration
113
+ >>> model = PatchTSMixerModel(configuration)
114
+
115
+ >>> # Accessing the model configuration
116
+ >>> configuration = model.config
117
+ ```"""
118
+
119
+ model_type = "patchtsmixer"
120
+ attribute_map = {
121
+ "hidden_size": "d_model",
122
+ "num_hidden_layers": "num_layers",
123
+ }
124
+
125
+ context_length: int = 32
126
+ patch_length: int = 8
127
+ num_input_channels: int = 1
128
+ patch_stride: int = 8
129
+ num_parallel_samples: int = 100
130
+ d_model: int = 8
131
+ expansion_factor: int = 2
132
+ num_layers: int = 3
133
+ dropout: float | int = 0.2
134
+ mode: str = "common_channel"
135
+ gated_attn: bool = True
136
+ norm_mlp: str = "LayerNorm"
137
+ self_attn: bool = False
138
+ self_attn_heads: int = 1
139
+ use_positional_encoding: bool = False
140
+ positional_encoding_type: str = "sincos"
141
+ scaling: str | bool | None = "std"
142
+ loss: str = "mse"
143
+ init_std: float = 0.02
144
+ norm_eps: float = 1e-5
145
+ mask_type: str = "random"
146
+ random_mask_ratio: float = 0.5
147
+ num_forecast_mask_patches: list[int] | tuple[int, ...] | int | None = (2,)
148
+ mask_value: int = 0
149
+ masked_loss: bool = True
150
+ channel_consistent_masking: bool = True
151
+ unmasked_channel_indices: list[int] | None = None
152
+ head_dropout: float | int = 0.2
153
+ distribution_output: str = "student_t"
154
+ prediction_length: int = 16
155
+ prediction_channel_indices: list | None = None
156
+ num_targets: int = 3
157
+ output_range: list | None = None
158
+ head_aggregation: str | None = "max_pool"
159
+
160
+ def __post_init__(self, **kwargs):
161
+ self.num_patches = (max(self.context_length, self.patch_length) - self.patch_length) // self.patch_stride + 1
162
+ self.patch_last = True
163
+ super().__post_init__(**kwargs)
164
+
165
+
166
+ __all__ = ["PatchTSMixerConfig"]
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/modeling_patchtsmixer.py ADDED
@@ -0,0 +1,2121 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2023 IBM and HuggingFace Inc. team. All Rights Reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ """PyTorch PatchTSMixer model."""
15
+
16
+ import math
17
+ from collections.abc import Callable
18
+ from dataclasses import dataclass
19
+
20
+ import torch
21
+ import torch.nn as nn
22
+
23
+ from transformers.modeling_utils import PreTrainedModel
24
+ from transformers.utils import ModelOutput
25
+
26
+ from ... import initialization as init
27
+ from ...modeling_flash_attention_utils import FlashAttentionKwargs
28
+ from ...modeling_utils import ALL_ATTENTION_FUNCTIONS
29
+ from ...processing_utils import Unpack
30
+ from ...time_series_utils import NegativeBinomialOutput, NormalOutput, StudentTOutput
31
+ from ...utils import TransformersKwargs, auto_docstring, logging
32
+ from .configuration_patchtsmixer import PatchTSMixerConfig
33
+
34
+
35
+ logger = logging.get_logger(__name__)
36
+
37
+
38
+ class PatchTSMixerGatedAttention(nn.Module):
39
+ """
40
+ Module that applies gated attention to input data.
41
+
42
+ Args:
43
+ in_size (`int`): The input size.
44
+ out_size (`int`): The output size.
45
+ """
46
+
47
+ def __init__(self, in_size: int, out_size: int):
48
+ super().__init__()
49
+ self.attn_layer = nn.Linear(in_size, out_size)
50
+ self.attn_softmax = nn.Softmax(dim=-1)
51
+
52
+ def forward(self, inputs):
53
+ attn_weight = self.attn_softmax(self.attn_layer(inputs))
54
+ inputs = inputs * attn_weight
55
+ return inputs
56
+
57
+
58
+ # Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTBatchNorm with PatchTST->PatchTSMixer
59
+ class PatchTSMixerBatchNorm(nn.Module):
60
+ """
61
+ Compute batch normalization over the sequence length (time) dimension.
62
+ """
63
+
64
+ def __init__(self, config: PatchTSMixerConfig):
65
+ super().__init__()
66
+ self.batchnorm = nn.BatchNorm1d(config.d_model, eps=config.norm_eps)
67
+
68
+ def forward(self, inputs: torch.Tensor):
69
+ """
70
+ Parameters:
71
+ inputs (`torch.Tensor` of shape `(batch_size, sequence_length, d_model)`):
72
+ input for Batch norm calculation
73
+ Returns:
74
+ `torch.Tensor` of shape `(batch_size, sequence_length, d_model)`
75
+ """
76
+ output = inputs.transpose(1, 2) # output: (batch_size, d_model, sequence_length)
77
+ output = self.batchnorm(output)
78
+ return output.transpose(1, 2)
79
+
80
+
81
+ class PatchTSMixerPositionalEncoding(nn.Module):
82
+ """
83
+ Class for positional encoding
84
+ """
85
+
86
+ def __init__(self, config: PatchTSMixerConfig):
87
+ super().__init__()
88
+ # positional encoding: [num_patches x d_model]
89
+ if config.use_positional_encoding:
90
+ self.position_enc = self._init_pe(config)
91
+ else:
92
+ self.position_enc = nn.Parameter(torch.zeros(config.num_patches, config.d_model))
93
+
94
+ @staticmethod
95
+ def _init_pe(config: PatchTSMixerConfig) -> nn.Parameter:
96
+ # Positional encoding
97
+ if config.positional_encoding_type == "random":
98
+ position_enc = nn.Parameter(torch.randn(config.num_patches, config.d_model), requires_grad=True)
99
+ elif config.positional_encoding_type == "sincos":
100
+ position_enc = torch.zeros(config.num_patches, config.d_model)
101
+ position = torch.arange(0, config.num_patches).unsqueeze(1)
102
+ div_term = torch.exp(torch.arange(0, config.d_model, 2) * -(math.log(10000.0) / config.d_model))
103
+ position_enc[:, 0::2] = torch.sin(position * div_term)
104
+ position_enc[:, 1::2] = torch.cos(position * div_term)
105
+ position_enc = position_enc - position_enc.mean()
106
+ position_enc = position_enc / (position_enc.std() * 10)
107
+ position_enc = nn.Parameter(position_enc, requires_grad=False)
108
+ else:
109
+ raise ValueError(
110
+ f"{config.positional_encoding_type} is not a valid positional encoder. Available types are 'random' and 'sincos'."
111
+ )
112
+ return position_enc
113
+
114
+ def forward(self, patch_input: torch.Tensor):
115
+ # hidden_state: [bs x num_channels x num_patches x d_model]
116
+ hidden_state = patch_input + self.position_enc
117
+ return hidden_state
118
+
119
+
120
+ class PatchTSMixerNormLayer(nn.Module):
121
+ """Normalization block
122
+
123
+ Args:
124
+ config (`PatchTSMixerConfig`):
125
+ Configuration.
126
+ """
127
+
128
+ def __init__(self, config: PatchTSMixerConfig):
129
+ super().__init__()
130
+
131
+ self.norm_mlp = config.norm_mlp
132
+
133
+ if "batch" in config.norm_mlp.lower():
134
+ self.norm = PatchTSMixerBatchNorm(config)
135
+ else:
136
+ self.norm = nn.LayerNorm(config.d_model, eps=config.norm_eps)
137
+
138
+ def forward(self, inputs: torch.Tensor):
139
+ """
140
+ Args:
141
+ inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`):
142
+ Input to the normalization layer.
143
+ Returns:
144
+ `torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`
145
+ """
146
+ if "batch" in self.norm_mlp.lower():
147
+ # reshape the data
148
+ inputs_reshaped = torch.reshape(
149
+ inputs,
150
+ (
151
+ inputs.shape[0] * inputs.shape[1],
152
+ inputs.shape[2],
153
+ inputs.shape[3],
154
+ ),
155
+ ) # inputs_reshaped: [batch_size*num_channels, num_patches, d_model]
156
+
157
+ # inputs_reshaped: [batch_size*num_channels, num_patches, d_model]
158
+ inputs_reshaped = self.norm(inputs_reshaped)
159
+
160
+ # put back data to the original shape
161
+ inputs = torch.reshape(inputs_reshaped, inputs.shape)
162
+
163
+ else:
164
+ inputs = self.norm(inputs)
165
+
166
+ return inputs
167
+
168
+
169
+ class PatchTSMixerMLP(nn.Module):
170
+ def __init__(self, in_features, out_features, config):
171
+ super().__init__()
172
+ num_hidden = in_features * config.expansion_factor
173
+ self.fc1 = nn.Linear(in_features, num_hidden)
174
+ self.dropout1 = nn.Dropout(config.dropout)
175
+ self.fc2 = nn.Linear(num_hidden, out_features)
176
+ self.dropout2 = nn.Dropout(config.dropout)
177
+
178
+ def forward(self, inputs: torch.Tensor):
179
+ """
180
+ Args:
181
+ inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`):
182
+ Input to the MLP layer.
183
+ Returns:
184
+ `torch.Tensor` of the same shape as `inputs`
185
+ """
186
+ inputs = self.dropout1(nn.functional.gelu(self.fc1(inputs)))
187
+ inputs = self.fc2(inputs)
188
+ inputs = self.dropout2(inputs)
189
+ return inputs
190
+
191
+
192
+ class PatchTSMixerChannelFeatureMixerBlock(nn.Module):
193
+ """This module mixes the features in the channel dimension.
194
+
195
+ Args:
196
+ config (`PatchTSMixerConfig`):
197
+ Configuration.
198
+ """
199
+
200
+ def __init__(self, config: PatchTSMixerConfig):
201
+ super().__init__()
202
+
203
+ self.norm = PatchTSMixerNormLayer(config)
204
+ self.gated_attn = config.gated_attn
205
+ self.mlp = PatchTSMixerMLP(
206
+ in_features=config.num_input_channels,
207
+ out_features=config.num_input_channels,
208
+ config=config,
209
+ )
210
+
211
+ if config.gated_attn:
212
+ self.gating_block = PatchTSMixerGatedAttention(
213
+ in_size=config.num_input_channels, out_size=config.num_input_channels
214
+ )
215
+
216
+ def forward(self, inputs: torch.Tensor):
217
+ """
218
+ Args:
219
+ inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`):
220
+ input to the MLP layer
221
+ Returns:
222
+ `torch.Tensor` of the same shape as `inputs`
223
+ """
224
+ residual = inputs
225
+ inputs = self.norm(inputs)
226
+
227
+ inputs = inputs.permute(0, 3, 2, 1)
228
+
229
+ if self.gated_attn:
230
+ inputs = self.gating_block(inputs)
231
+
232
+ inputs = self.mlp(inputs)
233
+
234
+ inputs = inputs.permute(0, 3, 2, 1)
235
+
236
+ out = inputs + residual
237
+ return out
238
+
239
+
240
+ # Copied from transformers.models.bert.modeling_bert.eager_attention_forward
241
+ def eager_attention_forward(
242
+ module: nn.Module,
243
+ query: torch.Tensor,
244
+ key: torch.Tensor,
245
+ value: torch.Tensor,
246
+ attention_mask: torch.Tensor | None,
247
+ scaling: float | None = None,
248
+ dropout: float = 0.0,
249
+ **kwargs: Unpack[TransformersKwargs],
250
+ ):
251
+ if scaling is None:
252
+ scaling = query.size(-1) ** -0.5
253
+
254
+ # Take the dot product between "query" and "key" to get the raw attention scores.
255
+ attn_weights = torch.matmul(query, key.transpose(2, 3)) * scaling
256
+
257
+ if attention_mask is not None:
258
+ attn_weights = attn_weights + attention_mask
259
+
260
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1)
261
+ attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)
262
+
263
+ attn_output = torch.matmul(attn_weights, value)
264
+ attn_output = attn_output.transpose(1, 2).contiguous()
265
+
266
+ return attn_output, attn_weights
267
+
268
+
269
+ # Copied from transformers.models.wav2vec2.modeling_wav2vec2.Wav2Vec2Attention with Wav2Vec2->PatchTSMixer
270
+ class PatchTSMixerAttention(nn.Module):
271
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
272
+
273
+ def __init__(
274
+ self,
275
+ embed_dim: int,
276
+ num_heads: int,
277
+ dropout: float = 0.0,
278
+ is_decoder: bool = False,
279
+ bias: bool = True,
280
+ is_causal: bool = False,
281
+ config: PatchTSMixerConfig | None = None,
282
+ ):
283
+ super().__init__()
284
+ self.embed_dim = embed_dim
285
+ self.num_heads = num_heads
286
+ self.dropout = dropout
287
+ self.head_dim = embed_dim // num_heads
288
+ self.config = config
289
+
290
+ if (self.head_dim * num_heads) != self.embed_dim:
291
+ raise ValueError(
292
+ f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim}"
293
+ f" and `num_heads`: {num_heads})."
294
+ )
295
+ self.scaling = self.head_dim**-0.5
296
+ self.is_decoder = is_decoder
297
+ self.is_causal = is_causal
298
+
299
+ self.k_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
300
+ self.v_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
301
+ self.q_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
302
+ self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
303
+
304
+ def forward(
305
+ self,
306
+ hidden_states: torch.Tensor,
307
+ key_value_states: torch.Tensor | None = None,
308
+ attention_mask: torch.Tensor | None = None,
309
+ output_attentions: bool | None = False,
310
+ # TODO: we need a refactor so that the different attention modules can get their specific kwargs
311
+ # ATM, we have mixed things encoder, decoder, and encoder-decoder attn
312
+ **kwargs: Unpack[FlashAttentionKwargs],
313
+ ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]:
314
+ """Input shape: Batch x Time x Channel"""
315
+
316
+ # if key_value_states are provided this layer is used as a cross-attention layer
317
+ # for the decoder
318
+ is_cross_attention = key_value_states is not None
319
+
320
+ # determine input shapes
321
+ input_shape = hidden_states.shape[:-1]
322
+
323
+ hidden_shape = (*input_shape, -1, self.head_dim)
324
+
325
+ # get query proj
326
+ query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
327
+
328
+ current_states = key_value_states if is_cross_attention else hidden_states
329
+ kv_shape = (*current_states.shape[:-1], -1, self.head_dim)
330
+ key_states = self.k_proj(current_states).view(kv_shape).transpose(1, 2)
331
+ value_states = self.v_proj(current_states).view(kv_shape).transpose(1, 2)
332
+
333
+ attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
334
+ self.config._attn_implementation, eager_attention_forward
335
+ )
336
+
337
+ attn_output, attn_weights = attention_interface(
338
+ self,
339
+ query_states,
340
+ key_states,
341
+ value_states,
342
+ attention_mask,
343
+ dropout=0.0 if not self.training else self.dropout,
344
+ scaling=self.scaling,
345
+ output_attentions=output_attentions,
346
+ **kwargs,
347
+ )
348
+
349
+ attn_output = attn_output.reshape(*input_shape, -1).contiguous()
350
+ attn_output = self.out_proj(attn_output)
351
+
352
+ return attn_output, attn_weights, None
353
+
354
+
355
+ class PatchMixerBlock(nn.Module):
356
+ """This module mixes the patch dimension.
357
+
358
+ Args:
359
+ config (`PatchTSMixerConfig`):
360
+ Configuration.
361
+ """
362
+
363
+ def __init__(self, config: PatchTSMixerConfig):
364
+ super().__init__()
365
+
366
+ self.norm = PatchTSMixerNormLayer(config)
367
+
368
+ self.self_attn = config.self_attn
369
+ self.gated_attn = config.gated_attn
370
+
371
+ self.mlp = PatchTSMixerMLP(
372
+ in_features=config.num_patches,
373
+ out_features=config.num_patches,
374
+ config=config,
375
+ )
376
+
377
+ if config.gated_attn:
378
+ self.gating_block = PatchTSMixerGatedAttention(in_size=config.num_patches, out_size=config.num_patches)
379
+
380
+ if config.self_attn:
381
+ self.self_attn_layer = PatchTSMixerAttention(
382
+ embed_dim=config.d_model,
383
+ num_heads=config.self_attn_heads,
384
+ dropout=config.dropout,
385
+ config=config,
386
+ )
387
+ self.norm_attn = PatchTSMixerNormLayer(config)
388
+
389
+ def forward(self, hidden_state):
390
+ """
391
+ Args:
392
+ hidden_state (`torch.Tensor`): Input tensor.
393
+
394
+ Returns:
395
+ `torch.Tensor`: Transformed tensor.
396
+ """
397
+ residual = hidden_state
398
+
399
+ hidden_state = self.norm(hidden_state)
400
+
401
+ if self.self_attn:
402
+ batch_size, n_vars, num_patches, d_model = hidden_state.shape
403
+ hidden_state_reshaped = hidden_state.reshape(batch_size * n_vars, num_patches, d_model)
404
+
405
+ x_attn, _, _ = self.self_attn_layer(hidden_state_reshaped, output_attentions=False)
406
+ x_attn = x_attn.reshape(batch_size, n_vars, num_patches, d_model)
407
+
408
+ # Transpose so that num_patches is the last dimension
409
+ hidden_state = hidden_state.transpose(2, 3)
410
+ hidden_state = self.mlp(hidden_state)
411
+
412
+ if self.gated_attn:
413
+ hidden_state = self.gating_block(hidden_state)
414
+
415
+ # Transpose back
416
+ hidden_state = hidden_state.transpose(2, 3)
417
+
418
+ if self.self_attn:
419
+ hidden_state = self.norm_attn(hidden_state + x_attn)
420
+
421
+ out = hidden_state + residual
422
+ return out
423
+
424
+
425
+ class FeatureMixerBlock(nn.Module):
426
+ """This module mixes the hidden feature dimension.
427
+
428
+ Args:
429
+ config (`PatchTSMixerConfig`):
430
+ Configuration.
431
+
432
+ """
433
+
434
+ def __init__(self, config: PatchTSMixerConfig):
435
+ super().__init__()
436
+
437
+ self.norm = PatchTSMixerNormLayer(config)
438
+
439
+ self.gated_attn = config.gated_attn
440
+
441
+ self.mlp = PatchTSMixerMLP(
442
+ in_features=config.d_model,
443
+ out_features=config.d_model,
444
+ config=config,
445
+ )
446
+
447
+ if config.gated_attn:
448
+ self.gating_block = PatchTSMixerGatedAttention(in_size=config.d_model, out_size=config.d_model)
449
+
450
+ def forward(self, hidden: torch.Tensor):
451
+ """
452
+ Args:
453
+ hidden (`torch.Tensor` of shape `(batch_size, num_patches, d_model)`):
454
+ Input tensor to the layer.
455
+
456
+ Returns:
457
+ `torch.Tensor`: Transformed tensor.
458
+ """
459
+ residual = hidden
460
+ hidden = self.norm(hidden)
461
+ hidden = self.mlp(hidden)
462
+
463
+ if self.gated_attn:
464
+ hidden = self.gating_block(hidden)
465
+
466
+ out = hidden + residual
467
+ return out
468
+
469
+
470
+ class PatchTSMixerLayer(nn.Module):
471
+ """
472
+ The `PatchTSMixer` layer that does all three kinds of mixing.
473
+
474
+ Args:
475
+ config (`PatchTSMixerConfig`):
476
+ Configuration.
477
+
478
+ """
479
+
480
+ def __init__(self, config: PatchTSMixerConfig):
481
+ super().__init__()
482
+
483
+ self.patch_mixer = PatchMixerBlock(config=config)
484
+ self.feature_mixer = FeatureMixerBlock(config=config)
485
+
486
+ self.mode = config.mode
487
+
488
+ if config.mode == "mix_channel":
489
+ self.channel_feature_mixer = PatchTSMixerChannelFeatureMixerBlock(config=config)
490
+
491
+ def forward(self, hidden: torch.Tensor):
492
+ """
493
+ Args:
494
+ hidden (`torch.Tensor` of shape `(batch_size, num_patches, d_model)`):
495
+ Input tensor to the layer.
496
+
497
+ Returns:
498
+ `torch.Tensor`: Transformed tensor.
499
+ """
500
+ if self.mode == "mix_channel":
501
+ hidden = self.channel_feature_mixer(hidden)
502
+
503
+ hidden = self.patch_mixer(hidden)
504
+ hidden = self.feature_mixer(hidden) # hidden: (batch_size x num_patches x d_model)
505
+ return hidden
506
+
507
+
508
+ class PatchTSMixerBlock(nn.Module):
509
+ """The main computing framework of the `PatchTSMixer` model.
510
+
511
+ Args:
512
+ config (`PatchTSMixerConfig`):
513
+ Configuration.
514
+ """
515
+
516
+ def __init__(self, config: PatchTSMixerConfig):
517
+ super().__init__()
518
+
519
+ num_layers = config.num_layers
520
+
521
+ self.mixers = nn.ModuleList([PatchTSMixerLayer(config=config) for _ in range(num_layers)])
522
+
523
+ def forward(self, hidden_state, output_hidden_states: bool = False):
524
+ """
525
+ Args:
526
+ hidden_state (`torch.Tensor`): The input tensor.
527
+ output_hidden_states (`bool`, *optional*, defaults to False.):
528
+ Whether to output the hidden states as well.
529
+
530
+ Returns:
531
+ `torch.Tensor`: The embedding. `list`: List of all hidden states if `output_hidden_states` is set to
532
+ `True`.
533
+ """
534
+ all_hidden_states = []
535
+
536
+ embedding = hidden_state
537
+
538
+ for mod in self.mixers:
539
+ embedding = mod(embedding)
540
+ if output_hidden_states:
541
+ all_hidden_states.append(embedding)
542
+
543
+ if output_hidden_states:
544
+ return embedding, all_hidden_states
545
+ else:
546
+ return embedding, None
547
+
548
+
549
+ class PatchTSMixerForPredictionHead(nn.Module):
550
+ """Prediction Head for Forecasting
551
+
552
+ Args:
553
+ config (`PatchTSMixerConfig`):
554
+ Configuration.
555
+ """
556
+
557
+ def __init__(self, config: PatchTSMixerConfig, distribution_output=None):
558
+ super().__init__()
559
+
560
+ self.prediction_channel_indices = config.prediction_channel_indices
561
+
562
+ if self.prediction_channel_indices is not None:
563
+ self.prediction_channel_indices.sort()
564
+
565
+ self.dropout_layer = nn.Dropout(config.head_dropout)
566
+ if distribution_output is None:
567
+ self.base_forecast_block = nn.Linear((config.num_patches * config.d_model), config.prediction_length)
568
+ else:
569
+ self.base_forecast_block = distribution_output.get_parameter_projection(
570
+ config.num_patches * config.d_model
571
+ )
572
+
573
+ self.flatten = nn.Flatten(start_dim=-2)
574
+
575
+ def forward(self, hidden_features):
576
+ """
577
+
578
+ Args:
579
+ hidden_features (`torch.Tensor` of shape `(batch_size, num_patch, d_model)` in `flatten` mode
580
+ or `(batch_size, n_vars, num_patch, d_model)` in `common_channel`/`mix_channel` mode.): Input hidden
581
+ features.
582
+
583
+ Returns:
584
+ `torch.Tensor` of shape `(batch_size, prediction_length, nvars)`.
585
+
586
+ """
587
+
588
+ hidden_features = self.flatten(hidden_features) # [batch_size x n_vars x num_patch * d_model]
589
+ hidden_features = self.dropout_layer(hidden_features) # [batch_size x n_vars x num_patch * d_model]
590
+ forecast = self.base_forecast_block(hidden_features) # [batch_size x n_vars x prediction_length]
591
+ if isinstance(forecast, tuple):
592
+ forecast = tuple(z.transpose(-1, -2) for z in forecast)
593
+ else:
594
+ forecast = forecast.transpose(-1, -2) # [batch_size x prediction_length x n_vars]
595
+
596
+ if self.prediction_channel_indices is not None:
597
+ if isinstance(forecast, tuple):
598
+ forecast = tuple(z[..., self.prediction_channel_indices] for z in forecast)
599
+ else:
600
+ forecast = forecast[..., self.prediction_channel_indices] # [batch_size x prediction_length x n_vars]
601
+
602
+ return forecast
603
+
604
+
605
+ class PatchTSMixerLinearHead(nn.Module):
606
+ """Linear head for Classification and Regression.
607
+
608
+ Args:
609
+ config (`PatchTSMixerConfig`):
610
+ Configuration.
611
+ """
612
+
613
+ def __init__(self, config: PatchTSMixerConfig, distribution_output=None):
614
+ super().__init__()
615
+
616
+ self.head_aggregation = config.head_aggregation
617
+ self.output_range = config.output_range
618
+
619
+ if config.head_aggregation is None:
620
+ mul_factor = config.num_patches
621
+ else:
622
+ mul_factor = 1
623
+ self.distribution_output = distribution_output
624
+ if distribution_output is None:
625
+ self.projection = nn.Linear(
626
+ config.d_model * config.num_input_channels * mul_factor,
627
+ config.num_targets,
628
+ )
629
+ else:
630
+ self.projection = distribution_output.get_parameter_projection(
631
+ config.d_model * config.num_input_channels * mul_factor
632
+ )
633
+
634
+ if config.head_aggregation is None:
635
+ self.flatten = nn.Flatten(start_dim=-3)
636
+ else:
637
+ self.flatten = nn.Flatten(start_dim=-2)
638
+
639
+ self.dropout = nn.Dropout(config.head_dropout)
640
+
641
+ def forward(self, hidden_features):
642
+ """
643
+ Args:
644
+ hidden_features (`torch.Tensor` of shape `(batch_size x num_patch x d_model)` in `flatten` mode
645
+ or `(batch_size x n_vars x num_patch x d_model)` in `common_channel`/`mix_channel` mode.): Input hidden
646
+ features.
647
+
648
+ Returns:
649
+ `torch.Tensor` of shape `(batch_size x num_targets)`.
650
+ """
651
+
652
+ # batch_size x d_model x num_patch or batch_size x n_vars x d_model x num_patch
653
+ hidden_features = hidden_features.transpose(-1, -2)
654
+ if self.head_aggregation == "use_last":
655
+ # batch_size x d_model (flatten) or # batch_size x n_vars x d_model (common_channel)
656
+ hidden_features = hidden_features[..., -1]
657
+ elif self.head_aggregation == "max_pool":
658
+ # batch_size x n_vars x d_model or batch_size x d_model
659
+ hidden_features = hidden_features.max(dim=-1).values
660
+ elif self.head_aggregation == "avg_pool":
661
+ # batch_size x n_vars x d_model or batch_size x d_model
662
+ hidden_features = hidden_features.mean(dim=-1)
663
+
664
+ if self.flatten:
665
+ hidden_features = self.flatten(hidden_features)
666
+ hidden_features = self.dropout(hidden_features)
667
+ hidden_features = self.projection(hidden_features) # batch_size x num_targets
668
+
669
+ if (self.distribution_output is None) and (self.output_range is not None):
670
+ hidden_features = (
671
+ torch.sigmoid(hidden_features) * (self.output_range[1] - self.output_range[0]) + self.output_range[0]
672
+ )
673
+ return hidden_features
674
+
675
+
676
+ @auto_docstring
677
+ class PatchTSMixerPreTrainedModel(PreTrainedModel):
678
+ # Weight initialization
679
+ config: PatchTSMixerConfig
680
+ base_model_prefix = "model"
681
+ main_input_name = "past_values"
682
+ input_modalities = ("time",)
683
+ supports_gradient_checkpointing = False
684
+
685
+ @torch.no_grad()
686
+ def _init_weights(self, module):
687
+ """Initialize weights"""
688
+ if isinstance(module, PatchTSMixerPositionalEncoding):
689
+ # initialize positional encoding
690
+ if self.config.positional_encoding_type == "random":
691
+ init.normal_(module.position_enc, mean=0.0, std=0.1)
692
+ elif isinstance(module, (nn.LayerNorm, nn.BatchNorm1d)):
693
+ init.zeros_(module.bias)
694
+ init.ones_(module.weight)
695
+ if getattr(module, "running_mean", None) is not None:
696
+ init.zeros_(module.running_mean)
697
+ init.ones_(module.running_var)
698
+ init.zeros_(module.num_batches_tracked)
699
+ elif isinstance(module, PatchTSMixerBatchNorm):
700
+ init.zeros_(module.batchnorm.bias)
701
+ init.ones_(module.batchnorm.weight)
702
+ elif isinstance(module, nn.Linear):
703
+ init.normal_(module.weight, mean=0.0, std=self.config.init_std)
704
+ if module.bias is not None:
705
+ init.zeros_(module.bias)
706
+
707
+
708
+ class PatchTSMixerPretrainHead(nn.Module):
709
+ """Pretraining head.
710
+
711
+ Args:
712
+ config (`PatchTSMixerConfig`):
713
+ Configuration.
714
+ """
715
+
716
+ def __init__(self, config: PatchTSMixerConfig):
717
+ super().__init__()
718
+
719
+ self.dropout_layer = nn.Dropout(config.head_dropout)
720
+ self.base_pt_block = nn.Linear(config.d_model, config.patch_length)
721
+
722
+ def forward(self, hidden_features):
723
+ """
724
+ Args:
725
+ hidden_features (`torch.Tensor` of shape `(batch_size x num_patch x d_model)` in `flatten` mode
726
+ or `(batch_size x n_vars x num_patch x d_model)` in `common_channel`/`mix_channel` mode.): Input hidden
727
+ features.
728
+
729
+ Returns:
730
+ `torch.Tensor` of shape `(batch_size x n_vars x num_patch x patch_length)`.
731
+ """
732
+
733
+ hidden_features = self.dropout_layer(hidden_features)
734
+ forecast = self.base_pt_block(hidden_features) # [batch_size x n_vars x num_patch x patch_length]
735
+ return forecast
736
+
737
+
738
+ # Copied from transformers.models.patchtst.modeling_patchtst.random_masking
739
+ def random_masking(
740
+ inputs: torch.Tensor,
741
+ mask_ratio: float,
742
+ unmasked_channel_indices: list | None = None,
743
+ channel_consistent_masking: bool = False,
744
+ mask_value: int = 0,
745
+ ):
746
+ """random_masking: Mask the input considering the control variables.
747
+
748
+ Args:
749
+ inputs (`torch.Tensor` of shape `(batch_size, num_channels, sequence_length, num_features)`):
750
+ The input tensor to mask.
751
+ mask_ratio (`float`):
752
+ Masking ratio applied to mask the input data during random pretraining. It is the number between 0 and 1.
753
+ unmasked_channel_indices (list, *optional*):
754
+ Indices of channels that will not be masked.
755
+ channel_consistent_masking (bool, *optional*, defaults to `False`):
756
+ When true, masking will be same across all channels of a timeseries. Otherwise, masking positions will vary
757
+ across channels.
758
+ mask_value (int, *optional*, defaults to 0):
759
+ Define the value of masked patches for pretraining.
760
+
761
+ Returns:
762
+ `tuple(torch.Tensor)`: inputs_mask, masked input, same shape as input Tensor and mask tensor of shape [bs x c x
763
+ n]
764
+ """
765
+ if mask_ratio < 0 or mask_ratio >= 1:
766
+ raise ValueError(f"Mask ratio {mask_ratio} has to be between 0 and 1.")
767
+
768
+ batch_size, num_channels, sequence_length, num_features = inputs.shape
769
+ device = inputs.device
770
+
771
+ len_keep = int(sequence_length * (1 - mask_ratio))
772
+
773
+ if channel_consistent_masking:
774
+ noise = torch.rand(batch_size, 1, sequence_length, device=device) # noise in [0, 1], bs x 1 x L
775
+ noise = noise.repeat(1, num_channels, 1) # bs x num_channels x time
776
+ else:
777
+ # noise in [0, 1], bs x num_channels x L
778
+ noise = torch.rand(batch_size, num_channels, sequence_length, device=device)
779
+
780
+ # mask: [bs x num_channels x num_patch]
781
+ mask = torch.ones(batch_size, num_channels, sequence_length, device=device)
782
+ mask[:, :, :len_keep] = 0
783
+
784
+ # sort noise for each sample
785
+ ids_shuffle = torch.argsort(noise, dim=-1) # ascend: small is keep, large is remove
786
+ ids_restore = torch.argsort(ids_shuffle, dim=-1) # ids_restore: [bs x num_channels x L]
787
+
788
+ mask = torch.gather(mask, dim=-1, index=ids_restore)
789
+ mask = mask.unsqueeze(-1).repeat(1, 1, 1, num_features) # mask: [bs x num_channels x num_patches x patch_length]
790
+ if unmasked_channel_indices is not None:
791
+ mask[:, unmasked_channel_indices, :, :] = 0
792
+
793
+ inputs_mask = inputs.masked_fill(mask.bool(), mask_value)
794
+ return inputs_mask, mask[..., 0]
795
+
796
+
797
+ # Copied from transformers.models.patchtst.modeling_patchtst.forecast_masking
798
+ def forecast_masking(
799
+ inputs: torch.Tensor,
800
+ num_forecast_mask_patches: list | int,
801
+ unmasked_channel_indices: list | None = None,
802
+ mask_value: int = 0,
803
+ ):
804
+ """Forecast masking that masks the last K patches where K is from the num_forecast_mask_patches.
805
+ If num_forecast_mask_patches is a list, samples in the batch will be randomly masked by numbers defined in the list.
806
+
807
+ Parameters:
808
+ inputs (`torch.Tensor`):
809
+ Input of shape `(bs, num_channels, num_patch, patch_length)`
810
+ num_forecast_mask_patches (`list`):
811
+ Number of patches to be masked at the end of each batch sample. e.g. 4 or [3, 5].
812
+ unmasked_channel_indices (`list`, *optional*):
813
+ Indices of channels that are not masked.
814
+ mask_value (`int`, *optional*, defaults to 0):
815
+ Values in the masked patches will be filled by `mask_value`.
816
+
817
+ Returns:
818
+ `tuple(torch.Tensor)`: inputs_mask, masked input, same shape as inputs Tensor and Mask tensor of shape `(bs,
819
+ num_channels , num_patch)` or `(bs, tsg1, tsg2, num_channels, num_patch)`
820
+ """
821
+
822
+ if isinstance(num_forecast_mask_patches, int):
823
+ num_forecast_mask_patches = [num_forecast_mask_patches]
824
+ forecast_mask_ratios = [1 for _ in num_forecast_mask_patches]
825
+
826
+ batch_size, num_channels, sequence_length, num_features = inputs.shape
827
+ mask = torch.zeros(batch_size, num_channels, sequence_length, device=inputs.device)
828
+
829
+ t_list = []
830
+ total_length = 0
831
+ total_ratio = sum(forecast_mask_ratios)
832
+
833
+ for patch_length, ratio in zip(num_forecast_mask_patches, forecast_mask_ratios):
834
+ if patch_length <= 0 or patch_length >= sequence_length:
835
+ raise ValueError(
836
+ f"num_forecast_mask_patches {patch_length} should be greater than 0 and less than total patches."
837
+ )
838
+ temp_len = int(batch_size * ratio / total_ratio)
839
+ t_list.append([patch_length, ratio, temp_len])
840
+ total_length += temp_len
841
+
842
+ t_list = sorted(t_list, key=lambda x: x[2])
843
+
844
+ if total_length < batch_size:
845
+ t_list[0][2] = t_list[0][2] + (batch_size - total_length)
846
+ elif total_length > batch_size:
847
+ t_list[-1][2] = t_list[-1][2] + (total_length - batch_size)
848
+
849
+ batch1 = 0
850
+ for patch_len, _, temp_len in t_list:
851
+ batch2 = batch1 + temp_len
852
+ mask[batch1:batch2, :, -patch_len:] = 1
853
+ batch1 = batch2
854
+
855
+ perm = torch.randperm(mask.shape[0])
856
+ mask = mask[perm]
857
+
858
+ mask = mask.unsqueeze(-1).repeat(1, 1, 1, num_features) # mask: [bs x num_channels x num_patch x patch_len]
859
+ if unmasked_channel_indices is not None:
860
+ mask[:, unmasked_channel_indices, :, :] = 0
861
+
862
+ inputs_mask = inputs.masked_fill(mask.bool(), mask_value)
863
+ return inputs_mask, mask[..., 0]
864
+
865
+
866
+ # Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTPatchify with PatchTST->PatchTSMixer
867
+ class PatchTSMixerPatchify(nn.Module):
868
+ """
869
+ A class to patchify the time series sequence into different patches
870
+
871
+ Returns:
872
+ `torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`
873
+ """
874
+
875
+ def __init__(self, config: PatchTSMixerConfig):
876
+ super().__init__()
877
+
878
+ self.sequence_length = config.context_length
879
+ self.patch_length = config.patch_length
880
+ self.patch_stride = config.patch_stride
881
+
882
+ if self.sequence_length <= self.patch_length:
883
+ raise ValueError(
884
+ f"Sequence length ({self.sequence_length}) has to be greater than the patch length ({self.patch_length})"
885
+ )
886
+
887
+ # get the number of patches
888
+ self.num_patches = (max(self.sequence_length, self.patch_length) - self.patch_length) // self.patch_stride + 1
889
+ new_sequence_length = self.patch_length + self.patch_stride * (self.num_patches - 1)
890
+ self.sequence_start = self.sequence_length - new_sequence_length
891
+
892
+ def forward(self, past_values: torch.Tensor):
893
+ """
894
+ Parameters:
895
+ past_values (`torch.Tensor` of shape `(batch_size, sequence_length, num_channels)`, *required*):
896
+ Input for patchification
897
+
898
+ Returns:
899
+ `torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`
900
+ """
901
+ sequence_length = past_values.shape[-2]
902
+ if sequence_length != self.sequence_length:
903
+ raise ValueError(
904
+ f"Input sequence length ({sequence_length}) doesn't match model configuration ({self.sequence_length})."
905
+ )
906
+ # output: [bs x new_sequence_length x num_channels]
907
+ output = past_values[:, self.sequence_start :, :]
908
+ # output: [bs x num_patches x num_input_channels x patch_length]
909
+ output = output.unfold(dimension=-2, size=self.patch_length, step=self.patch_stride)
910
+ # output: [bs x num_input_channels x num_patches x patch_length]
911
+ output = output.transpose(-2, -3).contiguous()
912
+ return output
913
+
914
+
915
+ # Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTMasking with PatchTST->PatchTSMixer
916
+ class PatchTSMixerMasking(nn.Module):
917
+ """
918
+ Class to perform random or forecast masking.
919
+
920
+ Parameters:
921
+ config (`PatchTSMixerConfig`): model config
922
+ Returns:
923
+ x_mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`)
924
+ Masked patched input
925
+ mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches)`)
926
+ Bool tensor indicating True on masked points
927
+ """
928
+
929
+ def __init__(self, config: PatchTSMixerConfig):
930
+ super().__init__()
931
+ self.random_mask_ratio = config.random_mask_ratio
932
+ self.channel_consistent_masking = config.channel_consistent_masking
933
+ self.mask_type = config.mask_type
934
+ self.num_forecast_mask_patches = config.num_forecast_mask_patches
935
+ self.unmasked_channel_indices = config.unmasked_channel_indices
936
+ self.mask_value = config.mask_value
937
+ if self.unmasked_channel_indices is not None:
938
+ self.unmasked_channel_indices = sorted(self.unmasked_channel_indices)
939
+
940
+ def forward(self, patch_input: torch.Tensor):
941
+ """
942
+ Parameters:
943
+ patch_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`, *required*):
944
+ Patch input
945
+
946
+ Return:
947
+ masked_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`)
948
+ Masked patched input
949
+ mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches)`)
950
+ Bool tensor indicating True on masked points
951
+
952
+ """
953
+ if self.mask_type == "random":
954
+ masked_input, mask = random_masking(
955
+ inputs=patch_input,
956
+ mask_ratio=self.random_mask_ratio,
957
+ unmasked_channel_indices=self.unmasked_channel_indices,
958
+ channel_consistent_masking=self.channel_consistent_masking,
959
+ mask_value=self.mask_value,
960
+ )
961
+ elif self.mask_type == "forecast":
962
+ masked_input, mask = forecast_masking(
963
+ inputs=patch_input,
964
+ num_forecast_mask_patches=self.num_forecast_mask_patches,
965
+ unmasked_channel_indices=self.unmasked_channel_indices,
966
+ mask_value=self.mask_value,
967
+ )
968
+ else:
969
+ raise ValueError(f"Invalid mask type {self.mask_type}.")
970
+
971
+ # mask: [bs x num_input_channels x num_patch]
972
+ mask = mask.bool()
973
+ return masked_input, mask
974
+
975
+
976
+ # Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTStdScaler with PatchTST->PatchTSMixer
977
+ class PatchTSMixerStdScaler(nn.Module):
978
+ """
979
+ Standardize features by calculating the mean and scaling along the first dimension, and then normalizes it by
980
+ subtracting from the mean and dividing by the standard deviation.
981
+ """
982
+
983
+ def __init__(self, config: PatchTSMixerConfig):
984
+ super().__init__()
985
+ self.dim = config.scaling_dim if hasattr(config, "scaling_dim") else 1
986
+ self.keepdim = config.keepdim if hasattr(config, "keepdim") else True
987
+ self.minimum_scale = config.minimum_scale if hasattr(config, "minimum_scale") else 1e-5
988
+
989
+ def forward(
990
+ self, data: torch.Tensor, observed_indicator: torch.Tensor
991
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
992
+ """
993
+ Parameters:
994
+ data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`):
995
+ input for Batch norm calculation
996
+ observed_indicator (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
997
+ Calculating the scale on the observed indicator.
998
+ Returns:
999
+ tuple of `torch.Tensor` of shapes
1000
+ (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`,
1001
+ `(batch_size, 1, num_input_channels)`)
1002
+ """
1003
+ denominator = observed_indicator.sum(self.dim, keepdim=self.keepdim)
1004
+ denominator = denominator.clamp_min(1.0)
1005
+ loc = (data * observed_indicator).sum(self.dim, keepdim=self.keepdim) / denominator
1006
+
1007
+ variance = (((data - loc) * observed_indicator) ** 2).sum(self.dim, keepdim=self.keepdim) / denominator
1008
+ scale = torch.sqrt(variance + self.minimum_scale)
1009
+ return (data - loc) / scale, loc, scale
1010
+
1011
+
1012
+ # Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTMeanScaler with PatchTST->PatchTSMixer
1013
+ class PatchTSMixerMeanScaler(nn.Module):
1014
+ """
1015
+ Computes a scaling factor as the weighted average absolute value along the first dimension, and scales the data
1016
+ accordingly.
1017
+ """
1018
+
1019
+ def __init__(self, config: PatchTSMixerConfig):
1020
+ super().__init__()
1021
+ self.dim = config.scaling_dim if hasattr(config, "scaling_dim") else 1
1022
+ self.keepdim = config.keepdim if hasattr(config, "keepdim") else True
1023
+ self.minimum_scale = config.minimum_scale if hasattr(config, "minimum_scale") else 1e-10
1024
+ self.default_scale = config.default_scale if hasattr(config, "default_scale") else None
1025
+
1026
+ def forward(
1027
+ self, data: torch.Tensor, observed_indicator: torch.Tensor
1028
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
1029
+ """
1030
+ Parameters:
1031
+ data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`):
1032
+ input for Batch norm calculation
1033
+ observed_indicator (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
1034
+ Calculating the scale on the observed indicator.
1035
+ Returns:
1036
+ tuple of `torch.Tensor` of shapes
1037
+ (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`,
1038
+ `(batch_size, 1, num_input_channels)`)
1039
+ """
1040
+ ts_sum = (data * observed_indicator).abs().sum(self.dim, keepdim=True)
1041
+ num_observed = observed_indicator.sum(self.dim, keepdim=True)
1042
+
1043
+ scale = ts_sum / torch.clamp(num_observed, min=1)
1044
+
1045
+ # If `default_scale` is provided, we use it, otherwise we use the scale
1046
+ # of the batch.
1047
+ if self.default_scale is None:
1048
+ batch_sum = ts_sum.sum(dim=0)
1049
+ batch_observations = torch.clamp(num_observed.sum(0), min=1)
1050
+ default_scale = torch.squeeze(batch_sum / batch_observations)
1051
+ else:
1052
+ default_scale = self.default_scale * torch.ones_like(scale)
1053
+
1054
+ # apply default scale where there are no observations
1055
+ scale = torch.where(num_observed > 0, scale, default_scale)
1056
+
1057
+ # ensure the scale is at least `self.minimum_scale`
1058
+ scale = torch.clamp(scale, min=self.minimum_scale)
1059
+ scaled_data = data / scale
1060
+
1061
+ if not self.keepdim:
1062
+ scale = scale.squeeze(dim=self.dim)
1063
+
1064
+ return scaled_data, torch.zeros_like(scale), scale
1065
+
1066
+
1067
+ # Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTNOPScaler with PatchTST->PatchTSMixer
1068
+ class PatchTSMixerNOPScaler(nn.Module):
1069
+ """
1070
+ Assigns a scaling factor equal to 1 along the first dimension, and therefore applies no scaling to the input data.
1071
+ """
1072
+
1073
+ def __init__(self, config: PatchTSMixerConfig):
1074
+ super().__init__()
1075
+ self.dim = config.scaling_dim if hasattr(config, "scaling_dim") else 1
1076
+ self.keepdim = config.keepdim if hasattr(config, "keepdim") else True
1077
+
1078
+ def forward(
1079
+ self, data: torch.Tensor, observed_indicator: torch.Tensor | None = None
1080
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
1081
+ """
1082
+ Parameters:
1083
+ data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`):
1084
+ input for Batch norm calculation
1085
+ Returns:
1086
+ tuple of `torch.Tensor` of shapes
1087
+ (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`,
1088
+ `(batch_size, 1, num_input_channels)`)
1089
+ """
1090
+ scale = torch.ones_like(data, requires_grad=False).mean(dim=self.dim, keepdim=self.keepdim)
1091
+ loc = torch.zeros_like(data, requires_grad=False).mean(dim=self.dim, keepdim=self.keepdim)
1092
+ return data, loc, scale
1093
+
1094
+
1095
+ @auto_docstring(
1096
+ custom_intro="""
1097
+ Base class for `PatchTSMixerEncoderOutput`, with potential hidden states.
1098
+ """
1099
+ )
1100
+ @dataclass
1101
+ class PatchTSMixerEncoderOutput(ModelOutput):
1102
+ r"""
1103
+ last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, d_model)`):
1104
+ Hidden-state at the output of the last layer of the model.
1105
+ hidden_states (`tuple(torch.FloatTensor)`, *optional*):
1106
+ Hidden-states of the model at the output of each layer.
1107
+ """
1108
+
1109
+ last_hidden_state: torch.FloatTensor | None = None
1110
+ hidden_states: tuple[torch.FloatTensor] | None = None
1111
+
1112
+
1113
+ class PatchTSMixerEncoder(PatchTSMixerPreTrainedModel):
1114
+ """
1115
+ Encoder for PatchTSMixer which inputs patched time-series and outputs patched embeddings.
1116
+
1117
+ Args:
1118
+ config (`PatchTSMixerConfig`):
1119
+ Configuration.
1120
+ """
1121
+
1122
+ def __init__(self, config: PatchTSMixerConfig):
1123
+ super().__init__(config)
1124
+
1125
+ self.return_dict = config.return_dict
1126
+
1127
+ self.patcher = nn.Linear(config.patch_length, config.d_model)
1128
+ if config.use_positional_encoding:
1129
+ self.positional_encoder = PatchTSMixerPositionalEncoding(config=config)
1130
+ else:
1131
+ self.positional_encoder = None
1132
+ self.mlp_mixer_encoder = PatchTSMixerBlock(config=config)
1133
+
1134
+ # Initialize weights and apply final processing
1135
+ self.post_init()
1136
+
1137
+ @auto_docstring
1138
+ def forward(
1139
+ self,
1140
+ past_values: torch.Tensor,
1141
+ output_hidden_states: bool | None = False,
1142
+ return_dict: bool | None = None,
1143
+ **kwargs,
1144
+ ) -> tuple | PatchTSMixerEncoderOutput:
1145
+ r"""
1146
+ past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
1147
+ Context values of the time series. For a pretraining task, this denotes the input time series to
1148
+ predict the masked portion. For a forecasting task, this denotes the history/past time series values.
1149
+ Similarly, for classification or regression tasks, it denotes the appropriate context values of the
1150
+ time series.
1151
+
1152
+ For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series,
1153
+ it is greater than 1.
1154
+
1155
+ Returns:
1156
+ `torch.FloatTensor` of shape `(batch_size, n_vars, num_patches, d_model)`
1157
+ """
1158
+
1159
+ return_dict = return_dict if return_dict is not None else self.return_dict
1160
+
1161
+ # flatten [bs x num_patch x d_model]. common_channel/mix_channel: [bs x n_vars x num_patch x d_model]
1162
+ patches = self.patcher(past_values)
1163
+
1164
+ # add positional encoder
1165
+ if self.positional_encoder is not None:
1166
+ patches = self.positional_encoder(patches)
1167
+
1168
+ last_hidden_state, hidden_states = self.mlp_mixer_encoder(patches, output_hidden_states=output_hidden_states)
1169
+
1170
+ if not return_dict:
1171
+ return tuple(
1172
+ v
1173
+ for v in [
1174
+ last_hidden_state,
1175
+ hidden_states,
1176
+ ]
1177
+ )
1178
+
1179
+ return PatchTSMixerEncoderOutput(last_hidden_state=last_hidden_state, hidden_states=hidden_states)
1180
+
1181
+
1182
+ @auto_docstring(
1183
+ custom_intro="""
1184
+ Base class for model's outputs, with potential hidden states.
1185
+ """
1186
+ )
1187
+ @dataclass
1188
+ class PatchTSMixerModelOutput(ModelOutput):
1189
+ r"""
1190
+ last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, d_model)`):
1191
+ Hidden-state at the output of the last layer of the model.
1192
+ hidden_states (`tuple(torch.FloatTensor)`, *optional*):
1193
+ Hidden-states of the model at the output of each layer.
1194
+ patch_input (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, patch_length)`):
1195
+ Patched input data to the model.
1196
+ mask (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches)`, *optional*):
1197
+ Bool Tensor indicating True in masked patches and False otherwise.
1198
+ loc (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*):
1199
+ Gives the mean of the context window per channel. Used for revin denorm outside the model, if revin
1200
+ enabled.
1201
+ scale (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*):
1202
+ Gives the std dev of the context window per channel. Used for revin denorm outside the model, if revin
1203
+ enabled.
1204
+ """
1205
+
1206
+ last_hidden_state: torch.FloatTensor | None = None
1207
+ hidden_states: tuple[torch.FloatTensor] | None = None
1208
+ patch_input: torch.FloatTensor | None = None
1209
+ mask: torch.FloatTensor | None = None
1210
+ loc: torch.FloatTensor | None = None
1211
+ scale: torch.FloatTensor | None = None
1212
+
1213
+
1214
+ @auto_docstring(
1215
+ custom_intro="""
1216
+ The PatchTSMixer Model for time-series forecasting.
1217
+ """
1218
+ )
1219
+ class PatchTSMixerModel(PatchTSMixerPreTrainedModel):
1220
+ def __init__(self, config: PatchTSMixerConfig, mask_input: bool = False):
1221
+ r"""
1222
+ mask_input (bool, *optional*, defaults to `False`):
1223
+ Whether to mask the input using the [`PatchTSMixerMasking`] module.
1224
+ """
1225
+ super().__init__(config)
1226
+
1227
+ self.return_dict = config.return_dict
1228
+ self.encoder = PatchTSMixerEncoder(config)
1229
+ self.patching = PatchTSMixerPatchify(config)
1230
+
1231
+ if mask_input is True:
1232
+ self.masking = PatchTSMixerMasking(config)
1233
+ else:
1234
+ self.masking = None
1235
+
1236
+ if config.scaling == "mean":
1237
+ self.scaler = PatchTSMixerMeanScaler(config)
1238
+ elif config.scaling == "std" or config.scaling is True:
1239
+ self.scaler = PatchTSMixerStdScaler(config)
1240
+ else:
1241
+ self.scaler = PatchTSMixerNOPScaler(config)
1242
+
1243
+ # Initialize weights and apply final processing
1244
+ self.post_init()
1245
+
1246
+ @auto_docstring
1247
+ def forward(
1248
+ self,
1249
+ past_values: torch.Tensor,
1250
+ observed_mask: torch.Tensor | None = None,
1251
+ output_hidden_states: bool | None = False,
1252
+ return_dict: bool | None = None,
1253
+ **kwargs,
1254
+ ) -> PatchTSMixerModelOutput:
1255
+ r"""
1256
+ past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
1257
+ Context values of the time series. For a pretraining task, this denotes the input time series to predict
1258
+ the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
1259
+ for classification or regression tasks, it denotes the appropriate context values of the time series.
1260
+
1261
+ For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
1262
+ greater than 1.
1263
+ observed_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
1264
+ Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
1265
+ in `[0, 1]`:
1266
+ - 1 for values that are **observed**,
1267
+ - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
1268
+ """
1269
+ return_dict = return_dict if return_dict is not None else self.return_dict
1270
+
1271
+ mask = None
1272
+ if observed_mask is None:
1273
+ observed_mask = torch.ones_like(past_values)
1274
+ scaled_past_values, loc, scale = self.scaler(past_values, observed_mask)
1275
+
1276
+ patched_x = self.patching(scaled_past_values) # [batch_size x num_input_channels x num_patch x patch_length
1277
+
1278
+ enc_input = patched_x
1279
+ if self.masking is not None:
1280
+ enc_input, mask = self.masking(patched_x)
1281
+ # enc_input: [batch_size x num_input_channels x num_patch x patch_length]
1282
+ # mask: [batch_size x num_input_channels x num_patch]
1283
+
1284
+ encoder_output = self.encoder(
1285
+ enc_input,
1286
+ output_hidden_states=output_hidden_states,
1287
+ return_dict=return_dict,
1288
+ )
1289
+
1290
+ if isinstance(encoder_output, tuple):
1291
+ encoder_output = PatchTSMixerEncoderOutput(*encoder_output)
1292
+
1293
+ if not return_dict:
1294
+ return tuple(
1295
+ v
1296
+ for v in [
1297
+ encoder_output.last_hidden_state,
1298
+ encoder_output.hidden_states,
1299
+ patched_x,
1300
+ mask,
1301
+ loc,
1302
+ scale,
1303
+ ]
1304
+ )
1305
+
1306
+ return PatchTSMixerModelOutput(
1307
+ last_hidden_state=encoder_output.last_hidden_state,
1308
+ hidden_states=encoder_output.hidden_states,
1309
+ patch_input=patched_x,
1310
+ mask=mask,
1311
+ loc=loc,
1312
+ scale=scale,
1313
+ )
1314
+
1315
+
1316
+ @auto_docstring(
1317
+ custom_intro="""
1318
+ Output type of [`PatchTSMixerForPreTrainingOutput`].
1319
+ """
1320
+ )
1321
+ @dataclass
1322
+ class PatchTSMixerForPreTrainingOutput(ModelOutput):
1323
+ r"""
1324
+ loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
1325
+ Total loss
1326
+ prediction_outputs (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, patch_length)`):
1327
+ Prediction output from the pretrain head.
1328
+ last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
1329
+ Backbone embeddings before passing through the head.
1330
+ hidden_states (`tuple(torch.FloatTensor)`, *optional*):
1331
+ Hidden-states of the model at the output of each layer.
1332
+ """
1333
+
1334
+ loss: torch.FloatTensor | None = None
1335
+ prediction_outputs: torch.FloatTensor | None = None
1336
+ last_hidden_state: torch.FloatTensor | None = None
1337
+ hidden_states: tuple[torch.FloatTensor] | None = None
1338
+
1339
+
1340
+ @auto_docstring(
1341
+ custom_intro="""
1342
+ `PatchTSMixer` for mask pretraining.
1343
+ """
1344
+ )
1345
+ class PatchTSMixerForPretraining(PatchTSMixerPreTrainedModel):
1346
+ def __init__(self, config: PatchTSMixerConfig):
1347
+ super().__init__(config)
1348
+ self.model = PatchTSMixerModel(config, mask_input=True)
1349
+ self.head = PatchTSMixerPretrainHead(config=config)
1350
+ self.masked_loss = config.masked_loss
1351
+ self.return_dict = config.return_dict
1352
+
1353
+ # Initialize weights and apply final processing
1354
+ self.post_init()
1355
+
1356
+ @auto_docstring
1357
+ def forward(
1358
+ self,
1359
+ past_values: torch.Tensor,
1360
+ observed_mask: torch.Tensor | None = None,
1361
+ output_hidden_states: bool | None = False,
1362
+ return_loss: bool = True,
1363
+ return_dict: bool | None = None,
1364
+ **kwargs,
1365
+ ) -> PatchTSMixerForPreTrainingOutput:
1366
+ r"""
1367
+ past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
1368
+ Context values of the time series. For a pretraining task, this denotes the input time series to predict
1369
+ the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
1370
+ for classification or regression tasks, it denotes the appropriate context values of the time series.
1371
+
1372
+ For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
1373
+ greater than 1.
1374
+ observed_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
1375
+ Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
1376
+ in `[0, 1]`:
1377
+ - 1 for values that are **observed**,
1378
+ - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
1379
+ return_loss (`bool`, *optional*):
1380
+ Whether to return the loss in the `forward` call.
1381
+ """
1382
+ return_dict = return_dict if return_dict is not None else self.return_dict
1383
+
1384
+ if self.masked_loss is True:
1385
+ loss = torch.nn.MSELoss(reduction="none")
1386
+ else:
1387
+ loss = torch.nn.MSELoss(reduction="mean")
1388
+
1389
+ # past_values: tensor [batch_size x context_length x num_input_channels]
1390
+ model_output = self.model(
1391
+ past_values,
1392
+ observed_mask=observed_mask,
1393
+ output_hidden_states=output_hidden_states,
1394
+ return_dict=return_dict,
1395
+ ) # x.last_hidden_state: [batch_size x nvars x num_patch x d_model]
1396
+ if isinstance(model_output, tuple):
1397
+ model_output = PatchTSMixerModelOutput(*model_output)
1398
+
1399
+ x_hat = self.head(model_output.last_hidden_state) # tensor [batch_size x nvars x num_patch x patch_length]
1400
+
1401
+ if return_loss is True:
1402
+ loss_val = loss(x_hat, model_output.patch_input)
1403
+ else:
1404
+ loss_val = None
1405
+
1406
+ # calculate masked_loss
1407
+ if self.masked_loss is True and loss_val is not None:
1408
+ loss_val = (loss_val.mean(dim=-1) * model_output.mask).sum() / (model_output.mask.sum() + 1e-10)
1409
+
1410
+ if not return_dict:
1411
+ return tuple(
1412
+ v
1413
+ for v in [
1414
+ loss_val,
1415
+ x_hat,
1416
+ model_output.last_hidden_state,
1417
+ model_output.hidden_states,
1418
+ ]
1419
+ )
1420
+
1421
+ return PatchTSMixerForPreTrainingOutput(
1422
+ loss=loss_val,
1423
+ prediction_outputs=x_hat, # tensor [batch_size x nvars x num_patch x patch_length]
1424
+ last_hidden_state=model_output.last_hidden_state, # x: [batch_size x nvars x num_patch x d_model]
1425
+ hidden_states=model_output.hidden_states,
1426
+ )
1427
+
1428
+
1429
+ @auto_docstring(
1430
+ custom_intro="""
1431
+ Output type of [`PatchTSMixerForPredictionOutput`].
1432
+ """
1433
+ )
1434
+ @dataclass
1435
+ class PatchTSMixerForPredictionOutput(ModelOutput):
1436
+ r"""
1437
+ loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
1438
+ Total loss.
1439
+ prediction_outputs (`torch.FloatTensor` of shape `(batch_size, prediction_length, num_input_channels)`):
1440
+ Prediction output from the forecast head.
1441
+ last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
1442
+ Backbone embeddings before passing through the head.
1443
+ hidden_states (`tuple(torch.FloatTensor)`, *optional*):
1444
+ Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
1445
+ loc (`torch.FloatTensor`, *optional* of shape `(batch_size, 1, num_input_channels)`):
1446
+ Input mean
1447
+ scale (`torch.FloatTensor`, *optional* of shape `(batch_size, 1, num_input_channels)`):
1448
+ Input std dev
1449
+ """
1450
+
1451
+ loss: torch.FloatTensor | None = None
1452
+ prediction_outputs: torch.FloatTensor | None = None
1453
+ last_hidden_state: torch.FloatTensor | None = None
1454
+ hidden_states: tuple[torch.FloatTensor] | None = None
1455
+ loc: torch.FloatTensor | None = None
1456
+ scale: torch.FloatTensor | None = None
1457
+
1458
+
1459
+ @auto_docstring(
1460
+ custom_intro="""
1461
+ Base class for time series model's predictions outputs that contains the sampled values from the chosen
1462
+ distribution.
1463
+ """
1464
+ )
1465
+ @dataclass
1466
+ class SamplePatchTSMixerPredictionOutput(ModelOutput):
1467
+ r"""
1468
+ sequences (`torch.FloatTensor` of shape `(batch_size, num_samples, prediction_length, number_channels)`):
1469
+ Sampled values from the chosen distribution.
1470
+ """
1471
+
1472
+ sequences: torch.FloatTensor | None = None
1473
+
1474
+
1475
+ @auto_docstring(
1476
+ custom_intro="""
1477
+ Base class for time series model's predictions outputs that contains the sampled values from the chosen
1478
+ distribution.
1479
+ """
1480
+ )
1481
+ @dataclass
1482
+ class SamplePatchTSMixerRegressionOutput(ModelOutput):
1483
+ r"""
1484
+ sequences (`torch.FloatTensor` of shape `(batch_size, num_samples, prediction_length, number_channels)`):
1485
+ Sampled values from the chosen distribution.
1486
+ """
1487
+
1488
+ sequences: torch.FloatTensor | None = None
1489
+
1490
+
1491
+ # Copied from transformers.models.time_series_transformer.modeling_time_series_transformer.nll
1492
+ def nll(input: torch.distributions.Distribution, target: torch.Tensor) -> torch.Tensor:
1493
+ """
1494
+ Computes the negative log likelihood loss from input distribution with respect to target.
1495
+ """
1496
+ return -input.log_prob(target)
1497
+
1498
+
1499
+ # Copied from transformers.models.time_series_transformer.modeling_time_series_transformer.weighted_average
1500
+ def weighted_average(input_tensor: torch.Tensor, weights: torch.Tensor | None = None, dim=None) -> torch.Tensor:
1501
+ """
1502
+ Computes the weighted average of a given tensor across a given `dim`, masking values associated with weight zero,
1503
+ meaning instead of `nan * 0 = nan` you will get `0 * 0 = 0`.
1504
+
1505
+ Args:
1506
+ input_tensor (`torch.FloatTensor`):
1507
+ Input tensor, of which the average must be computed.
1508
+ weights (`torch.FloatTensor`, *optional*):
1509
+ Weights tensor, of the same shape as `input_tensor`.
1510
+ dim (`int`, *optional*):
1511
+ The dim along which to average `input_tensor`.
1512
+
1513
+ Returns:
1514
+ `torch.FloatTensor`: The tensor with values averaged along the specified `dim`.
1515
+ """
1516
+ if weights is not None:
1517
+ weighted_tensor = torch.where(weights != 0, input_tensor * weights, torch.zeros_like(input_tensor))
1518
+ sum_weights = torch.clamp(weights.sum(dim=dim) if dim else weights.sum(), min=1.0)
1519
+ return (weighted_tensor.sum(dim=dim) if dim else weighted_tensor.sum()) / sum_weights
1520
+ else:
1521
+ return input_tensor.mean(dim=dim)
1522
+
1523
+
1524
+ class PatchTSMixerForPrediction(PatchTSMixerPreTrainedModel):
1525
+ r"""
1526
+ `PatchTSMixer` for forecasting application.
1527
+
1528
+ Args:
1529
+ config (`PatchTSMixerConfig`):
1530
+ Configuration.
1531
+
1532
+ Returns:
1533
+ `None`.
1534
+ """
1535
+
1536
+ def __init__(self, config: PatchTSMixerConfig):
1537
+ super().__init__(config)
1538
+ self.loss = config.loss
1539
+ self.return_dict = config.return_dict
1540
+ self.prediction_channel_indices = config.prediction_channel_indices
1541
+ self.num_parallel_samples = config.num_parallel_samples
1542
+
1543
+ if config.loss == "mse":
1544
+ self.distribution_output = None
1545
+ else:
1546
+ dim = config.prediction_length
1547
+ distribution_output_map = {
1548
+ "student_t": StudentTOutput,
1549
+ "normal": NormalOutput,
1550
+ "negative_binomial": NegativeBinomialOutput,
1551
+ }
1552
+ output_class = distribution_output_map.get(config.distribution_output)
1553
+ if output_class is not None:
1554
+ self.distribution_output = output_class(dim=dim)
1555
+ else:
1556
+ raise ValueError(f"Unknown distribution output {config.distribution_output}")
1557
+
1558
+ self.model = PatchTSMixerModel(config)
1559
+ self.head = PatchTSMixerForPredictionHead(
1560
+ config=config,
1561
+ distribution_output=self.distribution_output,
1562
+ )
1563
+
1564
+ # Initialize weights and apply final processing
1565
+ self.post_init()
1566
+
1567
+ @auto_docstring
1568
+ def forward(
1569
+ self,
1570
+ past_values: torch.Tensor,
1571
+ observed_mask: torch.Tensor | None = None,
1572
+ future_values: torch.Tensor | None = None,
1573
+ output_hidden_states: bool | None = False,
1574
+ return_loss: bool = True,
1575
+ return_dict: bool | None = None,
1576
+ **kwargs,
1577
+ ) -> PatchTSMixerForPredictionOutput:
1578
+ r"""
1579
+ past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
1580
+ Context values of the time series. For a pretraining task, this denotes the input time series to predict
1581
+ the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
1582
+ for classification or regression tasks, it denotes the appropriate context values of the time series.
1583
+
1584
+ For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
1585
+ greater than 1.
1586
+ observed_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
1587
+ Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
1588
+ in `[0, 1]`:
1589
+ - 1 for values that are **observed**,
1590
+ - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
1591
+ future_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting,:
1592
+ `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*):
1593
+ Target values of the time series, that serve as labels for the model. The `future_values` is what the
1594
+ Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT
1595
+ required for a pretraining task.
1596
+
1597
+ For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want
1598
+ to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter,
1599
+ pass the target data with all channels, as channel Filtering for both prediction and target will be
1600
+ manually applied before the loss computation.
1601
+ return_loss (`bool`, *optional*):
1602
+ Whether to return the loss in the `forward` call.
1603
+ """
1604
+ if self.loss == "mse":
1605
+ loss = nn.MSELoss(reduction="mean")
1606
+ elif self.loss == "nll":
1607
+ loss = nll
1608
+ else:
1609
+ raise ValueError("Invalid loss function: Allowed values: mse and nll")
1610
+
1611
+ return_dict = return_dict if return_dict is not None else self.return_dict
1612
+
1613
+ # past_values: tensor [batch_size x context_length x num_input_channels]
1614
+ model_output = self.model(
1615
+ past_values,
1616
+ observed_mask=observed_mask,
1617
+ output_hidden_states=output_hidden_states,
1618
+ return_dict=return_dict,
1619
+ ) # model_output: [batch_size x nvars x num_patch x d_model]
1620
+ if isinstance(model_output, tuple):
1621
+ model_output = PatchTSMixerModelOutput(*model_output)
1622
+
1623
+ # tensor [batch_size x prediction_length x num_input_channels]
1624
+ y_hat = self.head(model_output.last_hidden_state)
1625
+
1626
+ loss_val = None
1627
+ if self.prediction_channel_indices is not None:
1628
+ if self.distribution_output:
1629
+ distribution = self.distribution_output.distribution(
1630
+ y_hat,
1631
+ loc=model_output.loc[..., self.prediction_channel_indices],
1632
+ scale=model_output.scale[..., self.prediction_channel_indices],
1633
+ )
1634
+ if future_values is not None and return_loss is True:
1635
+ loss_val = loss(
1636
+ distribution,
1637
+ future_values[..., self.prediction_channel_indices],
1638
+ )
1639
+ # take average of the loss
1640
+ loss_val = weighted_average(loss_val)
1641
+ else:
1642
+ y_hat = (
1643
+ y_hat * model_output.scale[..., self.prediction_channel_indices]
1644
+ + model_output.loc[..., self.prediction_channel_indices]
1645
+ )
1646
+ if future_values is not None and return_loss is True:
1647
+ loss_val = loss(y_hat, future_values[..., self.prediction_channel_indices])
1648
+ else:
1649
+ if self.distribution_output:
1650
+ distribution = self.distribution_output.distribution(
1651
+ y_hat, loc=model_output.loc, scale=model_output.scale
1652
+ )
1653
+ if future_values is not None and return_loss is True:
1654
+ loss_val = loss(distribution, future_values)
1655
+ loss_val = weighted_average(loss_val)
1656
+ else:
1657
+ y_hat = y_hat * model_output.scale + model_output.loc
1658
+ if future_values is not None and return_loss is True:
1659
+ loss_val = loss(y_hat, future_values)
1660
+
1661
+ if self.prediction_channel_indices is not None:
1662
+ loc = model_output.loc[..., self.prediction_channel_indices]
1663
+ scale = model_output.scale[..., self.prediction_channel_indices]
1664
+ else:
1665
+ loc = model_output.loc
1666
+ scale = model_output.scale
1667
+
1668
+ if not return_dict:
1669
+ return tuple(
1670
+ v
1671
+ for v in [
1672
+ loss_val,
1673
+ y_hat,
1674
+ model_output.last_hidden_state,
1675
+ model_output.hidden_states,
1676
+ loc,
1677
+ scale,
1678
+ ]
1679
+ )
1680
+
1681
+ return PatchTSMixerForPredictionOutput(
1682
+ loss=loss_val,
1683
+ prediction_outputs=y_hat, # tensor [batch_size x prediction_length x num_input_channels]
1684
+ last_hidden_state=model_output.last_hidden_state, # x: [batch_size x nvars x num_patch x d_model]
1685
+ hidden_states=model_output.hidden_states,
1686
+ loc=loc,
1687
+ scale=scale,
1688
+ )
1689
+
1690
+ @torch.no_grad()
1691
+ def generate(
1692
+ self,
1693
+ past_values: torch.Tensor,
1694
+ observed_mask: torch.Tensor | None = None,
1695
+ ) -> SamplePatchTSMixerPredictionOutput:
1696
+ """
1697
+ Generate sequences of sample predictions from a model with a probability distribution head.
1698
+
1699
+ Args:
1700
+ past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
1701
+ Past values of the time series that serves as context in order to predict the future.
1702
+
1703
+ observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
1704
+ Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
1705
+ in `[0, 1]`:
1706
+
1707
+ - 1 for values that are **observed**,
1708
+ - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
1709
+
1710
+ Return:
1711
+ [`SamplePatchTSMixerPredictionOutput`] where the outputs `sequences` tensor will have shape `(batch_size,
1712
+ number of samples, prediction_length, num_input_channels)`.
1713
+ """
1714
+ # get number of samples
1715
+ num_parallel_samples = self.num_parallel_samples
1716
+
1717
+ # get model output
1718
+ outputs = self(
1719
+ past_values=past_values,
1720
+ future_values=None,
1721
+ observed_mask=observed_mask,
1722
+ output_hidden_states=False,
1723
+ )
1724
+
1725
+ # get distribution
1726
+
1727
+ distribution = self.distribution_output.distribution(
1728
+ outputs.prediction_outputs, loc=outputs.loc, scale=outputs.scale
1729
+ )
1730
+
1731
+ # get samples: list of [batch_size x prediction_length x num_channels]
1732
+ samples = [distribution.sample() for _ in range(num_parallel_samples)]
1733
+
1734
+ # stack tensors
1735
+ samples = torch.stack(samples, dim=1) # [batch_size x num_samples x prediction_length x num_channels]
1736
+ return SamplePatchTSMixerPredictionOutput(sequences=samples)
1737
+
1738
+
1739
+ @auto_docstring(
1740
+ custom_intro="""
1741
+ Output type of [`PatchTSMixerForTimeSeriesClassificationOutput`].
1742
+ """
1743
+ )
1744
+ @dataclass
1745
+ class PatchTSMixerForTimeSeriesClassificationOutput(ModelOutput):
1746
+ r"""
1747
+ loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
1748
+ Total loss.
1749
+ prediction_outputs (`torch.FloatTensor` of shape `(batch_size, num_labels)`):
1750
+ Prediction output from the classification head.
1751
+ last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
1752
+ Backbone embeddings before passing through the head.
1753
+ hidden_states (`tuple(torch.FloatTensor)`, *optional*):
1754
+ Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
1755
+ """
1756
+
1757
+ loss: torch.FloatTensor | None = None
1758
+ prediction_outputs: torch.FloatTensor | None = None
1759
+ last_hidden_state: torch.FloatTensor | None = None
1760
+ hidden_states: tuple[torch.FloatTensor] | None = None
1761
+
1762
+
1763
+ class PatchTSMixerForTimeSeriesClassification(PatchTSMixerPreTrainedModel):
1764
+ r"""
1765
+ `PatchTSMixer` for classification application.
1766
+
1767
+ Args:
1768
+ config (`PatchTSMixerConfig`):
1769
+ Configuration.
1770
+
1771
+ Returns:
1772
+ `None`.
1773
+ """
1774
+
1775
+ def __init__(self, config: PatchTSMixerConfig):
1776
+ super().__init__(config)
1777
+
1778
+ self.model = PatchTSMixerModel(config)
1779
+ self.head = PatchTSMixerLinearHead(
1780
+ config=config,
1781
+ )
1782
+ self.return_dict = config.return_dict
1783
+ if config.scaling in ["std", "mean", True]:
1784
+ self.inject_scale = InjectScalerStatistics4D(d_model=config.d_model, num_patches=config.num_patches)
1785
+ else:
1786
+ self.inject_scale = None
1787
+
1788
+ # Initialize weights and apply final processing
1789
+ self.post_init()
1790
+
1791
+ @auto_docstring
1792
+ def forward(
1793
+ self,
1794
+ past_values: torch.Tensor,
1795
+ target_values: torch.Tensor | None = None,
1796
+ output_hidden_states: bool | None = False,
1797
+ return_loss: bool = True,
1798
+ return_dict: bool | None = None,
1799
+ **kwargs,
1800
+ ) -> PatchTSMixerForTimeSeriesClassificationOutput:
1801
+ r"""
1802
+ past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
1803
+ Context values of the time series. For a pretraining task, this denotes the input time series to predict
1804
+ the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
1805
+ for classification or regression tasks, it denotes the appropriate context values of the time series.
1806
+
1807
+ For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
1808
+ greater than 1.
1809
+ target_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting,
1810
+ `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*):
1811
+ Target
1812
+ values of the time series, that serve as labels for the model. The `target_values` is what the
1813
+ Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT
1814
+ required for a pretraining task.
1815
+
1816
+ For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want
1817
+ to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter,
1818
+ pass the target data with all channels, as channel Filtering for both prediction and target will be
1819
+ manually applied before the loss computation.
1820
+
1821
+ For a classification task, it has a shape of `(batch_size,)`.
1822
+
1823
+ For a regression task, it has a shape of `(batch_size, num_targets)`.
1824
+ return_loss (`bool`, *optional*):
1825
+ Whether to return the loss in the `forward` call.
1826
+ """
1827
+
1828
+ loss = torch.nn.CrossEntropyLoss()
1829
+
1830
+ return_dict = return_dict if return_dict is not None else self.return_dict
1831
+
1832
+ model_output = self.model(
1833
+ past_values,
1834
+ output_hidden_states=output_hidden_states,
1835
+ return_dict=return_dict,
1836
+ ) # x: [batch_size x nvars x num_patch x d_model]
1837
+ if isinstance(model_output, tuple):
1838
+ model_output = PatchTSMixerModelOutput(*model_output)
1839
+
1840
+ if self.inject_scale is not None:
1841
+ model_output.last_hidden_state = self.inject_scale(
1842
+ model_output.last_hidden_state,
1843
+ loc=model_output.loc,
1844
+ scale=model_output.scale,
1845
+ ) # x: [batch_size x nvars x num_patch x d_model]
1846
+
1847
+ y_hat = self.head(model_output.last_hidden_state) # tensor [batch_size x n_labels]
1848
+
1849
+ if target_values is not None and return_loss is True:
1850
+ loss_val = loss(y_hat, target_values)
1851
+ else:
1852
+ loss_val = None
1853
+
1854
+ if not return_dict:
1855
+ return tuple(
1856
+ v
1857
+ for v in [
1858
+ loss_val,
1859
+ y_hat,
1860
+ model_output.last_hidden_state,
1861
+ model_output.hidden_states,
1862
+ ]
1863
+ )
1864
+
1865
+ return PatchTSMixerForTimeSeriesClassificationOutput(
1866
+ loss=loss_val,
1867
+ prediction_outputs=y_hat, # tensor [batch_size x n_labels]
1868
+ last_hidden_state=model_output.last_hidden_state, # x: [batch_size x nvars x num_patch x d_model]
1869
+ hidden_states=model_output.hidden_states,
1870
+ )
1871
+
1872
+
1873
+ @auto_docstring(
1874
+ custom_intro="""
1875
+ Output type of [`PatchTSMixerForRegressionOutput`].
1876
+ """
1877
+ )
1878
+ @dataclass
1879
+ class PatchTSMixerForRegressionOutput(ModelOutput):
1880
+ r"""
1881
+ loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
1882
+ Total loss.
1883
+ regression_outputs (`torch.FloatTensor` of shape `(batch_size, num_targets)`):
1884
+ Prediction output from the regression head.
1885
+ last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
1886
+ Backbone embeddings before passing through the head.
1887
+ hidden_states (`tuple(torch.FloatTensor)`, *optional*):
1888
+ Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
1889
+ """
1890
+
1891
+ loss: torch.FloatTensor | None = None
1892
+ regression_outputs: torch.FloatTensor | None = None
1893
+ last_hidden_state: torch.FloatTensor | None = None
1894
+ hidden_states: tuple[torch.FloatTensor] | None = None
1895
+
1896
+
1897
+ class InjectScalerStatistics4D(nn.Module):
1898
+ def __init__(self, d_model: int, num_patches: int, expansion: int = 2):
1899
+ super().__init__()
1900
+
1901
+ self.inverse_trans_expansion = nn.Linear(d_model + 2, expansion * d_model)
1902
+ self.inverse_trans_compression = nn.Linear(expansion * d_model, d_model)
1903
+ self.map_scale_expansion = nn.Linear(2, 2 * expansion)
1904
+ self.map_scale_compression = nn.Linear(2 * expansion, 2)
1905
+ self.num_patches = num_patches
1906
+
1907
+ def forward(self, inputs: torch.Tensor, loc: torch.Tensor, scale: torch.Tensor):
1908
+ """
1909
+ Args:
1910
+ inputs (`torch.Tensor` of shape `(batch_size, num_input_channels, num_patch, d_model)`)
1911
+ loc (`torch.Tensor` of shape `(batch_size, 1, num_input_channels)`)
1912
+ scale (`torch.Tensor` of shape `(batch_size, 1, num_input_channels)`)
1913
+ Returns:
1914
+ `torch.Tensor` of shape `(batch_size, num_input_channels, num_patch, d_model)`
1915
+ """
1916
+
1917
+ mean = loc.transpose(-1, -2) # [batch_size x n_channels x 1 ]
1918
+ mean = mean.unsqueeze(-2) # [batch_size x n_channels x 1 x 1]
1919
+ mean = mean.repeat(1, 1, self.num_patches, 1) # [batch_size x n_channels x num_patch x 1]
1920
+
1921
+ stdev = scale.transpose(-1, -2) # [batch_size x n_channels x 1 ]
1922
+ stdev = stdev.unsqueeze(-2) # [batch_size x n_channels x 1 x 1]
1923
+ stdev = stdev.repeat(1, 1, self.num_patches, 1) # [batch_size x n_channels x num_patch x 1]
1924
+
1925
+ concat_stats = torch.cat([mean, stdev], dim=-1) # [batch_size x n_channels x num_patch x 2]
1926
+
1927
+ concat_stats = self.map_scale_expansion(concat_stats) # [batch_size x n_channels x num_patch x (2*expansion)]
1928
+ concat_stats = self.map_scale_compression(concat_stats) # [batch_size x n_channels x num_patch x 2]
1929
+
1930
+ inputs = torch.cat([inputs, concat_stats], dim=-1) # [batch_size x channels x num_patch x d_model+2]
1931
+ inputs = self.inverse_trans_expansion(inputs) # [batch_size x channels x num_patch x (expansion*d_model)]
1932
+ inputs = self.inverse_trans_compression(inputs) # [batch_size x channels x num_patch x d_model]
1933
+
1934
+ return inputs
1935
+
1936
+
1937
+ @auto_docstring(
1938
+ custom_intro="""
1939
+ `PatchTSMixer` for regression application.
1940
+ """
1941
+ )
1942
+ class PatchTSMixerForRegression(PatchTSMixerPreTrainedModel):
1943
+ def __init__(self, config: PatchTSMixerConfig):
1944
+ super().__init__(config)
1945
+
1946
+ self.model = PatchTSMixerModel(config)
1947
+
1948
+ self.loss = config.loss
1949
+ self.distribution_output = config.distribution_output
1950
+
1951
+ self.return_dict = config.return_dict
1952
+ self.num_parallel_samples = config.num_parallel_samples
1953
+
1954
+ if config.loss == "mse":
1955
+ self.distribution_output = None
1956
+ else:
1957
+ distribution_output_map = {
1958
+ "student_t": StudentTOutput,
1959
+ "normal": NormalOutput,
1960
+ "negative_binomial": NegativeBinomialOutput,
1961
+ }
1962
+ output_class = distribution_output_map.get(config.distribution_output)
1963
+ if output_class is not None:
1964
+ self.distribution_output = output_class(dim=config.num_targets)
1965
+ else:
1966
+ raise ValueError(f"Unknown distribution output {config.distribution_output}")
1967
+
1968
+ if config.scaling in ["std", "mean", True]:
1969
+ self.inject_scale = InjectScalerStatistics4D(d_model=config.d_model, num_patches=config.num_patches)
1970
+ else:
1971
+ self.inject_scale = None
1972
+
1973
+ self.head = PatchTSMixerLinearHead(
1974
+ config=config,
1975
+ distribution_output=self.distribution_output,
1976
+ )
1977
+
1978
+ # Initialize weights and apply final processing
1979
+ self.post_init()
1980
+
1981
+ @auto_docstring
1982
+ def forward(
1983
+ self,
1984
+ past_values: torch.Tensor,
1985
+ target_values: torch.Tensor | None = None,
1986
+ output_hidden_states: bool | None = False,
1987
+ return_loss: bool = True,
1988
+ return_dict: bool | None = None,
1989
+ **kwargs,
1990
+ ) -> PatchTSMixerForRegressionOutput:
1991
+ r"""
1992
+ past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
1993
+ Context values of the time series. For a pretraining task, this denotes the input time series to predict
1994
+ the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
1995
+ for classification or regression tasks, it denotes the appropriate context values of the time series.
1996
+
1997
+ For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
1998
+ greater than 1.
1999
+ target_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting,
2000
+ `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*):
2001
+ Target values of the time series, that serve as labels for the model. The `target_values` is what the
2002
+ Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT
2003
+ required for a pretraining task.
2004
+
2005
+ For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want
2006
+ to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter,
2007
+ pass the target data with all channels, as channel Filtering for both prediction and target will be
2008
+ manually applied before the loss computation.
2009
+
2010
+ For a classification task, it has a shape of `(batch_size,)`.
2011
+
2012
+ For a regression task, it has a shape of `(batch_size, num_targets)`.
2013
+ return_loss (`bool`, *optional*):
2014
+ Whether to return the loss in the `forward` call.
2015
+ """
2016
+
2017
+ if self.loss == "mse":
2018
+ loss = nn.MSELoss(reduction="mean")
2019
+ elif self.loss == "nll":
2020
+ loss = nll
2021
+ else:
2022
+ raise ValueError("Invalid loss function: Allowed values: mse and nll")
2023
+
2024
+ return_dict = return_dict if return_dict is not None else self.return_dict
2025
+ model_output = self.model(
2026
+ past_values,
2027
+ output_hidden_states=output_hidden_states,
2028
+ return_dict=return_dict,
2029
+ ) # model_output: [batch_size x nvars x num_patch x d_model]
2030
+ if isinstance(model_output, tuple):
2031
+ model_output = PatchTSMixerModelOutput(*model_output)
2032
+
2033
+ if self.inject_scale is not None:
2034
+ model_output.last_hidden_state = self.inject_scale(
2035
+ model_output.last_hidden_state,
2036
+ loc=model_output.loc,
2037
+ scale=model_output.scale,
2038
+ ) # x: [batch_size x nvars x num_patch x d_model]
2039
+
2040
+ y_hat = self.head(model_output.last_hidden_state) # [batch_size x num_targets]
2041
+
2042
+ if target_values is not None and return_loss is True:
2043
+ if self.distribution_output:
2044
+ if self.distribution_output == "negative_binomial" and torch.any(target_values < 0):
2045
+ raise Exception("target_values cannot be negative for negative_binomial distribution.")
2046
+ distribution = self.distribution_output.distribution(y_hat)
2047
+ # y_hat should be a 2-tuple, each with dimension [bs, num_targets]
2048
+ y_hat = tuple(item.view(-1, self.config.num_targets) for item in y_hat)
2049
+ loss_val = loss(distribution, target_values)
2050
+ # take average of the loss
2051
+ loss_val = weighted_average(loss_val)
2052
+ else:
2053
+ loss_val = loss(y_hat, target_values)
2054
+ else:
2055
+ loss_val = None
2056
+
2057
+ if not return_dict:
2058
+ return tuple(
2059
+ v
2060
+ for v in [
2061
+ loss_val,
2062
+ y_hat,
2063
+ model_output.last_hidden_state,
2064
+ model_output.hidden_states,
2065
+ ]
2066
+ )
2067
+
2068
+ return PatchTSMixerForRegressionOutput(
2069
+ loss=loss_val,
2070
+ regression_outputs=y_hat, # tensor [batch_size x num_targets]
2071
+ last_hidden_state=model_output.last_hidden_state, # [batch_size x nvars x num_patch x d_model]
2072
+ hidden_states=model_output.hidden_states,
2073
+ )
2074
+
2075
+ @torch.no_grad()
2076
+ def generate(
2077
+ self,
2078
+ past_values: torch.Tensor,
2079
+ ) -> SamplePatchTSMixerRegressionOutput:
2080
+ """
2081
+ Generate sequences of sample predictions from a model with a probability distribution head.
2082
+
2083
+ Args:
2084
+ past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
2085
+ Past values of the time series that serves as context in order to predict the target values.
2086
+
2087
+ Return:
2088
+ [`SamplePatchTSMixerRegressionOutput`] where the outputs `sequences` tensor will have shape `(batch_size,
2089
+ number of samples, num_targets)`.
2090
+ """
2091
+ # get number of samples
2092
+ num_parallel_samples = self.num_parallel_samples
2093
+
2094
+ # get model output
2095
+ outputs = self(
2096
+ past_values=past_values,
2097
+ target_values=None,
2098
+ output_hidden_states=False,
2099
+ )
2100
+
2101
+ # get distribution
2102
+ distribution = self.distribution_output.distribution(outputs.regression_outputs)
2103
+
2104
+ # get samples
2105
+ samples = [
2106
+ distribution.sample() for _ in range(num_parallel_samples)
2107
+ ] # samples: list of [batch_size x num_targets]
2108
+ # stack tensors
2109
+ # [batch_size x num_samples x num_targets]
2110
+ samples = torch.stack(samples, dim=1).view(-1, num_parallel_samples, self.config.num_targets)
2111
+ return SamplePatchTSMixerRegressionOutput(sequences=samples)
2112
+
2113
+
2114
+ __all__ = [
2115
+ "PatchTSMixerPreTrainedModel",
2116
+ "PatchTSMixerModel",
2117
+ "PatchTSMixerForPretraining",
2118
+ "PatchTSMixerForPrediction",
2119
+ "PatchTSMixerForTimeSeriesClassification",
2120
+ "PatchTSMixerForRegression",
2121
+ ]
LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/timesformer/configuration_timesformer.py ADDED
@@ -0,0 +1,66 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2022 The HuggingFace Inc. team. All rights reserved.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ """TimeSformer model configuration"""
15
+
16
+ from huggingface_hub.dataclasses import strict
17
+
18
+ from ...configuration_utils import PreTrainedConfig
19
+ from ...utils import auto_docstring
20
+
21
+
22
+ @auto_docstring(checkpoint="facebook/timesformer-base-finetuned-k600")
23
+ @strict
24
+ class TimesformerConfig(PreTrainedConfig):
25
+ r"""
26
+ num_frames (`int`, *optional*, defaults to 8):
27
+ The number of frames in each video.
28
+ attention_type (`str`, *optional*, defaults to `"divided_space_time"`):
29
+ The attention type to use. Must be one of `"divided_space_time"`, `"space_only"`, `"joint_space_time"`.
30
+
31
+ Example:
32
+
33
+ ```python
34
+ >>> from transformers import TimesformerConfig, TimesformerModel
35
+
36
+ >>> # Initializing a TimeSformer timesformer-base style configuration
37
+ >>> configuration = TimesformerConfig()
38
+
39
+ >>> # Initializing a model from the configuration
40
+ >>> model = TimesformerModel(configuration)
41
+
42
+ >>> # Accessing the model configuration
43
+ >>> configuration = model.config
44
+ ```"""
45
+
46
+ model_type = "timesformer"
47
+
48
+ image_size: int | list[int] | tuple[int, int] = 224
49
+ patch_size: int | list[int] | tuple[int, int] = 16
50
+ num_channels: int = 3
51
+ num_frames: int = 8
52
+ hidden_size: int = 768
53
+ num_hidden_layers: int = 12
54
+ num_attention_heads: int = 12
55
+ intermediate_size: int = 3072
56
+ hidden_act: str = "gelu"
57
+ hidden_dropout_prob: float | int = 0.0
58
+ attention_probs_dropout_prob: float | int = 0.0
59
+ initializer_range: float = 0.02
60
+ layer_norm_eps: float = 1e-6
61
+ qkv_bias: bool = True
62
+ attention_type: str = "divided_space_time"
63
+ drop_path_rate: int = 0
64
+
65
+
66
+ __all__ = ["TimesformerConfig"]