diff --git a/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 b/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 new file mode 100644 index 0000000000000000000000000000000000000000..5f95c48a556a2bdfb988337940b83e6696ba4b48 --- /dev/null +++ b/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 @@ -0,0 +1,36 @@ +[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 +[ckpt] runs/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0010000.pt step=10000 +[decode] steps128_c1024_t1p45 generated 8/256 +[decode] steps128_c1024_t1p45 generated 16/256 +[decode] steps128_c1024_t1p45 generated 24/256 +[decode] steps128_c1024_t1p45 generated 32/256 +[decode] steps128_c1024_t1p45 generated 40/256 +[decode] steps128_c1024_t1p45 generated 48/256 +[decode] steps128_c1024_t1p45 generated 56/256 +[decode] steps128_c1024_t1p45 generated 64/256 +[decode] steps128_c1024_t1p45 generated 72/256 +[decode] steps128_c1024_t1p45 generated 80/256 +[decode] steps128_c1024_t1p45 generated 88/256 +[decode] steps128_c1024_t1p45 generated 96/256 +[decode] steps128_c1024_t1p45 generated 104/256 +[decode] steps128_c1024_t1p45 generated 112/256 +[decode] steps128_c1024_t1p45 generated 120/256 +[decode] steps128_c1024_t1p45 generated 128/256 +[decode] steps128_c1024_t1p45 generated 136/256 +[decode] steps128_c1024_t1p45 generated 144/256 +[decode] steps128_c1024_t1p45 generated 152/256 +[decode] steps128_c1024_t1p45 generated 160/256 +[decode] steps128_c1024_t1p45 generated 168/256 +[decode] steps128_c1024_t1p45 generated 176/256 +[decode] steps128_c1024_t1p45 generated 184/256 +[decode] steps128_c1024_t1p45 generated 192/256 +[decode] steps128_c1024_t1p45 generated 200/256 +[decode] steps128_c1024_t1p45 generated 208/256 +[decode] steps128_c1024_t1p45 generated 216/256 +[decode] steps128_c1024_t1p45 generated 224/256 +[decode] steps128_c1024_t1p45 generated 232/256 +[decode] steps128_c1024_t1p45 generated 240/256 +[decode] steps128_c1024_t1p45 generated 248/256 +[decode] steps128_c1024_t1p45 generated 256/256 +[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} +[watch-classic-1k] 2026-05-23_20:33:13 done step_0010000 diff --git a/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 b/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 new file mode 100644 index 0000000000000000000000000000000000000000..9bcac00856a73a014d1bc4c3cf4e156b13eaffa0 --- /dev/null +++ b/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 @@ -0,0 +1,36 @@ +[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 +[ckpt] runs/lta_lm1b_classic_dirichlet_len512_gbs512_4gpu_20k_save1k_20260523/step_0012000.pt step=12000 +[decode] steps128_c1024_t1p45 generated 8/256 +[decode] steps128_c1024_t1p45 generated 16/256 +[decode] steps128_c1024_t1p45 generated 24/256 +[decode] steps128_c1024_t1p45 generated 32/256 +[decode] steps128_c1024_t1p45 generated 40/256 +[decode] steps128_c1024_t1p45 generated 48/256 +[decode] steps128_c1024_t1p45 generated 56/256 +[decode] steps128_c1024_t1p45 generated 64/256 +[decode] steps128_c1024_t1p45 generated 72/256 +[decode] steps128_c1024_t1p45 generated 80/256 +[decode] steps128_c1024_t1p45 generated 88/256 +[decode] steps128_c1024_t1p45 generated 96/256 +[decode] steps128_c1024_t1p45 generated 104/256 +[decode] steps128_c1024_t1p45 generated 112/256 +[decode] steps128_c1024_t1p45 generated 120/256 +[decode] steps128_c1024_t1p45 generated 128/256 +[decode] steps128_c1024_t1p45 generated 136/256 +[decode] steps128_c1024_t1p45 generated 144/256 +[decode] steps128_c1024_t1p45 generated 152/256 +[decode] steps128_c1024_t1p45 generated 160/256 +[decode] steps128_c1024_t1p45 generated 168/256 +[decode] steps128_c1024_t1p45 generated 176/256 +[decode] steps128_c1024_t1p45 generated 184/256 +[decode] steps128_c1024_t1p45 generated 192/256 +[decode] steps128_c1024_t1p45 generated 200/256 +[decode] steps128_c1024_t1p45 generated 208/256 +[decode] steps128_c1024_t1p45 generated 216/256 +[decode] steps128_c1024_t1p45 generated 224/256 +[decode] steps128_c1024_t1p45 generated 232/256 +[decode] steps128_c1024_t1p45 generated 240/256 +[decode] steps128_c1024_t1p45 generated 248/256 +[decode] steps128_c1024_t1p45 generated 256/256 +[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} +[watch-classic-1k] 2026-05-23_21:17:11 done step_0012000 diff --git a/LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/processed_lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523.txt b/LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/processed_lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523.txt new file mode 100644 index 0000000000000000000000000000000000000000..e27fb1ed2a0cb7c990650f313292b08ed5e75da8 --- /dev/null +++ b/LTA_openwebtext_dualt/logs/lm1b_classic_dirichlet_every1k_infer_watch/processed_lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523.txt @@ -0,0 +1,5 @@ +runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0001000.pt +runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0002000.pt +runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0003000.pt +runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0004000.pt +runs/lta_lm1b_classic_dirichlet_len256_gbs512_4gpu_10k_save1k_20260523/step_0005000.pt diff --git a/LTA_openwebtext_dualt/logs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717.log.nohup b/LTA_openwebtext_dualt/logs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717.log.nohup new file mode 100644 index 0000000000000000000000000000000000000000..1f43e1302c87eb7a65eb92be123500a7aa4f4a22 --- /dev/null +++ b/LTA_openwebtext_dualt/logs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717.log.nohup @@ -0,0 +1,205 @@ +[launch] method=categorical_fullvocab_c1024_fullycoupled host=di-20260411014000-djqhq time=2026-05-13T18:47:17+00:00 +[launch] cwd=/e2e-data/evad-tech-vla/wanghan58/workspace/LTA_openwebtext_dualt +[launch] run_name=lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717 +[launch] save_dir=runs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717 +[launch] log_file=logs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717.log +NCCL version 2.25.1+cuda12.8 +{ + "device": "cuda:0", + "rank": 0, + "world_size": 4, + "samples": "wrapped_stream", + "vocab_size": 30522, + "tokenizer_vocab_size": 30522, + "save_dir": "runs/lta_lm1b_c1024_fullycoupled_4gpu_10k_basin_20260513_184717", + "batch_size": 64, + "grad_accum": 2, + "effective_batch_size": 512, + "global_batch_size": 512, + "lr_schedule": "constant_warmup", + "optimizer": "adamw", + "warmup_steps": 2500, + "min_lr": 6e-05, + "weight_decay": 0.0, + "adamw_param_groups": "nanogpt", + "adam_beta1": 0.9, + "adam_beta2": 0.999, + "adam_eps": 1e-08, + "muon_momentum": 0.95, + "muon_ns_steps": 5, + "muon_update_scale": 1.0, + "ema_decay": 0.0, + "ema_start_step": 0, + "model_type": "ddit", + "dual_t": true, + "corrupt_t_mode": "same", + "corrupt_min_t": 0.0, + "corrupt_max_t": 1.0, + "prefix_block_prob": 0.0, + "prefix_block_len": 128, + "mask_ratio_floor_schedule": "none", + "dirichlet_endpoint_mode": "categorical_dual_t", + "dirichlet_semantic_t_mode": "same", + "dirichlet_semantic_t_value": 0.0, + "endpoint_sequence_random_prob_alpha": 0.0, + "categorical_wrong_from_full_vocab": true, + "categorical_wrong_from_batch_valid_tokens": false, + "mask_mixture_original_prob": 0.0, + "mask_mixture_lowk_prob": 0.0, + "mask_mixture_lowcorrupt_prob": 0.0, + "mask_mixture_block_prob": 0.0, + "mask_mixture_all_prob": 0.0, + "mask_mixture_lowk_clean_tokens": "1,2,4,8,16,32,64", + "mask_mixture_lowcorrupt_tokens": "1,2,4,8,16,32,64", + "mask_mixture_block_tokens": "64,128", + "simplex_bridge_sampler": "dirichlet", + "logistic_normal_sigma_min": 0.18, + "logistic_normal_sigma_max": 2.2, + "logistic_normal_tau_min": 0.65, + "logistic_normal_tau_max": 1.15, + "torch_compile": false, + "compile_mode": "max-autotune", + "state_format": "prob", + "target_loss": "hard_ce", + "meanflow_weight": 0.0, + "rollout_train_prob": 0.0, + "rollout_train_steps": 1, + "rollout_train_infer_steps": 64, + "rollout_train_temp": 1.45, + "rollout_train_max_gamma": 1.0, + "rollout_train_corrupt_only": true, + "rollout_train_samplewise": false, + "rollout_train_compute_always": false, + "bridge_noise_init": "logistic_normal", + "noise_sigma": -1.0, + "allow_tf32": true, + "activation_checkpointing": false, + "activation_checkpoint_interval": 1, + "activation_checkpoint_scope": "block", + "ddp_static_graph": false, + "ddp_gradient_as_bucket_view": true, + "blocking_data_transfer": false, + "dataloader_prefetch_factor": 2, + "full_train_stats": false, + "record_pad_truncate": false, + "record_add_eos": false, + "record_add_special_tokens": false, + "record_pad_token": "pad", + "record_shuffle_buffer": 10000, + "wrap": true, + "wrap_mode": "stream", + "wrap_record_buffer_size": 200, + "owt_cached_chunks": false, + "owt_chunk_cache_dir": "", + "owt_chunk_cache_rebuild": false, + "owt_chunk_cache_write_batch": 4096, + "owt_exact_repeat_per_chunk": 0, + "online_chunk_shuffle": false, + "online_chunk_shuffle_buffer": 10000, + "openwebtext_split": "all", + "detokenizer": "auto", + "resolved_detokenizer": "lm1b", + "num_workers": 0, + "latest_every": 1000, + "resume_path": "" +} +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 +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 diff --git a/LTA_openwebtext_dualt/logs/lta_owt_bert_absrope_adaln_dirichlet_len1024_Cv_to_2v_mask1_sameT_gbs512_b4x4_1m_save1k_watch_20260525.log b/LTA_openwebtext_dualt/logs/lta_owt_bert_absrope_adaln_dirichlet_len1024_Cv_to_2v_mask1_sameT_gbs512_b4x4_1m_save1k_watch_20260525.log new file mode 100644 index 0000000000000000000000000000000000000000..ea503225c4b4fe44ea61a5ddd19dbf8dc0f54b13 --- /dev/null +++ b/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 +1,341 @@ + +***************************************** +Setting OMP_NUM_THREADS environment variable for each process to be 1 in default, to avoid your system being overloaded, please further tune the variable for optimal performance in your application as needed. +***************************************** +NCCL version 2.25.1+cuda12.8 +{ + "device": "cuda:0", + "rank": 0, + "world_size": 4, + "samples": "wrapped_stream", + "vocab_size": 30522, + "tokenizer_vocab_size": 30522, + "save_dir": "runs/lta_owt_bert_absrope_adaln_dirichlet_len1024_Cv_to_2v_mask1_sameT_gbs512_b4x4_1m_save1k_watch_20260525", + "max_len": 1024, + "effective_model_max_len": 1024, + "batch_size": 4, + "grad_accum": 32, + "effective_batch_size": 512, + "global_batch_size": 512, + "lr_schedule": "constant_warmup", + "optimizer": "adamw", + "epochs": 0.0, + "steps_per_epoch": 0, + "total_steps": 1000000, + "warmup_steps": 2500, + "warmup_epochs": -1.0, + "min_lr": 6e-05, + "weight_decay": 0.0, + "output_weight_decay": -1.0, + "adamw_param_groups": "nanogpt", + "adam_beta1": 0.9, + "adam_beta2": 0.999, + "adam_eps": 1e-08, + "muon_impl": "legacy", + "muon_momentum": 0.95, + "muon_ns_steps": 5, + "muon_update_scale": 1.0, + "muon_nesterov": false, + "muon_width_scale": false, + "muon_grouping": "", + "muon_param_count": 0, + "muon_adam_param_count": 0, + "muon_param_names": [], + "muon_adam_param_names": [], + "muon_effective_nesterov": false, + "muon_effective_width_scale": false, + "muon_effective_weight_decay": 0.0, + "muon_adam_fallback_nesterov": false, + "muon_adam_fallback_weight_decay": 0.0, + "ema_decay": 0.0, + "ema_start_step": 0, + "model_type": "ddit", + "ddit_mlp_type": "gelu", + "elf_num_time_tokens": 4, + "elf_num_model_mode_tokens": 0, + "abs_pos_embed": true, + "qk_norm": true, + "output_bias": false, + "output_init_std": -1.0, + "norm_type": "rmsnorm", + "target_loss": "hard_ce", + "linear_soft_target_power": 1.0, + "linear_soft_target_min_conf": 0.0, + "linear_soft_target_max_conf": 1.0, + "t_sampling_mode": "uniform", + "t_sampling_power": 1.0, + "t_sampling_eps": 0.0001, + "t_sampling_logit_mean": -1.5, + "t_sampling_logit_std": 0.8, + "t_sampling_gumbel_loc": 2.2, + "t_sampling_gumbel_scale": 0.8, + "dual_t": true, + "corrupt_t_mode": "same", + "corrupt_min_t": 0.0, + "corrupt_max_t": 1.0, + "prefix_block_prob": 0.0, + "prefix_block_len": 128, + "block_ar_two_stream": false, + "block_ar_block_len": 128, + "mask_ratio_floor_schedule": "none", + "dirichlet_endpoint_mode": "categorical_dual_t", + "dirichlet_semantic_t_mode": "same", + "dirichlet_semantic_t_value": 0.0, + "dirichlet_semantic_t_curve": "linear", + "dirichlet_semantic_t_power": 1.0, + "dirichlet_support_t_curve": "linear", + "dirichlet_support_t_power": 1.0, + "endpoint_sequence_random_prob_alpha": 0.0, + "categorical_wrong_from_full_vocab": true, + "categorical_wrong_from_batch_valid_tokens": false, + "categorical_wrong_basin_token_ids": "", + "categorical_wrong_basin_prob": 0.0, + "categorical_wrong_unigram_prob": 0.0, + "categorical_wrong_uniform_prob": 0.0, + "categorical_wrong_prob_floor": 0.0, + "categorical_gold_prob_floor": 0.0, + "categorical_gold_prob_ceil": 1.0, + "categorical_wrong_corpus_unigram_path": "", + "categorical_wrong_corpus_unigram_alpha": 1.0, + "categorical_wrong_basin_shared_prob": 0.0, + "categorical_wrong_unigram_shared_prob": 0.0, + "mask_mixture_original_prob": 0.0, + "mask_mixture_lowk_prob": 0.0, + "mask_mixture_lowcorrupt_prob": 0.0, + "mask_mixture_block_prob": 0.0, + "mask_mixture_all_prob": 0.0, + "mask_mixture_lowk_clean_tokens": "1,2,4,8,16,32,64", + "mask_mixture_lowcorrupt_tokens": "1,2,4,8,16,32,64", + "mask_mixture_block_tokens": "64,128", + "simplex_bridge_sampler": "dirichlet", + "logistic_normal_sigma_min": 0.18, + "logistic_normal_sigma_max": 2.2, + "logistic_normal_tau_min": 0.65, + "logistic_normal_tau_max": 1.15, + "torch_compile": false, + "compile_mode": "max-autotune", + "state_format": "prob", + "meanflow_weight": 0.0, + "rollout_train_prob": 0.0, + "rollout_train_steps": 1, + "rollout_train_steps_min": -1, + "rollout_train_infer_steps": 64, + "rollout_train_time_mode": "fixed_steps", + "rollout_train_s_dist": "uniform", + "rollout_train_s_min_frac": 0.0, + "rollout_train_s_max_frac": 0.125, + "rollout_train_s_beta_alpha": 2.0, + "rollout_train_s_beta_beta": 6.0, + "rollout_train_temp": 1.0, + "rollout_train_max_gamma": 1.0, + "rollout_train_rule": "flowmap", + "rollout_train_corrupt_only": true, + "rollout_train_samplewise": false, + "rollout_train_compute_always": false, + "rollout_train_keep_grad": false, + "rollout_train_sync_t": false, + "rollout_train_state_mix_mode": "final", + "rollout_train_state_mix_alpha": 0.5, + "bridge_noise_init": "logistic_normal", + "noise_sigma": -1.0, + "allow_tf32": true, + "activation_checkpointing": false, + "activation_checkpoint_interval": 1, + "activation_checkpoint_scope": "block", + "ddp_static_graph": false, + "ddp_gradient_as_bucket_view": true, + "blocking_data_transfer": false, + "dataloader_prefetch_factor": 2, + "full_train_stats": false, + "tokenized_hf": false, + "tokenized_pad_token": "pad", + "elf_conditional_hf": false, + "record_pad_truncate": false, + "record_add_eos": false, + "record_add_special_tokens": false, + "record_pad_token": "pad", + "record_shuffle_buffer": 10000, + "wrap": true, + "wrap_mode": "stream", + "wrap_record_buffer_size": 200, + "owt_cached_chunks": false, + "owt_chunk_cache_dir": "", + "owt_chunk_cache_rebuild": false, + "owt_chunk_cache_write_batch": 4096, + "owt_exact_repeat_per_chunk": 0, + "online_chunk_shuffle": false, + "online_chunk_shuffle_buffer": 10000, + "openwebtext_split": "train_minus_100k", + "detokenizer": "auto", + "resolved_detokenizer": null, + "num_workers": 0, + "latest_every": 1000, + "resume_path": "" +} +step=100 micro_steps=3200 elapsed=249.2s lr=1.212000e-05 loss=10.1760 loss_recon=10.1760 loss_meanflow=0.0000 mean_model_t=0.4997 mean_corrupt_t=0.4997 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.3481 corrupt_frac=1.0000 acc_corrupt=0.3481 loss_corrupt=10.1760 wrong_frac=0.5004 init_acc_corrupt=0.4996 acc_corrupt_t_0p0_0p2=0.0496 corrupt_frac_t_0p0_0p2=0.3376 acc_corrupt_t_0p2_0p4=0.1893 corrupt_frac_t_0p2_0p4=0.3320 acc_corrupt_t_0p8_1p0=0.6683 corrupt_frac_t_0p8_1p0=0.3432 out_w_norm=0.9225 out_g_norm=0.6925 acc_corrupt_t_0p6_0p8=0.4933 corrupt_frac_t_0p6_0p8=0.3421 acc_corrupt_t_0p4_0p6=0.3393 corrupt_frac_t_0p4_0p6=0.3380 loss_all=9.7138 init_gold_top10=0.3188 init_gold_top100=0.3206 +step=200 micro_steps=6400 elapsed=248.2s lr=2.412000e-05 loss=8.8807 loss_recon=8.8807 loss_meanflow=0.0000 mean_model_t=0.5019 mean_corrupt_t=0.5019 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.0485 corrupt_frac=1.0000 acc_corrupt=0.0485 loss_corrupt=8.8807 wrong_frac=0.4982 init_acc_corrupt=0.5018 acc_corrupt_t_0p0_0p2=0.0436 corrupt_frac_t_0p0_0p2=0.3396 acc_corrupt_t_0p4_0p6=0.0459 corrupt_frac_t_0p4_0p6=0.3392 acc_corrupt_t_0p6_0p8=0.0515 corrupt_frac_t_0p6_0p8=0.3421 out_w_norm=6.8523 out_g_norm=1.3715 acc_corrupt_t_0p2_0p4=0.0436 corrupt_frac_t_0p2_0p4=0.3354 acc_corrupt_t_0p8_1p0=0.0579 corrupt_frac_t_0p8_1p0=0.3377 loss_all=8.0119 init_gold_top10=0.5413 init_gold_top100=0.5425 +step=300 micro_steps=9600 elapsed=248.3s lr=3.612000e-05 loss=7.3930 loss_recon=7.3930 loss_meanflow=0.0000 mean_model_t=0.4951 mean_corrupt_t=0.4951 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.0520 corrupt_frac=1.0000 acc_corrupt=0.0520 loss_corrupt=7.3930 wrong_frac=0.5051 init_acc_corrupt=0.4949 acc_corrupt_t_0p2_0p4=0.0463 corrupt_frac_t_0p2_0p4=0.3356 acc_corrupt_t_0p6_0p8=0.0565 corrupt_frac_t_0p6_0p8=0.3379 out_w_norm=12.6071 out_g_norm=0.8735 acc_corrupt_t_0p4_0p6=0.0506 corrupt_frac_t_0p4_0p6=0.3421 acc_corrupt_t_0p8_1p0=0.0637 corrupt_frac_t_0p8_1p0=0.3355 acc_corrupt_t_0p0_0p2=0.0435 corrupt_frac_t_0p0_0p2=0.3427 loss_all=7.2145 init_gold_top10=0.7412 init_gold_top100=0.7422 +step=400 micro_steps=12800 elapsed=248.1s lr=4.812000e-05 loss=6.5192 loss_recon=6.5192 loss_meanflow=0.0000 mean_model_t=0.4988 mean_corrupt_t=0.4988 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.1602 corrupt_frac=1.0000 acc_corrupt=0.1602 loss_corrupt=6.5192 wrong_frac=0.5010 init_acc_corrupt=0.4990 acc_corrupt_t_0p4_0p6=0.1576 corrupt_frac_t_0p4_0p6=0.3416 acc_corrupt_t_0p8_1p0=0.2668 corrupt_frac_t_0p8_1p0=0.3396 out_w_norm=15.3591 out_g_norm=0.7304 acc_corrupt_t_0p2_0p4=0.1080 corrupt_frac_t_0p2_0p4=0.3391 acc_corrupt_t_0p6_0p8=0.2112 corrupt_frac_t_0p6_0p8=0.3371 acc_corrupt_t_0p0_0p2=0.0589 corrupt_frac_t_0p0_0p2=0.3379 loss_all=5.8103 init_gold_top10=0.5020 init_gold_top100=0.5044 +step=500 micro_steps=16000 elapsed=248.4s lr=6.012000e-05 loss=4.3971 loss_recon=4.3971 loss_meanflow=0.0000 mean_model_t=0.4997 mean_corrupt_t=0.4997 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.4498 corrupt_frac=1.0000 acc_corrupt=0.4498 loss_corrupt=4.3971 wrong_frac=0.5002 init_acc_corrupt=0.4998 acc_corrupt_t_0p0_0p2=0.0961 corrupt_frac_t_0p0_0p2=0.3383 acc_corrupt_t_0p2_0p4=0.2668 corrupt_frac_t_0p2_0p4=0.3341 acc_corrupt_t_0p4_0p6=0.4461 corrupt_frac_t_0p4_0p6=0.3349 out_w_norm=18.6346 out_g_norm=0.4967 acc_corrupt_t_0p6_0p8=0.6255 corrupt_frac_t_0p6_0p8=0.3443 acc_corrupt_t_0p8_1p0=0.8125 corrupt_frac_t_0p8_1p0=0.3424 loss_all=3.1059 init_gold_top10=0.6406 init_gold_top100=0.6423 +step=600 micro_steps=19200 elapsed=248.5s lr=7.212000e-05 loss=4.0138 loss_recon=4.0138 loss_meanflow=0.0000 mean_model_t=0.5010 mean_corrupt_t=0.5010 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.4983 corrupt_frac=1.0000 acc_corrupt=0.4983 loss_corrupt=4.0138 wrong_frac=0.4990 init_acc_corrupt=0.5010 acc_corrupt_t_0p0_0p2=0.1044 corrupt_frac_t_0p0_0p2=0.3321 acc_corrupt_t_0p2_0p4=0.3008 corrupt_frac_t_0p2_0p4=0.3488 acc_corrupt_t_0p4_0p6=0.4976 corrupt_frac_t_0p4_0p6=0.3406 acc_corrupt_t_0p8_1p0=0.8893 corrupt_frac_t_0p8_1p0=0.3433 out_w_norm=20.7167 out_g_norm=0.5580 acc_corrupt_t_0p6_0p8=0.6944 corrupt_frac_t_0p6_0p8=0.3394 loss_all=4.8030 init_gold_top10=0.4062 init_gold_top100=0.4077 +step=700 micro_steps=22400 elapsed=248.2s lr=8.412000e-05 loss=3.9537 loss_recon=3.9537 loss_meanflow=0.0000 mean_model_t=0.4985 mean_corrupt_t=0.4985 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5018 corrupt_frac=1.0000 acc_corrupt=0.5018 loss_corrupt=3.9537 wrong_frac=0.5014 init_acc_corrupt=0.4986 acc_corrupt_t_0p2_0p4=0.3049 corrupt_frac_t_0p2_0p4=0.3357 acc_corrupt_t_0p6_0p8=0.6987 corrupt_frac_t_0p6_0p8=0.3427 acc_corrupt_t_0p8_1p0=0.8964 corrupt_frac_t_0p8_1p0=0.3397 out_w_norm=21.7515 out_g_norm=0.5065 acc_corrupt_t_0p0_0p2=0.1125 corrupt_frac_t_0p0_0p2=0.3458 acc_corrupt_t_0p4_0p6=0.5031 corrupt_frac_t_0p4_0p6=0.3367 loss_all=4.4399 init_gold_top10=0.4092 init_gold_top100=0.4104 +step=800 micro_steps=25600 elapsed=247.9s lr=9.612000e-05 loss=3.8627 loss_recon=3.8627 loss_meanflow=0.0000 mean_model_t=0.4977 mean_corrupt_t=0.4977 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5052 corrupt_frac=1.0000 acc_corrupt=0.5052 loss_corrupt=3.8627 wrong_frac=0.5024 init_acc_corrupt=0.4977 acc_corrupt_t_0p0_0p2=0.1189 corrupt_frac_t_0p0_0p2=0.3425 acc_corrupt_t_0p2_0p4=0.3115 corrupt_frac_t_0p2_0p4=0.3392 acc_corrupt_t_0p4_0p6=0.5078 corrupt_frac_t_0p4_0p6=0.3418 out_w_norm=22.9196 out_g_norm=0.5135 acc_corrupt_t_0p6_0p8=0.7040 corrupt_frac_t_0p6_0p8=0.3384 acc_corrupt_t_0p8_1p0=0.8980 corrupt_frac_t_0p8_1p0=0.3355 loss_all=3.9068 init_gold_top10=0.4790 init_gold_top100=0.4800 +step=900 micro_steps=28800 elapsed=246.4s lr=1.081200e-04 loss=3.6869 loss_recon=3.6869 loss_meanflow=0.0000 mean_model_t=0.5051 mean_corrupt_t=0.5051 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5195 corrupt_frac=1.0000 acc_corrupt=0.5195 loss_corrupt=3.6869 wrong_frac=0.4951 init_acc_corrupt=0.5049 acc_corrupt_t_0p0_0p2=0.1210 corrupt_frac_t_0p0_0p2=0.3346 acc_corrupt_t_0p2_0p4=0.3184 corrupt_frac_t_0p2_0p4=0.3350 acc_corrupt_t_0p6_0p8=0.7106 corrupt_frac_t_0p6_0p8=0.3428 out_w_norm=24.2278 out_g_norm=0.4823 acc_corrupt_t_0p4_0p6=0.5184 corrupt_frac_t_0p4_0p6=0.3490 acc_corrupt_t_0p8_1p0=0.9050 corrupt_frac_t_0p8_1p0=0.3358 loss_all=1.9924 init_gold_top10=0.7144 init_gold_top100=0.7151 +step=1000 micro_steps=32000 elapsed=247.9s lr=1.201200e-04 loss=3.6043 loss_recon=3.6043 loss_meanflow=0.0000 mean_model_t=0.4991 mean_corrupt_t=0.4991 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5218 corrupt_frac=1.0000 acc_corrupt=0.5218 loss_corrupt=3.6043 wrong_frac=0.5009 init_acc_corrupt=0.4991 acc_corrupt_t_0p0_0p2=0.1252 corrupt_frac_t_0p0_0p2=0.3374 acc_corrupt_t_0p2_0p4=0.3337 corrupt_frac_t_0p2_0p4=0.3361 acc_corrupt_t_0p4_0p6=0.5299 corrupt_frac_t_0p4_0p6=0.3381 acc_corrupt_t_0p6_0p8=0.7203 corrupt_frac_t_0p6_0p8=0.3358 out_w_norm=25.4683 out_g_norm=0.5254 acc_corrupt_t_0p8_1p0=0.9046 corrupt_frac_t_0p8_1p0=0.3344 loss_all=5.2641 init_gold_top10=0.2729 init_gold_top100=0.2747 +step=1100 micro_steps=35200 elapsed=488.3s lr=1.321200e-04 loss=3.4953 loss_recon=3.4953 loss_meanflow=0.0000 mean_model_t=0.5000 mean_corrupt_t=0.5000 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5296 corrupt_frac=1.0000 acc_corrupt=0.5296 loss_corrupt=3.4953 wrong_frac=0.4999 init_acc_corrupt=0.5001 acc_corrupt_t_0p0_0p2=0.1288 corrupt_frac_t_0p0_0p2=0.3361 acc_corrupt_t_0p2_0p4=0.3405 corrupt_frac_t_0p2_0p4=0.3411 acc_corrupt_t_0p8_1p0=0.9114 corrupt_frac_t_0p8_1p0=0.3380 out_w_norm=26.8087 out_g_norm=0.4416 acc_corrupt_t_0p6_0p8=0.7280 corrupt_frac_t_0p6_0p8=0.3421 acc_corrupt_t_0p4_0p6=0.5384 corrupt_frac_t_0p4_0p6=0.3452 loss_all=4.7770 init_gold_top10=0.3320 init_gold_top100=0.3340 +step=1200 micro_steps=38400 elapsed=335.5s lr=1.441200e-04 loss=3.4037 loss_recon=3.4037 loss_meanflow=0.0000 mean_model_t=0.5057 mean_corrupt_t=0.5057 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5381 corrupt_frac=1.0000 acc_corrupt=0.5381 loss_corrupt=3.4037 wrong_frac=0.4940 init_acc_corrupt=0.5060 acc_corrupt_t_0p0_0p2=0.1293 corrupt_frac_t_0p0_0p2=0.3386 acc_corrupt_t_0p4_0p6=0.5446 corrupt_frac_t_0p4_0p6=0.3441 acc_corrupt_t_0p8_1p0=0.9125 corrupt_frac_t_0p8_1p0=0.3428 out_w_norm=28.3073 out_g_norm=0.4268 acc_corrupt_t_0p2_0p4=0.3447 corrupt_frac_t_0p2_0p4=0.3299 acc_corrupt_t_0p6_0p8=0.7335 corrupt_frac_t_0p6_0p8=0.3453 loss_all=2.6712 init_gold_top10=0.5635 init_gold_top100=0.5647 +step=1300 micro_steps=41600 elapsed=248.2s lr=1.561200e-04 loss=3.3806 loss_recon=3.3806 loss_meanflow=0.0000 mean_model_t=0.4999 mean_corrupt_t=0.4999 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5368 corrupt_frac=1.0000 acc_corrupt=0.5368 loss_corrupt=3.3806 wrong_frac=0.5001 init_acc_corrupt=0.5000 acc_corrupt_t_0p0_0p2=0.1337 corrupt_frac_t_0p0_0p2=0.3341 acc_corrupt_t_0p2_0p4=0.3498 corrupt_frac_t_0p2_0p4=0.3342 acc_corrupt_t_0p4_0p6=0.5497 corrupt_frac_t_0p4_0p6=0.3443 out_w_norm=29.9264 out_g_norm=0.4352 acc_corrupt_t_0p6_0p8=0.7372 corrupt_frac_t_0p6_0p8=0.3400 acc_corrupt_t_0p8_1p0=0.9136 corrupt_frac_t_0p8_1p0=0.3367 loss_all=3.6206 init_gold_top10=0.4763 init_gold_top100=0.4775 +step=1400 micro_steps=44800 elapsed=247.9s lr=1.681200e-04 loss=3.3366 loss_recon=3.3366 loss_meanflow=0.0000 mean_model_t=0.5021 mean_corrupt_t=0.5021 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5403 corrupt_frac=1.0000 acc_corrupt=0.5403 loss_corrupt=3.3366 wrong_frac=0.4980 init_acc_corrupt=0.5020 acc_corrupt_t_0p2_0p4=0.3509 corrupt_frac_t_0p2_0p4=0.3319 acc_corrupt_t_0p4_0p6=0.5510 corrupt_frac_t_0p4_0p6=0.3464 out_w_norm=31.7308 out_g_norm=0.3559 acc_corrupt_t_0p0_0p2=0.1339 corrupt_frac_t_0p0_0p2=0.3388 acc_corrupt_t_0p8_1p0=0.9161 corrupt_frac_t_0p8_1p0=0.3444 acc_corrupt_t_0p6_0p8=0.7406 corrupt_frac_t_0p6_0p8=0.3405 loss_all=4.6098 init_gold_top10=0.3313 init_gold_top100=0.3323 +step=1500 micro_steps=48000 elapsed=246.1s lr=1.801200e-04 loss=3.3254 loss_recon=3.3254 loss_meanflow=0.0000 mean_model_t=0.4988 mean_corrupt_t=0.4988 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5395 corrupt_frac=1.0000 acc_corrupt=0.5395 loss_corrupt=3.3254 wrong_frac=0.5013 init_acc_corrupt=0.4987 acc_corrupt_t_0p0_0p2=0.1370 corrupt_frac_t_0p0_0p2=0.3388 acc_corrupt_t_0p2_0p4=0.3511 corrupt_frac_t_0p2_0p4=0.3384 acc_corrupt_t_0p4_0p6=0.5571 corrupt_frac_t_0p4_0p6=0.3339 out_w_norm=33.7679 out_g_norm=0.3680 acc_corrupt_t_0p8_1p0=0.9172 corrupt_frac_t_0p8_1p0=0.3386 acc_corrupt_t_0p6_0p8=0.7393 corrupt_frac_t_0p6_0p8=0.3361 loss_all=3.2991 init_gold_top10=0.4912 init_gold_top100=0.4929 +step=1600 micro_steps=51200 elapsed=248.5s lr=1.921200e-04 loss=3.2735 loss_recon=3.2735 loss_meanflow=0.0000 mean_model_t=0.5031 mean_corrupt_t=0.5031 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5456 corrupt_frac=1.0000 acc_corrupt=0.5456 loss_corrupt=3.2735 wrong_frac=0.4969 init_acc_corrupt=0.5031 acc_corrupt_t_0p4_0p6=0.5580 corrupt_frac_t_0p4_0p6=0.3346 acc_corrupt_t_0p6_0p8=0.7452 corrupt_frac_t_0p6_0p8=0.3400 acc_corrupt_t_0p8_1p0=0.9152 corrupt_frac_t_0p8_1p0=0.3403 out_w_norm=35.8378 out_g_norm=0.3416 acc_corrupt_t_0p2_0p4=0.3557 corrupt_frac_t_0p2_0p4=0.3425 acc_corrupt_t_0p0_0p2=0.1380 corrupt_frac_t_0p0_0p2=0.3379 loss_all=4.0486 init_gold_top10=0.4185 init_gold_top100=0.4211 +step=1700 micro_steps=54400 elapsed=248.2s lr=2.041200e-04 loss=3.2503 loss_recon=3.2503 loss_meanflow=0.0000 mean_model_t=0.4992 mean_corrupt_t=0.4992 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5449 corrupt_frac=1.0000 acc_corrupt=0.5449 loss_corrupt=3.2503 wrong_frac=0.5006 init_acc_corrupt=0.4995 acc_corrupt_t_0p0_0p2=0.1373 corrupt_frac_t_0p0_0p2=0.3393 acc_corrupt_t_0p2_0p4=0.3576 corrupt_frac_t_0p2_0p4=0.3378 acc_corrupt_t_0p8_1p0=0.9181 corrupt_frac_t_0p8_1p0=0.3383 out_w_norm=37.9119 out_g_norm=0.2909 acc_corrupt_t_0p4_0p6=0.5623 corrupt_frac_t_0p4_0p6=0.3370 acc_corrupt_t_0p6_0p8=0.7470 corrupt_frac_t_0p6_0p8=0.3361 loss_all=3.8740 init_gold_top10=0.4106 init_gold_top100=0.4131 +step=1800 micro_steps=57600 elapsed=247.8s lr=2.161200e-04 loss=3.2161 loss_recon=3.2161 loss_meanflow=0.0000 mean_model_t=0.5004 mean_corrupt_t=0.5004 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5481 corrupt_frac=1.0000 acc_corrupt=0.5481 loss_corrupt=3.2161 wrong_frac=0.4996 init_acc_corrupt=0.5004 acc_corrupt_t_0p2_0p4=0.3607 corrupt_frac_t_0p2_0p4=0.3391 acc_corrupt_t_0p4_0p6=0.5650 corrupt_frac_t_0p4_0p6=0.3376 acc_corrupt_t_0p8_1p0=0.9205 corrupt_frac_t_0p8_1p0=0.3343 out_w_norm=39.9268 out_g_norm=0.2822 acc_corrupt_t_0p0_0p2=0.1406 corrupt_frac_t_0p0_0p2=0.3391 acc_corrupt_t_0p6_0p8=0.7526 corrupt_frac_t_0p6_0p8=0.3376 loss_all=2.6583 init_gold_top10=0.5962 init_gold_top100=0.5977 +step=1900 micro_steps=60800 elapsed=247.8s lr=2.281200e-04 loss=3.1648 loss_recon=3.1648 loss_meanflow=0.0000 mean_model_t=0.5061 mean_corrupt_t=0.5061 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5544 corrupt_frac=1.0000 acc_corrupt=0.5544 loss_corrupt=3.1648 wrong_frac=0.4941 init_acc_corrupt=0.5059 acc_corrupt_t_0p2_0p4=0.3626 corrupt_frac_t_0p2_0p4=0.3370 acc_corrupt_t_0p6_0p8=0.7544 corrupt_frac_t_0p6_0p8=0.3422 acc_corrupt_t_0p8_1p0=0.9200 corrupt_frac_t_0p8_1p0=0.3429 out_w_norm=41.8541 out_g_norm=0.2738 acc_corrupt_t_0p0_0p2=0.1400 corrupt_frac_t_0p0_0p2=0.3341 acc_corrupt_t_0p4_0p6=0.5694 corrupt_frac_t_0p4_0p6=0.3361 loss_all=4.7066 init_gold_top10=0.3164 init_gold_top100=0.3191 +step=2000 micro_steps=64000 elapsed=247.8s lr=2.401200e-04 loss=3.1913 loss_recon=3.1913 loss_meanflow=0.0000 mean_model_t=0.4982 mean_corrupt_t=0.4982 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5486 corrupt_frac=1.0000 acc_corrupt=0.5486 loss_corrupt=3.1913 wrong_frac=0.5018 init_acc_corrupt=0.4982 acc_corrupt_t_0p0_0p2=0.1391 corrupt_frac_t_0p0_0p2=0.3415 acc_corrupt_t_0p4_0p6=0.5699 corrupt_frac_t_0p4_0p6=0.3399 acc_corrupt_t_0p6_0p8=0.7570 corrupt_frac_t_0p6_0p8=0.3446 acc_corrupt_t_0p8_1p0=0.9219 corrupt_frac_t_0p8_1p0=0.3391 out_w_norm=43.7528 out_g_norm=0.2778 acc_corrupt_t_0p2_0p4=0.3622 corrupt_frac_t_0p2_0p4=0.3445 loss_all=2.8599 init_gold_top10=0.5359 init_gold_top100=0.5371 +step=2100 micro_steps=67200 elapsed=468.3s lr=2.521200e-04 loss=3.1438 loss_recon=3.1438 loss_meanflow=0.0000 mean_model_t=0.5000 mean_corrupt_t=0.5000 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5527 corrupt_frac=1.0000 acc_corrupt=0.5527 loss_corrupt=3.1438 wrong_frac=0.5000 init_acc_corrupt=0.5000 acc_corrupt_t_0p0_0p2=0.1410 corrupt_frac_t_0p0_0p2=0.3368 acc_corrupt_t_0p6_0p8=0.7597 corrupt_frac_t_0p6_0p8=0.3358 acc_corrupt_t_0p8_1p0=0.9225 corrupt_frac_t_0p8_1p0=0.3352 out_w_norm=45.5969 out_g_norm=0.2513 acc_corrupt_t_0p4_0p6=0.5745 corrupt_frac_t_0p4_0p6=0.3395 acc_corrupt_t_0p2_0p4=0.3644 corrupt_frac_t_0p2_0p4=0.3411 loss_all=5.0482 init_gold_top10=0.2610 init_gold_top100=0.2627 +step=2200 micro_steps=70400 elapsed=337.0s lr=2.641200e-04 loss=3.1195 loss_recon=3.1195 loss_meanflow=0.0000 mean_model_t=0.5021 mean_corrupt_t=0.5021 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5562 corrupt_frac=1.0000 acc_corrupt=0.5562 loss_corrupt=3.1195 wrong_frac=0.4979 init_acc_corrupt=0.5021 acc_corrupt_t_0p0_0p2=0.1432 corrupt_frac_t_0p0_0p2=0.3343 acc_corrupt_t_0p2_0p4=0.3669 corrupt_frac_t_0p2_0p4=0.3383 acc_corrupt_t_0p4_0p6=0.5739 corrupt_frac_t_0p4_0p6=0.3379 acc_corrupt_t_0p8_1p0=0.9250 corrupt_frac_t_0p8_1p0=0.3403 out_w_norm=47.3888 out_g_norm=0.2440 acc_corrupt_t_0p6_0p8=0.7618 corrupt_frac_t_0p6_0p8=0.3363 loss_all=1.7950 init_gold_top10=0.7090 init_gold_top100=0.7100 +step=2300 micro_steps=73600 elapsed=246.9s lr=2.761200e-04 loss=3.1104 loss_recon=3.1104 loss_meanflow=0.0000 mean_model_t=0.4985 mean_corrupt_t=0.4985 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5547 corrupt_frac=1.0000 acc_corrupt=0.5547 loss_corrupt=3.1104 wrong_frac=0.5016 init_acc_corrupt=0.4985 acc_corrupt_t_0p2_0p4=0.3685 corrupt_frac_t_0p2_0p4=0.3436 acc_corrupt_t_0p4_0p6=0.5784 corrupt_frac_t_0p4_0p6=0.3360 acc_corrupt_t_0p8_1p0=0.9257 corrupt_frac_t_0p8_1p0=0.3385 out_w_norm=49.1642 out_g_norm=0.2195 acc_corrupt_t_0p0_0p2=0.1426 corrupt_frac_t_0p0_0p2=0.3390 acc_corrupt_t_0p6_0p8=0.7648 corrupt_frac_t_0p6_0p8=0.3377 loss_all=3.0261 init_gold_top10=0.4915 init_gold_top100=0.4922 +step=2400 micro_steps=76800 elapsed=246.9s lr=2.881200e-04 loss=3.0872 loss_recon=3.0872 loss_meanflow=0.0000 mean_model_t=0.4990 mean_corrupt_t=0.4990 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5569 corrupt_frac=1.0000 acc_corrupt=0.5569 loss_corrupt=3.0872 wrong_frac=0.5009 init_acc_corrupt=0.4991 acc_corrupt_t_0p0_0p2=0.1428 corrupt_frac_t_0p0_0p2=0.3371 acc_corrupt_t_0p8_1p0=0.9256 corrupt_frac_t_0p8_1p0=0.3395 out_w_norm=50.9163 out_g_norm=0.2348 acc_corrupt_t_0p2_0p4=0.3697 corrupt_frac_t_0p2_0p4=0.3401 acc_corrupt_t_0p4_0p6=0.5791 corrupt_frac_t_0p4_0p6=0.3378 acc_corrupt_t_0p6_0p8=0.7665 corrupt_frac_t_0p6_0p8=0.3422 loss_all=3.8287 init_gold_top10=0.3652 init_gold_top100=0.3667 +step=2500 micro_steps=80000 elapsed=246.8s lr=3.000000e-04 loss=3.0590 loss_recon=3.0590 loss_meanflow=0.0000 mean_model_t=0.5005 mean_corrupt_t=0.5005 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5599 corrupt_frac=1.0000 acc_corrupt=0.5599 loss_corrupt=3.0590 wrong_frac=0.4995 init_acc_corrupt=0.5006 acc_corrupt_t_0p0_0p2=0.1428 corrupt_frac_t_0p0_0p2=0.3375 acc_corrupt_t_0p4_0p6=0.5834 corrupt_frac_t_0p4_0p6=0.3427 acc_corrupt_t_0p8_1p0=0.9301 corrupt_frac_t_0p8_1p0=0.3436 out_w_norm=52.6698 out_g_norm=0.2119 acc_corrupt_t_0p2_0p4=0.3704 corrupt_frac_t_0p2_0p4=0.3386 acc_corrupt_t_0p6_0p8=0.7696 corrupt_frac_t_0p6_0p8=0.3391 loss_all=5.9741 init_gold_top10=0.1738 init_gold_top100=0.1760 +step=2600 micro_steps=83200 elapsed=246.7s lr=3.000000e-04 loss=3.0383 loss_recon=3.0383 loss_meanflow=0.0000 mean_model_t=0.5016 mean_corrupt_t=0.5016 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5626 corrupt_frac=1.0000 acc_corrupt=0.5626 loss_corrupt=3.0383 wrong_frac=0.4985 init_acc_corrupt=0.5015 acc_corrupt_t_0p0_0p2=0.1458 corrupt_frac_t_0p0_0p2=0.3388 acc_corrupt_t_0p2_0p4=0.3700 corrupt_frac_t_0p2_0p4=0.3419 acc_corrupt_t_0p8_1p0=0.9301 corrupt_frac_t_0p8_1p0=0.3396 out_w_norm=54.4178 out_g_norm=0.2087 acc_corrupt_t_0p4_0p6=0.5863 corrupt_frac_t_0p4_0p6=0.3326 acc_corrupt_t_0p6_0p8=0.7702 corrupt_frac_t_0p6_0p8=0.3355 loss_all=4.7436 init_gold_top10=0.2754 init_gold_top100=0.2793 +step=2700 micro_steps=86400 elapsed=246.7s lr=3.000000e-04 loss=3.0074 loss_recon=3.0074 loss_meanflow=0.0000 mean_model_t=0.5010 mean_corrupt_t=0.5010 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5650 corrupt_frac=1.0000 acc_corrupt=0.5650 loss_corrupt=3.0074 wrong_frac=0.4989 init_acc_corrupt=0.5011 acc_corrupt_t_0p4_0p6=0.5910 corrupt_frac_t_0p4_0p6=0.3401 acc_corrupt_t_0p8_1p0=0.9301 corrupt_frac_t_0p8_1p0=0.3389 out_w_norm=56.1186 out_g_norm=0.1931 acc_corrupt_t_0p2_0p4=0.3728 corrupt_frac_t_0p2_0p4=0.3395 acc_corrupt_t_0p0_0p2=0.1451 corrupt_frac_t_0p0_0p2=0.3400 acc_corrupt_t_0p6_0p8=0.7753 corrupt_frac_t_0p6_0p8=0.3388 loss_all=1.5063 init_gold_top10=0.7212 init_gold_top100=0.7224 +step=2800 micro_steps=89600 elapsed=246.7s lr=3.000000e-04 loss=2.9977 loss_recon=2.9977 loss_meanflow=0.0000 mean_model_t=0.5000 mean_corrupt_t=0.5000 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5657 corrupt_frac=1.0000 acc_corrupt=0.5657 loss_corrupt=2.9977 wrong_frac=0.5002 init_acc_corrupt=0.4998 acc_corrupt_t_0p2_0p4=0.3789 corrupt_frac_t_0p2_0p4=0.3447 acc_corrupt_t_0p4_0p6=0.5938 corrupt_frac_t_0p4_0p6=0.3356 acc_corrupt_t_0p8_1p0=0.9319 corrupt_frac_t_0p8_1p0=0.3399 out_w_norm=57.6794 out_g_norm=0.2051 acc_corrupt_t_0p0_0p2=0.1443 corrupt_frac_t_0p0_0p2=0.3355 acc_corrupt_t_0p6_0p8=0.7781 corrupt_frac_t_0p6_0p8=0.3386 loss_all=0.5263 init_gold_top10=0.8708 init_gold_top100=0.8713 +step=2900 micro_steps=92800 elapsed=246.8s lr=3.000000e-04 loss=3.0064 loss_recon=3.0064 loss_meanflow=0.0000 mean_model_t=0.4963 mean_corrupt_t=0.4963 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5639 corrupt_frac=1.0000 acc_corrupt=0.5639 loss_corrupt=3.0064 wrong_frac=0.5036 init_acc_corrupt=0.4964 acc_corrupt_t_0p0_0p2=0.1482 corrupt_frac_t_0p0_0p2=0.3418 acc_corrupt_t_0p2_0p4=0.3803 corrupt_frac_t_0p2_0p4=0.3400 acc_corrupt_t_0p4_0p6=0.5956 corrupt_frac_t_0p4_0p6=0.3397 acc_corrupt_t_0p8_1p0=0.9345 corrupt_frac_t_0p8_1p0=0.3383 out_w_norm=59.1790 out_g_norm=0.1948 acc_corrupt_t_0p6_0p8=0.7820 corrupt_frac_t_0p6_0p8=0.3448 loss_all=3.5409 init_gold_top10=0.4199 init_gold_top100=0.4226 +step=3000 micro_steps=96000 elapsed=247.1s lr=3.000000e-04 loss=2.9819 loss_recon=2.9819 loss_meanflow=0.0000 mean_model_t=0.4984 mean_corrupt_t=0.4984 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5675 corrupt_frac=1.0000 acc_corrupt=0.5675 loss_corrupt=2.9819 wrong_frac=0.5017 init_acc_corrupt=0.4984 acc_corrupt_t_0p0_0p2=0.1461 corrupt_frac_t_0p0_0p2=0.3438 acc_corrupt_t_0p6_0p8=0.7844 corrupt_frac_t_0p6_0p8=0.3462 acc_corrupt_t_0p8_1p0=0.9340 corrupt_frac_t_0p8_1p0=0.3447 out_w_norm=60.6061 out_g_norm=0.1776 acc_corrupt_t_0p2_0p4=0.3821 corrupt_frac_t_0p2_0p4=0.3303 acc_corrupt_t_0p4_0p6=0.5973 corrupt_frac_t_0p4_0p6=0.3356 loss_all=2.7528 init_gold_top10=0.5759 init_gold_top100=0.5769 +step=3100 micro_steps=99200 elapsed=457.7s lr=3.000000e-04 loss=2.9494 loss_recon=2.9494 loss_meanflow=0.0000 mean_model_t=0.4989 mean_corrupt_t=0.4989 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5701 corrupt_frac=1.0000 acc_corrupt=0.5701 loss_corrupt=2.9494 wrong_frac=0.5010 init_acc_corrupt=0.4990 acc_corrupt_t_0p0_0p2=0.1489 corrupt_frac_t_0p0_0p2=0.3398 acc_corrupt_t_0p4_0p6=0.6002 corrupt_frac_t_0p4_0p6=0.3372 out_w_norm=61.9549 out_g_norm=0.1739 acc_corrupt_t_0p6_0p8=0.7842 corrupt_frac_t_0p6_0p8=0.3406 acc_corrupt_t_0p8_1p0=0.9361 corrupt_frac_t_0p8_1p0=0.3351 acc_corrupt_t_0p2_0p4=0.3853 corrupt_frac_t_0p2_0p4=0.3395 loss_all=4.0096 init_gold_top10=0.3740 init_gold_top100=0.3757 +step=3200 micro_steps=102400 elapsed=355.8s lr=3.000000e-04 loss=2.9174 loss_recon=2.9174 loss_meanflow=0.0000 mean_model_t=0.5016 mean_corrupt_t=0.5016 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5738 corrupt_frac=1.0000 acc_corrupt=0.5738 loss_corrupt=2.9174 wrong_frac=0.4984 init_acc_corrupt=0.5016 acc_corrupt_t_0p0_0p2=0.1490 corrupt_frac_t_0p0_0p2=0.3407 acc_corrupt_t_0p2_0p4=0.3880 corrupt_frac_t_0p2_0p4=0.3369 acc_corrupt_t_0p4_0p6=0.6053 corrupt_frac_t_0p4_0p6=0.3376 acc_corrupt_t_0p6_0p8=0.7870 corrupt_frac_t_0p6_0p8=0.3373 out_w_norm=63.2355 out_g_norm=0.1684 acc_corrupt_t_0p8_1p0=0.9361 corrupt_frac_t_0p8_1p0=0.3386 loss_all=4.2434 init_gold_top10=0.3120 init_gold_top100=0.3137 +step=3300 micro_steps=105600 elapsed=247.2s lr=3.000000e-04 loss=2.9444 loss_recon=2.9444 loss_meanflow=0.0000 mean_model_t=0.4982 mean_corrupt_t=0.4982 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5703 corrupt_frac=1.0000 acc_corrupt=0.5703 loss_corrupt=2.9444 wrong_frac=0.5020 init_acc_corrupt=0.4980 acc_corrupt_t_0p0_0p2=0.1473 corrupt_frac_t_0p0_0p2=0.3406 acc_corrupt_t_0p2_0p4=0.3843 corrupt_frac_t_0p2_0p4=0.3361 acc_corrupt_t_0p6_0p8=0.7885 corrupt_frac_t_0p6_0p8=0.3370 out_w_norm=64.4579 out_g_norm=0.1632 acc_corrupt_t_0p4_0p6=0.6033 corrupt_frac_t_0p4_0p6=0.3385 acc_corrupt_t_0p8_1p0=0.9369 corrupt_frac_t_0p8_1p0=0.3368 loss_all=5.2116 init_gold_top10=0.2288 init_gold_top100=0.2317 +step=3400 micro_steps=108800 elapsed=246.9s lr=3.000000e-04 loss=2.9099 loss_recon=2.9099 loss_meanflow=0.0000 mean_model_t=0.5002 mean_corrupt_t=0.5002 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5747 corrupt_frac=1.0000 acc_corrupt=0.5747 loss_corrupt=2.9099 wrong_frac=0.4998 init_acc_corrupt=0.5002 acc_corrupt_t_0p0_0p2=0.1510 corrupt_frac_t_0p0_0p2=0.3421 acc_corrupt_t_0p4_0p6=0.6062 corrupt_frac_t_0p4_0p6=0.3423 acc_corrupt_t_0p6_0p8=0.7921 corrupt_frac_t_0p6_0p8=0.3354 out_w_norm=65.6783 out_g_norm=0.1636 acc_corrupt_t_0p8_1p0=0.9383 corrupt_frac_t_0p8_1p0=0.3404 acc_corrupt_t_0p2_0p4=0.3867 corrupt_frac_t_0p2_0p4=0.3379 loss_all=3.9734 init_gold_top10=0.3418 init_gold_top100=0.3440 +step=3500 micro_steps=112000 elapsed=246.9s lr=3.000000e-04 loss=2.8990 loss_recon=2.8990 loss_meanflow=0.0000 mean_model_t=0.4993 mean_corrupt_t=0.4993 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5744 corrupt_frac=1.0000 acc_corrupt=0.5744 loss_corrupt=2.8990 wrong_frac=0.5009 init_acc_corrupt=0.4991 acc_corrupt_t_0p0_0p2=0.1503 corrupt_frac_t_0p0_0p2=0.3377 acc_corrupt_t_0p4_0p6=0.6065 corrupt_frac_t_0p4_0p6=0.3412 acc_corrupt_t_0p6_0p8=0.7921 corrupt_frac_t_0p6_0p8=0.3372 out_w_norm=66.9028 out_g_norm=0.1687 acc_corrupt_t_0p2_0p4=0.3882 corrupt_frac_t_0p2_0p4=0.3360 acc_corrupt_t_0p8_1p0=0.9365 corrupt_frac_t_0p8_1p0=0.3354 loss_all=2.3108 init_gold_top10=0.5696 init_gold_top100=0.5710 +step=3600 micro_steps=115200 elapsed=247.1s lr=3.000000e-04 loss=2.8946 loss_recon=2.8946 loss_meanflow=0.0000 mean_model_t=0.5009 mean_corrupt_t=0.5009 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5758 corrupt_frac=1.0000 acc_corrupt=0.5758 loss_corrupt=2.8946 wrong_frac=0.4990 init_acc_corrupt=0.5010 acc_corrupt_t_0p8_1p0=0.9378 corrupt_frac_t_0p8_1p0=0.3410 out_w_norm=68.0967 out_g_norm=0.1641 acc_corrupt_t_0p2_0p4=0.3897 corrupt_frac_t_0p2_0p4=0.3367 acc_corrupt_t_0p4_0p6=0.6071 corrupt_frac_t_0p4_0p6=0.3372 acc_corrupt_t_0p0_0p2=0.1502 corrupt_frac_t_0p0_0p2=0.3395 acc_corrupt_t_0p6_0p8=0.7927 corrupt_frac_t_0p6_0p8=0.3361 loss_all=3.1538 init_gold_top10=0.4495 init_gold_top100=0.4509 +step=3700 micro_steps=118400 elapsed=247.1s lr=3.000000e-04 loss=2.8661 loss_recon=2.8661 loss_meanflow=0.0000 mean_model_t=0.5038 mean_corrupt_t=0.5038 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5796 corrupt_frac=1.0000 acc_corrupt=0.5796 loss_corrupt=2.8661 wrong_frac=0.4960 init_acc_corrupt=0.5040 acc_corrupt_t_0p0_0p2=0.1489 corrupt_frac_t_0p0_0p2=0.3381 acc_corrupt_t_0p4_0p6=0.6074 corrupt_frac_t_0p4_0p6=0.3416 acc_corrupt_t_0p8_1p0=0.9397 corrupt_frac_t_0p8_1p0=0.3415 out_w_norm=69.2614 out_g_norm=0.1536 acc_corrupt_t_0p6_0p8=0.7943 corrupt_frac_t_0p6_0p8=0.3374 acc_corrupt_t_0p2_0p4=0.3875 corrupt_frac_t_0p2_0p4=0.3362 loss_all=2.7075 init_gold_top10=0.5378 init_gold_top100=0.5391 +step=3800 micro_steps=121600 elapsed=247.2s lr=3.000000e-04 loss=2.8988 loss_recon=2.8988 loss_meanflow=0.0000 mean_model_t=0.4993 mean_corrupt_t=0.4993 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5754 corrupt_frac=1.0000 acc_corrupt=0.5754 loss_corrupt=2.8988 wrong_frac=0.5007 init_acc_corrupt=0.4993 acc_corrupt_t_0p0_0p2=0.1488 corrupt_frac_t_0p0_0p2=0.3414 acc_corrupt_t_0p2_0p4=0.3879 corrupt_frac_t_0p2_0p4=0.3286 acc_corrupt_t_0p4_0p6=0.6076 corrupt_frac_t_0p4_0p6=0.3381 acc_corrupt_t_0p6_0p8=0.7944 corrupt_frac_t_0p6_0p8=0.3462 out_w_norm=70.3765 out_g_norm=0.1566 acc_corrupt_t_0p8_1p0=0.9388 corrupt_frac_t_0p8_1p0=0.3407 loss_all=3.0592 init_gold_top10=0.4521 init_gold_top100=0.4534 +step=3900 micro_steps=124800 elapsed=247.1s lr=3.000000e-04 loss=2.8688 loss_recon=2.8688 loss_meanflow=0.0000 mean_model_t=0.5014 mean_corrupt_t=0.5014 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5789 corrupt_frac=1.0000 acc_corrupt=0.5789 loss_corrupt=2.8688 wrong_frac=0.4986 init_acc_corrupt=0.5014 acc_corrupt_t_0p0_0p2=0.1495 corrupt_frac_t_0p0_0p2=0.3408 acc_corrupt_t_0p6_0p8=0.7960 corrupt_frac_t_0p6_0p8=0.3426 out_w_norm=71.4988 out_g_norm=0.1492 acc_corrupt_t_0p2_0p4=0.3901 corrupt_frac_t_0p2_0p4=0.3383 acc_corrupt_t_0p4_0p6=0.6137 corrupt_frac_t_0p4_0p6=0.3369 acc_corrupt_t_0p8_1p0=0.9406 corrupt_frac_t_0p8_1p0=0.3436 loss_all=3.4048 init_gold_top10=0.4609 init_gold_top100=0.4624 +step=4000 micro_steps=128000 elapsed=246.7s lr=3.000000e-04 loss=2.8544 loss_recon=2.8544 loss_meanflow=0.0000 mean_model_t=0.5015 mean_corrupt_t=0.5015 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5799 corrupt_frac=1.0000 acc_corrupt=0.5799 loss_corrupt=2.8544 wrong_frac=0.4985 init_acc_corrupt=0.5016 acc_corrupt_t_0p0_0p2=0.1523 corrupt_frac_t_0p0_0p2=0.3387 acc_corrupt_t_0p2_0p4=0.3921 corrupt_frac_t_0p2_0p4=0.3391 acc_corrupt_t_0p6_0p8=0.7962 corrupt_frac_t_0p6_0p8=0.3365 out_w_norm=72.5892 out_g_norm=0.1451 acc_corrupt_t_0p4_0p6=0.6126 corrupt_frac_t_0p4_0p6=0.3376 acc_corrupt_t_0p8_1p0=0.9415 corrupt_frac_t_0p8_1p0=0.3387 loss_all=3.2201 init_gold_top10=0.4409 init_gold_top100=0.4429 +step=4100 micro_steps=131200 elapsed=477.3s lr=3.000000e-04 loss=2.8306 loss_recon=2.8306 loss_meanflow=0.0000 mean_model_t=0.5044 mean_corrupt_t=0.5044 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5831 corrupt_frac=1.0000 acc_corrupt=0.5831 loss_corrupt=2.8306 wrong_frac=0.4957 init_acc_corrupt=0.5043 acc_corrupt_t_0p0_0p2=0.1521 corrupt_frac_t_0p0_0p2=0.3404 acc_corrupt_t_0p4_0p6=0.6147 corrupt_frac_t_0p4_0p6=0.3437 acc_corrupt_t_0p6_0p8=0.7970 corrupt_frac_t_0p6_0p8=0.3405 out_w_norm=73.6508 out_g_norm=0.1412 acc_corrupt_t_0p2_0p4=0.3921 corrupt_frac_t_0p2_0p4=0.3368 acc_corrupt_t_0p8_1p0=0.9404 corrupt_frac_t_0p8_1p0=0.3441 loss_all=0.8856 init_gold_top10=0.7737 init_gold_top100=0.7751 +step=4200 micro_steps=134400 elapsed=357.5s lr=3.000000e-04 loss=2.8273 loss_recon=2.8273 loss_meanflow=0.0000 mean_model_t=0.5041 mean_corrupt_t=0.5041 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5837 corrupt_frac=1.0000 acc_corrupt=0.5837 loss_corrupt=2.8273 wrong_frac=0.4960 init_acc_corrupt=0.5041 acc_corrupt_t_0p0_0p2=0.1511 corrupt_frac_t_0p0_0p2=0.3414 acc_corrupt_t_0p4_0p6=0.6152 corrupt_frac_t_0p4_0p6=0.3360 acc_corrupt_t_0p6_0p8=0.7984 corrupt_frac_t_0p6_0p8=0.3428 out_w_norm=74.7083 out_g_norm=0.1416 acc_corrupt_t_0p2_0p4=0.3945 corrupt_frac_t_0p2_0p4=0.3342 acc_corrupt_t_0p8_1p0=0.9389 corrupt_frac_t_0p8_1p0=0.3380 loss_all=3.3635 init_gold_top10=0.4641 init_gold_top100=0.4663 +step=4300 micro_steps=137600 elapsed=246.1s lr=3.000000e-04 loss=2.8711 loss_recon=2.8711 loss_meanflow=0.0000 mean_model_t=0.4976 mean_corrupt_t=0.4976 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5771 corrupt_frac=1.0000 acc_corrupt=0.5771 loss_corrupt=2.8711 wrong_frac=0.5025 init_acc_corrupt=0.4975 acc_corrupt_t_0p0_0p2=0.1523 corrupt_frac_t_0p0_0p2=0.3380 acc_corrupt_t_0p2_0p4=0.3924 corrupt_frac_t_0p2_0p4=0.3370 acc_corrupt_t_0p4_0p6=0.6165 corrupt_frac_t_0p4_0p6=0.3388 out_w_norm=75.7727 out_g_norm=0.1412 acc_corrupt_t_0p6_0p8=0.7988 corrupt_frac_t_0p6_0p8=0.3360 acc_corrupt_t_0p8_1p0=0.9420 corrupt_frac_t_0p8_1p0=0.3380 loss_all=4.0927 init_gold_top10=0.3137 init_gold_top100=0.3157 +step=4400 micro_steps=140800 elapsed=247.2s lr=3.000000e-04 loss=2.8268 loss_recon=2.8268 loss_meanflow=0.0000 mean_model_t=0.5036 mean_corrupt_t=0.5036 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5838 corrupt_frac=1.0000 acc_corrupt=0.5838 loss_corrupt=2.8268 wrong_frac=0.4965 init_acc_corrupt=0.5035 acc_corrupt_t_0p2_0p4=0.3954 corrupt_frac_t_0p2_0p4=0.3386 acc_corrupt_t_0p4_0p6=0.6187 corrupt_frac_t_0p4_0p6=0.3373 acc_corrupt_t_0p8_1p0=0.9411 corrupt_frac_t_0p8_1p0=0.3466 out_w_norm=76.8194 out_g_norm=0.1348 acc_corrupt_t_0p0_0p2=0.1518 corrupt_frac_t_0p0_0p2=0.3364 acc_corrupt_t_0p6_0p8=0.8012 corrupt_frac_t_0p6_0p8=0.3330 loss_all=1.9395 init_gold_top10=0.5818 init_gold_top100=0.5830 +step=4500 micro_steps=144000 elapsed=345.3s lr=3.000000e-04 loss=2.8555 loss_recon=2.8555 loss_meanflow=0.0000 mean_model_t=0.4970 mean_corrupt_t=0.4970 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5790 corrupt_frac=1.0000 acc_corrupt=0.5790 loss_corrupt=2.8555 wrong_frac=0.5029 init_acc_corrupt=0.4972 acc_corrupt_t_0p0_0p2=0.1521 corrupt_frac_t_0p0_0p2=0.3436 acc_corrupt_t_0p2_0p4=0.3977 corrupt_frac_t_0p2_0p4=0.3374 acc_corrupt_t_0p6_0p8=0.7999 corrupt_frac_t_0p6_0p8=0.3360 out_w_norm=77.8495 out_g_norm=0.1352 acc_corrupt_t_0p8_1p0=0.9405 corrupt_frac_t_0p8_1p0=0.3354 acc_corrupt_t_0p4_0p6=0.6186 corrupt_frac_t_0p4_0p6=0.3352 loss_all=4.0323 init_gold_top10=0.3994 init_gold_top100=0.4009 +step=4600 micro_steps=147200 elapsed=650.8s lr=3.000000e-04 loss=2.8642 loss_recon=2.8642 loss_meanflow=0.0000 mean_model_t=0.4962 mean_corrupt_t=0.4962 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5779 corrupt_frac=1.0000 acc_corrupt=0.5779 loss_corrupt=2.8642 wrong_frac=0.5037 init_acc_corrupt=0.4963 acc_corrupt_t_0p0_0p2=0.1505 corrupt_frac_t_0p0_0p2=0.3410 acc_corrupt_t_0p2_0p4=0.3957 corrupt_frac_t_0p2_0p4=0.3382 acc_corrupt_t_0p4_0p6=0.6159 corrupt_frac_t_0p4_0p6=0.3435 acc_corrupt_t_0p8_1p0=0.9404 corrupt_frac_t_0p8_1p0=0.3379 out_w_norm=78.8896 out_g_norm=0.1328 acc_corrupt_t_0p6_0p8=0.7995 corrupt_frac_t_0p6_0p8=0.3443 loss_all=1.5796 init_gold_top10=0.6882 init_gold_top100=0.6890 +step=4700 micro_steps=150400 elapsed=648.9s lr=3.000000e-04 loss=2.8173 loss_recon=2.8173 loss_meanflow=0.0000 mean_model_t=0.5024 mean_corrupt_t=0.5024 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5839 corrupt_frac=1.0000 acc_corrupt=0.5839 loss_corrupt=2.8173 wrong_frac=0.4975 init_acc_corrupt=0.5025 acc_corrupt_t_0p0_0p2=0.1537 corrupt_frac_t_0p0_0p2=0.3351 acc_corrupt_t_0p6_0p8=0.7995 corrupt_frac_t_0p6_0p8=0.3383 acc_corrupt_t_0p8_1p0=0.9408 corrupt_frac_t_0p8_1p0=0.3410 out_w_norm=79.9421 out_g_norm=0.1296 acc_corrupt_t_0p2_0p4=0.3946 corrupt_frac_t_0p2_0p4=0.3384 acc_corrupt_t_0p4_0p6=0.6163 corrupt_frac_t_0p4_0p6=0.3314 loss_all=2.2082 init_gold_top10=0.6035 init_gold_top100=0.6042 +step=4800 micro_steps=153600 elapsed=699.2s lr=3.000000e-04 loss=2.8238 loss_recon=2.8238 loss_meanflow=0.0000 mean_model_t=0.5005 mean_corrupt_t=0.5005 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5831 corrupt_frac=1.0000 acc_corrupt=0.5831 loss_corrupt=2.8238 wrong_frac=0.4994 init_acc_corrupt=0.5006 acc_corrupt_t_0p0_0p2=0.1511 corrupt_frac_t_0p0_0p2=0.3394 acc_corrupt_t_0p6_0p8=0.8032 corrupt_frac_t_0p6_0p8=0.3432 acc_corrupt_t_0p8_1p0=0.9426 corrupt_frac_t_0p8_1p0=0.3379 out_w_norm=80.9696 out_g_norm=0.1273 acc_corrupt_t_0p4_0p6=0.6199 corrupt_frac_t_0p4_0p6=0.3436 acc_corrupt_t_0p2_0p4=0.3958 corrupt_frac_t_0p2_0p4=0.3388 loss_all=3.3717 init_gold_top10=0.4116 init_gold_top100=0.4136 +step=4900 micro_steps=156800 elapsed=641.2s lr=3.000000e-04 loss=2.8187 loss_recon=2.8187 loss_meanflow=0.0000 mean_model_t=0.4987 mean_corrupt_t=0.4987 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5822 corrupt_frac=1.0000 acc_corrupt=0.5822 loss_corrupt=2.8187 wrong_frac=0.5013 init_acc_corrupt=0.4987 acc_corrupt_t_0p2_0p4=0.4002 corrupt_frac_t_0p2_0p4=0.3402 acc_corrupt_t_0p4_0p6=0.6204 corrupt_frac_t_0p4_0p6=0.3418 acc_corrupt_t_0p8_1p0=0.9444 corrupt_frac_t_0p8_1p0=0.3319 out_w_norm=81.9679 out_g_norm=0.1272 acc_corrupt_t_0p0_0p2=0.1515 corrupt_frac_t_0p0_0p2=0.3390 acc_corrupt_t_0p6_0p8=0.8021 corrupt_frac_t_0p6_0p8=0.3373 loss_all=2.5765 init_gold_top10=0.5508 init_gold_top100=0.5515 +step=5000 micro_steps=160000 elapsed=703.3s lr=3.000000e-04 loss=2.8073 loss_recon=2.8073 loss_meanflow=0.0000 mean_model_t=0.5018 mean_corrupt_t=0.5018 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5849 corrupt_frac=1.0000 acc_corrupt=0.5849 loss_corrupt=2.8073 wrong_frac=0.4981 init_acc_corrupt=0.5020 acc_corrupt_t_0p0_0p2=0.1536 corrupt_frac_t_0p0_0p2=0.3384 acc_corrupt_t_0p2_0p4=0.3968 corrupt_frac_t_0p2_0p4=0.3369 acc_corrupt_t_0p4_0p6=0.6207 corrupt_frac_t_0p4_0p6=0.3386 out_w_norm=82.9627 out_g_norm=0.1251 acc_corrupt_t_0p6_0p8=0.8030 corrupt_frac_t_0p6_0p8=0.3427 acc_corrupt_t_0p8_1p0=0.9422 corrupt_frac_t_0p8_1p0=0.3417 loss_all=2.9923 init_gold_top10=0.4417 init_gold_top100=0.4434 +step=5100 micro_steps=163200 elapsed=790.8s lr=3.000000e-04 loss=2.8464 loss_recon=2.8464 loss_meanflow=0.0000 mean_model_t=0.4973 mean_corrupt_t=0.4973 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5798 corrupt_frac=1.0000 acc_corrupt=0.5798 loss_corrupt=2.8464 wrong_frac=0.5028 init_acc_corrupt=0.4972 acc_corrupt_t_0p2_0p4=0.3944 corrupt_frac_t_0p2_0p4=0.3397 acc_corrupt_t_0p4_0p6=0.6219 corrupt_frac_t_0p4_0p6=0.3395 acc_corrupt_t_0p6_0p8=0.8050 corrupt_frac_t_0p6_0p8=0.3356 out_w_norm=83.9483 out_g_norm=0.1236 acc_corrupt_t_0p0_0p2=0.1507 corrupt_frac_t_0p0_0p2=0.3421 acc_corrupt_t_0p8_1p0=0.9432 corrupt_frac_t_0p8_1p0=0.3364 loss_all=3.4634 init_gold_top10=0.3970 init_gold_top100=0.3997 +step=5200 micro_steps=166400 elapsed=750.4s lr=3.000000e-04 loss=2.8599 loss_recon=2.8599 loss_meanflow=0.0000 mean_model_t=0.4930 mean_corrupt_t=0.4930 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5768 corrupt_frac=1.0000 acc_corrupt=0.5768 loss_corrupt=2.8599 wrong_frac=0.5070 init_acc_corrupt=0.4930 acc_corrupt_t_0p2_0p4=0.3956 corrupt_frac_t_0p2_0p4=0.3377 acc_corrupt_t_0p4_0p6=0.6210 corrupt_frac_t_0p4_0p6=0.3344 acc_corrupt_t_0p6_0p8=0.8035 corrupt_frac_t_0p6_0p8=0.3346 acc_corrupt_t_0p8_1p0=0.9426 corrupt_frac_t_0p8_1p0=0.3374 out_w_norm=84.9409 out_g_norm=0.1307 acc_corrupt_t_0p0_0p2=0.1502 corrupt_frac_t_0p0_0p2=0.3418 loss_all=1.7800 init_gold_top10=0.6040 init_gold_top100=0.6052 +step=5300 micro_steps=169600 elapsed=664.4s lr=3.000000e-04 loss=2.7718 loss_recon=2.7718 loss_meanflow=0.0000 mean_model_t=0.5028 mean_corrupt_t=0.5028 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5886 corrupt_frac=1.0000 acc_corrupt=0.5886 loss_corrupt=2.7718 wrong_frac=0.4974 init_acc_corrupt=0.5026 acc_corrupt_t_0p0_0p2=0.1568 corrupt_frac_t_0p0_0p2=0.3403 acc_corrupt_t_0p8_1p0=0.9435 corrupt_frac_t_0p8_1p0=0.3405 out_w_norm=85.9215 out_g_norm=0.1322 acc_corrupt_t_0p2_0p4=0.4003 corrupt_frac_t_0p2_0p4=0.3420 acc_corrupt_t_0p4_0p6=0.6215 corrupt_frac_t_0p4_0p6=0.3346 acc_corrupt_t_0p6_0p8=0.8076 corrupt_frac_t_0p6_0p8=0.3377 loss_all=3.7497 init_gold_top10=0.4092 init_gold_top100=0.4114 +step=5400 micro_steps=172800 elapsed=698.3s lr=3.000000e-04 loss=2.8065 loss_recon=2.8065 loss_meanflow=0.0000 mean_model_t=0.4986 mean_corrupt_t=0.4986 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5839 corrupt_frac=1.0000 acc_corrupt=0.5839 loss_corrupt=2.8065 wrong_frac=0.5015 init_acc_corrupt=0.4985 acc_corrupt_t_0p0_0p2=0.1550 corrupt_frac_t_0p0_0p2=0.3329 acc_corrupt_t_0p6_0p8=0.8049 corrupt_frac_t_0p6_0p8=0.3390 acc_corrupt_t_0p8_1p0=0.9453 corrupt_frac_t_0p8_1p0=0.3350 out_w_norm=86.8806 out_g_norm=0.1183 acc_corrupt_t_0p2_0p4=0.3961 corrupt_frac_t_0p2_0p4=0.3386 acc_corrupt_t_0p4_0p6=0.6248 corrupt_frac_t_0p4_0p6=0.3484 loss_all=2.8566 init_gold_top10=0.5085 init_gold_top100=0.5103 +step=5500 micro_steps=176000 elapsed=649.3s lr=3.000000e-04 loss=2.7956 loss_recon=2.7956 loss_meanflow=0.0000 mean_model_t=0.5000 mean_corrupt_t=0.5000 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5854 corrupt_frac=1.0000 acc_corrupt=0.5854 loss_corrupt=2.7956 wrong_frac=0.5000 init_acc_corrupt=0.5000 acc_corrupt_t_0p0_0p2=0.1565 corrupt_frac_t_0p0_0p2=0.3412 acc_corrupt_t_0p2_0p4=0.3973 corrupt_frac_t_0p2_0p4=0.3380 acc_corrupt_t_0p4_0p6=0.6215 corrupt_frac_t_0p4_0p6=0.3377 acc_corrupt_t_0p6_0p8=0.8079 corrupt_frac_t_0p6_0p8=0.3422 out_w_norm=87.8354 out_g_norm=0.1181 acc_corrupt_t_0p8_1p0=0.9444 corrupt_frac_t_0p8_1p0=0.3421 loss_all=2.6139 init_gold_top10=0.6626 init_gold_top100=0.6646 +step=5600 micro_steps=179200 elapsed=659.7s lr=3.000000e-04 loss=2.7978 loss_recon=2.7978 loss_meanflow=0.0000 mean_model_t=0.4985 mean_corrupt_t=0.4985 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5846 corrupt_frac=1.0000 acc_corrupt=0.5846 loss_corrupt=2.7978 wrong_frac=0.5014 init_acc_corrupt=0.4986 acc_corrupt_t_0p0_0p2=0.1546 corrupt_frac_t_0p0_0p2=0.3401 acc_corrupt_t_0p4_0p6=0.6244 corrupt_frac_t_0p4_0p6=0.3384 acc_corrupt_t_0p8_1p0=0.9444 corrupt_frac_t_0p8_1p0=0.3404 out_w_norm=88.8011 out_g_norm=0.1202 acc_corrupt_t_0p2_0p4=0.4008 corrupt_frac_t_0p2_0p4=0.3361 acc_corrupt_t_0p6_0p8=0.8058 corrupt_frac_t_0p6_0p8=0.3367 loss_all=2.1226 init_gold_top10=0.5999 init_gold_top100=0.6028 +step=5700 micro_steps=182400 elapsed=650.3s lr=3.000000e-04 loss=2.7698 loss_recon=2.7698 loss_meanflow=0.0000 mean_model_t=0.5037 mean_corrupt_t=0.5037 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5889 corrupt_frac=1.0000 acc_corrupt=0.5889 loss_corrupt=2.7698 wrong_frac=0.4965 init_acc_corrupt=0.5035 acc_corrupt_t_0p0_0p2=0.1526 corrupt_frac_t_0p0_0p2=0.3377 acc_corrupt_t_0p2_0p4=0.3971 corrupt_frac_t_0p2_0p4=0.3354 acc_corrupt_t_0p6_0p8=0.8048 corrupt_frac_t_0p6_0p8=0.3362 acc_corrupt_t_0p8_1p0=0.9446 corrupt_frac_t_0p8_1p0=0.3429 out_w_norm=89.7963 out_g_norm=0.1228 acc_corrupt_t_0p4_0p6=0.6244 corrupt_frac_t_0p4_0p6=0.3360 loss_all=2.3382 init_gold_top10=0.5327 init_gold_top100=0.5347 +step=5800 micro_steps=185600 elapsed=653.3s lr=3.000000e-04 loss=2.7884 loss_recon=2.7884 loss_meanflow=0.0000 mean_model_t=0.4995 mean_corrupt_t=0.4995 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5857 corrupt_frac=1.0000 acc_corrupt=0.5857 loss_corrupt=2.7884 wrong_frac=0.5007 init_acc_corrupt=0.4993 acc_corrupt_t_0p0_0p2=0.1550 corrupt_frac_t_0p0_0p2=0.3427 acc_corrupt_t_0p2_0p4=0.3999 corrupt_frac_t_0p2_0p4=0.3339 acc_corrupt_t_0p6_0p8=0.8095 corrupt_frac_t_0p6_0p8=0.3310 acc_corrupt_t_0p8_1p0=0.9437 corrupt_frac_t_0p8_1p0=0.3392 out_w_norm=90.8248 out_g_norm=0.1149 acc_corrupt_t_0p4_0p6=0.6220 corrupt_frac_t_0p4_0p6=0.3395 loss_all=2.9683 init_gold_top10=0.4565 init_gold_top100=0.4583 +step=5900 micro_steps=188800 elapsed=657.8s lr=3.000000e-04 loss=2.7772 loss_recon=2.7772 loss_meanflow=0.0000 mean_model_t=0.5011 mean_corrupt_t=0.5011 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5877 corrupt_frac=1.0000 acc_corrupt=0.5877 loss_corrupt=2.7772 wrong_frac=0.4989 init_acc_corrupt=0.5011 acc_corrupt_t_0p0_0p2=0.1569 corrupt_frac_t_0p0_0p2=0.3382 acc_corrupt_t_0p2_0p4=0.3993 corrupt_frac_t_0p2_0p4=0.3406 acc_corrupt_t_0p8_1p0=0.9448 corrupt_frac_t_0p8_1p0=0.3389 out_w_norm=91.8026 out_g_norm=0.1141 acc_corrupt_t_0p4_0p6=0.6263 corrupt_frac_t_0p4_0p6=0.3380 acc_corrupt_t_0p6_0p8=0.8082 corrupt_frac_t_0p6_0p8=0.3480 loss_all=1.7243 init_gold_top10=0.6489 init_gold_top100=0.6506 +step=6000 micro_steps=192000 elapsed=648.4s lr=3.000000e-04 loss=2.8014 loss_recon=2.8014 loss_meanflow=0.0000 mean_model_t=0.4967 mean_corrupt_t=0.4967 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5837 corrupt_frac=1.0000 acc_corrupt=0.5837 loss_corrupt=2.8014 wrong_frac=0.5034 init_acc_corrupt=0.4966 acc_corrupt_t_0p2_0p4=0.3990 corrupt_frac_t_0p2_0p4=0.3366 acc_corrupt_t_0p6_0p8=0.8086 corrupt_frac_t_0p6_0p8=0.3330 out_w_norm=92.7573 out_g_norm=0.1185 acc_corrupt_t_0p0_0p2=0.1536 corrupt_frac_t_0p0_0p2=0.3375 acc_corrupt_t_0p4_0p6=0.6266 corrupt_frac_t_0p4_0p6=0.3398 acc_corrupt_t_0p8_1p0=0.9451 corrupt_frac_t_0p8_1p0=0.3334 loss_all=0.9648 init_gold_top10=0.7515 init_gold_top100=0.7515 +step=6100 micro_steps=195200 elapsed=822.3s lr=3.000000e-04 loss=2.7990 loss_recon=2.7990 loss_meanflow=0.0000 mean_model_t=0.4977 mean_corrupt_t=0.4977 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5843 corrupt_frac=1.0000 acc_corrupt=0.5843 loss_corrupt=2.7990 wrong_frac=0.5020 init_acc_corrupt=0.4980 acc_corrupt_t_0p2_0p4=0.3987 corrupt_frac_t_0p2_0p4=0.3379 acc_corrupt_t_0p4_0p6=0.6271 corrupt_frac_t_0p4_0p6=0.3385 acc_corrupt_t_0p6_0p8=0.8091 corrupt_frac_t_0p6_0p8=0.3306 out_w_norm=93.6974 out_g_norm=0.1097 acc_corrupt_t_0p8_1p0=0.9443 corrupt_frac_t_0p8_1p0=0.3433 acc_corrupt_t_0p0_0p2=0.1543 corrupt_frac_t_0p0_0p2=0.3394 loss_all=2.1978 init_gold_top10=0.5244 init_gold_top100=0.5249 +step=6200 micro_steps=198400 elapsed=693.0s lr=3.000000e-04 loss=2.7466 loss_recon=2.7466 loss_meanflow=0.0000 mean_model_t=0.5039 mean_corrupt_t=0.5039 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5911 corrupt_frac=1.0000 acc_corrupt=0.5911 loss_corrupt=2.7466 wrong_frac=0.4959 init_acc_corrupt=0.5041 acc_corrupt_t_0p0_0p2=0.1542 corrupt_frac_t_0p0_0p2=0.3378 acc_corrupt_t_0p4_0p6=0.6246 corrupt_frac_t_0p4_0p6=0.3401 acc_corrupt_t_0p6_0p8=0.8090 corrupt_frac_t_0p6_0p8=0.3480 acc_corrupt_t_0p8_1p0=0.9454 corrupt_frac_t_0p8_1p0=0.3421 out_w_norm=94.6323 out_g_norm=0.1092 acc_corrupt_t_0p2_0p4=0.4015 corrupt_frac_t_0p2_0p4=0.3421 loss_all=1.7015 init_gold_top10=0.6582 init_gold_top100=0.6594 +step=6300 micro_steps=201600 elapsed=673.1s lr=3.000000e-04 loss=2.7727 loss_recon=2.7727 loss_meanflow=0.0000 mean_model_t=0.4997 mean_corrupt_t=0.4997 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5869 corrupt_frac=1.0000 acc_corrupt=0.5869 loss_corrupt=2.7727 wrong_frac=0.5003 init_acc_corrupt=0.4997 acc_corrupt_t_0p2_0p4=0.4016 corrupt_frac_t_0p2_0p4=0.3368 acc_corrupt_t_0p6_0p8=0.8092 corrupt_frac_t_0p6_0p8=0.3344 acc_corrupt_t_0p8_1p0=0.9447 corrupt_frac_t_0p8_1p0=0.3360 out_w_norm=95.5519 out_g_norm=0.1082 acc_corrupt_t_0p4_0p6=0.6269 corrupt_frac_t_0p4_0p6=0.3364 acc_corrupt_t_0p0_0p2=0.1530 corrupt_frac_t_0p0_0p2=0.3385 loss_all=2.3158 init_gold_top10=0.5325 init_gold_top100=0.5337 +step=6400 micro_steps=204800 elapsed=642.5s lr=3.000000e-04 loss=2.7292 loss_recon=2.7292 loss_meanflow=0.0000 mean_model_t=0.5027 mean_corrupt_t=0.5027 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5927 corrupt_frac=1.0000 acc_corrupt=0.5927 loss_corrupt=2.7292 wrong_frac=0.4973 init_acc_corrupt=0.5028 acc_corrupt_t_0p2_0p4=0.4060 corrupt_frac_t_0p2_0p4=0.3385 acc_corrupt_t_0p6_0p8=0.8110 corrupt_frac_t_0p6_0p8=0.3404 acc_corrupt_t_0p8_1p0=0.9447 corrupt_frac_t_0p8_1p0=0.3407 out_w_norm=96.4876 out_g_norm=0.1081 acc_corrupt_t_0p0_0p2=0.1565 corrupt_frac_t_0p0_0p2=0.3393 acc_corrupt_t_0p4_0p6=0.6277 corrupt_frac_t_0p4_0p6=0.3412 loss_all=1.8264 init_gold_top10=0.6467 init_gold_top100=0.6477 +step=6500 micro_steps=208000 elapsed=669.7s lr=3.000000e-04 loss=2.7348 loss_recon=2.7348 loss_meanflow=0.0000 mean_model_t=0.5032 mean_corrupt_t=0.5032 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5919 corrupt_frac=1.0000 acc_corrupt=0.5919 loss_corrupt=2.7348 wrong_frac=0.4967 init_acc_corrupt=0.5033 acc_corrupt_t_0p4_0p6=0.6312 corrupt_frac_t_0p4_0p6=0.3392 acc_corrupt_t_0p6_0p8=0.8096 corrupt_frac_t_0p6_0p8=0.3403 acc_corrupt_t_0p8_1p0=0.9463 corrupt_frac_t_0p8_1p0=0.3453 out_w_norm=97.4048 out_g_norm=0.1070 acc_corrupt_t_0p2_0p4=0.4034 corrupt_frac_t_0p2_0p4=0.3398 acc_corrupt_t_0p0_0p2=0.1526 corrupt_frac_t_0p0_0p2=0.3320 loss_all=1.3079 init_gold_top10=0.6763 init_gold_top100=0.6772 +step=6600 micro_steps=211200 elapsed=663.4s lr=3.000000e-04 loss=2.7890 loss_recon=2.7890 loss_meanflow=0.0000 mean_model_t=0.4947 mean_corrupt_t=0.4947 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5843 corrupt_frac=1.0000 acc_corrupt=0.5843 loss_corrupt=2.7890 wrong_frac=0.5053 init_acc_corrupt=0.4947 acc_corrupt_t_0p4_0p6=0.6282 corrupt_frac_t_0p4_0p6=0.3394 acc_corrupt_t_0p6_0p8=0.8100 corrupt_frac_t_0p6_0p8=0.3361 acc_corrupt_t_0p8_1p0=0.9450 corrupt_frac_t_0p8_1p0=0.3362 out_w_norm=98.3273 out_g_norm=0.1054 acc_corrupt_t_0p0_0p2=0.1560 corrupt_frac_t_0p0_0p2=0.3369 acc_corrupt_t_0p2_0p4=0.4028 corrupt_frac_t_0p2_0p4=0.3375 loss_all=2.1752 init_gold_top10=0.5371 init_gold_top100=0.5381 +step=6700 micro_steps=214400 elapsed=686.2s lr=3.000000e-04 loss=2.7276 loss_recon=2.7276 loss_meanflow=0.0000 mean_model_t=0.5038 mean_corrupt_t=0.5038 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5926 corrupt_frac=1.0000 acc_corrupt=0.5926 loss_corrupt=2.7276 wrong_frac=0.4961 init_acc_corrupt=0.5039 acc_corrupt_t_0p4_0p6=0.6293 corrupt_frac_t_0p4_0p6=0.3378 acc_corrupt_t_0p6_0p8=0.8115 corrupt_frac_t_0p6_0p8=0.3389 acc_corrupt_t_0p8_1p0=0.9461 corrupt_frac_t_0p8_1p0=0.3387 out_w_norm=99.2396 out_g_norm=0.1029 acc_corrupt_t_0p2_0p4=0.4024 corrupt_frac_t_0p2_0p4=0.3414 acc_corrupt_t_0p0_0p2=0.1541 corrupt_frac_t_0p0_0p2=0.3385 loss_all=2.4560 init_gold_top10=0.5300 init_gold_top100=0.5308 +step=6800 micro_steps=217600 elapsed=645.9s lr=3.000000e-04 loss=2.7419 loss_recon=2.7419 loss_meanflow=0.0000 mean_model_t=0.5020 mean_corrupt_t=0.5020 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5914 corrupt_frac=1.0000 acc_corrupt=0.5914 loss_corrupt=2.7419 wrong_frac=0.4980 init_acc_corrupt=0.5020 acc_corrupt_t_0p0_0p2=0.1524 corrupt_frac_t_0p0_0p2=0.3339 acc_corrupt_t_0p6_0p8=0.8119 corrupt_frac_t_0p6_0p8=0.3430 out_w_norm=100.1610 out_g_norm=0.1056 acc_corrupt_t_0p2_0p4=0.4035 corrupt_frac_t_0p2_0p4=0.3385 acc_corrupt_t_0p4_0p6=0.6293 corrupt_frac_t_0p4_0p6=0.3366 acc_corrupt_t_0p8_1p0=0.9455 corrupt_frac_t_0p8_1p0=0.3359 loss_all=1.2217 init_gold_top10=0.7466 init_gold_top100=0.7478 +step=6900 micro_steps=220800 elapsed=668.3s lr=3.000000e-04 loss=2.7614 loss_recon=2.7614 loss_meanflow=0.0000 mean_model_t=0.4996 mean_corrupt_t=0.4996 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5888 corrupt_frac=1.0000 acc_corrupt=0.5888 loss_corrupt=2.7614 wrong_frac=0.5006 init_acc_corrupt=0.4994 acc_corrupt_t_0p0_0p2=0.1541 corrupt_frac_t_0p0_0p2=0.3337 acc_corrupt_t_0p2_0p4=0.4034 corrupt_frac_t_0p2_0p4=0.3379 acc_corrupt_t_0p4_0p6=0.6288 corrupt_frac_t_0p4_0p6=0.3425 out_w_norm=101.0784 out_g_norm=0.1013 acc_corrupt_t_0p6_0p8=0.8135 corrupt_frac_t_0p6_0p8=0.3349 acc_corrupt_t_0p8_1p0=0.9455 corrupt_frac_t_0p8_1p0=0.3352 loss_all=1.0593 init_gold_top10=0.7419 init_gold_top100=0.7424 +step=7000 micro_steps=224000 elapsed=642.5s lr=3.000000e-04 loss=2.7354 loss_recon=2.7354 loss_meanflow=0.0000 mean_model_t=0.5019 mean_corrupt_t=0.5019 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5917 corrupt_frac=1.0000 acc_corrupt=0.5917 loss_corrupt=2.7354 wrong_frac=0.4981 init_acc_corrupt=0.5020 acc_corrupt_t_0p6_0p8=0.8110 corrupt_frac_t_0p6_0p8=0.3387 acc_corrupt_t_0p8_1p0=0.9466 corrupt_frac_t_0p8_1p0=0.3381 out_w_norm=101.9820 out_g_norm=0.1035 acc_corrupt_t_0p0_0p2=0.1545 corrupt_frac_t_0p0_0p2=0.3337 acc_corrupt_t_0p2_0p4=0.4047 corrupt_frac_t_0p2_0p4=0.3394 acc_corrupt_t_0p4_0p6=0.6284 corrupt_frac_t_0p4_0p6=0.3385 loss_all=4.7305 init_gold_top10=0.2310 init_gold_top100=0.2329 +step=7100 micro_steps=227200 elapsed=832.9s lr=3.000000e-04 loss=2.7888 loss_recon=2.7888 loss_meanflow=0.0000 mean_model_t=0.4954 mean_corrupt_t=0.4954 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5842 corrupt_frac=1.0000 acc_corrupt=0.5842 loss_corrupt=2.7888 wrong_frac=0.5047 init_acc_corrupt=0.4954 acc_corrupt_t_0p0_0p2=0.1551 corrupt_frac_t_0p0_0p2=0.3428 acc_corrupt_t_0p8_1p0=0.9472 corrupt_frac_t_0p8_1p0=0.3392 out_w_norm=102.8788 out_g_norm=0.1050 acc_corrupt_t_0p4_0p6=0.6300 corrupt_frac_t_0p4_0p6=0.3380 acc_corrupt_t_0p6_0p8=0.8100 corrupt_frac_t_0p6_0p8=0.3394 acc_corrupt_t_0p2_0p4=0.4037 corrupt_frac_t_0p2_0p4=0.3442 loss_all=2.2535 init_gold_top10=0.5654 init_gold_top100=0.5669 +step=7200 micro_steps=230400 elapsed=705.2s lr=3.000000e-04 loss=2.7658 loss_recon=2.7658 loss_meanflow=0.0000 mean_model_t=0.4970 mean_corrupt_t=0.4970 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5874 corrupt_frac=1.0000 acc_corrupt=0.5874 loss_corrupt=2.7658 wrong_frac=0.5031 init_acc_corrupt=0.4969 acc_corrupt_t_0p0_0p2=0.1548 corrupt_frac_t_0p0_0p2=0.3379 acc_corrupt_t_0p6_0p8=0.8139 corrupt_frac_t_0p6_0p8=0.3366 out_w_norm=103.8074 out_g_norm=0.1015 acc_corrupt_t_0p2_0p4=0.4056 corrupt_frac_t_0p2_0p4=0.3393 acc_corrupt_t_0p4_0p6=0.6299 corrupt_frac_t_0p4_0p6=0.3432 acc_corrupt_t_0p8_1p0=0.9462 corrupt_frac_t_0p8_1p0=0.3370 loss_all=1.2655 init_gold_top10=0.6638 init_gold_top100=0.6643 +step=7300 micro_steps=233600 elapsed=665.5s lr=3.000000e-04 loss=2.7764 loss_recon=2.7764 loss_meanflow=0.0000 mean_model_t=0.4964 mean_corrupt_t=0.4964 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5859 corrupt_frac=1.0000 acc_corrupt=0.5859 loss_corrupt=2.7764 wrong_frac=0.5037 init_acc_corrupt=0.4963 acc_corrupt_t_0p0_0p2=0.1565 corrupt_frac_t_0p0_0p2=0.3432 acc_corrupt_t_0p2_0p4=0.4030 corrupt_frac_t_0p2_0p4=0.3435 acc_corrupt_t_0p4_0p6=0.6297 corrupt_frac_t_0p4_0p6=0.3367 acc_corrupt_t_0p6_0p8=0.8123 corrupt_frac_t_0p6_0p8=0.3375 out_w_norm=104.7568 out_g_norm=0.1008 acc_corrupt_t_0p8_1p0=0.9467 corrupt_frac_t_0p8_1p0=0.3349 loss_all=1.6779 init_gold_top10=0.6226 init_gold_top100=0.6235 +step=7400 micro_steps=236800 elapsed=652.2s lr=3.000000e-04 loss=2.7703 loss_recon=2.7703 loss_meanflow=0.0000 mean_model_t=0.4977 mean_corrupt_t=0.4977 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5875 corrupt_frac=1.0000 acc_corrupt=0.5875 loss_corrupt=2.7703 wrong_frac=0.5022 init_acc_corrupt=0.4978 acc_corrupt_t_0p0_0p2=0.1535 corrupt_frac_t_0p0_0p2=0.3410 acc_corrupt_t_0p2_0p4=0.4033 corrupt_frac_t_0p2_0p4=0.3365 acc_corrupt_t_0p8_1p0=0.9466 corrupt_frac_t_0p8_1p0=0.3365 out_w_norm=105.6638 out_g_norm=0.1008 acc_corrupt_t_0p4_0p6=0.6308 corrupt_frac_t_0p4_0p6=0.3422 acc_corrupt_t_0p6_0p8=0.8138 corrupt_frac_t_0p6_0p8=0.3421 loss_all=3.6087 init_gold_top10=0.3962 init_gold_top100=0.3977 +step=7500 micro_steps=240000 elapsed=658.8s lr=3.000000e-04 loss=2.7644 loss_recon=2.7644 loss_meanflow=0.0000 mean_model_t=0.4973 mean_corrupt_t=0.4973 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5877 corrupt_frac=1.0000 acc_corrupt=0.5877 loss_corrupt=2.7644 wrong_frac=0.5026 init_acc_corrupt=0.4974 acc_corrupt_t_0p4_0p6=0.6319 corrupt_frac_t_0p4_0p6=0.3363 acc_corrupt_t_0p8_1p0=0.9464 corrupt_frac_t_0p8_1p0=0.3404 out_w_norm=106.5782 out_g_norm=0.1032 acc_corrupt_t_0p0_0p2=0.1541 corrupt_frac_t_0p0_0p2=0.3409 acc_corrupt_t_0p2_0p4=0.4063 corrupt_frac_t_0p2_0p4=0.3417 acc_corrupt_t_0p6_0p8=0.8127 corrupt_frac_t_0p6_0p8=0.3358 loss_all=2.5969 init_gold_top10=0.4861 init_gold_top100=0.4878 +step=7600 micro_steps=243200 elapsed=660.5s lr=3.000000e-04 loss=2.7281 loss_recon=2.7281 loss_meanflow=0.0000 mean_model_t=0.5007 mean_corrupt_t=0.5007 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5919 corrupt_frac=1.0000 acc_corrupt=0.5919 loss_corrupt=2.7281 wrong_frac=0.4995 init_acc_corrupt=0.5005 acc_corrupt_t_0p2_0p4=0.4071 corrupt_frac_t_0p2_0p4=0.3439 acc_corrupt_t_0p4_0p6=0.6341 corrupt_frac_t_0p4_0p6=0.3393 acc_corrupt_t_0p8_1p0=0.9467 corrupt_frac_t_0p8_1p0=0.3381 out_w_norm=107.4733 out_g_norm=0.0970 acc_corrupt_t_0p0_0p2=0.1522 corrupt_frac_t_0p0_0p2=0.3355 acc_corrupt_t_0p6_0p8=0.8156 corrupt_frac_t_0p6_0p8=0.3367 loss_all=2.8031 init_gold_top10=0.4363 init_gold_top100=0.4380 +step=7700 micro_steps=246400 elapsed=653.6s lr=3.000000e-04 loss=2.7333 loss_recon=2.7333 loss_meanflow=0.0000 mean_model_t=0.4998 mean_corrupt_t=0.4998 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5910 corrupt_frac=1.0000 acc_corrupt=0.5910 loss_corrupt=2.7333 wrong_frac=0.5004 init_acc_corrupt=0.4996 acc_corrupt_t_0p0_0p2=0.1559 corrupt_frac_t_0p0_0p2=0.3431 acc_corrupt_t_0p2_0p4=0.4088 corrupt_frac_t_0p2_0p4=0.3395 acc_corrupt_t_0p6_0p8=0.8150 corrupt_frac_t_0p6_0p8=0.3402 out_w_norm=108.3466 out_g_norm=0.0984 acc_corrupt_t_0p4_0p6=0.6325 corrupt_frac_t_0p4_0p6=0.3405 acc_corrupt_t_0p8_1p0=0.9485 corrupt_frac_t_0p8_1p0=0.3444 loss_all=3.0474 init_gold_top10=0.4146 init_gold_top100=0.4165 +step=7800 micro_steps=249600 elapsed=657.4s lr=3.000000e-04 loss=2.6996 loss_recon=2.6996 loss_meanflow=0.0000 mean_model_t=0.5056 mean_corrupt_t=0.5056 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5967 corrupt_frac=1.0000 acc_corrupt=0.5967 loss_corrupt=2.6996 wrong_frac=0.4945 init_acc_corrupt=0.5056 acc_corrupt_t_0p0_0p2=0.1557 corrupt_frac_t_0p0_0p2=0.3372 acc_corrupt_t_0p2_0p4=0.4054 corrupt_frac_t_0p2_0p4=0.3419 acc_corrupt_t_0p4_0p6=0.6342 corrupt_frac_t_0p4_0p6=0.3353 acc_corrupt_t_0p8_1p0=0.9475 corrupt_frac_t_0p8_1p0=0.3407 out_w_norm=109.2132 out_g_norm=0.1038 acc_corrupt_t_0p6_0p8=0.8157 corrupt_frac_t_0p6_0p8=0.3368 loss_all=2.0452 init_gold_top10=0.6318 init_gold_top100=0.6326 +step=7900 micro_steps=252800 elapsed=649.5s lr=3.000000e-04 loss=2.6981 loss_recon=2.6981 loss_meanflow=0.0000 mean_model_t=0.5049 mean_corrupt_t=0.5049 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5961 corrupt_frac=1.0000 acc_corrupt=0.5961 loss_corrupt=2.6981 wrong_frac=0.4949 init_acc_corrupt=0.5051 acc_corrupt_t_0p0_0p2=0.1545 corrupt_frac_t_0p0_0p2=0.3356 acc_corrupt_t_0p6_0p8=0.8143 corrupt_frac_t_0p6_0p8=0.3404 acc_corrupt_t_0p8_1p0=0.9466 corrupt_frac_t_0p8_1p0=0.3407 out_w_norm=110.0928 out_g_norm=0.0945 acc_corrupt_t_0p2_0p4=0.4070 corrupt_frac_t_0p2_0p4=0.3300 acc_corrupt_t_0p4_0p6=0.6336 corrupt_frac_t_0p4_0p6=0.3425 loss_all=2.7830 init_gold_top10=0.5002 init_gold_top100=0.5024 +step=8000 micro_steps=256000 elapsed=662.0s lr=3.000000e-04 loss=2.7355 loss_recon=2.7355 loss_meanflow=0.0000 mean_model_t=0.4981 mean_corrupt_t=0.4981 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5906 corrupt_frac=1.0000 acc_corrupt=0.5906 loss_corrupt=2.7355 wrong_frac=0.5016 init_acc_corrupt=0.4984 acc_corrupt_t_0p2_0p4=0.4086 corrupt_frac_t_0p2_0p4=0.3303 acc_corrupt_t_0p4_0p6=0.6334 corrupt_frac_t_0p4_0p6=0.3397 out_w_norm=110.9532 out_g_norm=0.0938 acc_corrupt_t_0p8_1p0=0.9470 corrupt_frac_t_0p8_1p0=0.3356 acc_corrupt_t_0p0_0p2=0.1551 corrupt_frac_t_0p0_0p2=0.3422 acc_corrupt_t_0p6_0p8=0.8159 corrupt_frac_t_0p6_0p8=0.3377 loss_all=2.3963 init_gold_top10=0.5017 init_gold_top100=0.5037 +step=8100 micro_steps=259200 elapsed=767.6s lr=3.000000e-04 loss=2.7255 loss_recon=2.7255 loss_meanflow=0.0000 mean_model_t=0.5003 mean_corrupt_t=0.5003 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5921 corrupt_frac=1.0000 acc_corrupt=0.5921 loss_corrupt=2.7255 wrong_frac=0.4999 init_acc_corrupt=0.5001 acc_corrupt_t_0p4_0p6=0.6332 corrupt_frac_t_0p4_0p6=0.3426 acc_corrupt_t_0p6_0p8=0.8146 corrupt_frac_t_0p6_0p8=0.3356 acc_corrupt_t_0p8_1p0=0.9471 corrupt_frac_t_0p8_1p0=0.3409 out_w_norm=111.8084 out_g_norm=0.0946 acc_corrupt_t_0p2_0p4=0.4059 corrupt_frac_t_0p2_0p4=0.3364 acc_corrupt_t_0p0_0p2=0.1555 corrupt_frac_t_0p0_0p2=0.3383 loss_all=3.1184 init_gold_top10=0.3708 init_gold_top100=0.3726 +step=8200 micro_steps=262400 elapsed=751.5s lr=3.000000e-04 loss=2.7114 loss_recon=2.7114 loss_meanflow=0.0000 mean_model_t=0.5030 mean_corrupt_t=0.5030 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5948 corrupt_frac=1.0000 acc_corrupt=0.5948 loss_corrupt=2.7114 wrong_frac=0.4970 init_acc_corrupt=0.5030 acc_corrupt_t_0p2_0p4=0.4094 corrupt_frac_t_0p2_0p4=0.3386 acc_corrupt_t_0p4_0p6=0.6324 corrupt_frac_t_0p4_0p6=0.3370 acc_corrupt_t_0p6_0p8=0.8152 corrupt_frac_t_0p6_0p8=0.3335 out_w_norm=112.6690 out_g_norm=0.0939 acc_corrupt_t_0p8_1p0=0.9468 corrupt_frac_t_0p8_1p0=0.3359 acc_corrupt_t_0p0_0p2=0.1563 corrupt_frac_t_0p0_0p2=0.3352 loss_all=1.7312 init_gold_top10=0.6360 init_gold_top100=0.6372 +step=8300 micro_steps=265600 elapsed=644.1s lr=3.000000e-04 loss=2.6946 loss_recon=2.6946 loss_meanflow=0.0000 mean_model_t=0.5026 mean_corrupt_t=0.5026 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5963 corrupt_frac=1.0000 acc_corrupt=0.5963 loss_corrupt=2.6946 wrong_frac=0.4972 init_acc_corrupt=0.5028 acc_corrupt_t_0p4_0p6=0.6379 corrupt_frac_t_0p4_0p6=0.3383 acc_corrupt_t_0p6_0p8=0.8165 corrupt_frac_t_0p6_0p8=0.3377 out_w_norm=113.5280 out_g_norm=0.0918 acc_corrupt_t_0p0_0p2=0.1590 corrupt_frac_t_0p0_0p2=0.3387 acc_corrupt_t_0p8_1p0=0.9478 corrupt_frac_t_0p8_1p0=0.3389 acc_corrupt_t_0p2_0p4=0.4093 corrupt_frac_t_0p2_0p4=0.3387 loss_all=1.6517 init_gold_top10=0.6252 init_gold_top100=0.6270 +step=8400 micro_steps=268800 elapsed=667.2s lr=3.000000e-04 loss=2.7335 loss_recon=2.7335 loss_meanflow=0.0000 mean_model_t=0.4984 mean_corrupt_t=0.4984 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5911 corrupt_frac=1.0000 acc_corrupt=0.5911 loss_corrupt=2.7335 wrong_frac=0.5017 init_acc_corrupt=0.4983 acc_corrupt_t_0p0_0p2=0.1594 corrupt_frac_t_0p0_0p2=0.3415 acc_corrupt_t_0p2_0p4=0.4075 corrupt_frac_t_0p2_0p4=0.3396 acc_corrupt_t_0p8_1p0=0.9480 corrupt_frac_t_0p8_1p0=0.3448 out_w_norm=114.3757 out_g_norm=0.0967 acc_corrupt_t_0p4_0p6=0.6344 corrupt_frac_t_0p4_0p6=0.3331 acc_corrupt_t_0p6_0p8=0.8156 corrupt_frac_t_0p6_0p8=0.3425 loss_all=2.3896 init_gold_top10=0.3992 init_gold_top100=0.4011 +step=8500 micro_steps=272000 elapsed=640.3s lr=3.000000e-04 loss=2.7414 loss_recon=2.7414 loss_meanflow=0.0000 mean_model_t=0.4958 mean_corrupt_t=0.4958 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5900 corrupt_frac=1.0000 acc_corrupt=0.5900 loss_corrupt=2.7414 wrong_frac=0.5042 init_acc_corrupt=0.4959 acc_corrupt_t_0p2_0p4=0.4107 corrupt_frac_t_0p2_0p4=0.3411 acc_corrupt_t_0p4_0p6=0.6365 corrupt_frac_t_0p4_0p6=0.3413 acc_corrupt_t_0p6_0p8=0.8156 corrupt_frac_t_0p6_0p8=0.3375 acc_corrupt_t_0p8_1p0=0.9487 corrupt_frac_t_0p8_1p0=0.3380 out_w_norm=115.2418 out_g_norm=0.0927 acc_corrupt_t_0p0_0p2=0.1585 corrupt_frac_t_0p0_0p2=0.3396 loss_all=4.3278 init_gold_top10=0.2913 init_gold_top100=0.2932 +step=8600 micro_steps=275200 elapsed=666.6s lr=3.000000e-04 loss=2.7268 loss_recon=2.7268 loss_meanflow=0.0000 mean_model_t=0.4993 mean_corrupt_t=0.4993 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5919 corrupt_frac=1.0000 acc_corrupt=0.5919 loss_corrupt=2.7268 wrong_frac=0.5007 init_acc_corrupt=0.4993 acc_corrupt_t_0p0_0p2=0.1562 corrupt_frac_t_0p0_0p2=0.3433 acc_corrupt_t_0p6_0p8=0.8167 corrupt_frac_t_0p6_0p8=0.3385 acc_corrupt_t_0p8_1p0=0.9487 corrupt_frac_t_0p8_1p0=0.3371 out_w_norm=116.0916 out_g_norm=0.0877 acc_corrupt_t_0p2_0p4=0.4063 corrupt_frac_t_0p2_0p4=0.3306 acc_corrupt_t_0p4_0p6=0.6347 corrupt_frac_t_0p4_0p6=0.3416 loss_all=4.3754 init_gold_top10=0.2520 init_gold_top100=0.2544 +step=8700 micro_steps=278400 elapsed=653.6s lr=3.000000e-04 loss=2.7047 loss_recon=2.7047 loss_meanflow=0.0000 mean_model_t=0.5018 mean_corrupt_t=0.5018 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5946 corrupt_frac=1.0000 acc_corrupt=0.5946 loss_corrupt=2.7047 wrong_frac=0.4984 init_acc_corrupt=0.5016 acc_corrupt_t_0p2_0p4=0.4066 corrupt_frac_t_0p2_0p4=0.3358 acc_corrupt_t_0p6_0p8=0.8154 corrupt_frac_t_0p6_0p8=0.3361 acc_corrupt_t_0p8_1p0=0.9486 corrupt_frac_t_0p8_1p0=0.3372 out_w_norm=116.9331 out_g_norm=0.0869 acc_corrupt_t_0p4_0p6=0.6357 corrupt_frac_t_0p4_0p6=0.3395 acc_corrupt_t_0p0_0p2=0.1571 corrupt_frac_t_0p0_0p2=0.3402 loss_all=2.7463 init_gold_top10=0.4114 init_gold_top100=0.4126 +step=8800 micro_steps=281600 elapsed=686.6s lr=3.000000e-04 loss=2.7116 loss_recon=2.7116 loss_meanflow=0.0000 mean_model_t=0.4994 mean_corrupt_t=0.4994 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5934 corrupt_frac=1.0000 acc_corrupt=0.5934 loss_corrupt=2.7116 wrong_frac=0.5005 init_acc_corrupt=0.4996 acc_corrupt_t_0p0_0p2=0.1569 corrupt_frac_t_0p0_0p2=0.3386 acc_corrupt_t_0p2_0p4=0.4114 corrupt_frac_t_0p2_0p4=0.3345 acc_corrupt_t_0p4_0p6=0.6379 corrupt_frac_t_0p4_0p6=0.3399 acc_corrupt_t_0p8_1p0=0.9489 corrupt_frac_t_0p8_1p0=0.3379 out_w_norm=117.7698 out_g_norm=0.0885 acc_corrupt_t_0p6_0p8=0.8189 corrupt_frac_t_0p6_0p8=0.3412 loss_all=2.5939 init_gold_top10=0.4475 init_gold_top100=0.4485 +step=8900 micro_steps=284800 elapsed=644.9s lr=3.000000e-04 loss=2.7466 loss_recon=2.7466 loss_meanflow=0.0000 mean_model_t=0.4955 mean_corrupt_t=0.4955 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5889 corrupt_frac=1.0000 acc_corrupt=0.5889 loss_corrupt=2.7466 wrong_frac=0.5043 init_acc_corrupt=0.4957 acc_corrupt_t_0p0_0p2=0.1574 corrupt_frac_t_0p0_0p2=0.3379 acc_corrupt_t_0p6_0p8=0.8168 corrupt_frac_t_0p6_0p8=0.3388 acc_corrupt_t_0p8_1p0=0.9487 corrupt_frac_t_0p8_1p0=0.3397 out_w_norm=118.6011 out_g_norm=0.0894 acc_corrupt_t_0p2_0p4=0.4070 corrupt_frac_t_0p2_0p4=0.3438 acc_corrupt_t_0p4_0p6=0.6379 corrupt_frac_t_0p4_0p6=0.3431 loss_all=3.0275 init_gold_top10=0.4407 init_gold_top100=0.4426 +step=9000 micro_steps=288000 elapsed=661.6s lr=3.000000e-04 loss=2.6913 loss_recon=2.6913 loss_meanflow=0.0000 mean_model_t=0.5037 mean_corrupt_t=0.5037 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5962 corrupt_frac=1.0000 acc_corrupt=0.5962 loss_corrupt=2.6913 wrong_frac=0.4964 init_acc_corrupt=0.5036 acc_corrupt_t_0p0_0p2=0.1570 corrupt_frac_t_0p0_0p2=0.3397 acc_corrupt_t_0p4_0p6=0.6352 corrupt_frac_t_0p4_0p6=0.3321 acc_corrupt_t_0p6_0p8=0.8173 corrupt_frac_t_0p6_0p8=0.3372 out_w_norm=119.4335 out_g_norm=0.0891 acc_corrupt_t_0p8_1p0=0.9495 corrupt_frac_t_0p8_1p0=0.3383 acc_corrupt_t_0p2_0p4=0.4065 corrupt_frac_t_0p2_0p4=0.3315 loss_all=2.1023 init_gold_top10=0.6228 init_gold_top100=0.6238 +step=9100 micro_steps=291200 elapsed=772.8s lr=3.000000e-04 loss=2.7351 loss_recon=2.7351 loss_meanflow=0.0000 mean_model_t=0.4980 mean_corrupt_t=0.4980 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5905 corrupt_frac=1.0000 acc_corrupt=0.5905 loss_corrupt=2.7351 wrong_frac=0.5019 init_acc_corrupt=0.4981 acc_corrupt_t_0p2_0p4=0.4075 corrupt_frac_t_0p2_0p4=0.3410 acc_corrupt_t_0p4_0p6=0.6359 corrupt_frac_t_0p4_0p6=0.3332 acc_corrupt_t_0p6_0p8=0.8169 corrupt_frac_t_0p6_0p8=0.3349 out_w_norm=120.2714 out_g_norm=0.0875 acc_corrupt_t_0p0_0p2=0.1552 corrupt_frac_t_0p0_0p2=0.3389 acc_corrupt_t_0p8_1p0=0.9482 corrupt_frac_t_0p8_1p0=0.3390 loss_all=3.7621 init_gold_top10=0.3994 init_gold_top100=0.4014 +step=9200 micro_steps=294400 elapsed=706.0s lr=3.000000e-04 loss=2.7077 loss_recon=2.7077 loss_meanflow=0.0000 mean_model_t=0.4989 mean_corrupt_t=0.4989 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5937 corrupt_frac=1.0000 acc_corrupt=0.5937 loss_corrupt=2.7077 wrong_frac=0.5010 init_acc_corrupt=0.4990 acc_corrupt_t_0p2_0p4=0.4100 corrupt_frac_t_0p2_0p4=0.3379 acc_corrupt_t_0p4_0p6=0.6382 corrupt_frac_t_0p4_0p6=0.3368 acc_corrupt_t_0p8_1p0=0.9477 corrupt_frac_t_0p8_1p0=0.3388 out_w_norm=121.1119 out_g_norm=0.0870 acc_corrupt_t_0p6_0p8=0.8161 corrupt_frac_t_0p6_0p8=0.3372 acc_corrupt_t_0p0_0p2=0.1587 corrupt_frac_t_0p0_0p2=0.3353 loss_all=3.0575 init_gold_top10=0.4058 init_gold_top100=0.4082 +step=9300 micro_steps=297600 elapsed=661.9s lr=3.000000e-04 loss=2.7200 loss_recon=2.7200 loss_meanflow=0.0000 mean_model_t=0.4990 mean_corrupt_t=0.4990 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5927 corrupt_frac=1.0000 acc_corrupt=0.5927 loss_corrupt=2.7200 wrong_frac=0.5009 init_acc_corrupt=0.4991 acc_corrupt_t_0p0_0p2=0.1564 corrupt_frac_t_0p0_0p2=0.3386 acc_corrupt_t_0p6_0p8=0.8203 corrupt_frac_t_0p6_0p8=0.3439 acc_corrupt_t_0p8_1p0=0.9491 corrupt_frac_t_0p8_1p0=0.3397 out_w_norm=121.9459 out_g_norm=0.0853 acc_corrupt_t_0p2_0p4=0.4085 corrupt_frac_t_0p2_0p4=0.3418 acc_corrupt_t_0p4_0p6=0.6361 corrupt_frac_t_0p4_0p6=0.3378 loss_all=3.0337 init_gold_top10=0.3738 init_gold_top100=0.3748 +step=9400 micro_steps=300800 elapsed=654.0s lr=3.000000e-04 loss=2.6920 loss_recon=2.6920 loss_meanflow=0.0000 mean_model_t=0.5020 mean_corrupt_t=0.5020 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5959 corrupt_frac=1.0000 acc_corrupt=0.5959 loss_corrupt=2.6920 wrong_frac=0.4979 init_acc_corrupt=0.5021 acc_corrupt_t_0p0_0p2=0.1552 corrupt_frac_t_0p0_0p2=0.3387 acc_corrupt_t_0p4_0p6=0.6360 corrupt_frac_t_0p4_0p6=0.3356 acc_corrupt_t_0p8_1p0=0.9491 corrupt_frac_t_0p8_1p0=0.3395 out_w_norm=122.7778 out_g_norm=0.0841 acc_corrupt_t_0p2_0p4=0.4079 corrupt_frac_t_0p2_0p4=0.3436 acc_corrupt_t_0p6_0p8=0.8187 corrupt_frac_t_0p6_0p8=0.3400 loss_all=2.5220 init_gold_top10=0.5081 init_gold_top100=0.5095 +step=9500 micro_steps=304000 elapsed=667.5s lr=3.000000e-04 loss=2.6650 loss_recon=2.6650 loss_meanflow=0.0000 mean_model_t=0.5031 mean_corrupt_t=0.5031 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5984 corrupt_frac=1.0000 acc_corrupt=0.5984 loss_corrupt=2.6650 wrong_frac=0.4968 init_acc_corrupt=0.5032 acc_corrupt_t_0p0_0p2=0.1574 corrupt_frac_t_0p0_0p2=0.3332 acc_corrupt_t_0p4_0p6=0.6358 corrupt_frac_t_0p4_0p6=0.3411 acc_corrupt_t_0p8_1p0=0.9497 corrupt_frac_t_0p8_1p0=0.3372 out_w_norm=123.5998 out_g_norm=0.0856 acc_corrupt_t_0p6_0p8=0.8183 corrupt_frac_t_0p6_0p8=0.3434 acc_corrupt_t_0p2_0p4=0.4111 corrupt_frac_t_0p2_0p4=0.3372 loss_all=5.0354 init_gold_top10=0.1895 init_gold_top100=0.1912 +step=9600 micro_steps=307200 elapsed=650.8s lr=3.000000e-04 loss=2.6948 loss_recon=2.6948 loss_meanflow=0.0000 mean_model_t=0.4993 mean_corrupt_t=0.4993 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5952 corrupt_frac=1.0000 acc_corrupt=0.5952 loss_corrupt=2.6948 wrong_frac=0.5005 init_acc_corrupt=0.4995 acc_corrupt_t_0p2_0p4=0.4104 corrupt_frac_t_0p2_0p4=0.3407 acc_corrupt_t_0p4_0p6=0.6372 corrupt_frac_t_0p4_0p6=0.3343 acc_corrupt_t_0p8_1p0=0.9480 corrupt_frac_t_0p8_1p0=0.3295 out_w_norm=124.4039 out_g_norm=0.0856 acc_corrupt_t_0p0_0p2=0.1589 corrupt_frac_t_0p0_0p2=0.3333 acc_corrupt_t_0p6_0p8=0.8176 corrupt_frac_t_0p6_0p8=0.3442 loss_all=4.5141 init_gold_top10=0.3186 init_gold_top100=0.3206 +step=9700 micro_steps=310400 elapsed=668.5s lr=3.000000e-04 loss=2.7597 loss_recon=2.7597 loss_meanflow=0.0000 mean_model_t=0.4932 mean_corrupt_t=0.4932 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5875 corrupt_frac=1.0000 acc_corrupt=0.5875 loss_corrupt=2.7597 wrong_frac=0.5069 init_acc_corrupt=0.4931 acc_corrupt_t_0p6_0p8=0.8165 corrupt_frac_t_0p6_0p8=0.3343 acc_corrupt_t_0p8_1p0=0.9496 corrupt_frac_t_0p8_1p0=0.3270 out_w_norm=125.2177 out_g_norm=0.0852 acc_corrupt_t_0p0_0p2=0.1577 corrupt_frac_t_0p0_0p2=0.3418 acc_corrupt_t_0p4_0p6=0.6344 corrupt_frac_t_0p4_0p6=0.3425 acc_corrupt_t_0p2_0p4=0.4091 corrupt_frac_t_0p2_0p4=0.3394 loss_all=4.2617 init_gold_top10=0.3176 init_gold_top100=0.3193 +step=9800 micro_steps=313600 elapsed=642.6s lr=3.000000e-04 loss=2.6706 loss_recon=2.6706 loss_meanflow=0.0000 mean_model_t=0.5048 mean_corrupt_t=0.5048 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5986 corrupt_frac=1.0000 acc_corrupt=0.5986 loss_corrupt=2.6706 wrong_frac=0.4952 init_acc_corrupt=0.5048 acc_corrupt_t_0p2_0p4=0.4088 corrupt_frac_t_0p2_0p4=0.3373 acc_corrupt_t_0p8_1p0=0.9494 corrupt_frac_t_0p8_1p0=0.3414 out_w_norm=126.0520 out_g_norm=0.0821 acc_corrupt_t_0p0_0p2=0.1589 corrupt_frac_t_0p0_0p2=0.3424 acc_corrupt_t_0p6_0p8=0.8176 corrupt_frac_t_0p6_0p8=0.3349 acc_corrupt_t_0p4_0p6=0.6358 corrupt_frac_t_0p4_0p6=0.3422 loss_all=0.7549 init_gold_top10=0.8105 init_gold_top100=0.8113 +step=9900 micro_steps=316800 elapsed=670.4s lr=3.000000e-04 loss=2.7014 loss_recon=2.7014 loss_meanflow=0.0000 mean_model_t=0.5007 mean_corrupt_t=0.5007 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5945 corrupt_frac=1.0000 acc_corrupt=0.5945 loss_corrupt=2.7014 wrong_frac=0.4994 init_acc_corrupt=0.5006 acc_corrupt_t_0p2_0p4=0.4065 corrupt_frac_t_0p2_0p4=0.3350 acc_corrupt_t_0p6_0p8=0.8191 corrupt_frac_t_0p6_0p8=0.3394 acc_corrupt_t_0p8_1p0=0.9501 corrupt_frac_t_0p8_1p0=0.3405 out_w_norm=126.8655 out_g_norm=0.0809 acc_corrupt_t_0p4_0p6=0.6404 corrupt_frac_t_0p4_0p6=0.3377 acc_corrupt_t_0p0_0p2=0.1573 corrupt_frac_t_0p0_0p2=0.3409 loss_all=3.0748 init_gold_top10=0.3967 init_gold_top100=0.3984 +step=10000 micro_steps=320000 elapsed=647.5s lr=3.000000e-04 loss=2.6882 loss_recon=2.6882 loss_meanflow=0.0000 mean_model_t=0.5006 mean_corrupt_t=0.5006 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5961 corrupt_frac=1.0000 acc_corrupt=0.5961 loss_corrupt=2.6882 wrong_frac=0.4991 init_acc_corrupt=0.5009 acc_corrupt_t_0p0_0p2=0.1590 corrupt_frac_t_0p0_0p2=0.3395 acc_corrupt_t_0p8_1p0=0.9491 corrupt_frac_t_0p8_1p0=0.3349 out_w_norm=127.6908 out_g_norm=0.0827 acc_corrupt_t_0p2_0p4=0.4089 corrupt_frac_t_0p2_0p4=0.3437 acc_corrupt_t_0p4_0p6=0.6388 corrupt_frac_t_0p4_0p6=0.3409 acc_corrupt_t_0p6_0p8=0.8190 corrupt_frac_t_0p6_0p8=0.3328 loss_all=2.1507 init_gold_top10=0.6318 init_gold_top100=0.6328 +step=10100 micro_steps=323200 elapsed=845.9s lr=3.000000e-04 loss=2.6763 loss_recon=2.6763 loss_meanflow=0.0000 mean_model_t=0.5007 mean_corrupt_t=0.5007 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5972 corrupt_frac=1.0000 acc_corrupt=0.5972 loss_corrupt=2.6763 wrong_frac=0.4993 init_acc_corrupt=0.5007 acc_corrupt_t_0p6_0p8=0.8179 corrupt_frac_t_0p6_0p8=0.3402 out_w_norm=128.5026 out_g_norm=0.0811 acc_corrupt_t_0p0_0p2=0.1596 corrupt_frac_t_0p0_0p2=0.3417 acc_corrupt_t_0p8_1p0=0.9495 corrupt_frac_t_0p8_1p0=0.3336 acc_corrupt_t_0p4_0p6=0.6389 corrupt_frac_t_0p4_0p6=0.3450 acc_corrupt_t_0p2_0p4=0.4109 corrupt_frac_t_0p2_0p4=0.3417 loss_all=1.8252 init_gold_top10=0.6191 init_gold_top100=0.6199 +step=10200 micro_steps=326400 elapsed=695.8s lr=3.000000e-04 loss=2.7018 loss_recon=2.7018 loss_meanflow=0.0000 mean_model_t=0.4974 mean_corrupt_t=0.4974 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5938 corrupt_frac=1.0000 acc_corrupt=0.5938 loss_corrupt=2.7018 wrong_frac=0.5024 init_acc_corrupt=0.4977 acc_corrupt_t_0p4_0p6=0.6387 corrupt_frac_t_0p4_0p6=0.3444 acc_corrupt_t_0p6_0p8=0.8210 corrupt_frac_t_0p6_0p8=0.3285 acc_corrupt_t_0p8_1p0=0.9504 corrupt_frac_t_0p8_1p0=0.3368 out_w_norm=129.3110 out_g_norm=0.0818 acc_corrupt_t_0p0_0p2=0.1590 corrupt_frac_t_0p0_0p2=0.3414 acc_corrupt_t_0p2_0p4=0.4126 corrupt_frac_t_0p2_0p4=0.3357 loss_all=1.2812 init_gold_top10=0.7019 init_gold_top100=0.7034 +step=10300 micro_steps=329600 elapsed=672.9s lr=3.000000e-04 loss=2.6639 loss_recon=2.6639 loss_meanflow=0.0000 mean_model_t=0.5029 mean_corrupt_t=0.5029 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5992 corrupt_frac=1.0000 acc_corrupt=0.5992 loss_corrupt=2.6639 wrong_frac=0.4970 init_acc_corrupt=0.5030 acc_corrupt_t_0p0_0p2=0.1580 corrupt_frac_t_0p0_0p2=0.3345 acc_corrupt_t_0p6_0p8=0.8225 corrupt_frac_t_0p6_0p8=0.3430 acc_corrupt_t_0p8_1p0=0.9492 corrupt_frac_t_0p8_1p0=0.3416 out_w_norm=130.1160 out_g_norm=0.0798 acc_corrupt_t_0p4_0p6=0.6403 corrupt_frac_t_0p4_0p6=0.3445 acc_corrupt_t_0p2_0p4=0.4120 corrupt_frac_t_0p2_0p4=0.3351 loss_all=2.0823 init_gold_top10=0.5723 init_gold_top100=0.5740 +step=10400 micro_steps=332800 elapsed=645.8s lr=3.000000e-04 loss=2.7074 loss_recon=2.7074 loss_meanflow=0.0000 mean_model_t=0.4975 mean_corrupt_t=0.4975 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5933 corrupt_frac=1.0000 acc_corrupt=0.5933 loss_corrupt=2.7074 wrong_frac=0.5023 init_acc_corrupt=0.4977 acc_corrupt_t_0p4_0p6=0.6371 corrupt_frac_t_0p4_0p6=0.3447 acc_corrupt_t_0p6_0p8=0.8208 corrupt_frac_t_0p6_0p8=0.3360 acc_corrupt_t_0p8_1p0=0.9496 corrupt_frac_t_0p8_1p0=0.3422 out_w_norm=130.9129 out_g_norm=0.0814 acc_corrupt_t_0p0_0p2=0.1576 corrupt_frac_t_0p0_0p2=0.3403 acc_corrupt_t_0p2_0p4=0.4107 corrupt_frac_t_0p2_0p4=0.3371 loss_all=2.3848 init_gold_top10=0.5127 init_gold_top100=0.5151 +step=10500 micro_steps=336000 elapsed=671.5s lr=3.000000e-04 loss=2.6893 loss_recon=2.6893 loss_meanflow=0.0000 mean_model_t=0.4992 mean_corrupt_t=0.4992 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5960 corrupt_frac=1.0000 acc_corrupt=0.5960 loss_corrupt=2.6893 wrong_frac=0.5008 init_acc_corrupt=0.4992 acc_corrupt_t_0p2_0p4=0.4128 corrupt_frac_t_0p2_0p4=0.3417 acc_corrupt_t_0p8_1p0=0.9501 corrupt_frac_t_0p8_1p0=0.3354 out_w_norm=131.7024 out_g_norm=0.0788 acc_corrupt_t_0p0_0p2=0.1582 corrupt_frac_t_0p0_0p2=0.3416 acc_corrupt_t_0p4_0p6=0.6411 corrupt_frac_t_0p4_0p6=0.3391 acc_corrupt_t_0p6_0p8=0.8218 corrupt_frac_t_0p6_0p8=0.3466 loss_all=2.9530 init_gold_top10=0.4424 init_gold_top100=0.4438 +step=10600 micro_steps=339200 elapsed=642.8s lr=3.000000e-04 loss=2.7234 loss_recon=2.7234 loss_meanflow=0.0000 mean_model_t=0.4945 mean_corrupt_t=0.4945 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5912 corrupt_frac=1.0000 acc_corrupt=0.5912 loss_corrupt=2.7234 wrong_frac=0.5055 init_acc_corrupt=0.4945 acc_corrupt_t_0p4_0p6=0.6397 corrupt_frac_t_0p4_0p6=0.3365 acc_corrupt_t_0p6_0p8=0.8223 corrupt_frac_t_0p6_0p8=0.3330 out_w_norm=132.4841 out_g_norm=0.0793 acc_corrupt_t_0p2_0p4=0.4122 corrupt_frac_t_0p2_0p4=0.3412 acc_corrupt_t_0p8_1p0=0.9501 corrupt_frac_t_0p8_1p0=0.3370 acc_corrupt_t_0p0_0p2=0.1617 corrupt_frac_t_0p0_0p2=0.3411 loss_all=3.0640 init_gold_top10=0.3984 init_gold_top100=0.3999 +step=10700 micro_steps=342400 elapsed=664.1s lr=3.000000e-04 loss=2.7026 loss_recon=2.7026 loss_meanflow=0.0000 mean_model_t=0.4975 mean_corrupt_t=0.4975 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5939 corrupt_frac=1.0000 acc_corrupt=0.5939 loss_corrupt=2.7026 wrong_frac=0.5027 init_acc_corrupt=0.4974 acc_corrupt_t_0p2_0p4=0.4116 corrupt_frac_t_0p2_0p4=0.3401 acc_corrupt_t_0p6_0p8=0.8195 corrupt_frac_t_0p6_0p8=0.3387 acc_corrupt_t_0p8_1p0=0.9504 corrupt_frac_t_0p8_1p0=0.3359 out_w_norm=133.2736 out_g_norm=0.0801 acc_corrupt_t_0p4_0p6=0.6401 corrupt_frac_t_0p4_0p6=0.3426 acc_corrupt_t_0p0_0p2=0.1599 corrupt_frac_t_0p0_0p2=0.3433 loss_all=1.6566 init_gold_top10=0.5759 init_gold_top100=0.5762 +step=10800 micro_steps=345600 elapsed=651.0s lr=3.000000e-04 loss=2.7030 loss_recon=2.7030 loss_meanflow=0.0000 mean_model_t=0.4957 mean_corrupt_t=0.4957 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5932 corrupt_frac=1.0000 acc_corrupt=0.5932 loss_corrupt=2.7030 wrong_frac=0.5043 init_acc_corrupt=0.4957 acc_corrupt_t_0p0_0p2=0.1599 corrupt_frac_t_0p0_0p2=0.3384 acc_corrupt_t_0p4_0p6=0.6402 corrupt_frac_t_0p4_0p6=0.3408 acc_corrupt_t_0p6_0p8=0.8192 corrupt_frac_t_0p6_0p8=0.3365 out_w_norm=134.0506 out_g_norm=0.0788 acc_corrupt_t_0p2_0p4=0.4123 corrupt_frac_t_0p2_0p4=0.3392 acc_corrupt_t_0p8_1p0=0.9492 corrupt_frac_t_0p8_1p0=0.3331 loss_all=1.0048 init_gold_top10=0.7144 init_gold_top100=0.7156 +step=10900 micro_steps=348800 elapsed=694.6s lr=3.000000e-04 loss=2.6890 loss_recon=2.6890 loss_meanflow=0.0000 mean_model_t=0.5002 mean_corrupt_t=0.5002 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5954 corrupt_frac=1.0000 acc_corrupt=0.5954 loss_corrupt=2.6890 wrong_frac=0.4999 init_acc_corrupt=0.5001 acc_corrupt_t_0p2_0p4=0.4125 corrupt_frac_t_0p2_0p4=0.3404 acc_corrupt_t_0p4_0p6=0.6423 corrupt_frac_t_0p4_0p6=0.3343 acc_corrupt_t_0p8_1p0=0.9487 corrupt_frac_t_0p8_1p0=0.3336 out_w_norm=134.8268 out_g_norm=0.0824 acc_corrupt_t_0p6_0p8=0.8199 corrupt_frac_t_0p6_0p8=0.3398 acc_corrupt_t_0p0_0p2=0.1568 corrupt_frac_t_0p0_0p2=0.3342 loss_all=1.2489 init_gold_top10=0.6880 init_gold_top100=0.6887 +step=11000 micro_steps=352000 elapsed=658.5s lr=3.000000e-04 loss=2.6673 loss_recon=2.6673 loss_meanflow=0.0000 mean_model_t=0.5023 mean_corrupt_t=0.5023 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5991 corrupt_frac=1.0000 acc_corrupt=0.5991 loss_corrupt=2.6673 wrong_frac=0.4976 init_acc_corrupt=0.5024 acc_corrupt_t_0p2_0p4=0.4142 corrupt_frac_t_0p2_0p4=0.3385 acc_corrupt_t_0p6_0p8=0.8216 corrupt_frac_t_0p6_0p8=0.3374 acc_corrupt_t_0p8_1p0=0.9494 corrupt_frac_t_0p8_1p0=0.3418 out_w_norm=135.6049 out_g_norm=0.0780 acc_corrupt_t_0p0_0p2=0.1593 corrupt_frac_t_0p0_0p2=0.3408 acc_corrupt_t_0p4_0p6=0.6423 corrupt_frac_t_0p4_0p6=0.3349 loss_all=4.5499 init_gold_top10=0.2642 init_gold_top100=0.2644 +step=11100 micro_steps=355200 elapsed=848.5s lr=3.000000e-04 loss=2.6619 loss_recon=2.6619 loss_meanflow=0.0000 mean_model_t=0.5038 mean_corrupt_t=0.5038 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5997 corrupt_frac=1.0000 acc_corrupt=0.5997 loss_corrupt=2.6619 wrong_frac=0.4962 init_acc_corrupt=0.5038 acc_corrupt_t_0p4_0p6=0.6386 corrupt_frac_t_0p4_0p6=0.3366 acc_corrupt_t_0p8_1p0=0.9501 corrupt_frac_t_0p8_1p0=0.3473 out_w_norm=136.3694 out_g_norm=0.0746 acc_corrupt_t_0p2_0p4=0.4092 corrupt_frac_t_0p2_0p4=0.3385 acc_corrupt_t_0p6_0p8=0.8230 corrupt_frac_t_0p6_0p8=0.3410 acc_corrupt_t_0p0_0p2=0.1598 corrupt_frac_t_0p0_0p2=0.3334 loss_all=0.8429 init_gold_top10=0.7468 init_gold_top100=0.7478 +step=11200 micro_steps=358400 elapsed=736.0s lr=3.000000e-04 loss=2.6817 loss_recon=2.6817 loss_meanflow=0.0000 mean_model_t=0.4990 mean_corrupt_t=0.4990 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5959 corrupt_frac=1.0000 acc_corrupt=0.5959 loss_corrupt=2.6817 wrong_frac=0.5009 init_acc_corrupt=0.4991 acc_corrupt_t_0p0_0p2=0.1584 corrupt_frac_t_0p0_0p2=0.3359 acc_corrupt_t_0p2_0p4=0.4098 corrupt_frac_t_0p2_0p4=0.3379 acc_corrupt_t_0p8_1p0=0.9492 corrupt_frac_t_0p8_1p0=0.3326 out_w_norm=137.1358 out_g_norm=0.0796 acc_corrupt_t_0p4_0p6=0.6409 corrupt_frac_t_0p4_0p6=0.3382 acc_corrupt_t_0p6_0p8=0.8220 corrupt_frac_t_0p6_0p8=0.3353 loss_all=1.7695 init_gold_top10=0.6304 init_gold_top100=0.6313 +step=11300 micro_steps=361600 elapsed=700.9s lr=3.000000e-04 loss=2.6979 loss_recon=2.6979 loss_meanflow=0.0000 mean_model_t=0.4977 mean_corrupt_t=0.4977 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5948 corrupt_frac=1.0000 acc_corrupt=0.5948 loss_corrupt=2.6979 wrong_frac=0.5024 init_acc_corrupt=0.4976 acc_corrupt_t_0p2_0p4=0.4107 corrupt_frac_t_0p2_0p4=0.3392 acc_corrupt_t_0p4_0p6=0.6447 corrupt_frac_t_0p4_0p6=0.3375 acc_corrupt_t_0p6_0p8=0.8207 corrupt_frac_t_0p6_0p8=0.3445 out_w_norm=137.9133 out_g_norm=0.0789 acc_corrupt_t_0p0_0p2=0.1590 corrupt_frac_t_0p0_0p2=0.3352 acc_corrupt_t_0p8_1p0=0.9509 corrupt_frac_t_0p8_1p0=0.3380 loss_all=3.1596 init_gold_top10=0.3938 init_gold_top100=0.3955 +step=11400 micro_steps=364800 elapsed=844.4s lr=3.000000e-04 loss=2.6698 loss_recon=2.6698 loss_meanflow=0.0000 mean_model_t=0.5007 mean_corrupt_t=0.5007 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5976 corrupt_frac=1.0000 acc_corrupt=0.5976 loss_corrupt=2.6698 wrong_frac=0.4994 init_acc_corrupt=0.5006 acc_corrupt_t_0p0_0p2=0.1599 corrupt_frac_t_0p0_0p2=0.3361 acc_corrupt_t_0p2_0p4=0.4118 corrupt_frac_t_0p2_0p4=0.3392 acc_corrupt_t_0p4_0p6=0.6420 corrupt_frac_t_0p4_0p6=0.3424 acc_corrupt_t_0p6_0p8=0.8220 corrupt_frac_t_0p6_0p8=0.3384 out_w_norm=138.6687 out_g_norm=0.0779 acc_corrupt_t_0p8_1p0=0.9497 corrupt_frac_t_0p8_1p0=0.3384 loss_all=3.0330 init_gold_top10=0.4612 init_gold_top100=0.4631 +step=11500 micro_steps=368000 elapsed=829.7s lr=3.000000e-04 loss=2.6705 loss_recon=2.6705 loss_meanflow=0.0000 mean_model_t=0.5006 mean_corrupt_t=0.5006 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5975 corrupt_frac=1.0000 acc_corrupt=0.5975 loss_corrupt=2.6705 wrong_frac=0.4994 init_acc_corrupt=0.5007 acc_corrupt_t_0p0_0p2=0.1598 corrupt_frac_t_0p0_0p2=0.3379 acc_corrupt_t_0p4_0p6=0.6393 corrupt_frac_t_0p4_0p6=0.3388 acc_corrupt_t_0p6_0p8=0.8223 corrupt_frac_t_0p6_0p8=0.3386 acc_corrupt_t_0p8_1p0=0.9494 corrupt_frac_t_0p8_1p0=0.3337 out_w_norm=139.4190 out_g_norm=0.0737 acc_corrupt_t_0p2_0p4=0.4131 corrupt_frac_t_0p2_0p4=0.3344 loss_all=2.8063 init_gold_top10=0.5381 init_gold_top100=0.5398 +step=11600 micro_steps=371200 elapsed=880.6s lr=3.000000e-04 loss=2.6477 loss_recon=2.6477 loss_meanflow=0.0000 mean_model_t=0.5043 mean_corrupt_t=0.5043 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.6013 corrupt_frac=1.0000 acc_corrupt=0.6013 loss_corrupt=2.6477 wrong_frac=0.4957 init_acc_corrupt=0.5043 acc_corrupt_t_0p2_0p4=0.4149 corrupt_frac_t_0p2_0p4=0.3371 acc_corrupt_t_0p6_0p8=0.8243 corrupt_frac_t_0p6_0p8=0.3362 acc_corrupt_t_0p8_1p0=0.9498 corrupt_frac_t_0p8_1p0=0.3428 out_w_norm=140.1790 out_g_norm=0.0741 acc_corrupt_t_0p0_0p2=0.1575 corrupt_frac_t_0p0_0p2=0.3371 acc_corrupt_t_0p4_0p6=0.6417 corrupt_frac_t_0p4_0p6=0.3387 loss_all=2.0381 init_gold_top10=0.5466 init_gold_top100=0.5479 +step=11700 micro_steps=374400 elapsed=790.2s lr=3.000000e-04 loss=2.6639 loss_recon=2.6639 loss_meanflow=0.0000 mean_model_t=0.5028 mean_corrupt_t=0.5028 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5996 corrupt_frac=1.0000 acc_corrupt=0.5996 loss_corrupt=2.6639 wrong_frac=0.4972 init_acc_corrupt=0.5028 acc_corrupt_t_0p0_0p2=0.1592 corrupt_frac_t_0p0_0p2=0.3406 acc_corrupt_t_0p2_0p4=0.4129 corrupt_frac_t_0p2_0p4=0.3377 acc_corrupt_t_0p8_1p0=0.9512 corrupt_frac_t_0p8_1p0=0.3437 out_w_norm=140.9233 out_g_norm=0.0740 acc_corrupt_t_0p4_0p6=0.6434 corrupt_frac_t_0p4_0p6=0.3426 acc_corrupt_t_0p6_0p8=0.8201 corrupt_frac_t_0p6_0p8=0.3429 loss_all=3.1155 init_gold_top10=0.4241 init_gold_top100=0.4263 +step=11800 micro_steps=377600 elapsed=708.9s lr=3.000000e-04 loss=2.6963 loss_recon=2.6963 loss_meanflow=0.0000 mean_model_t=0.4957 mean_corrupt_t=0.4957 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5938 corrupt_frac=1.0000 acc_corrupt=0.5938 loss_corrupt=2.6963 wrong_frac=0.5042 init_acc_corrupt=0.4958 acc_corrupt_t_0p2_0p4=0.4125 corrupt_frac_t_0p2_0p4=0.3450 acc_corrupt_t_0p6_0p8=0.8218 corrupt_frac_t_0p6_0p8=0.3425 acc_corrupt_t_0p8_1p0=0.9497 corrupt_frac_t_0p8_1p0=0.3338 out_w_norm=141.6821 out_g_norm=0.0742 acc_corrupt_t_0p0_0p2=0.1574 corrupt_frac_t_0p0_0p2=0.3379 acc_corrupt_t_0p4_0p6=0.6420 corrupt_frac_t_0p4_0p6=0.3408 loss_all=2.9546 init_gold_top10=0.4146 init_gold_top100=0.4153 +step=11900 micro_steps=380800 elapsed=641.2s lr=3.000000e-04 loss=2.7300 loss_recon=2.7300 loss_meanflow=0.0000 mean_model_t=0.4943 mean_corrupt_t=0.4943 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5907 corrupt_frac=1.0000 acc_corrupt=0.5907 loss_corrupt=2.7300 wrong_frac=0.5056 init_acc_corrupt=0.4944 acc_corrupt_t_0p0_0p2=0.1574 corrupt_frac_t_0p0_0p2=0.3446 acc_corrupt_t_0p6_0p8=0.8208 corrupt_frac_t_0p6_0p8=0.3371 acc_corrupt_t_0p8_1p0=0.9492 corrupt_frac_t_0p8_1p0=0.3372 out_w_norm=142.4417 out_g_norm=0.0749 acc_corrupt_t_0p2_0p4=0.4095 corrupt_frac_t_0p2_0p4=0.3404 acc_corrupt_t_0p4_0p6=0.6418 corrupt_frac_t_0p4_0p6=0.3398 loss_all=4.3424 init_gold_top10=0.3503 init_gold_top100=0.3521 +step=12000 micro_steps=384000 elapsed=912.2s lr=3.000000e-04 loss=2.6658 loss_recon=2.6658 loss_meanflow=0.0000 mean_model_t=0.4997 mean_corrupt_t=0.4997 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5984 corrupt_frac=1.0000 acc_corrupt=0.5984 loss_corrupt=2.6658 wrong_frac=0.5003 init_acc_corrupt=0.4997 acc_corrupt_t_0p0_0p2=0.1602 corrupt_frac_t_0p0_0p2=0.3420 acc_corrupt_t_0p2_0p4=0.4152 corrupt_frac_t_0p2_0p4=0.3415 acc_corrupt_t_0p6_0p8=0.8247 corrupt_frac_t_0p6_0p8=0.3395 out_w_norm=143.2062 out_g_norm=0.0747 acc_corrupt_t_0p8_1p0=0.9496 corrupt_frac_t_0p8_1p0=0.3359 acc_corrupt_t_0p4_0p6=0.6423 corrupt_frac_t_0p4_0p6=0.3398 loss_all=4.5359 init_gold_top10=0.2490 init_gold_top100=0.2524 +step=12100 micro_steps=387200 elapsed=1108.8s lr=3.000000e-04 loss=2.6854 loss_recon=2.6854 loss_meanflow=0.0000 mean_model_t=0.4992 mean_corrupt_t=0.4992 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5965 corrupt_frac=1.0000 acc_corrupt=0.5965 loss_corrupt=2.6854 wrong_frac=0.5006 init_acc_corrupt=0.4994 acc_corrupt_t_0p0_0p2=0.1602 corrupt_frac_t_0p0_0p2=0.3420 acc_corrupt_t_0p4_0p6=0.6448 corrupt_frac_t_0p4_0p6=0.3405 acc_corrupt_t_0p8_1p0=0.9516 corrupt_frac_t_0p8_1p0=0.3324 out_w_norm=143.9685 out_g_norm=0.0725 acc_corrupt_t_0p2_0p4=0.4130 corrupt_frac_t_0p2_0p4=0.3386 acc_corrupt_t_0p6_0p8=0.8219 corrupt_frac_t_0p6_0p8=0.3355 loss_all=2.9383 init_gold_top10=0.4868 init_gold_top100=0.4878 +step=12200 micro_steps=390400 elapsed=1077.3s lr=3.000000e-04 loss=2.6768 loss_recon=2.6768 loss_meanflow=0.0000 mean_model_t=0.5000 mean_corrupt_t=0.5000 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5973 corrupt_frac=1.0000 acc_corrupt=0.5973 loss_corrupt=2.6768 wrong_frac=0.4999 init_acc_corrupt=0.5001 acc_corrupt_t_0p2_0p4=0.4118 corrupt_frac_t_0p2_0p4=0.3399 acc_corrupt_t_0p4_0p6=0.6437 corrupt_frac_t_0p4_0p6=0.3350 acc_corrupt_t_0p8_1p0=0.9503 corrupt_frac_t_0p8_1p0=0.3389 out_w_norm=144.7137 out_g_norm=0.0721 acc_corrupt_t_0p0_0p2=0.1584 corrupt_frac_t_0p0_0p2=0.3404 acc_corrupt_t_0p6_0p8=0.8252 corrupt_frac_t_0p6_0p8=0.3418 loss_all=1.6161 init_gold_top10=0.6577 init_gold_top100=0.6587 +step=12300 micro_steps=393600 elapsed=1058.8s lr=3.000000e-04 loss=2.6903 loss_recon=2.6903 loss_meanflow=0.0000 mean_model_t=0.4964 mean_corrupt_t=0.4964 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5953 corrupt_frac=1.0000 acc_corrupt=0.5953 loss_corrupt=2.6903 wrong_frac=0.5036 init_acc_corrupt=0.4964 acc_corrupt_t_0p2_0p4=0.4165 corrupt_frac_t_0p2_0p4=0.3411 acc_corrupt_t_0p4_0p6=0.6425 corrupt_frac_t_0p4_0p6=0.3390 out_w_norm=145.4439 out_g_norm=0.0727 acc_corrupt_t_0p8_1p0=0.9508 corrupt_frac_t_0p8_1p0=0.3416 acc_corrupt_t_0p6_0p8=0.8264 corrupt_frac_t_0p6_0p8=0.3338 acc_corrupt_t_0p0_0p2=0.1569 corrupt_frac_t_0p0_0p2=0.3348 loss_all=1.6735 init_gold_top10=0.5957 init_gold_top100=0.5969 +step=12400 micro_steps=396800 elapsed=1063.2s lr=3.000000e-04 loss=2.6521 loss_recon=2.6521 loss_meanflow=0.0000 mean_model_t=0.5024 mean_corrupt_t=0.5024 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.6003 corrupt_frac=1.0000 acc_corrupt=0.6003 loss_corrupt=2.6521 wrong_frac=0.4978 init_acc_corrupt=0.5022 acc_corrupt_t_0p0_0p2=0.1564 corrupt_frac_t_0p0_0p2=0.3344 acc_corrupt_t_0p2_0p4=0.4137 corrupt_frac_t_0p2_0p4=0.3311 acc_corrupt_t_0p4_0p6=0.6435 corrupt_frac_t_0p4_0p6=0.3418 acc_corrupt_t_0p8_1p0=0.9500 corrupt_frac_t_0p8_1p0=0.3453 out_w_norm=146.1728 out_g_norm=0.0721 acc_corrupt_t_0p6_0p8=0.8226 corrupt_frac_t_0p6_0p8=0.3400 loss_all=2.3955 init_gold_top10=0.4827 init_gold_top100=0.4844 +step=12500 micro_steps=400000 elapsed=1059.4s lr=3.000000e-04 loss=2.6353 loss_recon=2.6353 loss_meanflow=0.0000 mean_model_t=0.5046 mean_corrupt_t=0.5046 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.6030 corrupt_frac=1.0000 acc_corrupt=0.6030 loss_corrupt=2.6353 wrong_frac=0.4954 init_acc_corrupt=0.5046 acc_corrupt_t_0p0_0p2=0.1598 corrupt_frac_t_0p0_0p2=0.3389 acc_corrupt_t_0p6_0p8=0.8223 corrupt_frac_t_0p6_0p8=0.3426 acc_corrupt_t_0p8_1p0=0.9511 corrupt_frac_t_0p8_1p0=0.3401 out_w_norm=146.8965 out_g_norm=0.0705 acc_corrupt_t_0p4_0p6=0.6463 corrupt_frac_t_0p4_0p6=0.3368 acc_corrupt_t_0p2_0p4=0.4151 corrupt_frac_t_0p2_0p4=0.3350 loss_all=2.5688 init_gold_top10=0.5767 init_gold_top100=0.5774 +step=12600 micro_steps=403200 elapsed=1061.6s lr=3.000000e-04 loss=2.6764 loss_recon=2.6764 loss_meanflow=0.0000 mean_model_t=0.4981 mean_corrupt_t=0.4981 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5965 corrupt_frac=1.0000 acc_corrupt=0.5965 loss_corrupt=2.6764 wrong_frac=0.5020 init_acc_corrupt=0.4980 acc_corrupt_t_0p4_0p6=0.6447 corrupt_frac_t_0p4_0p6=0.3433 acc_corrupt_t_0p8_1p0=0.9504 corrupt_frac_t_0p8_1p0=0.3406 out_w_norm=147.6190 out_g_norm=0.0719 acc_corrupt_t_0p0_0p2=0.1603 corrupt_frac_t_0p0_0p2=0.3413 acc_corrupt_t_0p6_0p8=0.8235 corrupt_frac_t_0p6_0p8=0.3415 acc_corrupt_t_0p2_0p4=0.4140 corrupt_frac_t_0p2_0p4=0.3375 loss_all=3.2488 init_gold_top10=0.4009 init_gold_top100=0.4016 +step=12700 micro_steps=406400 elapsed=693.7s lr=3.000000e-04 loss=2.6789 loss_recon=2.6789 loss_meanflow=0.0000 mean_model_t=0.4985 mean_corrupt_t=0.4985 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5968 corrupt_frac=1.0000 acc_corrupt=0.5968 loss_corrupt=2.6789 wrong_frac=0.5015 init_acc_corrupt=0.4985 acc_corrupt_t_0p0_0p2=0.1578 corrupt_frac_t_0p0_0p2=0.3376 acc_corrupt_t_0p4_0p6=0.6480 corrupt_frac_t_0p4_0p6=0.3363 acc_corrupt_t_0p6_0p8=0.8242 corrupt_frac_t_0p6_0p8=0.3385 out_w_norm=148.3549 out_g_norm=0.0713 acc_corrupt_t_0p2_0p4=0.4146 corrupt_frac_t_0p2_0p4=0.3390 acc_corrupt_t_0p8_1p0=0.9519 corrupt_frac_t_0p8_1p0=0.3351 loss_all=1.0905 init_gold_top10=0.7327 init_gold_top100=0.7334 +step=12800 micro_steps=409600 elapsed=981.2s lr=3.000000e-04 loss=2.6677 loss_recon=2.6677 loss_meanflow=0.0000 mean_model_t=0.4988 mean_corrupt_t=0.4988 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5977 corrupt_frac=1.0000 acc_corrupt=0.5977 loss_corrupt=2.6677 wrong_frac=0.5012 init_acc_corrupt=0.4988 acc_corrupt_t_0p2_0p4=0.4161 corrupt_frac_t_0p2_0p4=0.3422 acc_corrupt_t_0p4_0p6=0.6448 corrupt_frac_t_0p4_0p6=0.3409 acc_corrupt_t_0p8_1p0=0.9507 corrupt_frac_t_0p8_1p0=0.3360 out_w_norm=149.0799 out_g_norm=0.0703 acc_corrupt_t_0p6_0p8=0.8234 corrupt_frac_t_0p6_0p8=0.3339 acc_corrupt_t_0p0_0p2=0.1607 corrupt_frac_t_0p0_0p2=0.3398 loss_all=3.7772 init_gold_top10=0.3591 init_gold_top100=0.3613 +step=12900 micro_steps=412800 elapsed=786.9s lr=3.000000e-04 loss=2.6457 loss_recon=2.6457 loss_meanflow=0.0000 mean_model_t=0.5026 mean_corrupt_t=0.5026 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.6009 corrupt_frac=1.0000 acc_corrupt=0.6009 loss_corrupt=2.6457 wrong_frac=0.4974 init_acc_corrupt=0.5026 acc_corrupt_t_0p0_0p2=0.1580 corrupt_frac_t_0p0_0p2=0.3374 acc_corrupt_t_0p6_0p8=0.8267 corrupt_frac_t_0p6_0p8=0.3406 out_w_norm=149.8111 out_g_norm=0.0698 acc_corrupt_t_0p2_0p4=0.4148 corrupt_frac_t_0p2_0p4=0.3436 acc_corrupt_t_0p4_0p6=0.6445 corrupt_frac_t_0p4_0p6=0.3406 acc_corrupt_t_0p8_1p0=0.9513 corrupt_frac_t_0p8_1p0=0.3409 loss_all=3.9694 init_gold_top10=0.3030 init_gold_top100=0.3054 +step=13000 micro_steps=416000 elapsed=686.7s lr=3.000000e-04 loss=2.6762 loss_recon=2.6762 loss_meanflow=0.0000 mean_model_t=0.4977 mean_corrupt_t=0.4977 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5964 corrupt_frac=1.0000 acc_corrupt=0.5964 loss_corrupt=2.6762 wrong_frac=0.5024 init_acc_corrupt=0.4976 acc_corrupt_t_0p4_0p6=0.6443 corrupt_frac_t_0p4_0p6=0.3388 acc_corrupt_t_0p6_0p8=0.8221 corrupt_frac_t_0p6_0p8=0.3388 acc_corrupt_t_0p8_1p0=0.9505 corrupt_frac_t_0p8_1p0=0.3375 out_w_norm=150.5468 out_g_norm=0.0708 acc_corrupt_t_0p0_0p2=0.1595 corrupt_frac_t_0p0_0p2=0.3390 acc_corrupt_t_0p2_0p4=0.4130 corrupt_frac_t_0p2_0p4=0.3360 loss_all=2.8862 init_gold_top10=0.4663 init_gold_top100=0.4678 +step=13100 micro_steps=419200 elapsed=797.0s lr=3.000000e-04 loss=2.6863 loss_recon=2.6863 loss_meanflow=0.0000 mean_model_t=0.4969 mean_corrupt_t=0.4969 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5952 corrupt_frac=1.0000 acc_corrupt=0.5952 loss_corrupt=2.6863 wrong_frac=0.5031 init_acc_corrupt=0.4969 acc_corrupt_t_0p0_0p2=0.1600 corrupt_frac_t_0p0_0p2=0.3472 acc_corrupt_t_0p2_0p4=0.4148 corrupt_frac_t_0p2_0p4=0.3394 acc_corrupt_t_0p6_0p8=0.8225 corrupt_frac_t_0p6_0p8=0.3444 acc_corrupt_t_0p8_1p0=0.9509 corrupt_frac_t_0p8_1p0=0.3376 out_w_norm=151.3018 out_g_norm=0.0691 acc_corrupt_t_0p4_0p6=0.6438 corrupt_frac_t_0p4_0p6=0.3334 loss_all=4.9974 init_gold_top10=0.2483 init_gold_top100=0.2498 +step=13200 micro_steps=422400 elapsed=701.5s lr=3.000000e-04 loss=2.6915 loss_recon=2.6915 loss_meanflow=0.0000 mean_model_t=0.4949 mean_corrupt_t=0.4949 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5938 corrupt_frac=1.0000 acc_corrupt=0.5938 loss_corrupt=2.6915 wrong_frac=0.5049 init_acc_corrupt=0.4951 acc_corrupt_t_0p0_0p2=0.1573 corrupt_frac_t_0p0_0p2=0.3415 acc_corrupt_t_0p6_0p8=0.8232 corrupt_frac_t_0p6_0p8=0.3359 out_w_norm=152.0209 out_g_norm=0.0683 acc_corrupt_t_0p2_0p4=0.4136 corrupt_frac_t_0p2_0p4=0.3426 acc_corrupt_t_0p4_0p6=0.6430 corrupt_frac_t_0p4_0p6=0.3433 acc_corrupt_t_0p8_1p0=0.9515 corrupt_frac_t_0p8_1p0=0.3369 loss_all=2.3652 init_gold_top10=0.5081 init_gold_top100=0.5095 +step=13300 micro_steps=425600 elapsed=667.8s lr=3.000000e-04 loss=2.6664 loss_recon=2.6664 loss_meanflow=0.0000 mean_model_t=0.4995 mean_corrupt_t=0.4995 mean_loss_t_weight=1.0000 linear_soft_target_mean_conf=0.0000 prior_center_loss_beta=0.0000 rollout_train_applied=0.0000 grad_enabled_before_rollout=1.0000 grad_enabled_after_rollout=1.0000 logits_requires_grad=1.0000 raw_loss_requires_grad=1.0000 acc_all=0.5982 corrupt_frac=1.0000 acc_corrupt=0.5982 loss_corrupt=2.6664 wrong_frac=0.5005 init_acc_corrupt=0.4996 acc_corrupt_t_0p4_0p6=0.6470 corrupt_frac_t_0p4_0p6=0.3389 acc_corrupt_t_0p6_0p8=0.8247 corrupt_frac_t_0p6_0p8=0.3397 out_w_norm=152.7259 out_g_norm=0.0682 acc_corrupt_t_0p0_0p2=0.1561 corrupt_frac_t_0p0_0p2=0.3374 acc_corrupt_t_0p2_0p4=0.4143 corrupt_frac_t_0p2_0p4=0.3347 acc_corrupt_t_0p8_1p0=0.9523 corrupt_frac_t_0p8_1p0=0.3359 loss_all=3.2159 init_gold_top10=0.4009 init_gold_top100=0.4026 +W0526 16:17:55.734000 1191028 torch/distributed/elastic/agent/server/api.py:719] Received 15 death signal, shutting down workers +W0526 16:17:55.735000 1191028 torch/distributed/elastic/multiprocessing/api.py:898] Sending process 1191116 closing signal SIGTERM +W0526 16:17:55.735000 1191028 torch/distributed/elastic/multiprocessing/api.py:898] Sending process 1191117 closing signal SIGTERM +W0526 16:17:55.736000 1191028 torch/distributed/elastic/multiprocessing/api.py:898] Sending process 1191118 closing signal SIGTERM +W0526 16:17:55.740000 1191028 torch/distributed/elastic/multiprocessing/api.py:898] Sending process 1191119 closing signal SIGTERM +Traceback (most recent call last): + File "", line 198, in _run_module_as_main + File "", line 88, in _run_code + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/run.py", line 922, in + main() + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/elastic/multiprocessing/errors/__init__.py", line 355, in wrapper + return f(*args, **kwargs) + ^^^^^^^^^^^^^^^^^^ + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/run.py", line 918, in main + run(args) + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/run.py", line 909, in run + elastic_launch( + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/launcher/api.py", line 139, in __call__ + return launch_agent(self._config, self._entrypoint, list(args)) + ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/launcher/api.py", line 261, in launch_agent + result = agent.run() + ^^^^^^^^^^^ + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/elastic/metrics/api.py", line 137, in wrapper + result = f(*args, **kwargs) + ^^^^^^^^^^^^^^^^^^ + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/elastic/agent/server/api.py", line 711, in run + result = self._invoke_run(role) + ^^^^^^^^^^^^^^^^^^^^^^ + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/elastic/agent/server/api.py", line 870, in _invoke_run + time.sleep(monitor_interval) + File "/usr/local/lib/python3.12/dist-packages/torch/distributed/elastic/multiprocessing/api.py", line 84, in _terminate_process_handler + raise SignalException(f"Process {os.getpid()} got signal: {sigval}", sigval=sigval) +torch.distributed.elastic.multiprocessing.api.SignalException: Process 1191028 got signal: 15 diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..7a254464cc78ccea32b3ded00513c44c4e4da412 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/annotated_doc-0.0.4.dist-info/licenses/LICENSE @@ -0,0 +1,21 @@ +The MIT License (MIT) + +Copyright (c) 2025 Sebastián Ramírez + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in +all copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN +THE SOFTWARE. diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/pygments/modeline.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/pygments/modeline.py new file mode 100644 index 0000000000000000000000000000000000000000..81ec15773dd6cdc40c72e7bda280ce33f645ef1a --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/pygments/modeline.py @@ -0,0 +1,43 @@ +""" + pygments.modeline + ~~~~~~~~~~~~~~~~~ + + A simple modeline parser (based on pymodeline). + + :copyright: Copyright 2006-present by the Pygments team, see AUTHORS. + :license: BSD, see LICENSE for details. +""" + +import re + +__all__ = ['get_filetype_from_buffer'] + + +modeline_re = re.compile(r''' + (?: vi | vim | ex ) (?: [<=>]? \d* )? : + .* (?: ft | filetype | syn | syntax ) = ( [^:\s]+ ) +''', re.VERBOSE) + + +def get_filetype_from_line(l): # noqa: E741 + m = modeline_re.search(l) + if m: + return m.group(1) + + +def get_filetype_from_buffer(buf, max_lines=5): + """ + Scan the buffer for modelines and return filetype if one is found. + """ + lines = buf.splitlines() + for line in lines[-1:-max_lines-1:-1]: + ret = get_filetype_from_line(line) + if ret: + return ret + for i in range(max_lines, -1, -1): + if i < len(lines): + ret = get_filetype_from_line(lines[i]) + if ret: + return ret + + return None diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/nt.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/nt.py new file mode 100644 index 0000000000000000000000000000000000000000..389551b223a761fa2f97e929b60bf3ca5baed94c --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/nt.py @@ -0,0 +1,163 @@ +import contextlib +import ctypes +import os + +from ctypes.wintypes import ( + BOOL, + CHAR, + DWORD, + HANDLE, + LONG, + LPWSTR, + MAX_PATH, + PDWORD, + ULONG, +) + +from shellingham._core import SHELL_NAMES + + +INVALID_HANDLE_VALUE = HANDLE(-1).value +ERROR_NO_MORE_FILES = 18 +ERROR_INSUFFICIENT_BUFFER = 122 +TH32CS_SNAPPROCESS = 2 +PROCESS_QUERY_LIMITED_INFORMATION = 0x1000 + + +kernel32 = ctypes.windll.kernel32 + + +def _check_handle(error_val=0): + def check(ret, func, args): + if ret == error_val: + raise ctypes.WinError() + return ret + + return check + + +def _check_expected(expected): + def check(ret, func, args): + if ret: + return True + code = ctypes.GetLastError() + if code == expected: + return False + raise ctypes.WinError(code) + + return check + + +class ProcessEntry32(ctypes.Structure): + _fields_ = ( + ("dwSize", DWORD), + ("cntUsage", DWORD), + ("th32ProcessID", DWORD), + ("th32DefaultHeapID", ctypes.POINTER(ULONG)), + ("th32ModuleID", DWORD), + ("cntThreads", DWORD), + ("th32ParentProcessID", DWORD), + ("pcPriClassBase", LONG), + ("dwFlags", DWORD), + ("szExeFile", CHAR * MAX_PATH), + ) + + +kernel32.CloseHandle.argtypes = [HANDLE] +kernel32.CloseHandle.restype = BOOL + +kernel32.CreateToolhelp32Snapshot.argtypes = [DWORD, DWORD] +kernel32.CreateToolhelp32Snapshot.restype = HANDLE +kernel32.CreateToolhelp32Snapshot.errcheck = _check_handle( # type: ignore + INVALID_HANDLE_VALUE, +) + +kernel32.Process32First.argtypes = [HANDLE, ctypes.POINTER(ProcessEntry32)] +kernel32.Process32First.restype = BOOL +kernel32.Process32First.errcheck = _check_expected( # type: ignore + ERROR_NO_MORE_FILES, +) + +kernel32.Process32Next.argtypes = [HANDLE, ctypes.POINTER(ProcessEntry32)] +kernel32.Process32Next.restype = BOOL +kernel32.Process32Next.errcheck = _check_expected( # type: ignore + ERROR_NO_MORE_FILES, +) + +kernel32.GetCurrentProcessId.argtypes = [] +kernel32.GetCurrentProcessId.restype = DWORD + +kernel32.OpenProcess.argtypes = [DWORD, BOOL, DWORD] +kernel32.OpenProcess.restype = HANDLE +kernel32.OpenProcess.errcheck = _check_handle( # type: ignore + INVALID_HANDLE_VALUE, +) + +kernel32.QueryFullProcessImageNameW.argtypes = [HANDLE, DWORD, LPWSTR, PDWORD] +kernel32.QueryFullProcessImageNameW.restype = BOOL +kernel32.QueryFullProcessImageNameW.errcheck = _check_expected( # type: ignore + ERROR_INSUFFICIENT_BUFFER, +) + + +@contextlib.contextmanager +def _handle(f, *args, **kwargs): + handle = f(*args, **kwargs) + try: + yield handle + finally: + kernel32.CloseHandle(handle) + + +def _iter_processes(): + f = kernel32.CreateToolhelp32Snapshot + with _handle(f, TH32CS_SNAPPROCESS, 0) as snap: + entry = ProcessEntry32() + entry.dwSize = ctypes.sizeof(entry) + ret = kernel32.Process32First(snap, entry) + while ret: + yield entry + ret = kernel32.Process32Next(snap, entry) + + +def _get_full_path(proch): + size = DWORD(MAX_PATH) + while True: + path_buff = ctypes.create_unicode_buffer("", size.value) + if kernel32.QueryFullProcessImageNameW(proch, 0, path_buff, size): + return path_buff.value + size.value *= 2 + + +def get_shell(pid=None, max_depth=10): + proc_map = { + proc.th32ProcessID: (proc.th32ParentProcessID, proc.szExeFile) + for proc in _iter_processes() + } + pid = pid or os.getpid() + + for _ in range(0, max_depth + 1): + try: + ppid, executable = proc_map[pid] + except KeyError: # No such process? Give up. + break + + # The executable name would be encoded with the current code page if + # we're in ANSI mode (usually). Try to decode it into str/unicode, + # replacing invalid characters to be safe (not thoeratically necessary, + # I think). Note that we need to use 'mbcs' instead of encoding + # settings from sys because this is from the Windows API, not Python + # internals (which those settings reflect). (pypa/pipenv#3382) + if isinstance(executable, bytes): + executable = executable.decode("mbcs", "replace") + + name = executable.rpartition(".")[0].lower() + if name not in SHELL_NAMES: + pid = ppid + continue + + key = PROCESS_QUERY_LIMITED_INFORMATION + with _handle(kernel32.OpenProcess, key, 0, pid) as proch: + return (name, _get_full_path(proch)) + + return None diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/_core.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/_core.py new file mode 100644 index 0000000000000000000000000000000000000000..adc49e6e7a9d3edf062c55e0078136899f78d30d --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/_core.py @@ -0,0 +1,3 @@ +import collections + +Process = collections.namedtuple("Process", "args pid ppid") diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/proc.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/proc.py new file mode 100644 index 0000000000000000000000000000000000000000..950f63228e5b328f82b70da8851ec60c6a2ff029 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/shellingham/posix/proc.py @@ -0,0 +1,83 @@ +import io +import os +import re +import sys + +from ._core import Process + +# FreeBSD: https://www.freebsd.org/cgi/man.cgi?query=procfs +# NetBSD: https://man.netbsd.org/NetBSD-9.3-STABLE/mount_procfs.8 +# DragonFlyBSD: https://www.dragonflybsd.org/cgi/web-man?command=procfs +BSD_STAT_PPID = 2 + +# See https://docs.kernel.org/filesystems/proc.html +LINUX_STAT_PPID = 3 + +STAT_PATTERN = re.compile(r"\(.+\)|\S+") + + +def detect_proc(): + """Detect /proc filesystem style. + + This checks the /proc/{pid} directory for possible formats. Returns one of + the following as str: + + * `stat`: Linux-style, i.e. ``/proc/{pid}/stat``. + * `status`: BSD-style, i.e. ``/proc/{pid}/status``. + """ + pid = os.getpid() + for name in ("stat", "status"): + if os.path.exists(os.path.join("/proc", str(pid), name)): + return name + raise ProcFormatError("unsupported proc format") + + +def _use_bsd_stat_format(): + try: + return os.uname().sysname.lower() in ("freebsd", "netbsd", "dragonfly") + except Exception: + return False + + +def _get_ppid(pid, name): + path = os.path.join("/proc", str(pid), name) + with io.open(path, encoding="ascii", errors="replace") as f: + parts = STAT_PATTERN.findall(f.read()) + # We only care about TTY and PPID -- both are numbers. + if _use_bsd_stat_format(): + return parts[BSD_STAT_PPID] + return parts[LINUX_STAT_PPID] + + +def _get_cmdline(pid): + path = os.path.join("/proc", str(pid), "cmdline") + encoding = sys.getfilesystemencoding() or "utf-8" + with io.open(path, encoding=encoding, errors="replace") as f: + # XXX: Command line arguments can be arbitrary byte sequences, not + # necessarily decodable. For Shellingham's purpose, however, we don't + # care. (pypa/pipenv#2820) + # cmdline appends an extra NULL at the end, hence the [:-1]. + return tuple(f.read().split("\0")[:-1]) + + +class ProcFormatError(EnvironmentError): + pass + + +def iter_process_parents(pid, max_depth=10): + """Try to look up the process tree via the /proc interface.""" + stat_name = detect_proc() + + # Inner generator function so we correctly throw an error eagerly if proc + # is not supported, rather than on the first call to the iterator. This + # allows the call site detects the correct implementation. + def _iter_process_parents(pid, max_depth): + for _ in range(max_depth): + ppid = _get_ppid(pid, stat_name) + args = _get_cmdline(pid) + yield Process(args=args, pid=pid, ppid=ppid) + if ppid == "0": + break + pid = ppid + + return _iter_process_parents(pid, max_depth) diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/image_processing_aria.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/image_processing_aria.py new file mode 100644 index 0000000000000000000000000000000000000000..e0abc55236f0499d63e5c774fec363bba893efe7 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/image_processing_aria.py @@ -0,0 +1,226 @@ +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# This file was automatically generated from src/transformers/models/aria/modular_aria.py. +# Do NOT edit this file manually as any edits will be overwritten by the generation of +# the file from the modular. If any change should be done, please apply the change to the +# modular_aria.py file directly. One of our CI enforces this. +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# Copyright 2024 The Rhymes-AI Teams Authors and The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import torch +from torchvision.transforms.v2 import functional as tvF + +from ...image_processing_backends import TorchvisionBackend +from ...image_processing_utils import BatchFeature, get_patch_output_size, select_best_resolution +from ...image_transforms import divide_to_patches +from ...image_utils import ChannelDimension, PILImageResampling, SizeDict, get_image_size +from ...processing_utils import ImagesKwargs, Unpack +from ...utils import TensorType, auto_docstring + + +class AriaImageProcessorKwargs(ImagesKwargs, total=False): + r""" + max_image_size (`int`, *optional*, defaults to `self.max_image_size`): + Maximum image size. Must be either 490 or 980. + min_image_size (`int`, *optional*, defaults to `self.min_image_size`): + Minimum image size. Images smaller than this in any dimension will be scaled up. + split_resolutions (`list[list[int]]`, *optional*, defaults to `self.split_resolutions`): + A list of possible resolutions as (height, width) pairs for splitting high-resolution images into patches. + split_image (`bool`, *optional*, defaults to `self.split_image`): + Whether to split the image into patches using the best matching resolution from `split_resolutions`. + """ + + max_image_size: int + min_image_size: int + split_resolutions: list[list[int]] + split_image: bool + + +@auto_docstring +class AriaImageProcessor(TorchvisionBackend): + model_input_names = ["pixel_values", "pixel_mask", "num_crops"] + valid_kwargs = AriaImageProcessorKwargs + + resample = PILImageResampling.BICUBIC + image_mean = [0.5, 0.5, 0.5] + image_std = [0.5, 0.5, 0.5] + max_image_size = 980 + min_image_size = 336 + split_image = False + split_resolutions = None + do_convert_rgb = True + do_rescale = True + do_normalize = True + + def __init__(self, **kwargs: Unpack[AriaImageProcessorKwargs]): + if kwargs.get("split_resolutions") is None: + 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 + kwargs["split_resolutions"] = [[el[0] * 490, el[1] * 490] for el in default_resolutions] + super().__init__(**kwargs) + + def _get_padding_size(self, original_resolution: tuple, target_resolution: tuple) -> list[int]: + """Get padding size for patching, returns [left, top, right, bottom] for tvF.pad.""" + original_height, original_width = original_resolution + target_height, target_width = target_resolution + paste_x, r_x = divmod(target_width - original_width, 2) + paste_y, r_y = divmod(target_height - original_height, 2) + return [paste_x, paste_y, paste_x + r_x, paste_y + r_y] + + def _resize_for_patching( + self, + image: "torch.Tensor", + target_resolution: tuple, + resample: "PILImageResampling | tvF.InterpolationMode | int | None", + ) -> "torch.Tensor": + """Resize an image to a target resolution while maintaining aspect ratio.""" + new_height, new_width = get_patch_output_size( + image, target_resolution, input_data_format=ChannelDimension.FIRST + ) + return self.resize(image, SizeDict(height=new_height, width=new_width), resample) + + def _pad_for_patching( + self, + image: "torch.Tensor", + target_resolution: tuple, + ) -> "torch.Tensor": + """Pad an image to a target resolution while maintaining aspect ratio.""" + new_resolution = get_patch_output_size(image, target_resolution, input_data_format=ChannelDimension.FIRST) + padding = self._get_padding_size(new_resolution, target_resolution) + return tvF.pad(image, padding=padding) + + def get_image_patches( + self, + image: "torch.Tensor", + grid_pinpoints: list[list[int]], + patch_size: int, + resample: "PILImageResampling | tvF.InterpolationMode | int | None", + ) -> list["torch.Tensor"]: + """ + Process an image with variable resolutions by dividing it into patches. + + Args: + image (`torch.Tensor`): + The input image to be processed (channels-first format). + grid_pinpoints (`list[list[int]]`): + A list of possible resolutions as (height, width) pairs. + patch_size (`int`): + Size of each square patch to divide the image into. + resample (`PILImageResampling | tvF.InterpolationMode | int | None`): + Resampling filter to use when resizing. + + Returns: + `list[torch.Tensor]`: A list of image patches in channels-first format. + """ + if not isinstance(grid_pinpoints, list): + raise TypeError("grid_pinpoints must be a list of possible resolutions.") + + image_size = get_image_size(image, channel_dim=ChannelDimension.FIRST) + best_resolution = select_best_resolution(image_size, grid_pinpoints) + resized_image = self._resize_for_patching(image, best_resolution, resample) + padded_image = self._pad_for_patching(resized_image, best_resolution) + patches = divide_to_patches(padded_image, patch_size=patch_size) + return patches + + def _preprocess( + self, + images: list["torch.Tensor"], + do_rescale: bool, + rescale_factor: float, + do_normalize: bool, + image_mean: float | list[float] | None, + image_std: float | list[float] | None, + disable_grouping: bool | None, + return_tensors: str | TensorType | None, + max_image_size: int = 980, + min_image_size: int = 336, + split_resolutions: list[list[int]] | None = None, + split_image: bool = False, + resample: "PILImageResampling | tvF.InterpolationMode | int | None" = None, + **kwargs, + ) -> BatchFeature: + if max_image_size not in [490, 980]: + raise ValueError("max_image_size must be either 490 or 980") + + pixel_masks = [] + processed_crops = [] + num_crops = None + + for image in images: + if split_image: + crop_images = self.get_image_patches(image, split_resolutions, max_image_size, resample) + else: + crop_images = [image] + + if num_crops is None or len(crop_images) > num_crops: + num_crops = len(crop_images) + + for crop_image in crop_images: + h, w = crop_image.shape[-2], crop_image.shape[-1] + scale = max_image_size / max(h, w) + if w >= h: + new_h = max(int(h * scale), min_image_size) + new_w = max_image_size + else: + new_h = max_image_size + new_w = max(int(w * scale), min_image_size) + + crop_image = self.resize(crop_image, SizeDict(height=new_h, width=new_w), resample) + + padding_bottom = max_image_size - new_h + padding_right = max_image_size - new_w + crop_image = tvF.pad(crop_image, [0, 0, padding_right, padding_bottom]) + + pixel_mask = torch.zeros((max_image_size, max_image_size), dtype=torch.bool) + pixel_mask[:new_h, :new_w] = True + pixel_masks.append(pixel_mask) + processed_crops.append(crop_image) + + stacked_images = torch.stack(processed_crops, dim=0) + stacked_images = self.rescale_and_normalize( + stacked_images, do_rescale, rescale_factor, do_normalize, image_mean, image_std + ) + stacked_masks = torch.stack(pixel_masks, dim=0) + + return BatchFeature( + data={ + "pixel_values": stacked_images, + "pixel_mask": stacked_masks, + "num_crops": num_crops, + }, + tensor_type=return_tensors, + ) + + def get_number_of_image_patches(self, height: int, width: int, images_kwargs=None): + """ + A utility that returns number of image patches for a given image size. + + Args: + height (`int`): + Height of the input image. + width (`int`): + Width of the input image. + images_kwargs (`dict`, *optional*): + Any kwargs to override defaults of the image processor. + + Returns: + `int`: Number of patches per image. + """ + split_image = images_kwargs.get("split_image", self.split_image) + max_image_size = images_kwargs.get("max_image_size", self.max_image_size) + + resized_height, resized_width = select_best_resolution((height, width), self.split_resolutions) + num_patches = 1 if not split_image else resized_height // max_image_size * resized_width // max_image_size + return num_patches + + +__all__ = ["AriaImageProcessor"] diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/processing_aria.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/processing_aria.py new file mode 100644 index 0000000000000000000000000000000000000000..8c9fa8188c81a9cf725d8938e33e5ec4a6be1a66 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/aria/processing_aria.py @@ -0,0 +1,177 @@ +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# This file was automatically generated from src/transformers/models/aria/modular_aria.py. +# Do NOT edit this file manually as any edits will be overwritten by the generation of +# the file from the modular. If any change should be done, please apply the change to the +# modular_aria.py file directly. One of our CI enforces this. +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# Copyright 2024 The Rhymes-AI Teams Authors and The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from ...image_processing_utils import BatchFeature +from ...image_utils import ImageInput +from ...processing_utils import ImagesKwargs, MultiModalData, ProcessingKwargs, ProcessorMixin, Unpack +from ...tokenization_python import PreTokenizedInput, TextInput +from ...utils import TensorType, auto_docstring +from ..auto import AutoTokenizer + + +class AriaImagesKwargs(ImagesKwargs, total=False): + """ + split_image (`bool`, *optional*, defaults to `False`): + Whether to split large images into multiple crops. When enabled, images exceeding the maximum size are + divided into overlapping crops that are processed separately and then combined. This allows processing + of very high-resolution images that exceed the model's input size limits. + max_image_size (`int`, *optional*, defaults to `980`): + Maximum image size (in pixels) for a single image crop. Images larger than this will be split into + multiple crops when `split_image=True`, or resized if splitting is disabled. This parameter controls + the maximum resolution of individual image patches processed by the model. + min_image_size (`int`, *optional*): + Minimum image size (in pixels) for a single image crop. Images smaller than this will be upscaled to + meet the minimum requirement. If not specified, images are processed at their original size (subject + to the maximum size constraint). + """ + + split_image: bool + max_image_size: int + min_image_size: int + + +class AriaProcessorKwargs(ProcessingKwargs, total=False): + images_kwargs: AriaImagesKwargs + + _defaults = { + "text_kwargs": { + "padding": False, + "return_mm_token_type_ids": False, + }, + "images_kwargs": { + "max_image_size": 980, + "split_image": False, + }, + "return_tensors": TensorType.PYTORCH, + } + + +@auto_docstring +class AriaProcessor(ProcessorMixin): + def __init__( + self, + image_processor=None, + tokenizer: AutoTokenizer | str = None, + chat_template: str | None = None, + size_conversion: dict[float | int, int] | None = None, + ): + r""" + size_conversion (`Dict`, *optional*): + A dictionary indicating size conversions for images. + """ + if size_conversion is None: + size_conversion = {490: 128, 980: 256} + self.size_conversion = {int(k): v for k, v in size_conversion.items()} + + self.image_token = tokenizer.image_token + self.image_token_id = tokenizer.image_token_id + if tokenizer is not None and tokenizer.pad_token is None: + tokenizer.pad_token = tokenizer.unk_token + + super().__init__(image_processor, tokenizer, chat_template=chat_template) + + @auto_docstring + def __call__( + self, + text: TextInput | PreTokenizedInput | list[TextInput] | list[PreTokenizedInput], + images: ImageInput | None = None, + **kwargs: Unpack[AriaProcessorKwargs], + ) -> BatchFeature: + r""" + Returns: + [`BatchFeature`]: A [`BatchFeature`] with the following fields: + - **input_ids** -- List of token ids to be fed to a model. Returned when `text` is not `None`. + - **attention_mask** -- List of indices specifying which tokens should be attended to by the model (when + `return_attention_mask=True` or if *"attention_mask"* is in `self.model_input_names` and if `text` is not + `None`). + - **pixel_values** -- Pixel values to be fed to a model. Returned when `images` is not `None`. + - **pixel_mask** -- Pixel mask to be fed to a model. Returned when `images` is not `None`. + """ + output_kwargs = self._merge_kwargs( + AriaProcessorKwargs, + tokenizer_init_kwargs=self.tokenizer.init_kwargs, + **kwargs, + ) + + if isinstance(text, str): + text = [text] + elif not isinstance(text, list) and not isinstance(text[0], str): + raise TypeError("Invalid input text. Please provide a string, or a list of strings") + + if images is not None: + image_inputs = self.image_processor(images, **output_kwargs["images_kwargs"]) + # expand the image_token according to the num_crops and tokens per image + tokens_per_image = self.size_conversion[image_inputs.pixel_values.shape[2]] + prompt_strings = [] + num_crops = image_inputs.pop("num_crops") * tokens_per_image + for sample in text: + sample = sample.replace(self.tokenizer.image_token, self.tokenizer.image_token * num_crops) + prompt_strings.append(sample) + + else: + image_inputs = {} + prompt_strings = text + + return_tensors = output_kwargs["text_kwargs"].pop("return_tensors", None) + return_mm_token_type_ids = output_kwargs["text_kwargs"].pop("return_mm_token_type_ids", False) + text_inputs = self.tokenizer(prompt_strings, **output_kwargs["text_kwargs"], return_tensors=None) + self._check_special_mm_tokens(prompt_strings, text_inputs, modalities=["image"]) + + if return_mm_token_type_ids: + text_inputs["mm_token_type_ids"] = self.create_mm_token_type_ids(text_inputs["input_ids"]) + return BatchFeature(data={**text_inputs, **image_inputs}, tensor_type=return_tensors) + + def _get_num_multimodal_tokens(self, image_sizes=None, **kwargs): + """ + Computes the number of placeholder tokens needed for multimodal inputs with the given sizes. + Args: + image_sizes (`list[list[int]]`, *optional*): + The input sizes formatted as (height, width) per each image. + Returns: + `MultiModalData`: A `MultiModalData` object holding number of tokens per each of the provided + input modalities, along with other useful data. + """ + + vision_data = {} + if image_sizes is not None: + images_kwargs = AriaProcessorKwargs._defaults.get("images_kwargs", {}) + images_kwargs.update(kwargs) + + max_size = images_kwargs.get("max_image_size", None) or self.image_processor.max_image_size + num_image_patches = [ + self.image_processor.get_number_of_image_patches(*image_size, images_kwargs) + for image_size in image_sizes + ] + num_image_tokens = [self.size_conversion[max_size] * num_patches for num_patches in num_image_patches] + vision_data.update({"num_image_tokens": num_image_tokens, "num_image_patches": num_image_patches}) + + return MultiModalData(**vision_data) + + @property + def model_input_names(self): + tokenizer_input_names = self.tokenizer.model_input_names + image_processor_input_names = self.image_processor.model_input_names + + # Remove `num_crops`, it is popped and used only when processing. Make a copy of list when removing + # otherwise `self.image_processor.model_input_names` is also modified + image_processor_input_names = [name for name in image_processor_input_names if name != "num_crops"] + return list(dict.fromkeys(tokenizer_input_names + image_processor_input_names)) + + +__all__ = ["AriaProcessor"] diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/__init__.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..10e4ae987f7dd7642a190d594be313a98f955713 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/__init__.py @@ -0,0 +1,28 @@ +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from typing import TYPE_CHECKING + +from ...utils import _LazyModule +from ...utils.import_utils import define_import_structure + + +if TYPE_CHECKING: + from .configuration_eomt_dinov3 import * + from .modeling_eomt_dinov3 import * +else: + import sys + + _file = globals()["__file__"] + sys.modules[__name__] = _LazyModule(__name__, _file, define_import_structure(_file), module_spec=__spec__) diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/configuration_eomt_dinov3.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/configuration_eomt_dinov3.py new file mode 100644 index 0000000000000000000000000000000000000000..93f6350f26d2e08c5faf36ae5079b451bc3becc7 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/configuration_eomt_dinov3.py @@ -0,0 +1,107 @@ +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# This file was automatically generated from src/transformers/models/eomt_dinov3/modular_eomt_dinov3.py. +# Do NOT edit this file manually as any edits will be overwritten by the generation of +# the file from the modular. If any change should be done, please apply the change to the +# modular_eomt_dinov3.py file directly. One of our CI enforces this. +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from huggingface_hub.dataclasses import strict + +from ...configuration_utils import PreTrainedConfig +from ...modeling_rope_utils import RopeParameters +from ...utils import auto_docstring + + +@auto_docstring(checkpoint="tue-mps/coco_panoptic_eomt_large_640_dinov3") +@strict +class EomtDinov3Config(PreTrainedConfig): + r""" + layerscale_value (`float`, *optional*, defaults to 1.0): + Initial value for the LayerScale parameter. + num_upscale_blocks (`int`, *optional*, defaults to 2): + Number of upsampling blocks used in the decoder or segmentation head. + num_blocks (`int`, *optional*, defaults to 4): + Number of feature blocks or stages in the architecture. + no_object_weight (`float`, *optional*, defaults to 0.1): + Loss weight for the "no object" class in panoptic/instance segmentation. + train_num_points (`int`, *optional*, defaults to 12544): + Number of points to sample for mask loss computation during training. + oversample_ratio (`float`, *optional*, defaults to 3.0): + Oversampling ratio used in point sampling for mask training. + importance_sample_ratio (`float`, *optional*, defaults to 0.75): + Ratio of points to sample based on importance during training. + num_queries (`int`, *optional*, defaults to 200): + Number of object queries in the Transformer. + num_register_tokens (`int`, *optional*, defaults to 4): + Number of learnable register tokens added to the transformer input. + query_bias (`bool`, *optional*, defaults to `True`): + Whether to use bias in query projection. + key_bias (`bool`, *optional*, defaults to `False`): + Whether to use bias in key projection. + value_bias (`bool`, *optional*, defaults to `True`): + Whether to use bias in value projection. + proj_bias (`bool`, *optional*, defaults to `True`): + Whether to use bias in output projection. + use_gated_mlp (`bool`, *optional*, defaults to `False`): + Whether to use gated MLP layers. + pos_embed_shift (`float`, *optional*): + Shift value for position embeddings. + pos_embed_jitter (`float`, *optional*): + Jitter value for position embeddings. + pos_embed_rescale (`float`, *optional*, defaults to 2.0): + Rescale value for position embeddings. + """ + + model_type = "eomt_dinov3" + + hidden_size: int = 1024 + num_hidden_layers: int = 24 + num_attention_heads: int = 16 + hidden_act: str = "gelu" + hidden_dropout_prob: float | int = 0.0 + initializer_range: float = 0.02 + layer_norm_eps: float = 1e-6 + image_size: int | list[int] | tuple[int, int] = 640 + patch_size: int | list[int] | tuple[int, int] = 16 + num_channels: int = 3 + layerscale_value: float = 1.0 + drop_path_rate: float | int = 0.0 + num_upscale_blocks: int = 2 + attention_dropout: float | int = 0.0 + num_blocks: int = 4 + no_object_weight: float = 0.1 + class_weight: float = 2.0 + mask_weight: float = 5.0 + dice_weight: float = 5.0 + train_num_points: int = 12544 + oversample_ratio: float = 3.0 + importance_sample_ratio: float = 0.75 + num_queries: int = 200 + num_register_tokens: int = 4 + default_theta = 100.0 + intermediate_size: int = 4096 + rope_parameters: RopeParameters | dict | None = None + query_bias: bool = True + key_bias: bool = False + value_bias: bool = True + proj_bias: bool = True + mlp_bias: bool = True + use_gated_mlp: bool = False + pos_embed_shift: float | None = None + pos_embed_jitter: float | None = None + pos_embed_rescale: float | None = 2.0 + + +__all__ = ["EomtDinov3Config"] diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modeling_eomt_dinov3.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modeling_eomt_dinov3.py new file mode 100644 index 0000000000000000000000000000000000000000..18d2c37f042c2f87ee6a50940163b44dd7b7b922 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modeling_eomt_dinov3.py @@ -0,0 +1,1374 @@ +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# This file was automatically generated from src/transformers/models/eomt_dinov3/modular_eomt_dinov3.py. +# Do NOT edit this file manually as any edits will be overwritten by the generation of +# the file from the modular. If any change should be done, please apply the change to the +# modular_eomt_dinov3.py file directly. One of our CI enforces this. +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +from collections.abc import Callable +from dataclasses import dataclass +from typing import Optional + +import numpy as np +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from ... import initialization as init +from ...activations import ACT2FN +from ...file_utils import ModelOutput, is_scipy_available, requires_backends +from ...modeling_layers import GradientCheckpointingLayer +from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel +from ...processing_utils import Unpack +from ...pytorch_utils import compile_compatible_method_lru_cache +from ...utils import TransformersKwargs, auto_docstring, is_accelerate_available +from ...utils.generic import maybe_autocast, merge_with_config_defaults +from ...utils.output_capturing import capture_outputs +from .configuration_eomt_dinov3 import EomtDinov3Config + + +if is_scipy_available(): + from scipy.optimize import linear_sum_assignment + +if is_accelerate_available(): + from accelerate import PartialState + from accelerate.utils import reduce + + +def rotate_half(x): + """Rotates half the hidden dims of the input.""" + x1 = x[..., : x.shape[-1] // 2] + x2 = x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + + +def eager_attention_forward( + module: nn.Module, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attention_mask: torch.Tensor | None, + scaling: float | None = None, + dropout: float = 0.0, + **kwargs: Unpack[TransformersKwargs], +): + if scaling is None: + scaling = query.size(-1) ** -0.5 + + # Take the dot product between "query" and "key" to get the raw attention scores. + attn_weights = torch.matmul(query, key.transpose(2, 3)) * scaling + + if attention_mask is not None: + attn_weights = attn_weights + attention_mask + + attn_weights = nn.functional.softmax(attn_weights, dim=-1) + attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) + + attn_output = torch.matmul(attn_weights, value) + attn_output = attn_output.transpose(1, 2).contiguous() + + return attn_output, attn_weights + + +def apply_rotary_pos_emb( + q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, **kwargs +) -> tuple[torch.Tensor, torch.Tensor]: + """Applies Rotary Position Embedding to the query and key tensors, but only to the patch tokens, + ignoring the prefix tokens (cls token and register tokens). + + Args: + q (`torch.Tensor`): The query tensor. + k (`torch.Tensor`): The key tensor. + cos (`torch.Tensor`): The cosine part of the rotary embedding. + sin (`torch.Tensor`): The sine part of the rotary embedding. + + Returns: + `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding. + """ + + num_tokens = q.shape[-2] + num_patches = sin.shape[-2] + num_prefix_tokens = num_tokens - num_patches # cls token + register tokens + + q_prefix_tokens, q_patches = q.split((num_prefix_tokens, num_patches), dim=-2) + k_prefix_tokens, k_patches = k.split((num_prefix_tokens, num_patches), dim=-2) + + # apply rope only to patch tokens + q_patches = (q_patches * cos) + (rotate_half(q_patches) * sin) + k_patches = (k_patches * cos) + (rotate_half(k_patches) * sin) + + q = torch.cat((q_prefix_tokens, q_patches), dim=-2) + k = torch.cat((k_prefix_tokens, k_patches), dim=-2) + + return q, k + + +class EomtDinov3Attention(nn.Module): + """ + Multi-headed attention compatible with ALL_ATTENTION_FUNCTIONS. + """ + + def __init__(self, config: EomtDinov3Config): + super().__init__() + self.config = config + self.embed_dim = config.hidden_size + self.num_heads = config.num_attention_heads + self.head_dim = self.embed_dim // self.num_heads + self.is_causal = False + + self.scaling = self.head_dim**-0.5 + self.is_causal = False + + self.dropout = config.attention_dropout + self.k_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.key_bias) + self.v_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.value_bias) + + self.q_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.query_bias) + self.o_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=config.proj_bias) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None]: + """Input shape: Batch x Time x Channel""" + + batch_size, patches, _ = hidden_states.size() + + query_states = self.q_proj(hidden_states) + key_states = self.k_proj(hidden_states) + value_states = self.v_proj(hidden_states) + + query_states = query_states.view(batch_size, patches, self.num_heads, self.head_dim).transpose(1, 2) + key_states = key_states.view(batch_size, patches, self.num_heads, self.head_dim).transpose(1, 2) + value_states = value_states.view(batch_size, patches, self.num_heads, self.head_dim).transpose(1, 2) + + cos, sin = position_embeddings + query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) + + attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface( + self.config._attn_implementation, eager_attention_forward + ) + + attn_output, attn_weights = attention_interface( + self, + query_states, + key_states, + value_states, + attention_mask, + dropout=0.0 if not self.training else self.dropout, + scaling=self.scaling, + **kwargs, + ) + + attn_output = attn_output.reshape(batch_size, patches, -1).contiguous() + attn_output = self.o_proj(attn_output) + + return attn_output, attn_weights + + +class EomtDinov3Embeddings(nn.Module): + """ + Construct the CLS token, mask token, position and patch embeddings. + """ + + def __init__(self, config: EomtDinov3Config): + super().__init__() + self.config = config + self.cls_token = nn.Parameter(torch.randn(1, 1, config.hidden_size)) + self.mask_token = nn.Parameter(torch.zeros(1, 1, config.hidden_size)) + self.register_tokens = nn.Parameter(torch.empty(1, config.num_register_tokens, config.hidden_size)) + self.patch_embeddings = nn.Conv2d( + config.num_channels, config.hidden_size, kernel_size=config.patch_size, stride=config.patch_size + ) + self.num_prefix_tokens = 1 + config.num_register_tokens + + def forward(self, pixel_values: torch.Tensor, bool_masked_pos: torch.Tensor | None = None) -> torch.Tensor: + batch_size = pixel_values.shape[0] + target_dtype = self.patch_embeddings.weight.dtype + + # (batch_size, num_channels, height, width) -> (batch_size, num_patches, hidden_size) + patch_embeddings = self.patch_embeddings(pixel_values.to(dtype=target_dtype)) + patch_embeddings = patch_embeddings.flatten(2).transpose(1, 2) + + if bool_masked_pos is not None: + mask_token = self.mask_token.to(patch_embeddings.dtype) + patch_embeddings = torch.where(bool_masked_pos.unsqueeze(-1), mask_token, patch_embeddings) + + # Add CLS and register tokens + cls_token = self.cls_token.expand(batch_size, -1, -1) + register_tokens = self.register_tokens.expand(batch_size, -1, -1) + embeddings = torch.cat([cls_token, register_tokens, patch_embeddings], dim=1) + + return embeddings + + +class EomtDinov3MLP(nn.Module): + def __init__(self, config): + super().__init__() + self.config = config + self.hidden_size = config.hidden_size + self.intermediate_size = config.intermediate_size + self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias) + self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.mlp_bias) + self.act_fn = ACT2FN[config.hidden_act] + + def forward(self, x): + return self.down_proj(self.act_fn(self.up_proj(x))) + + +class EomtDinov3GatedMLP(nn.Module): + def __init__(self, config): + super().__init__() + self.config = config + self.hidden_size = config.hidden_size + self.intermediate_size = config.intermediate_size + self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias) + self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=config.mlp_bias) + self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=config.mlp_bias) + self.act_fn = ACT2FN[config.hidden_act] + + def forward(self, x): + down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) + return down_proj + + +class EomtDinov3DropPath(nn.Module): + """Stochastic depth (DropPath) per sample, for residual blocks. + + Identity when ``drop_prob`` is 0 or outside training. See `Deep Networks with Stochastic Depth + `_. + """ + + def __init__(self, drop_prob: float = 0.0) -> None: + super().__init__() + self.drop_prob = drop_prob + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + if self.drop_prob == 0.0 or not self.training: + return hidden_states + keep_prob = 1 - self.drop_prob + shape = (hidden_states.shape[0],) + (1,) * (hidden_states.ndim - 1) + random_tensor = torch.rand(shape, dtype=hidden_states.dtype, device=hidden_states.device) + random_tensor = torch.floor(random_tensor + keep_prob) + return hidden_states.div(keep_prob) * random_tensor + + def extra_repr(self) -> str: + return f"p={self.drop_prob}" + + +class EomtDinov3Layer(GradientCheckpointingLayer): + """This corresponds to the Block class in the original implementation.""" + + def __init__(self, config: EomtDinov3Config): + super().__init__() + + self.norm1 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) + self.attention = EomtDinov3Attention(config) + self.layer_scale1 = EomtDinov3LayerScale(config) + self.drop_path = EomtDinov3DropPath(config.drop_path_rate) if config.drop_path_rate > 0.0 else nn.Identity() + + self.norm2 = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) + + if config.use_gated_mlp: + self.mlp = EomtDinov3GatedMLP(config) + else: + self.mlp = EomtDinov3MLP(config) + self.layer_scale2 = EomtDinov3LayerScale(config) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> torch.Tensor: + # Attention with residual connection + residual = hidden_states + hidden_states = self.norm1(hidden_states) + hidden_states, _ = self.attention( + hidden_states, + attention_mask=attention_mask, + position_embeddings=position_embeddings, + **kwargs, + ) + hidden_states = self.layer_scale1(hidden_states) + hidden_states = self.drop_path(hidden_states) + residual + + # MLP with residual connection + residual = hidden_states + hidden_states = self.norm2(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = self.layer_scale2(hidden_states) + hidden_states = self.drop_path(hidden_states) + residual + + return hidden_states + + +class EomtDinov3LayerScale(nn.Module): + def __init__(self, config) -> None: + super().__init__() + self.lambda1 = nn.Parameter(config.layerscale_value * torch.ones(config.hidden_size)) + + def forward(self, hidden_state: torch.Tensor) -> torch.Tensor: + return hidden_state * self.lambda1 + + +@compile_compatible_method_lru_cache(maxsize=32) +def get_patches_center_coordinates( + num_patches_h: int, num_patches_w: int, dtype: torch.dtype, device: torch.device +) -> torch.Tensor: + """ + Computes the 2D coordinates of the centers of image patches, normalized to the range [-1, +1]. + The center of each patch is exactly halfway between its top-left and bottom-right corners. + + Args: + num_patches_h (int): Number of patches along the vertical (height) axis. + num_patches_w (int): Number of patches along the horizontal (width) axis. + dtype (torch.dtype): The desired data type of the returned tensor. + + Returns: + torch.Tensor: A tensor of shape (height * width, 2), where each row contains the (y, x) + coordinates of a patch center, normalized to [-1, +1]. + """ + coords_h = torch.arange(0.5, num_patches_h, dtype=dtype, device=device) + coords_w = torch.arange(0.5, num_patches_w, dtype=dtype, device=device) + coords_h = coords_h / num_patches_h + coords_w = coords_w / num_patches_w + # (height, width, 2) -> (height * width, 2) + coords = torch.stack(torch.meshgrid(coords_h, coords_w, indexing="ij"), dim=-1) + coords = coords.flatten(0, 1) + # Shift range [0, 1] to [-1, +1] + coords = 2.0 * coords - 1.0 + return coords + + +def augment_patches_center_coordinates( + coords: torch.Tensor, + shift: float | None = None, + jitter: float | None = None, + rescale: float | None = None, +) -> torch.Tensor: + # Shift coords by adding a uniform value in [-shift, shift] + if shift is not None: + shift_hw = torch.empty((1, 2), device=coords.device, dtype=coords.dtype) + shift_hw = shift_hw.uniform_(-shift, shift) + coords = coords + shift_hw + + # Jitter coords by multiplying the range [-1, 1] by a log-uniform value in [1/jitter, jitter] + if jitter is not None: + jitter_range = np.log(jitter) + jitter_hw = torch.empty((1, 2), device=coords.device, dtype=coords.dtype) + jitter_hw = jitter_hw.uniform_(-jitter_range, jitter_range).exp() + coords = coords * jitter_hw + + # Rescale coords by multiplying the range [-1, 1] by a log-uniform value in [1/rescale, rescale] + if rescale is not None: + rescale_range = np.log(rescale) + rescale_hw = torch.empty(1, device=coords.device, dtype=coords.dtype) + rescale_hw = rescale_hw.uniform_(-rescale_range, rescale_range).exp() + coords = coords * rescale_hw + + return coords + + +class EomtDinov3RotaryEmbedding(nn.Module): + inv_freq: Tensor + + def __init__(self, config: EomtDinov3Config, device=None): + super().__init__() + self.config = config + + self.rope_type = self.config.rope_parameters["rope_type"] + rope_init_fn: Callable = self.compute_default_rope_parameters + if self.rope_type != "default": + raise ValueError("`EomtDinov3` only supports `default` RoPE! Please check your `rope_type`") + inv_freq, self.attention_scaling = rope_init_fn(self.config, device) + + self.register_buffer("inv_freq", inv_freq, persistent=False) + self.register_buffer("original_inv_freq", inv_freq.clone(), persistent=False) + + def forward(self, pixel_values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + _, _, height, width = pixel_values.shape + num_patches_h = height // self.config.patch_size + num_patches_w = width // self.config.patch_size + + device = pixel_values.device + device_type = device.type if isinstance(device.type, str) and device.type != "mps" else "cpu" + + with maybe_autocast(device_type=device_type, enabled=False): # Force float32 + # Although we could precompute static patch_coords from image_size and patch_size in the config, + # the model was trained with random_scale, so it can process images of varying sizes. + # Therefore, it's better to compute patch_coords dynamically (with lru_cache). + patch_coords = get_patches_center_coordinates( + num_patches_h, num_patches_w, dtype=torch.float32, device=device + ) + if self.training: + patch_coords = augment_patches_center_coordinates( + patch_coords, + shift=self.config.pos_embed_shift, + jitter=self.config.pos_embed_jitter, + rescale=self.config.pos_embed_rescale, + ) + + # (height * width, 2, head_dim / 4) -> (height * width, head_dim / 2) -> (height * width, head_dim) + angles = 2 * math.pi * patch_coords[:, :, None] * self.inv_freq[None, None, :] + angles = angles.flatten(1, 2) + angles = angles.tile(2) + + cos = torch.cos(angles) + sin = torch.sin(angles) + + dtype = pixel_values.dtype + return cos.to(dtype=dtype), sin.to(dtype=dtype) + + @staticmethod + def compute_default_rope_parameters( + config: EomtDinov3Config | None = None, + device: Optional["torch.device"] = None, + seq_len: int | None = None, + ) -> torch.Tensor: + """ + Computes the inverse frequencies according to the original RoPE implementation + Args: + config ([`~transformers.PreTrainedConfig`]): + The model configuration. + device (`torch.device`): + The device to use for initialization of the inverse frequencies. + seq_len (`int`, *optional*): + The current sequence length. Unused for this type of RoPE. + Returns: + Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the + post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE). + """ + base = config.rope_parameters["rope_theta"] + head_dim = config.hidden_size // config.num_attention_heads + + attention_factor = 1.0 # Unused in this type of RoPE + + # Compute the inverse frequencies + inv_freq = 1 / base ** torch.arange(0, 1, 4 / head_dim, dtype=torch.float32, device=device) + return inv_freq, attention_factor + + +# Adapted from https://github.com/facebookresearch/detectron2/blob/main/projects/PointRend/point_rend/point_features.py +def sample_point( + input_features: torch.Tensor, point_coordinates: torch.Tensor, add_dim=False, **kwargs +) -> torch.Tensor: + """ + A wrapper around `torch.nn.functional.grid_sample` to support 3D point_coordinates tensors. + + Args: + input_features (`torch.Tensor` of shape (batch_size, channels, height, width)): + A tensor that contains features map on a height * width grid + point_coordinates (`torch.Tensor` of shape (batch_size, num_points, 2) or (batch_size, grid_height, grid_width,: + 2)): + A tensor that contains [0, 1] * [0, 1] normalized point coordinates + add_dim (`bool`): + boolean value to keep track of added dimension + + Returns: + point_features (`torch.Tensor` of shape (batch_size, channels, num_points) or (batch_size, channels, + height_grid, width_grid): + A tensor that contains features for points in `point_coordinates`. + """ + if point_coordinates.dim() == 3: + add_dim = True + point_coordinates = point_coordinates.unsqueeze(2) + + # use nn.function.grid_sample to get features for points in `point_coordinates` via bilinear interpolation + point_features = torch.nn.functional.grid_sample(input_features, 2.0 * point_coordinates - 1.0, **kwargs) + if add_dim: + point_features = point_features.squeeze(3) + + return point_features + + +def pair_wise_dice_loss(inputs: Tensor, labels: Tensor) -> Tensor: + """ + A pair wise version of the dice loss, see `dice_loss` for usage. + + Args: + inputs (`torch.Tensor`): + A tensor representing a mask + labels (`torch.Tensor`): + A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs + (0 for the negative class and 1 for the positive class). + + Returns: + `torch.Tensor`: The computed loss between each pairs. + """ + inputs = inputs.sigmoid().flatten(1) + numerator = 2 * torch.matmul(inputs, labels.T) + # using broadcasting to get a [num_queries, NUM_CLASSES] matrix + denominator = inputs.sum(-1)[:, None] + labels.sum(-1)[None, :] + loss = 1 - (numerator + 1) / (denominator + 1) + return loss + + +def pair_wise_sigmoid_cross_entropy_loss(inputs: torch.Tensor, labels: torch.Tensor) -> torch.Tensor: + r""" + A pair wise version of the cross entropy loss, see `sigmoid_cross_entropy_loss` for usage. + + Args: + inputs (`torch.Tensor`): + A tensor representing a mask. + labels (`torch.Tensor`): + A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs + (0 for the negative class and 1 for the positive class). + + Returns: + loss (`torch.Tensor`): The computed loss between each pairs. + """ + + height_and_width = inputs.shape[1] + + criterion = nn.BCEWithLogitsLoss(reduction="none") + cross_entropy_loss_pos = criterion(inputs, torch.ones_like(inputs)) + cross_entropy_loss_neg = criterion(inputs, torch.zeros_like(inputs)) + + loss_pos = torch.matmul(cross_entropy_loss_pos / height_and_width, labels.T) + loss_neg = torch.matmul(cross_entropy_loss_neg / height_and_width, (1 - labels).T) + loss = loss_pos + loss_neg + return loss + + +# Adapted from https://github.com/facebookresearch/EomtDinov3/blob/main/eomt_dinov3/modeling/matcher.py +class EomtDinov3HungarianMatcher(nn.Module): + """This class computes an assignment between the labels and the predictions of the network. + + For efficiency reasons, the labels don't include the no_object. Because of this, in general, there are more + predictions than labels. In this case, we do a 1-to-1 matching of the best predictions, while the others are + un-matched (and thus treated as non-objects). + """ + + def __init__( + self, cost_class: float = 1.0, cost_mask: float = 1.0, cost_dice: float = 1.0, num_points: int = 12544 + ): + """Creates the matcher + + Params: + cost_class (`float`, *optional*, defaults to 1.0): + Relative weight of the classification error in the matching cost. + cost_mask (`float`, *optional*, defaults to 1.0): + This is the relative weight of the focal loss of the binary mask in the matching cost. + cost_dice (`float`, *optional*, defaults to 1.0): + This is the relative weight of the dice loss of the binary mask in the matching cost. + num_points (`int`, *optional*, defaults to 12544): + No. of points to sample on which the mask loss will be calculated. The same set of K points are + uniformly sampled for all prediction and ground truth masks to construct the cost matrix for bipartite + matching. + """ + super().__init__() + if cost_class == 0 and cost_mask == 0 and cost_dice == 0: + raise ValueError("All costs can't be 0") + + self.num_points = num_points + self.cost_class = cost_class + self.cost_mask = cost_mask + self.cost_dice = cost_dice + + @torch.no_grad() + def forward( + self, + masks_queries_logits: torch.Tensor, + class_queries_logits: torch.Tensor, + mask_labels: torch.Tensor, + class_labels: torch.Tensor, + ) -> list[tuple[Tensor]]: + """ + Params: + masks_queries_logits (`torch.Tensor`): + A tensor of dim `batch_size, num_queries, num_labels` with the classification logits. + class_queries_logits (`torch.Tensor`): + A tensor of dim `batch_size, num_queries, height, width` with the predicted masks. + class_labels (`torch.Tensor`): + A tensor of dim `num_target_boxes` (where num_target_boxes is the number of ground-truth objects in the + target) containing the class labels. + mask_labels (`torch.Tensor`): + A tensor of dim `num_target_boxes, height, width` containing the target masks. + + Returns: + matched_indices (`list[tuple[Tensor]]`): A list of size batch_size, containing tuples of (index_i, index_j) + where: + - index_i is the indices of the selected predictions (in order) + - index_j is the indices of the corresponding selected labels (in order) + For each batch element, it holds: + len(index_i) = len(index_j) = min(num_queries, num_target_boxes). + """ + indices: list[tuple[np.array]] = [] + + # iterate through batch size + batch_size = masks_queries_logits.shape[0] + for i in range(batch_size): + pred_probs = class_queries_logits[i].softmax(-1) + pred_mask = masks_queries_logits[i] + + # 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. + cost_class = -pred_probs[:, class_labels[i]] + target_mask = mask_labels[i].to(pred_mask) + target_mask = target_mask[:, None] + pred_mask = pred_mask[:, None] + + # Sample ground truth and predicted masks + point_coordinates = torch.rand(1, self.num_points, 2, device=pred_mask.device) + + target_coordinates = point_coordinates.repeat(target_mask.shape[0], 1, 1) + target_mask = sample_point(target_mask, target_coordinates, align_corners=False).squeeze(1) + + pred_coordinates = point_coordinates.repeat(pred_mask.shape[0], 1, 1) + pred_mask = sample_point(pred_mask, pred_coordinates, align_corners=False).squeeze(1) + + # compute the cross entropy loss between each mask pairs -> shape (num_queries, num_labels) + cost_mask = pair_wise_sigmoid_cross_entropy_loss(pred_mask, target_mask) + # Compute the dice loss between each mask pairs -> shape (num_queries, num_labels) + cost_dice = pair_wise_dice_loss(pred_mask, target_mask) + # final cost matrix + cost_matrix = self.cost_mask * cost_mask + self.cost_class * cost_class + self.cost_dice * cost_dice + # eliminate infinite values in cost_matrix to avoid the error ``ValueError: cost matrix is infeasible`` + cost_matrix = torch.minimum(cost_matrix, torch.tensor(1e10)) + cost_matrix = torch.maximum(cost_matrix, torch.tensor(-1e10)) + cost_matrix = torch.nan_to_num(cost_matrix, 0) + # do the assignment using the hungarian algorithm in scipy + assigned_indices: tuple[np.array] = linear_sum_assignment(cost_matrix.cpu()) + indices.append(assigned_indices) + + # It could be stacked in one tensor + matched_indices = [ + (torch.as_tensor(i, dtype=torch.int64), torch.as_tensor(j, dtype=torch.int64)) for i, j in indices + ] + return matched_indices + + +def dice_loss(inputs: Tensor, labels: Tensor, num_masks: int) -> Tensor: + r""" + Compute the DICE loss, similar to generalized IOU for masks as follows: + + $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x \cap y }{x \cup y + 1}} $$ + + In practice, since `labels` is a binary mask, (only 0s and 1s), dice can be computed as follow + + $$ \mathcal{L}_{\text{dice}(x, y) = 1 - \frac{2 * x * y }{x + y + 1}} $$ + + Args: + inputs (`torch.Tensor`): + A tensor representing a mask. + labels (`torch.Tensor`): + A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs + (0 for the negative class and 1 for the positive class). + num_masks (`int`): + The number of masks present in the current batch, used for normalization. + + Returns: + `torch.Tensor`: The computed loss. + """ + probs = inputs.sigmoid().flatten(1) + numerator = 2 * (probs * labels).sum(-1) + denominator = probs.sum(-1) + labels.sum(-1) + loss = 1 - (numerator + 1) / (denominator + 1) + loss = loss.sum() / num_masks + return loss + + +def sigmoid_cross_entropy_loss(inputs: torch.Tensor, labels: torch.Tensor, num_masks: int) -> torch.Tensor: + r""" + Args: + inputs (`torch.Tensor`): + A float tensor of arbitrary shape. + labels (`torch.Tensor`): + A tensor with the same shape as inputs. Stores the binary classification labels for each element in inputs + (0 for the negative class and 1 for the positive class). + + Returns: + loss (`torch.Tensor`): The computed loss. + """ + criterion = nn.BCEWithLogitsLoss(reduction="none") + cross_entropy_loss = criterion(inputs, labels) + + loss = cross_entropy_loss.mean(1).sum() / num_masks + return loss + + +# Adapted from https://github.com/facebookresearch/EomtDinov3/blob/main/eomt_dinov3/modeling/criterion.py +class EomtDinov3Loss(nn.Module): + def __init__(self, config: EomtDinov3Config, weight_dict: dict[str, float]): + """ + The EomtDinov3 Loss. The loss is computed very similar to DETR. The process happens in two steps: 1) we + compute hungarian assignment between ground truth masks and the outputs of the model 2) we supervise each pair + of matched ground-truth / prediction (supervise class and mask) + + Args: + config (`EomtDinov3Config`): + The configuration for EomtDinov3 model also containing loss calculation specific parameters. + weight_dict (`dict[str, float]`): + A dictionary of weights to be applied to the different losses. + """ + super().__init__() + requires_backends(self, ["scipy"]) + self.num_labels = config.num_labels + self.weight_dict = weight_dict + + # Weight to apply to the null class + self.eos_coef = config.no_object_weight + empty_weight = torch.ones(self.num_labels + 1) + empty_weight[-1] = self.eos_coef + self.register_buffer("empty_weight", empty_weight) + + # pointwise mask loss parameters + self.num_points = config.train_num_points + self.oversample_ratio = config.oversample_ratio + self.importance_sample_ratio = config.importance_sample_ratio + + self.matcher = EomtDinov3HungarianMatcher( + cost_class=config.class_weight, + cost_dice=config.dice_weight, + cost_mask=config.mask_weight, + num_points=self.num_points, + ) + + def _max_by_axis(self, sizes: list[list[int]]) -> list[int]: + maxes = sizes[0] + for sublist in sizes[1:]: + for index, item in enumerate(sublist): + maxes[index] = max(maxes[index], item) + return maxes + + # Adapted from nested_tensor_from_tensor_list() in original implementation + def _pad_images_to_max_in_batch(self, tensors: list[Tensor]) -> tuple[Tensor, Tensor]: + # get the maximum size in the batch + max_size = self._max_by_axis([list(tensor.shape) for tensor in tensors]) + # compute final size + batch_shape = [len(tensors)] + max_size + batch_size, _, height, width = batch_shape + dtype = tensors[0].dtype + device = tensors[0].device + padded_tensors = torch.zeros(batch_shape, dtype=dtype, device=device) + padding_masks = torch.ones((batch_size, height, width), dtype=torch.bool, device=device) + # pad the tensors to the size of the biggest one + for tensor, padded_tensor, padding_mask in zip(tensors, padded_tensors, padding_masks): + padded_tensor[: tensor.shape[0], : tensor.shape[1], : tensor.shape[2]].copy_(tensor) + padding_mask[: tensor.shape[1], : tensor.shape[2]] = False + + return padded_tensors, padding_masks + + def loss_labels( + self, class_queries_logits: Tensor, class_labels: list[Tensor], indices: tuple[np.array] + ) -> dict[str, Tensor]: + """Compute the losses related to the labels using cross entropy. + + Args: + class_queries_logits (`torch.Tensor`): + A tensor of shape `batch_size, num_queries, num_labels` + class_labels (`list[torch.Tensor]`): + List of class labels of shape `(labels)`. + indices (`tuple[np.array])`: + The indices computed by the Hungarian matcher. + + Returns: + `dict[str, Tensor]`: A dict of `torch.Tensor` containing the following key: + - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels. + """ + pred_logits = class_queries_logits + batch_size, num_queries, _ = pred_logits.shape + criterion = nn.CrossEntropyLoss(weight=self.empty_weight) + idx = self._get_predictions_permutation_indices(indices) # shape of (batch_size, num_queries) + target_classes_o = torch.cat( + [target[j] for target, (_, j) in zip(class_labels, indices)] + ) # shape of (batch_size, num_queries) + target_classes = torch.full( + (batch_size, num_queries), fill_value=self.num_labels, dtype=torch.int64, device=pred_logits.device + ) + target_classes[idx] = target_classes_o + # Permute target_classes (batch_size, num_queries, num_labels) -> (batch_size, num_labels, num_queries) + pred_logits_transposed = pred_logits.transpose(1, 2) + loss_ce = criterion(pred_logits_transposed, target_classes) + losses = {"loss_cross_entropy": loss_ce} + return losses + + def loss_masks( + self, + masks_queries_logits: torch.Tensor, + mask_labels: list[torch.Tensor], + indices: tuple[np.array], + num_masks: int, + ) -> dict[str, torch.Tensor]: + """Compute the losses related to the masks using sigmoid_cross_entropy_loss and dice loss. + + Args: + masks_queries_logits (`torch.Tensor`): + A tensor of shape `(batch_size, num_queries, height, width)`. + mask_labels (`torch.Tensor`): + List of mask labels of shape `(labels, height, width)`. + indices (`tuple[np.array])`: + The indices computed by the Hungarian matcher. + num_masks (`int)`: + The number of masks, used for normalization. + + Returns: + losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing two keys: + - **loss_mask** -- The loss computed using sigmoid cross entropy loss on the predicted and ground truth. + masks. + - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth, + masks. + """ + src_idx = self._get_predictions_permutation_indices(indices) + tgt_idx = self._get_targets_permutation_indices(indices) + # shape (batch_size * num_queries, height, width) + pred_masks = masks_queries_logits[src_idx] + # shape (batch_size, num_queries, height, width) + # pad all and stack the targets to the num_labels dimension + target_masks, _ = self._pad_images_to_max_in_batch(mask_labels) + target_masks = target_masks[tgt_idx] + + # No need to upsample predictions as we are using normalized coordinates + pred_masks = pred_masks[:, None] + target_masks = target_masks[:, None] + + # Sample point coordinates + with torch.no_grad(): + point_coordinates = self.sample_points_using_uncertainty( + pred_masks, + lambda logits: self.calculate_uncertainty(logits), + self.num_points, + self.oversample_ratio, + self.importance_sample_ratio, + ) + + point_labels = sample_point(target_masks, point_coordinates, align_corners=False).squeeze(1) + + point_logits = sample_point(pred_masks, point_coordinates, align_corners=False).squeeze(1) + + losses = { + "loss_mask": sigmoid_cross_entropy_loss(point_logits, point_labels, num_masks), + "loss_dice": dice_loss(point_logits, point_labels, num_masks), + } + + del pred_masks + del target_masks + return losses + + def _get_predictions_permutation_indices(self, indices): + # Permute predictions following indices + batch_indices = torch.cat([torch.full_like(src, i) for i, (src, _) in enumerate(indices)]) + predictions_indices = torch.cat([src for (src, _) in indices]) + return batch_indices, predictions_indices + + def _get_targets_permutation_indices(self, indices): + # Permute labels following indices + batch_indices = torch.cat([torch.full_like(tgt, i) for i, (_, tgt) in enumerate(indices)]) + target_indices = torch.cat([tgt for (_, tgt) in indices]) + return batch_indices, target_indices + + def calculate_uncertainty(self, logits: torch.Tensor) -> torch.Tensor: + """ + In EomtDinov3 paper, uncertainty is estimated as L1 distance between 0.0 and the logit prediction in 'logits' + for the foreground class in `classes`. + + Args: + logits (`torch.Tensor`): + 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: + the number of foreground classes. The values are logits. + + Returns: + scores (`torch.Tensor`): A tensor of shape (R, 1, ...) that contains uncertainty scores with the most + uncertain locations having the highest uncertainty score. + """ + uncertainty_scores = -(torch.abs(logits)) + return uncertainty_scores + + def sample_points_using_uncertainty( + self, + logits: torch.Tensor, + uncertainty_function, + num_points: int, + oversample_ratio: int, + importance_sample_ratio: float, + ) -> torch.Tensor: + """ + This function is meant for sampling points in [0, 1] * [0, 1] coordinate space based on their uncertainty. The + uncertainty is calculated for each point using the passed `uncertainty function` that takes points logit + prediction as input. + + Args: + logits (`float`): + Logit predictions for P points. + uncertainty_function: + A function that takes logit predictions for P points and returns their uncertainties. + num_points (`int`): + The number of points P to sample. + oversample_ratio (`int`): + Oversampling parameter. + importance_sample_ratio (`float`): + Ratio of points that are sampled via importance sampling. + + Returns: + point_coordinates (`torch.Tensor`): + Coordinates for P sampled points. + """ + + num_boxes = logits.shape[0] + num_points_sampled = int(num_points * oversample_ratio) + + # Get random point coordinates + point_coordinates = torch.rand(num_boxes, num_points_sampled, 2, device=logits.device) + # Get sampled prediction value for the point coordinates + point_logits = sample_point(logits, point_coordinates, align_corners=False) + # Calculate the uncertainties based on the sampled prediction values of the points + point_uncertainties = uncertainty_function(point_logits) + + num_uncertain_points = int(importance_sample_ratio * num_points) + num_random_points = num_points - num_uncertain_points + + idx = torch.topk(point_uncertainties[:, 0, :], k=num_uncertain_points, dim=1)[1] + shift = num_points_sampled * torch.arange(num_boxes, dtype=torch.long, device=logits.device) + idx += shift[:, None] + point_coordinates = point_coordinates.view(-1, 2)[idx.view(-1), :].view(num_boxes, num_uncertain_points, 2) + + if num_random_points > 0: + point_coordinates = torch.cat( + [point_coordinates, torch.rand(num_boxes, num_random_points, 2, device=logits.device)], + dim=1, + ) + return point_coordinates + + def forward( + self, + masks_queries_logits: torch.Tensor, + class_queries_logits: torch.Tensor, + mask_labels: list[torch.Tensor], + class_labels: list[torch.Tensor], + auxiliary_predictions: dict[str, torch.Tensor] | None = None, + ) -> dict[str, torch.Tensor]: + """ + This performs the loss computation. + + Args: + masks_queries_logits (`torch.Tensor`): + A tensor of shape `(batch_size, num_queries, height, width)`. + class_queries_logits (`torch.Tensor`): + A tensor of shape `(batch_size, num_queries, num_labels)`. + mask_labels (`torch.Tensor`): + List of mask labels of shape `(labels, height, width)`. + class_labels (`list[torch.Tensor]`): + List of class labels of shape `(labels)`. + auxiliary_predictions (`dict[str, torch.Tensor]`, *optional*): + if `use_auxiliary_loss` was set to `true` in [`EomtDinov3Config`], then it contains the logits from + the inner layers of the EomtDinov3MaskedAttentionDecoder. + + Returns: + losses (`dict[str, Tensor]`): A dict of `torch.Tensor` containing three keys: + - **loss_cross_entropy** -- The loss computed using cross entropy on the predicted and ground truth labels. + - **loss_mask** -- The loss computed using sigmoid cross_entropy loss on the predicted and ground truth + masks. + - **loss_dice** -- The loss computed using dice loss on the predicted on the predicted and ground truth + masks. + if `use_auxiliary_loss` was set to `true` in [`EomtDinov3Config`], the dictionary contains additional + losses for each auxiliary predictions. + """ + + # retrieve the matching between the outputs of the last layer and the labels + indices = self.matcher(masks_queries_logits, class_queries_logits, mask_labels, class_labels) + # compute the average number of target masks for normalization purposes + num_masks = self.get_num_masks(class_labels, device=class_labels[0].device) + # get all the losses + losses: dict[str, Tensor] = { + **self.loss_masks(masks_queries_logits, mask_labels, indices, num_masks), + **self.loss_labels(class_queries_logits, class_labels, indices), + } + # in case of auxiliary losses, we repeat this process with the output of each intermediate layer. + if auxiliary_predictions is not None: + for idx, aux_outputs in enumerate(auxiliary_predictions): + masks_queries_logits = aux_outputs["masks_queries_logits"] + class_queries_logits = aux_outputs["class_queries_logits"] + loss_dict = self.forward(masks_queries_logits, class_queries_logits, mask_labels, class_labels) + loss_dict = {f"{key}_{idx}": value for key, value in loss_dict.items()} + losses.update(loss_dict) + + return losses + + def get_num_masks(self, class_labels: torch.Tensor, device: torch.device) -> torch.Tensor: + """ + Computes the average number of target masks across the batch, for normalization purposes. + """ + num_masks = sum(len(classes) for classes in class_labels) + num_masks = torch.as_tensor(num_masks, dtype=torch.float, device=device) + world_size = 1 + if is_accelerate_available(): + if PartialState._shared_state != {}: + num_masks = reduce(num_masks) + world_size = PartialState().num_processes + + num_masks = torch.clamp(num_masks / world_size, min=1) + return num_masks + + +@auto_docstring( + custom_intro=""" + Class for outputs of [`EomtDinov3ForUniversalSegmentationOutput`]. + + This output can be directly passed to [`~EomtDinov3ImageProcessor.post_process_semantic_segmentation`] or + [`~EomtDinov3ImageProcessor.post_process_instance_segmentation`] or + [`~EomtDinov3ImageProcessor.post_process_panoptic_segmentation`] to compute final segmentation maps. Please, see + [`~EomtDinov3ImageProcessor] for details regarding usage. + """ +) +@dataclass +class EomtDinov3ForUniversalSegmentationOutput(ModelOutput): + r""" + loss (`torch.Tensor`, *optional*): + The computed loss, returned when labels are present. + class_queries_logits (`torch.FloatTensor`): + A tensor of shape `(batch_size, num_queries, num_labels + 1)` representing the proposed classes for each + query. Note the `+ 1` is needed because we incorporate the null class. + masks_queries_logits (`torch.FloatTensor`): + A tensor of shape `(batch_size, num_queries, height, width)` representing the proposed masks for each + query. + last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`): + Last hidden states (final feature map) of the last layer. + hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`): + Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each stage) of + shape `(batch_size, sequence_length, hidden_size)`. Hidden-states all layers of the model. + attentions (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`): + Tuple of `tuple(torch.FloatTensor)` (one for each layer) of shape `(batch_size, num_heads, sequence_length, + sequence_length)`. Self and Cross Attentions weights from transformer decoder. + patch_offsets (`list[torch.Tensor]`, *optional*): + list of tuples indicating the image index and start and end positions of patches for semantic segmentation. + """ + + loss: torch.FloatTensor | None = None + class_queries_logits: torch.FloatTensor | None = None + masks_queries_logits: torch.FloatTensor | None = None + last_hidden_state: torch.FloatTensor | None = None + hidden_states: tuple[torch.FloatTensor] | None = None + attentions: tuple[torch.FloatTensor] | None = None + patch_offsets: list[torch.Tensor] | None = None + + +@auto_docstring +class EomtDinov3PreTrainedModel(PreTrainedModel): + """ + An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained + models. + """ + + config: EomtDinov3Config + base_model_prefix = "eomt_dinov3" + main_input_name = "pixel_values" + input_modalities = ("image",) + supports_gradient_checkpointing = False + _no_split_modules = ["EomtDinov3Layer"] + _supports_sdpa = True + _can_record_outputs = { + "hidden_states": EomtDinov3Layer, + "attentions": EomtDinov3Attention, + } + config_class = EomtDinov3Config + + @torch.no_grad() + def _init_weights(self, module: nn.Module) -> None: + super()._init_weights(module) + std = self.config.initializer_range + if isinstance(module, EomtDinov3LayerScale): + if hasattr(module, "lambda1"): + init.constant_(module.lambda1, self.config.layerscale_value) + elif isinstance(module, EomtDinov3Embeddings): + init.trunc_normal_(module.cls_token, mean=0.0, std=std) + init.zeros_(module.register_tokens) + elif isinstance(module, EomtDinov3Loss): + empty_weight = torch.ones(module.num_labels + 1) + empty_weight[-1] = module.eos_coef + init.copy_(module.empty_weight, empty_weight) + elif isinstance(module, EomtDinov3ForUniversalSegmentation): + init.ones_(module.attn_mask_probs) + + +class EomtDinov3LayerNorm2d(nn.LayerNorm): + def __init__(self, num_channels, eps=1e-6, affine=True): + super().__init__(num_channels, eps=eps, elementwise_affine=affine) + + def forward(self, hidden_state: torch.Tensor) -> torch.Tensor: + hidden_state = hidden_state.permute(0, 2, 3, 1) + hidden_state = F.layer_norm(hidden_state, self.normalized_shape, self.weight, self.bias, self.eps) + hidden_state = hidden_state.permute(0, 3, 1, 2) + return hidden_state + + +class EomtDinov3ScaleLayer(nn.Module): + def __init__(self, config: EomtDinov3Config): + super().__init__() + hidden_size = config.hidden_size + self.conv1 = nn.ConvTranspose2d(hidden_size, hidden_size, kernel_size=2, stride=2) + self.activation = ACT2FN[config.hidden_act] + self.conv2 = nn.Conv2d( + hidden_size, + hidden_size, + kernel_size=3, + padding=1, + groups=hidden_size, + bias=False, + ) + + self.layernorm2d = EomtDinov3LayerNorm2d(hidden_size) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = self.conv1(hidden_states) + hidden_states = self.activation(hidden_states) + hidden_states = self.conv2(hidden_states) + hidden_states = self.layernorm2d(hidden_states) + return hidden_states + + +class EomtDinov3ScaleBlock(nn.Module): + def __init__(self, config: EomtDinov3Config): + super().__init__() + self.num_blocks = config.num_upscale_blocks + self.block = nn.ModuleList([EomtDinov3ScaleLayer(config) for _ in range(self.num_blocks)]) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + for block in self.block: + hidden_states = block(hidden_states) + return hidden_states + + +class EomtDinov3MaskHead(nn.Module): + def __init__(self, config: EomtDinov3Config): + super().__init__() + + hidden_size = config.hidden_size + self.fc1 = nn.Linear(hidden_size, hidden_size) + self.fc2 = nn.Linear(hidden_size, hidden_size) + self.fc3 = nn.Linear(hidden_size, hidden_size) + self.activation = ACT2FN[config.hidden_act] + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = self.activation(self.fc1(hidden_states)) + hidden_states = self.activation(self.fc2(hidden_states)) + hidden_states = self.fc3(hidden_states) + return hidden_states + + +@auto_docstring( + custom_intro=""" + The EoMT-DINOv3 model with head on top for instance/semantic/panoptic segmentation. + """, +) +class EomtDinov3ForUniversalSegmentation(EomtDinov3PreTrainedModel): + main_input_name = "pixel_values" + + def __init__(self, config: EomtDinov3Config): + super().__init__(config) + self.config = config + self.num_hidden_layers = config.num_hidden_layers + self.embeddings = EomtDinov3Embeddings(config) + self.layernorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps) + + self.query = nn.Embedding(config.num_queries, config.hidden_size) + self.layers = nn.ModuleList([EomtDinov3Layer(config) for _ in range(config.num_hidden_layers)]) + + self.upscale_block = EomtDinov3ScaleBlock(config) + self.mask_head = EomtDinov3MaskHead(config) + + self.class_predictor = nn.Linear(config.hidden_size, config.num_labels + 1) + + self.grid_size = (config.image_size // config.patch_size, config.image_size // config.patch_size) + self.weight_dict: dict[str, float] = { + "loss_cross_entropy": config.class_weight, + "loss_mask": config.mask_weight, + "loss_dice": config.dice_weight, + } + + self.criterion = EomtDinov3Loss(config=config, weight_dict=self.weight_dict) + + self.register_buffer("attn_mask_probs", torch.ones(config.num_blocks)) + + self.num_prefix_tokens = 1 + config.num_register_tokens + self.dropout = nn.Dropout(config.hidden_dropout_prob) + self.embeddings.register_parameter("mask_token", None) + + self.rope_embeddings = EomtDinov3RotaryEmbedding(config) + + self.post_init() + + def get_loss_dict( + self, + masks_queries_logits: Tensor, + class_queries_logits: Tensor, + mask_labels: Tensor, + class_labels: Tensor, + auxiliary_predictions: dict[str, Tensor], + ) -> dict[str, Tensor]: + loss_dict: dict[str, Tensor] = self.criterion( + masks_queries_logits=masks_queries_logits, + class_queries_logits=class_queries_logits, + mask_labels=mask_labels, + class_labels=class_labels, + auxiliary_predictions=auxiliary_predictions, + ) + + # weight each loss by `self.weight_dict[]` including auxiliary losses + for key, weight in self.weight_dict.items(): + for loss_key, loss in loss_dict.items(): + if key in loss_key: + loss *= weight + + return loss_dict + + def get_loss(self, loss_dict: dict[str, Tensor]) -> Tensor: + return sum(loss_dict.values()) + + @merge_with_config_defaults + @capture_outputs + @auto_docstring + def forward( + self, + pixel_values: Tensor, + mask_labels: list[Tensor] | None = None, + class_labels: list[Tensor] | None = None, + patch_offsets: list[Tensor] | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> EomtDinov3ForUniversalSegmentationOutput: + r""" + mask_labels (`list[torch.Tensor]`, *optional*): + list of mask labels of shape `(num_labels, height, width)` to be fed to a model + class_labels (`list[torch.LongTensor]`, *optional*): + list of target class labels of shape `(num_labels, height, width)` to be fed to a model. They identify the + labels of `mask_labels`, e.g. the label of `mask_labels[i][j]` if `class_labels[i][j]`. + patch_offsets (`list[torch.Tensor]`, *optional*): + list of tuples indicating the image index and start and end positions of patches for semantic segmentation. + """ + masks_queries_logits_per_layer, class_queries_logits_per_layer = (), () + + hidden_states = self.dropout(self.embeddings(pixel_values)) + position_embeddings = self.rope_embeddings(pixel_values.to(hidden_states.dtype)) + attention_mask = None + + for idx, layer_module in enumerate(self.layers): + if idx == self.num_hidden_layers - self.config.num_blocks: + query = self.query.weight[None, :, :].expand(hidden_states.shape[0], -1, -1).to(hidden_states.device) + hidden_states = torch.cat((query, hidden_states), dim=1) + + if idx >= self.num_hidden_layers - self.config.num_blocks and ( + self.training or self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks] > 0 + ): + norm_hidden_states = self.layernorm(hidden_states) + masks_queries_logits, class_queries_logits = self.predict(norm_hidden_states) + + masks_queries_logits_per_layer += (masks_queries_logits,) + class_queries_logits_per_layer += (class_queries_logits,) + + attention_mask = torch.ones( + hidden_states.shape[0], + hidden_states.shape[1], + hidden_states.shape[1], + device=hidden_states.device, + dtype=torch.bool, + ) + + interpolated_logits = F.interpolate(masks_queries_logits, size=self.grid_size, mode="bilinear") + interpolated_logits = interpolated_logits.view( + interpolated_logits.size(0), interpolated_logits.size(1), -1 + ) + + num_query_tokens = self.config.num_queries + encoder_start_tokens = num_query_tokens + self.num_prefix_tokens + + # Set attention mask for queries to focus on encoder tokens based on interpolated logits + attention_mask[:, :num_query_tokens, encoder_start_tokens:] = interpolated_logits > 0 + + # Disable attention mask for random query tokens. + attention_mask = self._disable_attention_mask( + attention_mask, + prob=self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks], + num_query_tokens=num_query_tokens, + encoder_start_tokens=encoder_start_tokens, + device=attention_mask.device, + ) + + # Expand attention mask to 4d mask. + attention_mask = attention_mask[:, None, ...].expand(-1, self.config.num_attention_heads, -1, -1) + dtype_min = torch.finfo(hidden_states.dtype).min + attention_mask = attention_mask.to(hidden_states.dtype).masked_fill(~attention_mask, dtype_min) + + hidden_states = layer_module( + hidden_states, + attention_mask=attention_mask, + position_embeddings=position_embeddings, + ) + + sequence_output = self.layernorm(hidden_states) + + masks_queries_logits, class_queries_logits = self.predict(sequence_output) + masks_queries_logits_per_layer += (masks_queries_logits,) + class_queries_logits_per_layer += (class_queries_logits,) + + loss = None + if mask_labels is not None and class_labels is not None: + loss = 0.0 + for masks_queries_logits, class_queries_logits in zip( + masks_queries_logits_per_layer, class_queries_logits_per_layer + ): + loss_dict = self.get_loss_dict( + masks_queries_logits=masks_queries_logits, + class_queries_logits=class_queries_logits, + mask_labels=mask_labels, + class_labels=class_labels, + auxiliary_predictions=None, + ) + loss += self.get_loss(loss_dict) + + return EomtDinov3ForUniversalSegmentationOutput( + loss=loss, + masks_queries_logits=masks_queries_logits, + class_queries_logits=class_queries_logits, + last_hidden_state=sequence_output, + patch_offsets=patch_offsets, + ) + + def get_input_embeddings(self): + return self.embeddings.patch_embeddings + + def predict(self, logits: torch.Tensor): + query_tokens = logits[:, : self.config.num_queries, :] + class_logits = self.class_predictor(query_tokens) + + prefix_tokens = logits[:, self.config.num_queries + self.embeddings.num_prefix_tokens :, :] + prefix_tokens = prefix_tokens.transpose(1, 2) + + prefix_tokens = prefix_tokens.reshape(prefix_tokens.shape[0], -1, *self.grid_size) + + query_tokens = self.mask_head(query_tokens) + prefix_tokens = self.upscale_block(prefix_tokens) + + mask_logits = torch.einsum("bqc, bchw -> bqhw", query_tokens, prefix_tokens) + + return mask_logits, class_logits + + @staticmethod + def _disable_attention_mask(attn_mask, prob, num_query_tokens, encoder_start_tokens, device): + if prob < 1: + # Generate random queries to disable based on the probs + random_queries = torch.rand(attn_mask.shape[0], num_query_tokens, device=device) > prob + + # Disable attention to the query tokens, considering the prefix tokens + attn_mask[:, :num_query_tokens, encoder_start_tokens:][random_queries] = 1 + + return attn_mask + + +__all__ = ["EomtDinov3PreTrainedModel", "EomtDinov3ForUniversalSegmentation"] diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modular_eomt_dinov3.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modular_eomt_dinov3.py new file mode 100644 index 0000000000000000000000000000000000000000..2b6920a3b75838da9c9abc247d5f38782a365ed4 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/eomt_dinov3/modular_eomt_dinov3.py @@ -0,0 +1,364 @@ +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""PyTorch EoMT model backed by DINOv3.""" + +from collections.abc import Callable +from typing import Optional + +import torch +import torch.nn.functional as F +from huggingface_hub.dataclasses import strict +from torch import Tensor, nn + +from ... import initialization as init +from ...modeling_rope_utils import RopeParameters +from ...modeling_utils import PreTrainedModel +from ...processing_utils import Unpack +from ...utils import ( + TransformersKwargs, + auto_docstring, +) +from ...utils.generic import merge_with_config_defaults +from ...utils.output_capturing import capture_outputs +from ..dinov3_vit.modeling_dinov3_vit import ( + DINOv3ViTAttention, + DINOv3ViTEmbeddings, + DINOv3ViTLayer, + DINOv3ViTLayerScale, + DINOv3ViTRopePositionEmbedding, +) +from ..eomt.configuration_eomt import EomtConfig +from ..eomt.modeling_eomt import ( + EomtForUniversalSegmentation, + EomtForUniversalSegmentationOutput, + EomtLoss, + EomtPreTrainedModel, +) + + +@auto_docstring(checkpoint="tue-mps/coco_panoptic_eomt_large_640_dinov3") +@strict +class EomtDinov3Config(EomtConfig): + r""" + layerscale_value (`float`, *optional*, defaults to 1.0): + Initial value for the LayerScale parameter. + num_upscale_blocks (`int`, *optional*, defaults to 2): + Number of upsampling blocks used in the decoder or segmentation head. + num_blocks (`int`, *optional*, defaults to 4): + Number of feature blocks or stages in the architecture. + no_object_weight (`float`, *optional*, defaults to 0.1): + Loss weight for the "no object" class in panoptic/instance segmentation. + train_num_points (`int`, *optional*, defaults to 12544): + Number of points to sample for mask loss computation during training. + oversample_ratio (`float`, *optional*, defaults to 3.0): + Oversampling ratio used in point sampling for mask training. + importance_sample_ratio (`float`, *optional*, defaults to 0.75): + Ratio of points to sample based on importance during training. + num_queries (`int`, *optional*, defaults to 200): + Number of object queries in the Transformer. + num_register_tokens (`int`, *optional*, defaults to 4): + Number of learnable register tokens added to the transformer input. + query_bias (`bool`, *optional*, defaults to `True`): + Whether to use bias in query projection. + key_bias (`bool`, *optional*, defaults to `False`): + Whether to use bias in key projection. + value_bias (`bool`, *optional*, defaults to `True`): + Whether to use bias in value projection. + proj_bias (`bool`, *optional*, defaults to `True`): + Whether to use bias in output projection. + use_gated_mlp (`bool`, *optional*, defaults to `False`): + Whether to use gated MLP layers. + pos_embed_shift (`float`, *optional*): + Shift value for position embeddings. + pos_embed_jitter (`float`, *optional*): + Jitter value for position embeddings. + pos_embed_rescale (`float`, *optional*, defaults to 2.0): + Rescale value for position embeddings. + """ + + model_type = "eomt_dinov3" + default_theta = 100.0 + + hidden_size: int = 1024 + num_hidden_layers: int = 24 + num_attention_heads: int = 16 + intermediate_size: int = 4096 + hidden_act: str = "gelu" + hidden_dropout_prob: float | int = 0.0 + initializer_range: float = 0.02 + layer_norm_eps: float = 1e-6 + image_size: int | list[int] | tuple[int, int] = 640 + patch_size: int | list[int] | tuple[int, int] = 16 + num_channels: int = 3 + layerscale_value: float = 1.0 + drop_path_rate: float | int = 0.0 + num_upscale_blocks: int = 2 + attention_dropout: float | int = 0.0 + num_blocks: int = 4 + no_object_weight: float = 0.1 + class_weight: float = 2.0 + mask_weight: float = 5.0 + dice_weight: float = 5.0 + train_num_points: int = 12544 + oversample_ratio: float = 3.0 + importance_sample_ratio: float = 0.75 + num_queries: int = 200 + num_register_tokens: int = 4 + rope_parameters: RopeParameters | dict | None = None + query_bias: bool = True + key_bias: bool = False + value_bias: bool = True + proj_bias: bool = True + mlp_bias: bool = True + use_gated_mlp: bool = False + pos_embed_shift: float | None = None + pos_embed_jitter: float | None = None + pos_embed_rescale: float | None = 2.0 + + mlp_ratio = AttributeError() + use_swiglu_ffn = AttributeError() + + +class EomtDinov3Attention(DINOv3ViTAttention): + pass + + +class EomtDinov3Embeddings(DINOv3ViTEmbeddings): + def __init__(self, config: EomtDinov3Config): + super().__init__(config) + self.num_prefix_tokens = 1 + config.num_register_tokens + + +class EomtDinov3Layer(DINOv3ViTLayer): + pass + + +class EomtDinov3LayerScale(DINOv3ViTLayerScale): + pass + + +class EomtDinov3RotaryEmbedding(DINOv3ViTRopePositionEmbedding): + inv_freq: Tensor + + def __init__(self, config: EomtDinov3Config, device=None): + nn.Module.__init__(self) + self.config = config + + self.rope_type = self.config.rope_parameters["rope_type"] + rope_init_fn: Callable = self.compute_default_rope_parameters + if self.rope_type != "default": + raise ValueError("`EomtDinov3` only supports `default` RoPE! Please check your `rope_type`") + inv_freq, self.attention_scaling = rope_init_fn(self.config, device) + + self.register_buffer("inv_freq", inv_freq, persistent=False) + self.register_buffer("original_inv_freq", inv_freq.clone(), persistent=False) + + @staticmethod + def compute_default_rope_parameters( + config: EomtDinov3Config | None = None, + device: Optional["torch.device"] = None, + seq_len: int | None = None, + ) -> torch.Tensor: + """ + Computes the inverse frequencies according to the original RoPE implementation + Args: + config ([`~transformers.PreTrainedConfig`]): + The model configuration. + device (`torch.device`): + The device to use for initialization of the inverse frequencies. + seq_len (`int`, *optional*): + The current sequence length. Unused for this type of RoPE. + Returns: + Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the + post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE). + """ + base = config.rope_parameters["rope_theta"] + head_dim = config.hidden_size // config.num_attention_heads + + attention_factor = 1.0 # Unused in this type of RoPE + + # Compute the inverse frequencies + inv_freq = 1 / base ** torch.arange(0, 1, 4 / head_dim, dtype=torch.float32, device=device) + return inv_freq, attention_factor + + +class EomtDinov3Loss(EomtLoss): + pass + + +class EomtDinov3ForUniversalSegmentationOutput(EomtForUniversalSegmentationOutput): + pass + + +class EomtDinov3PreTrainedModel(EomtPreTrainedModel): + config_class = EomtDinov3Config + base_model_prefix = "eomt_dinov3" + _no_split_modules = ["EomtDinov3Layer"] + _can_record_outputs = { + "hidden_states": EomtDinov3Layer, + "attentions": EomtDinov3Attention, + } + + def _init_weights(self, module: nn.Module) -> None: + PreTrainedModel._init_weights(module) + std = self.config.initializer_range + if isinstance(module, EomtDinov3LayerScale): + if hasattr(module, "lambda1"): + init.constant_(module.lambda1, self.config.layerscale_value) + elif isinstance(module, EomtDinov3Embeddings): + init.trunc_normal_(module.cls_token, mean=0.0, std=std) + init.zeros_(module.register_tokens) + elif isinstance(module, EomtDinov3Loss): + empty_weight = torch.ones(module.num_labels + 1) + empty_weight[-1] = module.eos_coef + init.copy_(module.empty_weight, empty_weight) + elif isinstance(module, EomtDinov3ForUniversalSegmentation): + init.ones_(module.attn_mask_probs) + + +@auto_docstring( + custom_intro=""" + The EoMT-DINOv3 model with head on top for instance/semantic/panoptic segmentation. + """, +) +class EomtDinov3ForUniversalSegmentation(EomtDinov3PreTrainedModel, EomtForUniversalSegmentation): + def __init__(self, config: EomtDinov3Config): + super().__init__(config) + + self.num_prefix_tokens = 1 + config.num_register_tokens + self.dropout = nn.Dropout(config.hidden_dropout_prob) + self.embeddings = EomtDinov3Embeddings(config) + self.embeddings.register_parameter("mask_token", None) + + self.rope_embeddings = EomtDinov3RotaryEmbedding(config) + self.layers = nn.ModuleList([EomtDinov3Layer(config) for _ in range(config.num_hidden_layers)]) + + self.post_init() + + # We redefine forward here because EoMT-DINOv3 uses DINOv3 backbone components (RoPE embeddings, layers) + # which require different integration than the base EoMT model that uses a separate encoder. + @merge_with_config_defaults + @capture_outputs + @auto_docstring + def forward( + self, + pixel_values: Tensor, + mask_labels: list[Tensor] | None = None, + class_labels: list[Tensor] | None = None, + patch_offsets: list[Tensor] | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> EomtDinov3ForUniversalSegmentationOutput: + r""" + mask_labels (`list[torch.Tensor]`, *optional*): + list of mask labels of shape `(num_labels, height, width)` to be fed to a model + class_labels (`list[torch.LongTensor]`, *optional*): + list of target class labels of shape `(num_labels, height, width)` to be fed to a model. They identify the + labels of `mask_labels`, e.g. the label of `mask_labels[i][j]` if `class_labels[i][j]`. + patch_offsets (`list[torch.Tensor]`, *optional*): + list of tuples indicating the image index and start and end positions of patches for semantic segmentation. + """ + masks_queries_logits_per_layer, class_queries_logits_per_layer = (), () + + hidden_states = self.dropout(self.embeddings(pixel_values)) + position_embeddings = self.rope_embeddings(pixel_values.to(hidden_states.dtype)) + attention_mask = None + + for idx, layer_module in enumerate(self.layers): + if idx == self.num_hidden_layers - self.config.num_blocks: + query = self.query.weight[None, :, :].expand(hidden_states.shape[0], -1, -1).to(hidden_states.device) + hidden_states = torch.cat((query, hidden_states), dim=1) + + if idx >= self.num_hidden_layers - self.config.num_blocks and ( + self.training or self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks] > 0 + ): + norm_hidden_states = self.layernorm(hidden_states) + masks_queries_logits, class_queries_logits = self.predict(norm_hidden_states) + + masks_queries_logits_per_layer += (masks_queries_logits,) + class_queries_logits_per_layer += (class_queries_logits,) + + attention_mask = torch.ones( + hidden_states.shape[0], + hidden_states.shape[1], + hidden_states.shape[1], + device=hidden_states.device, + dtype=torch.bool, + ) + + interpolated_logits = F.interpolate(masks_queries_logits, size=self.grid_size, mode="bilinear") + interpolated_logits = interpolated_logits.view( + interpolated_logits.size(0), interpolated_logits.size(1), -1 + ) + + num_query_tokens = self.config.num_queries + encoder_start_tokens = num_query_tokens + self.num_prefix_tokens + + # Set attention mask for queries to focus on encoder tokens based on interpolated logits + attention_mask[:, :num_query_tokens, encoder_start_tokens:] = interpolated_logits > 0 + + # Disable attention mask for random query tokens. + attention_mask = self._disable_attention_mask( + attention_mask, + prob=self.attn_mask_probs[idx - self.num_hidden_layers + self.config.num_blocks], + num_query_tokens=num_query_tokens, + encoder_start_tokens=encoder_start_tokens, + device=attention_mask.device, + ) + + # Expand attention mask to 4d mask. + attention_mask = attention_mask[:, None, ...].expand(-1, self.config.num_attention_heads, -1, -1) + dtype_min = torch.finfo(hidden_states.dtype).min + attention_mask = attention_mask.to(hidden_states.dtype).masked_fill(~attention_mask, dtype_min) + + hidden_states = layer_module( + hidden_states, + attention_mask=attention_mask, + position_embeddings=position_embeddings, + ) + + sequence_output = self.layernorm(hidden_states) + + masks_queries_logits, class_queries_logits = self.predict(sequence_output) + masks_queries_logits_per_layer += (masks_queries_logits,) + class_queries_logits_per_layer += (class_queries_logits,) + + loss = None + if mask_labels is not None and class_labels is not None: + loss = 0.0 + for masks_queries_logits, class_queries_logits in zip( + masks_queries_logits_per_layer, class_queries_logits_per_layer + ): + loss_dict = self.get_loss_dict( + masks_queries_logits=masks_queries_logits, + class_queries_logits=class_queries_logits, + mask_labels=mask_labels, + class_labels=class_labels, + auxiliary_predictions=None, + ) + loss += self.get_loss(loss_dict) + + return EomtDinov3ForUniversalSegmentationOutput( + loss=loss, + masks_queries_logits=masks_queries_logits, + class_queries_logits=class_queries_logits, + last_hidden_state=sequence_output, + patch_offsets=patch_offsets, + ) + + +__all__ = [ + "EomtDinov3Config", + "EomtDinov3PreTrainedModel", + "EomtDinov3ForUniversalSegmentation", +] diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/__init__.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..285c1970308a47827806fca349d130703f40a2c8 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/__init__.py @@ -0,0 +1,27 @@ +# Copyright 2024 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import TYPE_CHECKING + +from ...utils import _LazyModule +from ...utils.import_utils import define_import_structure + + +if TYPE_CHECKING: + from .configuration_patchtsmixer import * + from .modeling_patchtsmixer import * +else: + import sys + + _file = globals()["__file__"] + sys.modules[__name__] = _LazyModule(__name__, _file, define_import_structure(_file), module_spec=__spec__) diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/configuration_patchtsmixer.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/configuration_patchtsmixer.py new file mode 100644 index 0000000000000000000000000000000000000000..92dd5a3d3247ccafac2a598e94c14ebc60413764 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/configuration_patchtsmixer.py @@ -0,0 +1,166 @@ +# Copyright 2023 IBM and HuggingFace Inc. team. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""PatchTSMixer model configuration""" + +from huggingface_hub.dataclasses import strict + +from ...configuration_utils import PreTrainedConfig +from ...utils import auto_docstring + + +@auto_docstring(checkpoint="ibm/patchtsmixer-etth1-pretrain") +@strict +class PatchTSMixerConfig(PreTrainedConfig): + r""" + context_length (`int`, *optional*, defaults to 32): + The context/history length for the input sequence. + patch_length (`int`, *optional*, defaults to 8): + The patch length for the input sequence. + patch_stride (`int`, *optional*, defaults to 8): + Determines the overlap between two consecutive patches. Set it to patch_length (or greater), if we want + non-overlapping patches. + num_parallel_samples (`int`, *optional*, defaults to 100): + The number of samples to generate in parallel for probabilistic forecast. + expansion_factor (`int`, *optional*, defaults to 2): + Expansion factor to use inside MLP. Recommended range is 2-5. Larger value indicates more complex model. + mode (`str`, *optional*, defaults to `"common_channel"`): + Mixer Mode. Determines how to process the channels. Allowed values: "common_channel", "mix_channel". In + "common_channel" mode, we follow Channel-independent modelling with no explicit channel-mixing. Channel + mixing happens in an implicit manner via shared weights across channels. (preferred first approach) In + "mix_channel" mode, we follow explicit channel-mixing in addition to patch and feature mixer. (preferred + approach when channel correlations are very important to model) + gated_attn (`bool`, *optional*, defaults to `True`): + Enable Gated Attention. + norm_mlp (`str`, *optional*, defaults to `"LayerNorm"`): + Normalization layer (BatchNorm or LayerNorm). + self_attn (`bool`, *optional*, defaults to `False`): + Enable Tiny self attention across patches. This can be enabled when the output of Vanilla PatchTSMixer with + gated attention is not satisfactory. Enabling this leads to explicit pair-wise attention and modelling + across patches. + self_attn_heads (`int`, *optional*, defaults to 1): + Number of self-attention heads. Works only when `self_attn` is set to `True`. + use_positional_encoding (`bool`, *optional*, defaults to `False`): + Enable the use of positional embedding for the tiny self-attention layers. Works only when `self_attn` is + set to `True`. + positional_encoding_type (`str`, *optional*, defaults to `"sincos"`): + Positional encodings. Options `"random"` and `"sincos"` are supported. Works only when + `use_positional_encoding` is set to `True` + scaling (`string` or `bool`, *optional*, defaults to `"std"`): + Whether to scale the input targets via "mean" scaler, "std" scaler or no scaler if `None`. If `True`, the + scaler is set to "mean". + loss (`string`, *optional*, defaults to `"mse"`): + The loss function for the model corresponding to the `distribution_output` head. For parametric + distributions it is the negative log likelihood ("nll") and for point estimates it is the mean squared + error "mse". + norm_eps (`float`, *optional*, defaults to 1e-05): + A value added to the denominator for numerical stability of normalization. + mask_type (`str`, *optional*, defaults to `"random"`): + Type of masking to use for Masked Pretraining mode. Allowed values are "random", "forecast". In Random + masking, points are masked randomly. In Forecast masking, points are masked towards the end. + random_mask_ratio (`float`, *optional*, defaults to 0.5): + Masking ratio to use when `mask_type` is `random`. Higher value indicates more masking. + num_forecast_mask_patches (`int` or `list`, *optional*, defaults to `[2]`): + Number of patches to be masked at the end of each batch sample. If it is an integer, all the samples in the + batch will have the same number of masked patches. If it is a list, samples in the batch will be randomly + masked by numbers defined in the list. This argument is only used for forecast pretraining. + mask_value (`float`, *optional*, defaults to `0.0`): + Mask value to use. + masked_loss (`bool`, *optional*, defaults to `True`): + Whether to compute pretraining loss only at the masked portions, or on the entire output. + channel_consistent_masking (`bool`, *optional*, defaults to `True`): + When true, masking will be same across all channels of a timeseries. Otherwise, masking positions will vary + across channels. + unmasked_channel_indices (`list`, *optional*): + Channels that are not masked during pretraining. + head_dropout (`float`, *optional*, defaults to 0.2): + The dropout probability the `PatchTSMixer` head. + distribution_output (`string`, *optional*, defaults to `"student_t"`): + The distribution emission head for the model when loss is "nll". Could be either "student_t", "normal" or + "negative_binomial". + prediction_length (`int`, *optional*, defaults to 16): + Number of time steps to forecast for a forecasting task. Also known as the Forecast Horizon. + prediction_channel_indices (`list`, *optional*): + List of channel indices to forecast. If None, forecast all channels. Target data is expected to have all + channels and we explicitly filter the channels in prediction and target before loss computation. + num_targets (`int`, *optional*, defaults to 3): + Number of targets (dimensionality of the regressed variable) for a regression task. + output_range (`list`, *optional*): + Output range to restrict for the regression task. Defaults to None. + head_aggregation (`str`, *optional*, defaults to `"max_pool"`): + Aggregation mode to enable for classification or regression task. Allowed values are `None`, "use_last", + "max_pool", "avg_pool". + + Example: + + ```python + >>> from transformers import PatchTSMixerConfig, PatchTSMixerModel + + >>> # Initializing a default PatchTSMixer configuration + >>> configuration = PatchTSMixerConfig() + + >>> # Randomly initializing a model (with random weights) from the configuration + >>> model = PatchTSMixerModel(configuration) + + >>> # Accessing the model configuration + >>> configuration = model.config + ```""" + + model_type = "patchtsmixer" + attribute_map = { + "hidden_size": "d_model", + "num_hidden_layers": "num_layers", + } + + context_length: int = 32 + patch_length: int = 8 + num_input_channels: int = 1 + patch_stride: int = 8 + num_parallel_samples: int = 100 + d_model: int = 8 + expansion_factor: int = 2 + num_layers: int = 3 + dropout: float | int = 0.2 + mode: str = "common_channel" + gated_attn: bool = True + norm_mlp: str = "LayerNorm" + self_attn: bool = False + self_attn_heads: int = 1 + use_positional_encoding: bool = False + positional_encoding_type: str = "sincos" + scaling: str | bool | None = "std" + loss: str = "mse" + init_std: float = 0.02 + norm_eps: float = 1e-5 + mask_type: str = "random" + random_mask_ratio: float = 0.5 + num_forecast_mask_patches: list[int] | tuple[int, ...] | int | None = (2,) + mask_value: int = 0 + masked_loss: bool = True + channel_consistent_masking: bool = True + unmasked_channel_indices: list[int] | None = None + head_dropout: float | int = 0.2 + distribution_output: str = "student_t" + prediction_length: int = 16 + prediction_channel_indices: list | None = None + num_targets: int = 3 + output_range: list | None = None + head_aggregation: str | None = "max_pool" + + def __post_init__(self, **kwargs): + self.num_patches = (max(self.context_length, self.patch_length) - self.patch_length) // self.patch_stride + 1 + self.patch_last = True + super().__post_init__(**kwargs) + + +__all__ = ["PatchTSMixerConfig"] diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/modeling_patchtsmixer.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/modeling_patchtsmixer.py new file mode 100644 index 0000000000000000000000000000000000000000..3146261e3204916fbeca067be90fcbd25de7d044 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/patchtsmixer/modeling_patchtsmixer.py @@ -0,0 +1,2121 @@ +# Copyright 2023 IBM and HuggingFace Inc. team. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""PyTorch PatchTSMixer model.""" + +import math +from collections.abc import Callable +from dataclasses import dataclass + +import torch +import torch.nn as nn + +from transformers.modeling_utils import PreTrainedModel +from transformers.utils import ModelOutput + +from ... import initialization as init +from ...modeling_flash_attention_utils import FlashAttentionKwargs +from ...modeling_utils import ALL_ATTENTION_FUNCTIONS +from ...processing_utils import Unpack +from ...time_series_utils import NegativeBinomialOutput, NormalOutput, StudentTOutput +from ...utils import TransformersKwargs, auto_docstring, logging +from .configuration_patchtsmixer import PatchTSMixerConfig + + +logger = logging.get_logger(__name__) + + +class PatchTSMixerGatedAttention(nn.Module): + """ + Module that applies gated attention to input data. + + Args: + in_size (`int`): The input size. + out_size (`int`): The output size. + """ + + def __init__(self, in_size: int, out_size: int): + super().__init__() + self.attn_layer = nn.Linear(in_size, out_size) + self.attn_softmax = nn.Softmax(dim=-1) + + def forward(self, inputs): + attn_weight = self.attn_softmax(self.attn_layer(inputs)) + inputs = inputs * attn_weight + return inputs + + +# Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTBatchNorm with PatchTST->PatchTSMixer +class PatchTSMixerBatchNorm(nn.Module): + """ + Compute batch normalization over the sequence length (time) dimension. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + self.batchnorm = nn.BatchNorm1d(config.d_model, eps=config.norm_eps) + + def forward(self, inputs: torch.Tensor): + """ + Parameters: + inputs (`torch.Tensor` of shape `(batch_size, sequence_length, d_model)`): + input for Batch norm calculation + Returns: + `torch.Tensor` of shape `(batch_size, sequence_length, d_model)` + """ + output = inputs.transpose(1, 2) # output: (batch_size, d_model, sequence_length) + output = self.batchnorm(output) + return output.transpose(1, 2) + + +class PatchTSMixerPositionalEncoding(nn.Module): + """ + Class for positional encoding + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + # positional encoding: [num_patches x d_model] + if config.use_positional_encoding: + self.position_enc = self._init_pe(config) + else: + self.position_enc = nn.Parameter(torch.zeros(config.num_patches, config.d_model)) + + @staticmethod + def _init_pe(config: PatchTSMixerConfig) -> nn.Parameter: + # Positional encoding + if config.positional_encoding_type == "random": + position_enc = nn.Parameter(torch.randn(config.num_patches, config.d_model), requires_grad=True) + elif config.positional_encoding_type == "sincos": + position_enc = torch.zeros(config.num_patches, config.d_model) + position = torch.arange(0, config.num_patches).unsqueeze(1) + div_term = torch.exp(torch.arange(0, config.d_model, 2) * -(math.log(10000.0) / config.d_model)) + position_enc[:, 0::2] = torch.sin(position * div_term) + position_enc[:, 1::2] = torch.cos(position * div_term) + position_enc = position_enc - position_enc.mean() + position_enc = position_enc / (position_enc.std() * 10) + position_enc = nn.Parameter(position_enc, requires_grad=False) + else: + raise ValueError( + f"{config.positional_encoding_type} is not a valid positional encoder. Available types are 'random' and 'sincos'." + ) + return position_enc + + def forward(self, patch_input: torch.Tensor): + # hidden_state: [bs x num_channels x num_patches x d_model] + hidden_state = patch_input + self.position_enc + return hidden_state + + +class PatchTSMixerNormLayer(nn.Module): + """Normalization block + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + self.norm_mlp = config.norm_mlp + + if "batch" in config.norm_mlp.lower(): + self.norm = PatchTSMixerBatchNorm(config) + else: + self.norm = nn.LayerNorm(config.d_model, eps=config.norm_eps) + + def forward(self, inputs: torch.Tensor): + """ + Args: + inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`): + Input to the normalization layer. + Returns: + `torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))` + """ + if "batch" in self.norm_mlp.lower(): + # reshape the data + inputs_reshaped = torch.reshape( + inputs, + ( + inputs.shape[0] * inputs.shape[1], + inputs.shape[2], + inputs.shape[3], + ), + ) # inputs_reshaped: [batch_size*num_channels, num_patches, d_model] + + # inputs_reshaped: [batch_size*num_channels, num_patches, d_model] + inputs_reshaped = self.norm(inputs_reshaped) + + # put back data to the original shape + inputs = torch.reshape(inputs_reshaped, inputs.shape) + + else: + inputs = self.norm(inputs) + + return inputs + + +class PatchTSMixerMLP(nn.Module): + def __init__(self, in_features, out_features, config): + super().__init__() + num_hidden = in_features * config.expansion_factor + self.fc1 = nn.Linear(in_features, num_hidden) + self.dropout1 = nn.Dropout(config.dropout) + self.fc2 = nn.Linear(num_hidden, out_features) + self.dropout2 = nn.Dropout(config.dropout) + + def forward(self, inputs: torch.Tensor): + """ + Args: + inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`): + Input to the MLP layer. + Returns: + `torch.Tensor` of the same shape as `inputs` + """ + inputs = self.dropout1(nn.functional.gelu(self.fc1(inputs))) + inputs = self.fc2(inputs) + inputs = self.dropout2(inputs) + return inputs + + +class PatchTSMixerChannelFeatureMixerBlock(nn.Module): + """This module mixes the features in the channel dimension. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + self.norm = PatchTSMixerNormLayer(config) + self.gated_attn = config.gated_attn + self.mlp = PatchTSMixerMLP( + in_features=config.num_input_channels, + out_features=config.num_input_channels, + config=config, + ) + + if config.gated_attn: + self.gating_block = PatchTSMixerGatedAttention( + in_size=config.num_input_channels, out_size=config.num_input_channels + ) + + def forward(self, inputs: torch.Tensor): + """ + Args: + inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`): + input to the MLP layer + Returns: + `torch.Tensor` of the same shape as `inputs` + """ + residual = inputs + inputs = self.norm(inputs) + + inputs = inputs.permute(0, 3, 2, 1) + + if self.gated_attn: + inputs = self.gating_block(inputs) + + inputs = self.mlp(inputs) + + inputs = inputs.permute(0, 3, 2, 1) + + out = inputs + residual + return out + + +# Copied from transformers.models.bert.modeling_bert.eager_attention_forward +def eager_attention_forward( + module: nn.Module, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attention_mask: torch.Tensor | None, + scaling: float | None = None, + dropout: float = 0.0, + **kwargs: Unpack[TransformersKwargs], +): + if scaling is None: + scaling = query.size(-1) ** -0.5 + + # Take the dot product between "query" and "key" to get the raw attention scores. + attn_weights = torch.matmul(query, key.transpose(2, 3)) * scaling + + if attention_mask is not None: + attn_weights = attn_weights + attention_mask + + attn_weights = nn.functional.softmax(attn_weights, dim=-1) + attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) + + attn_output = torch.matmul(attn_weights, value) + attn_output = attn_output.transpose(1, 2).contiguous() + + return attn_output, attn_weights + + +# Copied from transformers.models.wav2vec2.modeling_wav2vec2.Wav2Vec2Attention with Wav2Vec2->PatchTSMixer +class PatchTSMixerAttention(nn.Module): + """Multi-headed attention from 'Attention Is All You Need' paper""" + + def __init__( + self, + embed_dim: int, + num_heads: int, + dropout: float = 0.0, + is_decoder: bool = False, + bias: bool = True, + is_causal: bool = False, + config: PatchTSMixerConfig | None = None, + ): + super().__init__() + self.embed_dim = embed_dim + self.num_heads = num_heads + self.dropout = dropout + self.head_dim = embed_dim // num_heads + self.config = config + + if (self.head_dim * num_heads) != self.embed_dim: + raise ValueError( + f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim}" + f" and `num_heads`: {num_heads})." + ) + self.scaling = self.head_dim**-0.5 + self.is_decoder = is_decoder + self.is_causal = is_causal + + self.k_proj = nn.Linear(embed_dim, embed_dim, bias=bias) + self.v_proj = nn.Linear(embed_dim, embed_dim, bias=bias) + self.q_proj = nn.Linear(embed_dim, embed_dim, bias=bias) + self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias) + + def forward( + self, + hidden_states: torch.Tensor, + key_value_states: torch.Tensor | None = None, + attention_mask: torch.Tensor | None = None, + output_attentions: bool | None = False, + # TODO: we need a refactor so that the different attention modules can get their specific kwargs + # ATM, we have mixed things encoder, decoder, and encoder-decoder attn + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + """Input shape: Batch x Time x Channel""" + + # if key_value_states are provided this layer is used as a cross-attention layer + # for the decoder + is_cross_attention = key_value_states is not None + + # determine input shapes + input_shape = hidden_states.shape[:-1] + + hidden_shape = (*input_shape, -1, self.head_dim) + + # get query proj + query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2) + + current_states = key_value_states if is_cross_attention else hidden_states + kv_shape = (*current_states.shape[:-1], -1, self.head_dim) + key_states = self.k_proj(current_states).view(kv_shape).transpose(1, 2) + value_states = self.v_proj(current_states).view(kv_shape).transpose(1, 2) + + attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface( + self.config._attn_implementation, eager_attention_forward + ) + + attn_output, attn_weights = attention_interface( + self, + query_states, + key_states, + value_states, + attention_mask, + dropout=0.0 if not self.training else self.dropout, + scaling=self.scaling, + output_attentions=output_attentions, + **kwargs, + ) + + attn_output = attn_output.reshape(*input_shape, -1).contiguous() + attn_output = self.out_proj(attn_output) + + return attn_output, attn_weights, None + + +class PatchMixerBlock(nn.Module): + """This module mixes the patch dimension. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + self.norm = PatchTSMixerNormLayer(config) + + self.self_attn = config.self_attn + self.gated_attn = config.gated_attn + + self.mlp = PatchTSMixerMLP( + in_features=config.num_patches, + out_features=config.num_patches, + config=config, + ) + + if config.gated_attn: + self.gating_block = PatchTSMixerGatedAttention(in_size=config.num_patches, out_size=config.num_patches) + + if config.self_attn: + self.self_attn_layer = PatchTSMixerAttention( + embed_dim=config.d_model, + num_heads=config.self_attn_heads, + dropout=config.dropout, + config=config, + ) + self.norm_attn = PatchTSMixerNormLayer(config) + + def forward(self, hidden_state): + """ + Args: + hidden_state (`torch.Tensor`): Input tensor. + + Returns: + `torch.Tensor`: Transformed tensor. + """ + residual = hidden_state + + hidden_state = self.norm(hidden_state) + + if self.self_attn: + batch_size, n_vars, num_patches, d_model = hidden_state.shape + hidden_state_reshaped = hidden_state.reshape(batch_size * n_vars, num_patches, d_model) + + x_attn, _, _ = self.self_attn_layer(hidden_state_reshaped, output_attentions=False) + x_attn = x_attn.reshape(batch_size, n_vars, num_patches, d_model) + + # Transpose so that num_patches is the last dimension + hidden_state = hidden_state.transpose(2, 3) + hidden_state = self.mlp(hidden_state) + + if self.gated_attn: + hidden_state = self.gating_block(hidden_state) + + # Transpose back + hidden_state = hidden_state.transpose(2, 3) + + if self.self_attn: + hidden_state = self.norm_attn(hidden_state + x_attn) + + out = hidden_state + residual + return out + + +class FeatureMixerBlock(nn.Module): + """This module mixes the hidden feature dimension. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + self.norm = PatchTSMixerNormLayer(config) + + self.gated_attn = config.gated_attn + + self.mlp = PatchTSMixerMLP( + in_features=config.d_model, + out_features=config.d_model, + config=config, + ) + + if config.gated_attn: + self.gating_block = PatchTSMixerGatedAttention(in_size=config.d_model, out_size=config.d_model) + + def forward(self, hidden: torch.Tensor): + """ + Args: + hidden (`torch.Tensor` of shape `(batch_size, num_patches, d_model)`): + Input tensor to the layer. + + Returns: + `torch.Tensor`: Transformed tensor. + """ + residual = hidden + hidden = self.norm(hidden) + hidden = self.mlp(hidden) + + if self.gated_attn: + hidden = self.gating_block(hidden) + + out = hidden + residual + return out + + +class PatchTSMixerLayer(nn.Module): + """ + The `PatchTSMixer` layer that does all three kinds of mixing. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + self.patch_mixer = PatchMixerBlock(config=config) + self.feature_mixer = FeatureMixerBlock(config=config) + + self.mode = config.mode + + if config.mode == "mix_channel": + self.channel_feature_mixer = PatchTSMixerChannelFeatureMixerBlock(config=config) + + def forward(self, hidden: torch.Tensor): + """ + Args: + hidden (`torch.Tensor` of shape `(batch_size, num_patches, d_model)`): + Input tensor to the layer. + + Returns: + `torch.Tensor`: Transformed tensor. + """ + if self.mode == "mix_channel": + hidden = self.channel_feature_mixer(hidden) + + hidden = self.patch_mixer(hidden) + hidden = self.feature_mixer(hidden) # hidden: (batch_size x num_patches x d_model) + return hidden + + +class PatchTSMixerBlock(nn.Module): + """The main computing framework of the `PatchTSMixer` model. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + num_layers = config.num_layers + + self.mixers = nn.ModuleList([PatchTSMixerLayer(config=config) for _ in range(num_layers)]) + + def forward(self, hidden_state, output_hidden_states: bool = False): + """ + Args: + hidden_state (`torch.Tensor`): The input tensor. + output_hidden_states (`bool`, *optional*, defaults to False.): + Whether to output the hidden states as well. + + Returns: + `torch.Tensor`: The embedding. `list`: List of all hidden states if `output_hidden_states` is set to + `True`. + """ + all_hidden_states = [] + + embedding = hidden_state + + for mod in self.mixers: + embedding = mod(embedding) + if output_hidden_states: + all_hidden_states.append(embedding) + + if output_hidden_states: + return embedding, all_hidden_states + else: + return embedding, None + + +class PatchTSMixerForPredictionHead(nn.Module): + """Prediction Head for Forecasting + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig, distribution_output=None): + super().__init__() + + self.prediction_channel_indices = config.prediction_channel_indices + + if self.prediction_channel_indices is not None: + self.prediction_channel_indices.sort() + + self.dropout_layer = nn.Dropout(config.head_dropout) + if distribution_output is None: + self.base_forecast_block = nn.Linear((config.num_patches * config.d_model), config.prediction_length) + else: + self.base_forecast_block = distribution_output.get_parameter_projection( + config.num_patches * config.d_model + ) + + self.flatten = nn.Flatten(start_dim=-2) + + def forward(self, hidden_features): + """ + + Args: + hidden_features (`torch.Tensor` of shape `(batch_size, num_patch, d_model)` in `flatten` mode + or `(batch_size, n_vars, num_patch, d_model)` in `common_channel`/`mix_channel` mode.): Input hidden + features. + + Returns: + `torch.Tensor` of shape `(batch_size, prediction_length, nvars)`. + + """ + + hidden_features = self.flatten(hidden_features) # [batch_size x n_vars x num_patch * d_model] + hidden_features = self.dropout_layer(hidden_features) # [batch_size x n_vars x num_patch * d_model] + forecast = self.base_forecast_block(hidden_features) # [batch_size x n_vars x prediction_length] + if isinstance(forecast, tuple): + forecast = tuple(z.transpose(-1, -2) for z in forecast) + else: + forecast = forecast.transpose(-1, -2) # [batch_size x prediction_length x n_vars] + + if self.prediction_channel_indices is not None: + if isinstance(forecast, tuple): + forecast = tuple(z[..., self.prediction_channel_indices] for z in forecast) + else: + forecast = forecast[..., self.prediction_channel_indices] # [batch_size x prediction_length x n_vars] + + return forecast + + +class PatchTSMixerLinearHead(nn.Module): + """Linear head for Classification and Regression. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig, distribution_output=None): + super().__init__() + + self.head_aggregation = config.head_aggregation + self.output_range = config.output_range + + if config.head_aggregation is None: + mul_factor = config.num_patches + else: + mul_factor = 1 + self.distribution_output = distribution_output + if distribution_output is None: + self.projection = nn.Linear( + config.d_model * config.num_input_channels * mul_factor, + config.num_targets, + ) + else: + self.projection = distribution_output.get_parameter_projection( + config.d_model * config.num_input_channels * mul_factor + ) + + if config.head_aggregation is None: + self.flatten = nn.Flatten(start_dim=-3) + else: + self.flatten = nn.Flatten(start_dim=-2) + + self.dropout = nn.Dropout(config.head_dropout) + + def forward(self, hidden_features): + """ + Args: + hidden_features (`torch.Tensor` of shape `(batch_size x num_patch x d_model)` in `flatten` mode + or `(batch_size x n_vars x num_patch x d_model)` in `common_channel`/`mix_channel` mode.): Input hidden + features. + + Returns: + `torch.Tensor` of shape `(batch_size x num_targets)`. + """ + + # batch_size x d_model x num_patch or batch_size x n_vars x d_model x num_patch + hidden_features = hidden_features.transpose(-1, -2) + if self.head_aggregation == "use_last": + # batch_size x d_model (flatten) or # batch_size x n_vars x d_model (common_channel) + hidden_features = hidden_features[..., -1] + elif self.head_aggregation == "max_pool": + # batch_size x n_vars x d_model or batch_size x d_model + hidden_features = hidden_features.max(dim=-1).values + elif self.head_aggregation == "avg_pool": + # batch_size x n_vars x d_model or batch_size x d_model + hidden_features = hidden_features.mean(dim=-1) + + if self.flatten: + hidden_features = self.flatten(hidden_features) + hidden_features = self.dropout(hidden_features) + hidden_features = self.projection(hidden_features) # batch_size x num_targets + + if (self.distribution_output is None) and (self.output_range is not None): + hidden_features = ( + torch.sigmoid(hidden_features) * (self.output_range[1] - self.output_range[0]) + self.output_range[0] + ) + return hidden_features + + +@auto_docstring +class PatchTSMixerPreTrainedModel(PreTrainedModel): + # Weight initialization + config: PatchTSMixerConfig + base_model_prefix = "model" + main_input_name = "past_values" + input_modalities = ("time",) + supports_gradient_checkpointing = False + + @torch.no_grad() + def _init_weights(self, module): + """Initialize weights""" + if isinstance(module, PatchTSMixerPositionalEncoding): + # initialize positional encoding + if self.config.positional_encoding_type == "random": + init.normal_(module.position_enc, mean=0.0, std=0.1) + elif isinstance(module, (nn.LayerNorm, nn.BatchNorm1d)): + init.zeros_(module.bias) + init.ones_(module.weight) + if getattr(module, "running_mean", None) is not None: + init.zeros_(module.running_mean) + init.ones_(module.running_var) + init.zeros_(module.num_batches_tracked) + elif isinstance(module, PatchTSMixerBatchNorm): + init.zeros_(module.batchnorm.bias) + init.ones_(module.batchnorm.weight) + elif isinstance(module, nn.Linear): + init.normal_(module.weight, mean=0.0, std=self.config.init_std) + if module.bias is not None: + init.zeros_(module.bias) + + +class PatchTSMixerPretrainHead(nn.Module): + """Pretraining head. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + self.dropout_layer = nn.Dropout(config.head_dropout) + self.base_pt_block = nn.Linear(config.d_model, config.patch_length) + + def forward(self, hidden_features): + """ + Args: + hidden_features (`torch.Tensor` of shape `(batch_size x num_patch x d_model)` in `flatten` mode + or `(batch_size x n_vars x num_patch x d_model)` in `common_channel`/`mix_channel` mode.): Input hidden + features. + + Returns: + `torch.Tensor` of shape `(batch_size x n_vars x num_patch x patch_length)`. + """ + + hidden_features = self.dropout_layer(hidden_features) + forecast = self.base_pt_block(hidden_features) # [batch_size x n_vars x num_patch x patch_length] + return forecast + + +# Copied from transformers.models.patchtst.modeling_patchtst.random_masking +def random_masking( + inputs: torch.Tensor, + mask_ratio: float, + unmasked_channel_indices: list | None = None, + channel_consistent_masking: bool = False, + mask_value: int = 0, +): + """random_masking: Mask the input considering the control variables. + + Args: + inputs (`torch.Tensor` of shape `(batch_size, num_channels, sequence_length, num_features)`): + The input tensor to mask. + mask_ratio (`float`): + Masking ratio applied to mask the input data during random pretraining. It is the number between 0 and 1. + unmasked_channel_indices (list, *optional*): + Indices of channels that will not be masked. + channel_consistent_masking (bool, *optional*, defaults to `False`): + When true, masking will be same across all channels of a timeseries. Otherwise, masking positions will vary + across channels. + mask_value (int, *optional*, defaults to 0): + Define the value of masked patches for pretraining. + + Returns: + `tuple(torch.Tensor)`: inputs_mask, masked input, same shape as input Tensor and mask tensor of shape [bs x c x + n] + """ + if mask_ratio < 0 or mask_ratio >= 1: + raise ValueError(f"Mask ratio {mask_ratio} has to be between 0 and 1.") + + batch_size, num_channels, sequence_length, num_features = inputs.shape + device = inputs.device + + len_keep = int(sequence_length * (1 - mask_ratio)) + + if channel_consistent_masking: + noise = torch.rand(batch_size, 1, sequence_length, device=device) # noise in [0, 1], bs x 1 x L + noise = noise.repeat(1, num_channels, 1) # bs x num_channels x time + else: + # noise in [0, 1], bs x num_channels x L + noise = torch.rand(batch_size, num_channels, sequence_length, device=device) + + # mask: [bs x num_channels x num_patch] + mask = torch.ones(batch_size, num_channels, sequence_length, device=device) + mask[:, :, :len_keep] = 0 + + # sort noise for each sample + ids_shuffle = torch.argsort(noise, dim=-1) # ascend: small is keep, large is remove + ids_restore = torch.argsort(ids_shuffle, dim=-1) # ids_restore: [bs x num_channels x L] + + mask = torch.gather(mask, dim=-1, index=ids_restore) + mask = mask.unsqueeze(-1).repeat(1, 1, 1, num_features) # mask: [bs x num_channels x num_patches x patch_length] + if unmasked_channel_indices is not None: + mask[:, unmasked_channel_indices, :, :] = 0 + + inputs_mask = inputs.masked_fill(mask.bool(), mask_value) + return inputs_mask, mask[..., 0] + + +# Copied from transformers.models.patchtst.modeling_patchtst.forecast_masking +def forecast_masking( + inputs: torch.Tensor, + num_forecast_mask_patches: list | int, + unmasked_channel_indices: list | None = None, + mask_value: int = 0, +): + """Forecast masking that masks the last K patches where K is from the num_forecast_mask_patches. + If num_forecast_mask_patches is a list, samples in the batch will be randomly masked by numbers defined in the list. + + Parameters: + inputs (`torch.Tensor`): + Input of shape `(bs, num_channels, num_patch, patch_length)` + num_forecast_mask_patches (`list`): + Number of patches to be masked at the end of each batch sample. e.g. 4 or [3, 5]. + unmasked_channel_indices (`list`, *optional*): + Indices of channels that are not masked. + mask_value (`int`, *optional*, defaults to 0): + Values in the masked patches will be filled by `mask_value`. + + Returns: + `tuple(torch.Tensor)`: inputs_mask, masked input, same shape as inputs Tensor and Mask tensor of shape `(bs, + num_channels , num_patch)` or `(bs, tsg1, tsg2, num_channels, num_patch)` + """ + + if isinstance(num_forecast_mask_patches, int): + num_forecast_mask_patches = [num_forecast_mask_patches] + forecast_mask_ratios = [1 for _ in num_forecast_mask_patches] + + batch_size, num_channels, sequence_length, num_features = inputs.shape + mask = torch.zeros(batch_size, num_channels, sequence_length, device=inputs.device) + + t_list = [] + total_length = 0 + total_ratio = sum(forecast_mask_ratios) + + for patch_length, ratio in zip(num_forecast_mask_patches, forecast_mask_ratios): + if patch_length <= 0 or patch_length >= sequence_length: + raise ValueError( + f"num_forecast_mask_patches {patch_length} should be greater than 0 and less than total patches." + ) + temp_len = int(batch_size * ratio / total_ratio) + t_list.append([patch_length, ratio, temp_len]) + total_length += temp_len + + t_list = sorted(t_list, key=lambda x: x[2]) + + if total_length < batch_size: + t_list[0][2] = t_list[0][2] + (batch_size - total_length) + elif total_length > batch_size: + t_list[-1][2] = t_list[-1][2] + (total_length - batch_size) + + batch1 = 0 + for patch_len, _, temp_len in t_list: + batch2 = batch1 + temp_len + mask[batch1:batch2, :, -patch_len:] = 1 + batch1 = batch2 + + perm = torch.randperm(mask.shape[0]) + mask = mask[perm] + + mask = mask.unsqueeze(-1).repeat(1, 1, 1, num_features) # mask: [bs x num_channels x num_patch x patch_len] + if unmasked_channel_indices is not None: + mask[:, unmasked_channel_indices, :, :] = 0 + + inputs_mask = inputs.masked_fill(mask.bool(), mask_value) + return inputs_mask, mask[..., 0] + + +# Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTPatchify with PatchTST->PatchTSMixer +class PatchTSMixerPatchify(nn.Module): + """ + A class to patchify the time series sequence into different patches + + Returns: + `torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)` + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + + self.sequence_length = config.context_length + self.patch_length = config.patch_length + self.patch_stride = config.patch_stride + + if self.sequence_length <= self.patch_length: + raise ValueError( + f"Sequence length ({self.sequence_length}) has to be greater than the patch length ({self.patch_length})" + ) + + # get the number of patches + self.num_patches = (max(self.sequence_length, self.patch_length) - self.patch_length) // self.patch_stride + 1 + new_sequence_length = self.patch_length + self.patch_stride * (self.num_patches - 1) + self.sequence_start = self.sequence_length - new_sequence_length + + def forward(self, past_values: torch.Tensor): + """ + Parameters: + past_values (`torch.Tensor` of shape `(batch_size, sequence_length, num_channels)`, *required*): + Input for patchification + + Returns: + `torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)` + """ + sequence_length = past_values.shape[-2] + if sequence_length != self.sequence_length: + raise ValueError( + f"Input sequence length ({sequence_length}) doesn't match model configuration ({self.sequence_length})." + ) + # output: [bs x new_sequence_length x num_channels] + output = past_values[:, self.sequence_start :, :] + # output: [bs x num_patches x num_input_channels x patch_length] + output = output.unfold(dimension=-2, size=self.patch_length, step=self.patch_stride) + # output: [bs x num_input_channels x num_patches x patch_length] + output = output.transpose(-2, -3).contiguous() + return output + + +# Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTMasking with PatchTST->PatchTSMixer +class PatchTSMixerMasking(nn.Module): + """ + Class to perform random or forecast masking. + + Parameters: + config (`PatchTSMixerConfig`): model config + Returns: + x_mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`) + Masked patched input + mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches)`) + Bool tensor indicating True on masked points + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + self.random_mask_ratio = config.random_mask_ratio + self.channel_consistent_masking = config.channel_consistent_masking + self.mask_type = config.mask_type + self.num_forecast_mask_patches = config.num_forecast_mask_patches + self.unmasked_channel_indices = config.unmasked_channel_indices + self.mask_value = config.mask_value + if self.unmasked_channel_indices is not None: + self.unmasked_channel_indices = sorted(self.unmasked_channel_indices) + + def forward(self, patch_input: torch.Tensor): + """ + Parameters: + patch_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`, *required*): + Patch input + + Return: + masked_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`) + Masked patched input + mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches)`) + Bool tensor indicating True on masked points + + """ + if self.mask_type == "random": + masked_input, mask = random_masking( + inputs=patch_input, + mask_ratio=self.random_mask_ratio, + unmasked_channel_indices=self.unmasked_channel_indices, + channel_consistent_masking=self.channel_consistent_masking, + mask_value=self.mask_value, + ) + elif self.mask_type == "forecast": + masked_input, mask = forecast_masking( + inputs=patch_input, + num_forecast_mask_patches=self.num_forecast_mask_patches, + unmasked_channel_indices=self.unmasked_channel_indices, + mask_value=self.mask_value, + ) + else: + raise ValueError(f"Invalid mask type {self.mask_type}.") + + # mask: [bs x num_input_channels x num_patch] + mask = mask.bool() + return masked_input, mask + + +# Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTStdScaler with PatchTST->PatchTSMixer +class PatchTSMixerStdScaler(nn.Module): + """ + Standardize features by calculating the mean and scaling along the first dimension, and then normalizes it by + subtracting from the mean and dividing by the standard deviation. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + self.dim = config.scaling_dim if hasattr(config, "scaling_dim") else 1 + self.keepdim = config.keepdim if hasattr(config, "keepdim") else True + self.minimum_scale = config.minimum_scale if hasattr(config, "minimum_scale") else 1e-5 + + def forward( + self, data: torch.Tensor, observed_indicator: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Parameters: + data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`): + input for Batch norm calculation + observed_indicator (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`): + Calculating the scale on the observed indicator. + Returns: + tuple of `torch.Tensor` of shapes + (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`, + `(batch_size, 1, num_input_channels)`) + """ + denominator = observed_indicator.sum(self.dim, keepdim=self.keepdim) + denominator = denominator.clamp_min(1.0) + loc = (data * observed_indicator).sum(self.dim, keepdim=self.keepdim) / denominator + + variance = (((data - loc) * observed_indicator) ** 2).sum(self.dim, keepdim=self.keepdim) / denominator + scale = torch.sqrt(variance + self.minimum_scale) + return (data - loc) / scale, loc, scale + + +# Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTMeanScaler with PatchTST->PatchTSMixer +class PatchTSMixerMeanScaler(nn.Module): + """ + Computes a scaling factor as the weighted average absolute value along the first dimension, and scales the data + accordingly. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + self.dim = config.scaling_dim if hasattr(config, "scaling_dim") else 1 + self.keepdim = config.keepdim if hasattr(config, "keepdim") else True + self.minimum_scale = config.minimum_scale if hasattr(config, "minimum_scale") else 1e-10 + self.default_scale = config.default_scale if hasattr(config, "default_scale") else None + + def forward( + self, data: torch.Tensor, observed_indicator: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Parameters: + data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`): + input for Batch norm calculation + observed_indicator (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`): + Calculating the scale on the observed indicator. + Returns: + tuple of `torch.Tensor` of shapes + (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`, + `(batch_size, 1, num_input_channels)`) + """ + ts_sum = (data * observed_indicator).abs().sum(self.dim, keepdim=True) + num_observed = observed_indicator.sum(self.dim, keepdim=True) + + scale = ts_sum / torch.clamp(num_observed, min=1) + + # If `default_scale` is provided, we use it, otherwise we use the scale + # of the batch. + if self.default_scale is None: + batch_sum = ts_sum.sum(dim=0) + batch_observations = torch.clamp(num_observed.sum(0), min=1) + default_scale = torch.squeeze(batch_sum / batch_observations) + else: + default_scale = self.default_scale * torch.ones_like(scale) + + # apply default scale where there are no observations + scale = torch.where(num_observed > 0, scale, default_scale) + + # ensure the scale is at least `self.minimum_scale` + scale = torch.clamp(scale, min=self.minimum_scale) + scaled_data = data / scale + + if not self.keepdim: + scale = scale.squeeze(dim=self.dim) + + return scaled_data, torch.zeros_like(scale), scale + + +# Copied from transformers.models.patchtst.modeling_patchtst.PatchTSTNOPScaler with PatchTST->PatchTSMixer +class PatchTSMixerNOPScaler(nn.Module): + """ + Assigns a scaling factor equal to 1 along the first dimension, and therefore applies no scaling to the input data. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__() + self.dim = config.scaling_dim if hasattr(config, "scaling_dim") else 1 + self.keepdim = config.keepdim if hasattr(config, "keepdim") else True + + def forward( + self, data: torch.Tensor, observed_indicator: torch.Tensor | None = None + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Parameters: + data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`): + input for Batch norm calculation + Returns: + tuple of `torch.Tensor` of shapes + (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`, + `(batch_size, 1, num_input_channels)`) + """ + scale = torch.ones_like(data, requires_grad=False).mean(dim=self.dim, keepdim=self.keepdim) + loc = torch.zeros_like(data, requires_grad=False).mean(dim=self.dim, keepdim=self.keepdim) + return data, loc, scale + + +@auto_docstring( + custom_intro=""" + Base class for `PatchTSMixerEncoderOutput`, with potential hidden states. + """ +) +@dataclass +class PatchTSMixerEncoderOutput(ModelOutput): + r""" + last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, d_model)`): + Hidden-state at the output of the last layer of the model. + hidden_states (`tuple(torch.FloatTensor)`, *optional*): + Hidden-states of the model at the output of each layer. + """ + + last_hidden_state: torch.FloatTensor | None = None + hidden_states: tuple[torch.FloatTensor] | None = None + + +class PatchTSMixerEncoder(PatchTSMixerPreTrainedModel): + """ + Encoder for PatchTSMixer which inputs patched time-series and outputs patched embeddings. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__(config) + + self.return_dict = config.return_dict + + self.patcher = nn.Linear(config.patch_length, config.d_model) + if config.use_positional_encoding: + self.positional_encoder = PatchTSMixerPositionalEncoding(config=config) + else: + self.positional_encoder = None + self.mlp_mixer_encoder = PatchTSMixerBlock(config=config) + + # Initialize weights and apply final processing + self.post_init() + + @auto_docstring + def forward( + self, + past_values: torch.Tensor, + output_hidden_states: bool | None = False, + return_dict: bool | None = None, + **kwargs, + ) -> tuple | PatchTSMixerEncoderOutput: + r""" + past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`): + Context values of the time series. For a pretraining task, this denotes the input time series to + predict the masked portion. For a forecasting task, this denotes the history/past time series values. + Similarly, for classification or regression tasks, it denotes the appropriate context values of the + time series. + + For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, + it is greater than 1. + + Returns: + `torch.FloatTensor` of shape `(batch_size, n_vars, num_patches, d_model)` + """ + + return_dict = return_dict if return_dict is not None else self.return_dict + + # flatten [bs x num_patch x d_model]. common_channel/mix_channel: [bs x n_vars x num_patch x d_model] + patches = self.patcher(past_values) + + # add positional encoder + if self.positional_encoder is not None: + patches = self.positional_encoder(patches) + + last_hidden_state, hidden_states = self.mlp_mixer_encoder(patches, output_hidden_states=output_hidden_states) + + if not return_dict: + return tuple( + v + for v in [ + last_hidden_state, + hidden_states, + ] + ) + + return PatchTSMixerEncoderOutput(last_hidden_state=last_hidden_state, hidden_states=hidden_states) + + +@auto_docstring( + custom_intro=""" + Base class for model's outputs, with potential hidden states. + """ +) +@dataclass +class PatchTSMixerModelOutput(ModelOutput): + r""" + last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, d_model)`): + Hidden-state at the output of the last layer of the model. + hidden_states (`tuple(torch.FloatTensor)`, *optional*): + Hidden-states of the model at the output of each layer. + patch_input (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, patch_length)`): + Patched input data to the model. + mask (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches)`, *optional*): + Bool Tensor indicating True in masked patches and False otherwise. + loc (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*): + Gives the mean of the context window per channel. Used for revin denorm outside the model, if revin + enabled. + scale (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*): + Gives the std dev of the context window per channel. Used for revin denorm outside the model, if revin + enabled. + """ + + last_hidden_state: torch.FloatTensor | None = None + hidden_states: tuple[torch.FloatTensor] | None = None + patch_input: torch.FloatTensor | None = None + mask: torch.FloatTensor | None = None + loc: torch.FloatTensor | None = None + scale: torch.FloatTensor | None = None + + +@auto_docstring( + custom_intro=""" + The PatchTSMixer Model for time-series forecasting. + """ +) +class PatchTSMixerModel(PatchTSMixerPreTrainedModel): + def __init__(self, config: PatchTSMixerConfig, mask_input: bool = False): + r""" + mask_input (bool, *optional*, defaults to `False`): + Whether to mask the input using the [`PatchTSMixerMasking`] module. + """ + super().__init__(config) + + self.return_dict = config.return_dict + self.encoder = PatchTSMixerEncoder(config) + self.patching = PatchTSMixerPatchify(config) + + if mask_input is True: + self.masking = PatchTSMixerMasking(config) + else: + self.masking = None + + if config.scaling == "mean": + self.scaler = PatchTSMixerMeanScaler(config) + elif config.scaling == "std" or config.scaling is True: + self.scaler = PatchTSMixerStdScaler(config) + else: + self.scaler = PatchTSMixerNOPScaler(config) + + # Initialize weights and apply final processing + self.post_init() + + @auto_docstring + def forward( + self, + past_values: torch.Tensor, + observed_mask: torch.Tensor | None = None, + output_hidden_states: bool | None = False, + return_dict: bool | None = None, + **kwargs, + ) -> PatchTSMixerModelOutput: + r""" + past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`): + Context values of the time series. For a pretraining task, this denotes the input time series to predict + the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly, + for classification or regression tasks, it denotes the appropriate context values of the time series. + + For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is + greater than 1. + observed_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*): + Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected + in `[0, 1]`: + - 1 for values that are **observed**, + - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros). + """ + return_dict = return_dict if return_dict is not None else self.return_dict + + mask = None + if observed_mask is None: + observed_mask = torch.ones_like(past_values) + scaled_past_values, loc, scale = self.scaler(past_values, observed_mask) + + patched_x = self.patching(scaled_past_values) # [batch_size x num_input_channels x num_patch x patch_length + + enc_input = patched_x + if self.masking is not None: + enc_input, mask = self.masking(patched_x) + # enc_input: [batch_size x num_input_channels x num_patch x patch_length] + # mask: [batch_size x num_input_channels x num_patch] + + encoder_output = self.encoder( + enc_input, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + + if isinstance(encoder_output, tuple): + encoder_output = PatchTSMixerEncoderOutput(*encoder_output) + + if not return_dict: + return tuple( + v + for v in [ + encoder_output.last_hidden_state, + encoder_output.hidden_states, + patched_x, + mask, + loc, + scale, + ] + ) + + return PatchTSMixerModelOutput( + last_hidden_state=encoder_output.last_hidden_state, + hidden_states=encoder_output.hidden_states, + patch_input=patched_x, + mask=mask, + loc=loc, + scale=scale, + ) + + +@auto_docstring( + custom_intro=""" + Output type of [`PatchTSMixerForPreTrainingOutput`]. + """ +) +@dataclass +class PatchTSMixerForPreTrainingOutput(ModelOutput): + r""" + loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`): + Total loss + prediction_outputs (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, patch_length)`): + Prediction output from the pretrain head. + last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`): + Backbone embeddings before passing through the head. + hidden_states (`tuple(torch.FloatTensor)`, *optional*): + Hidden-states of the model at the output of each layer. + """ + + loss: torch.FloatTensor | None = None + prediction_outputs: torch.FloatTensor | None = None + last_hidden_state: torch.FloatTensor | None = None + hidden_states: tuple[torch.FloatTensor] | None = None + + +@auto_docstring( + custom_intro=""" + `PatchTSMixer` for mask pretraining. + """ +) +class PatchTSMixerForPretraining(PatchTSMixerPreTrainedModel): + def __init__(self, config: PatchTSMixerConfig): + super().__init__(config) + self.model = PatchTSMixerModel(config, mask_input=True) + self.head = PatchTSMixerPretrainHead(config=config) + self.masked_loss = config.masked_loss + self.return_dict = config.return_dict + + # Initialize weights and apply final processing + self.post_init() + + @auto_docstring + def forward( + self, + past_values: torch.Tensor, + observed_mask: torch.Tensor | None = None, + output_hidden_states: bool | None = False, + return_loss: bool = True, + return_dict: bool | None = None, + **kwargs, + ) -> PatchTSMixerForPreTrainingOutput: + r""" + past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`): + Context values of the time series. For a pretraining task, this denotes the input time series to predict + the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly, + for classification or regression tasks, it denotes the appropriate context values of the time series. + + For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is + greater than 1. + observed_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*): + Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected + in `[0, 1]`: + - 1 for values that are **observed**, + - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros). + return_loss (`bool`, *optional*): + Whether to return the loss in the `forward` call. + """ + return_dict = return_dict if return_dict is not None else self.return_dict + + if self.masked_loss is True: + loss = torch.nn.MSELoss(reduction="none") + else: + loss = torch.nn.MSELoss(reduction="mean") + + # past_values: tensor [batch_size x context_length x num_input_channels] + model_output = self.model( + past_values, + observed_mask=observed_mask, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) # x.last_hidden_state: [batch_size x nvars x num_patch x d_model] + if isinstance(model_output, tuple): + model_output = PatchTSMixerModelOutput(*model_output) + + x_hat = self.head(model_output.last_hidden_state) # tensor [batch_size x nvars x num_patch x patch_length] + + if return_loss is True: + loss_val = loss(x_hat, model_output.patch_input) + else: + loss_val = None + + # calculate masked_loss + if self.masked_loss is True and loss_val is not None: + loss_val = (loss_val.mean(dim=-1) * model_output.mask).sum() / (model_output.mask.sum() + 1e-10) + + if not return_dict: + return tuple( + v + for v in [ + loss_val, + x_hat, + model_output.last_hidden_state, + model_output.hidden_states, + ] + ) + + return PatchTSMixerForPreTrainingOutput( + loss=loss_val, + prediction_outputs=x_hat, # tensor [batch_size x nvars x num_patch x patch_length] + last_hidden_state=model_output.last_hidden_state, # x: [batch_size x nvars x num_patch x d_model] + hidden_states=model_output.hidden_states, + ) + + +@auto_docstring( + custom_intro=""" + Output type of [`PatchTSMixerForPredictionOutput`]. + """ +) +@dataclass +class PatchTSMixerForPredictionOutput(ModelOutput): + r""" + loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`): + Total loss. + prediction_outputs (`torch.FloatTensor` of shape `(batch_size, prediction_length, num_input_channels)`): + Prediction output from the forecast head. + last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`): + Backbone embeddings before passing through the head. + hidden_states (`tuple(torch.FloatTensor)`, *optional*): + Hidden-states of the model at the output of each layer plus the optional initial embedding outputs. + loc (`torch.FloatTensor`, *optional* of shape `(batch_size, 1, num_input_channels)`): + Input mean + scale (`torch.FloatTensor`, *optional* of shape `(batch_size, 1, num_input_channels)`): + Input std dev + """ + + loss: torch.FloatTensor | None = None + prediction_outputs: torch.FloatTensor | None = None + last_hidden_state: torch.FloatTensor | None = None + hidden_states: tuple[torch.FloatTensor] | None = None + loc: torch.FloatTensor | None = None + scale: torch.FloatTensor | None = None + + +@auto_docstring( + custom_intro=""" + Base class for time series model's predictions outputs that contains the sampled values from the chosen + distribution. + """ +) +@dataclass +class SamplePatchTSMixerPredictionOutput(ModelOutput): + r""" + sequences (`torch.FloatTensor` of shape `(batch_size, num_samples, prediction_length, number_channels)`): + Sampled values from the chosen distribution. + """ + + sequences: torch.FloatTensor | None = None + + +@auto_docstring( + custom_intro=""" + Base class for time series model's predictions outputs that contains the sampled values from the chosen + distribution. + """ +) +@dataclass +class SamplePatchTSMixerRegressionOutput(ModelOutput): + r""" + sequences (`torch.FloatTensor` of shape `(batch_size, num_samples, prediction_length, number_channels)`): + Sampled values from the chosen distribution. + """ + + sequences: torch.FloatTensor | None = None + + +# Copied from transformers.models.time_series_transformer.modeling_time_series_transformer.nll +def nll(input: torch.distributions.Distribution, target: torch.Tensor) -> torch.Tensor: + """ + Computes the negative log likelihood loss from input distribution with respect to target. + """ + return -input.log_prob(target) + + +# Copied from transformers.models.time_series_transformer.modeling_time_series_transformer.weighted_average +def weighted_average(input_tensor: torch.Tensor, weights: torch.Tensor | None = None, dim=None) -> torch.Tensor: + """ + Computes the weighted average of a given tensor across a given `dim`, masking values associated with weight zero, + meaning instead of `nan * 0 = nan` you will get `0 * 0 = 0`. + + Args: + input_tensor (`torch.FloatTensor`): + Input tensor, of which the average must be computed. + weights (`torch.FloatTensor`, *optional*): + Weights tensor, of the same shape as `input_tensor`. + dim (`int`, *optional*): + The dim along which to average `input_tensor`. + + Returns: + `torch.FloatTensor`: The tensor with values averaged along the specified `dim`. + """ + if weights is not None: + weighted_tensor = torch.where(weights != 0, input_tensor * weights, torch.zeros_like(input_tensor)) + sum_weights = torch.clamp(weights.sum(dim=dim) if dim else weights.sum(), min=1.0) + return (weighted_tensor.sum(dim=dim) if dim else weighted_tensor.sum()) / sum_weights + else: + return input_tensor.mean(dim=dim) + + +class PatchTSMixerForPrediction(PatchTSMixerPreTrainedModel): + r""" + `PatchTSMixer` for forecasting application. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + + Returns: + `None`. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__(config) + self.loss = config.loss + self.return_dict = config.return_dict + self.prediction_channel_indices = config.prediction_channel_indices + self.num_parallel_samples = config.num_parallel_samples + + if config.loss == "mse": + self.distribution_output = None + else: + dim = config.prediction_length + distribution_output_map = { + "student_t": StudentTOutput, + "normal": NormalOutput, + "negative_binomial": NegativeBinomialOutput, + } + output_class = distribution_output_map.get(config.distribution_output) + if output_class is not None: + self.distribution_output = output_class(dim=dim) + else: + raise ValueError(f"Unknown distribution output {config.distribution_output}") + + self.model = PatchTSMixerModel(config) + self.head = PatchTSMixerForPredictionHead( + config=config, + distribution_output=self.distribution_output, + ) + + # Initialize weights and apply final processing + self.post_init() + + @auto_docstring + def forward( + self, + past_values: torch.Tensor, + observed_mask: torch.Tensor | None = None, + future_values: torch.Tensor | None = None, + output_hidden_states: bool | None = False, + return_loss: bool = True, + return_dict: bool | None = None, + **kwargs, + ) -> PatchTSMixerForPredictionOutput: + r""" + past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`): + Context values of the time series. For a pretraining task, this denotes the input time series to predict + the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly, + for classification or regression tasks, it denotes the appropriate context values of the time series. + + For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is + greater than 1. + observed_mask (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*): + Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected + in `[0, 1]`: + - 1 for values that are **observed**, + - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros). + future_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting,: + `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*): + Target values of the time series, that serve as labels for the model. The `future_values` is what the + Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT + required for a pretraining task. + + For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want + to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter, + pass the target data with all channels, as channel Filtering for both prediction and target will be + manually applied before the loss computation. + return_loss (`bool`, *optional*): + Whether to return the loss in the `forward` call. + """ + if self.loss == "mse": + loss = nn.MSELoss(reduction="mean") + elif self.loss == "nll": + loss = nll + else: + raise ValueError("Invalid loss function: Allowed values: mse and nll") + + return_dict = return_dict if return_dict is not None else self.return_dict + + # past_values: tensor [batch_size x context_length x num_input_channels] + model_output = self.model( + past_values, + observed_mask=observed_mask, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) # model_output: [batch_size x nvars x num_patch x d_model] + if isinstance(model_output, tuple): + model_output = PatchTSMixerModelOutput(*model_output) + + # tensor [batch_size x prediction_length x num_input_channels] + y_hat = self.head(model_output.last_hidden_state) + + loss_val = None + if self.prediction_channel_indices is not None: + if self.distribution_output: + distribution = self.distribution_output.distribution( + y_hat, + loc=model_output.loc[..., self.prediction_channel_indices], + scale=model_output.scale[..., self.prediction_channel_indices], + ) + if future_values is not None and return_loss is True: + loss_val = loss( + distribution, + future_values[..., self.prediction_channel_indices], + ) + # take average of the loss + loss_val = weighted_average(loss_val) + else: + y_hat = ( + y_hat * model_output.scale[..., self.prediction_channel_indices] + + model_output.loc[..., self.prediction_channel_indices] + ) + if future_values is not None and return_loss is True: + loss_val = loss(y_hat, future_values[..., self.prediction_channel_indices]) + else: + if self.distribution_output: + distribution = self.distribution_output.distribution( + y_hat, loc=model_output.loc, scale=model_output.scale + ) + if future_values is not None and return_loss is True: + loss_val = loss(distribution, future_values) + loss_val = weighted_average(loss_val) + else: + y_hat = y_hat * model_output.scale + model_output.loc + if future_values is not None and return_loss is True: + loss_val = loss(y_hat, future_values) + + if self.prediction_channel_indices is not None: + loc = model_output.loc[..., self.prediction_channel_indices] + scale = model_output.scale[..., self.prediction_channel_indices] + else: + loc = model_output.loc + scale = model_output.scale + + if not return_dict: + return tuple( + v + for v in [ + loss_val, + y_hat, + model_output.last_hidden_state, + model_output.hidden_states, + loc, + scale, + ] + ) + + return PatchTSMixerForPredictionOutput( + loss=loss_val, + prediction_outputs=y_hat, # tensor [batch_size x prediction_length x num_input_channels] + last_hidden_state=model_output.last_hidden_state, # x: [batch_size x nvars x num_patch x d_model] + hidden_states=model_output.hidden_states, + loc=loc, + scale=scale, + ) + + @torch.no_grad() + def generate( + self, + past_values: torch.Tensor, + observed_mask: torch.Tensor | None = None, + ) -> SamplePatchTSMixerPredictionOutput: + """ + Generate sequences of sample predictions from a model with a probability distribution head. + + Args: + past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`): + Past values of the time series that serves as context in order to predict the future. + + observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*): + Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected + in `[0, 1]`: + + - 1 for values that are **observed**, + - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros). + + Return: + [`SamplePatchTSMixerPredictionOutput`] where the outputs `sequences` tensor will have shape `(batch_size, + number of samples, prediction_length, num_input_channels)`. + """ + # get number of samples + num_parallel_samples = self.num_parallel_samples + + # get model output + outputs = self( + past_values=past_values, + future_values=None, + observed_mask=observed_mask, + output_hidden_states=False, + ) + + # get distribution + + distribution = self.distribution_output.distribution( + outputs.prediction_outputs, loc=outputs.loc, scale=outputs.scale + ) + + # get samples: list of [batch_size x prediction_length x num_channels] + samples = [distribution.sample() for _ in range(num_parallel_samples)] + + # stack tensors + samples = torch.stack(samples, dim=1) # [batch_size x num_samples x prediction_length x num_channels] + return SamplePatchTSMixerPredictionOutput(sequences=samples) + + +@auto_docstring( + custom_intro=""" + Output type of [`PatchTSMixerForTimeSeriesClassificationOutput`]. + """ +) +@dataclass +class PatchTSMixerForTimeSeriesClassificationOutput(ModelOutput): + r""" + loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`): + Total loss. + prediction_outputs (`torch.FloatTensor` of shape `(batch_size, num_labels)`): + Prediction output from the classification head. + last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`): + Backbone embeddings before passing through the head. + hidden_states (`tuple(torch.FloatTensor)`, *optional*): + Hidden-states of the model at the output of each layer plus the optional initial embedding outputs. + """ + + loss: torch.FloatTensor | None = None + prediction_outputs: torch.FloatTensor | None = None + last_hidden_state: torch.FloatTensor | None = None + hidden_states: tuple[torch.FloatTensor] | None = None + + +class PatchTSMixerForTimeSeriesClassification(PatchTSMixerPreTrainedModel): + r""" + `PatchTSMixer` for classification application. + + Args: + config (`PatchTSMixerConfig`): + Configuration. + + Returns: + `None`. + """ + + def __init__(self, config: PatchTSMixerConfig): + super().__init__(config) + + self.model = PatchTSMixerModel(config) + self.head = PatchTSMixerLinearHead( + config=config, + ) + self.return_dict = config.return_dict + if config.scaling in ["std", "mean", True]: + self.inject_scale = InjectScalerStatistics4D(d_model=config.d_model, num_patches=config.num_patches) + else: + self.inject_scale = None + + # Initialize weights and apply final processing + self.post_init() + + @auto_docstring + def forward( + self, + past_values: torch.Tensor, + target_values: torch.Tensor | None = None, + output_hidden_states: bool | None = False, + return_loss: bool = True, + return_dict: bool | None = None, + **kwargs, + ) -> PatchTSMixerForTimeSeriesClassificationOutput: + r""" + past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`): + Context values of the time series. For a pretraining task, this denotes the input time series to predict + the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly, + for classification or regression tasks, it denotes the appropriate context values of the time series. + + For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is + greater than 1. + target_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting, + `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*): + Target + values of the time series, that serve as labels for the model. The `target_values` is what the + Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT + required for a pretraining task. + + For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want + to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter, + pass the target data with all channels, as channel Filtering for both prediction and target will be + manually applied before the loss computation. + + For a classification task, it has a shape of `(batch_size,)`. + + For a regression task, it has a shape of `(batch_size, num_targets)`. + return_loss (`bool`, *optional*): + Whether to return the loss in the `forward` call. + """ + + loss = torch.nn.CrossEntropyLoss() + + return_dict = return_dict if return_dict is not None else self.return_dict + + model_output = self.model( + past_values, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) # x: [batch_size x nvars x num_patch x d_model] + if isinstance(model_output, tuple): + model_output = PatchTSMixerModelOutput(*model_output) + + if self.inject_scale is not None: + model_output.last_hidden_state = self.inject_scale( + model_output.last_hidden_state, + loc=model_output.loc, + scale=model_output.scale, + ) # x: [batch_size x nvars x num_patch x d_model] + + y_hat = self.head(model_output.last_hidden_state) # tensor [batch_size x n_labels] + + if target_values is not None and return_loss is True: + loss_val = loss(y_hat, target_values) + else: + loss_val = None + + if not return_dict: + return tuple( + v + for v in [ + loss_val, + y_hat, + model_output.last_hidden_state, + model_output.hidden_states, + ] + ) + + return PatchTSMixerForTimeSeriesClassificationOutput( + loss=loss_val, + prediction_outputs=y_hat, # tensor [batch_size x n_labels] + last_hidden_state=model_output.last_hidden_state, # x: [batch_size x nvars x num_patch x d_model] + hidden_states=model_output.hidden_states, + ) + + +@auto_docstring( + custom_intro=""" + Output type of [`PatchTSMixerForRegressionOutput`]. + """ +) +@dataclass +class PatchTSMixerForRegressionOutput(ModelOutput): + r""" + loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`): + Total loss. + regression_outputs (`torch.FloatTensor` of shape `(batch_size, num_targets)`): + Prediction output from the regression head. + last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`): + Backbone embeddings before passing through the head. + hidden_states (`tuple(torch.FloatTensor)`, *optional*): + Hidden-states of the model at the output of each layer plus the optional initial embedding outputs. + """ + + loss: torch.FloatTensor | None = None + regression_outputs: torch.FloatTensor | None = None + last_hidden_state: torch.FloatTensor | None = None + hidden_states: tuple[torch.FloatTensor] | None = None + + +class InjectScalerStatistics4D(nn.Module): + def __init__(self, d_model: int, num_patches: int, expansion: int = 2): + super().__init__() + + self.inverse_trans_expansion = nn.Linear(d_model + 2, expansion * d_model) + self.inverse_trans_compression = nn.Linear(expansion * d_model, d_model) + self.map_scale_expansion = nn.Linear(2, 2 * expansion) + self.map_scale_compression = nn.Linear(2 * expansion, 2) + self.num_patches = num_patches + + def forward(self, inputs: torch.Tensor, loc: torch.Tensor, scale: torch.Tensor): + """ + Args: + inputs (`torch.Tensor` of shape `(batch_size, num_input_channels, num_patch, d_model)`) + loc (`torch.Tensor` of shape `(batch_size, 1, num_input_channels)`) + scale (`torch.Tensor` of shape `(batch_size, 1, num_input_channels)`) + Returns: + `torch.Tensor` of shape `(batch_size, num_input_channels, num_patch, d_model)` + """ + + mean = loc.transpose(-1, -2) # [batch_size x n_channels x 1 ] + mean = mean.unsqueeze(-2) # [batch_size x n_channels x 1 x 1] + mean = mean.repeat(1, 1, self.num_patches, 1) # [batch_size x n_channels x num_patch x 1] + + stdev = scale.transpose(-1, -2) # [batch_size x n_channels x 1 ] + stdev = stdev.unsqueeze(-2) # [batch_size x n_channels x 1 x 1] + stdev = stdev.repeat(1, 1, self.num_patches, 1) # [batch_size x n_channels x num_patch x 1] + + concat_stats = torch.cat([mean, stdev], dim=-1) # [batch_size x n_channels x num_patch x 2] + + concat_stats = self.map_scale_expansion(concat_stats) # [batch_size x n_channels x num_patch x (2*expansion)] + concat_stats = self.map_scale_compression(concat_stats) # [batch_size x n_channels x num_patch x 2] + + inputs = torch.cat([inputs, concat_stats], dim=-1) # [batch_size x channels x num_patch x d_model+2] + inputs = self.inverse_trans_expansion(inputs) # [batch_size x channels x num_patch x (expansion*d_model)] + inputs = self.inverse_trans_compression(inputs) # [batch_size x channels x num_patch x d_model] + + return inputs + + +@auto_docstring( + custom_intro=""" + `PatchTSMixer` for regression application. + """ +) +class PatchTSMixerForRegression(PatchTSMixerPreTrainedModel): + def __init__(self, config: PatchTSMixerConfig): + super().__init__(config) + + self.model = PatchTSMixerModel(config) + + self.loss = config.loss + self.distribution_output = config.distribution_output + + self.return_dict = config.return_dict + self.num_parallel_samples = config.num_parallel_samples + + if config.loss == "mse": + self.distribution_output = None + else: + distribution_output_map = { + "student_t": StudentTOutput, + "normal": NormalOutput, + "negative_binomial": NegativeBinomialOutput, + } + output_class = distribution_output_map.get(config.distribution_output) + if output_class is not None: + self.distribution_output = output_class(dim=config.num_targets) + else: + raise ValueError(f"Unknown distribution output {config.distribution_output}") + + if config.scaling in ["std", "mean", True]: + self.inject_scale = InjectScalerStatistics4D(d_model=config.d_model, num_patches=config.num_patches) + else: + self.inject_scale = None + + self.head = PatchTSMixerLinearHead( + config=config, + distribution_output=self.distribution_output, + ) + + # Initialize weights and apply final processing + self.post_init() + + @auto_docstring + def forward( + self, + past_values: torch.Tensor, + target_values: torch.Tensor | None = None, + output_hidden_states: bool | None = False, + return_loss: bool = True, + return_dict: bool | None = None, + **kwargs, + ) -> PatchTSMixerForRegressionOutput: + r""" + past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`): + Context values of the time series. For a pretraining task, this denotes the input time series to predict + the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly, + for classification or regression tasks, it denotes the appropriate context values of the time series. + + For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is + greater than 1. + target_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting, + `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*): + Target values of the time series, that serve as labels for the model. The `target_values` is what the + Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT + required for a pretraining task. + + For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want + to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter, + pass the target data with all channels, as channel Filtering for both prediction and target will be + manually applied before the loss computation. + + For a classification task, it has a shape of `(batch_size,)`. + + For a regression task, it has a shape of `(batch_size, num_targets)`. + return_loss (`bool`, *optional*): + Whether to return the loss in the `forward` call. + """ + + if self.loss == "mse": + loss = nn.MSELoss(reduction="mean") + elif self.loss == "nll": + loss = nll + else: + raise ValueError("Invalid loss function: Allowed values: mse and nll") + + return_dict = return_dict if return_dict is not None else self.return_dict + model_output = self.model( + past_values, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) # model_output: [batch_size x nvars x num_patch x d_model] + if isinstance(model_output, tuple): + model_output = PatchTSMixerModelOutput(*model_output) + + if self.inject_scale is not None: + model_output.last_hidden_state = self.inject_scale( + model_output.last_hidden_state, + loc=model_output.loc, + scale=model_output.scale, + ) # x: [batch_size x nvars x num_patch x d_model] + + y_hat = self.head(model_output.last_hidden_state) # [batch_size x num_targets] + + if target_values is not None and return_loss is True: + if self.distribution_output: + if self.distribution_output == "negative_binomial" and torch.any(target_values < 0): + raise Exception("target_values cannot be negative for negative_binomial distribution.") + distribution = self.distribution_output.distribution(y_hat) + # y_hat should be a 2-tuple, each with dimension [bs, num_targets] + y_hat = tuple(item.view(-1, self.config.num_targets) for item in y_hat) + loss_val = loss(distribution, target_values) + # take average of the loss + loss_val = weighted_average(loss_val) + else: + loss_val = loss(y_hat, target_values) + else: + loss_val = None + + if not return_dict: + return tuple( + v + for v in [ + loss_val, + y_hat, + model_output.last_hidden_state, + model_output.hidden_states, + ] + ) + + return PatchTSMixerForRegressionOutput( + loss=loss_val, + regression_outputs=y_hat, # tensor [batch_size x num_targets] + last_hidden_state=model_output.last_hidden_state, # [batch_size x nvars x num_patch x d_model] + hidden_states=model_output.hidden_states, + ) + + @torch.no_grad() + def generate( + self, + past_values: torch.Tensor, + ) -> SamplePatchTSMixerRegressionOutput: + """ + Generate sequences of sample predictions from a model with a probability distribution head. + + Args: + past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`): + Past values of the time series that serves as context in order to predict the target values. + + Return: + [`SamplePatchTSMixerRegressionOutput`] where the outputs `sequences` tensor will have shape `(batch_size, + number of samples, num_targets)`. + """ + # get number of samples + num_parallel_samples = self.num_parallel_samples + + # get model output + outputs = self( + past_values=past_values, + target_values=None, + output_hidden_states=False, + ) + + # get distribution + distribution = self.distribution_output.distribution(outputs.regression_outputs) + + # get samples + samples = [ + distribution.sample() for _ in range(num_parallel_samples) + ] # samples: list of [batch_size x num_targets] + # stack tensors + # [batch_size x num_samples x num_targets] + samples = torch.stack(samples, dim=1).view(-1, num_parallel_samples, self.config.num_targets) + return SamplePatchTSMixerRegressionOutput(sequences=samples) + + +__all__ = [ + "PatchTSMixerPreTrainedModel", + "PatchTSMixerModel", + "PatchTSMixerForPretraining", + "PatchTSMixerForPrediction", + "PatchTSMixerForTimeSeriesClassification", + "PatchTSMixerForRegression", +] diff --git a/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/timesformer/configuration_timesformer.py b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/timesformer/configuration_timesformer.py new file mode 100644 index 0000000000000000000000000000000000000000..554d0e02531e394e21f83d72af4665e1b0728bb2 --- /dev/null +++ b/LTA_openwebtext_dualt/mini_owt_logdirichlet/.venv_qwen35_uv/lib/python3.12/site-packages/transformers/models/timesformer/configuration_timesformer.py @@ -0,0 +1,66 @@ +# Copyright 2022 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""TimeSformer model configuration""" + +from huggingface_hub.dataclasses import strict + +from ...configuration_utils import PreTrainedConfig +from ...utils import auto_docstring + + +@auto_docstring(checkpoint="facebook/timesformer-base-finetuned-k600") +@strict +class TimesformerConfig(PreTrainedConfig): + r""" + num_frames (`int`, *optional*, defaults to 8): + The number of frames in each video. + attention_type (`str`, *optional*, defaults to `"divided_space_time"`): + The attention type to use. Must be one of `"divided_space_time"`, `"space_only"`, `"joint_space_time"`. + + Example: + + ```python + >>> from transformers import TimesformerConfig, TimesformerModel + + >>> # Initializing a TimeSformer timesformer-base style configuration + >>> configuration = TimesformerConfig() + + >>> # Initializing a model from the configuration + >>> model = TimesformerModel(configuration) + + >>> # Accessing the model configuration + >>> configuration = model.config + ```""" + + model_type = "timesformer" + + image_size: int | list[int] | tuple[int, int] = 224 + patch_size: int | list[int] | tuple[int, int] = 16 + num_channels: int = 3 + num_frames: int = 8 + hidden_size: int = 768 + num_hidden_layers: int = 12 + num_attention_heads: int = 12 + intermediate_size: int = 3072 + hidden_act: str = "gelu" + hidden_dropout_prob: float | int = 0.0 + attention_probs_dropout_prob: float | int = 0.0 + initializer_range: float = 0.02 + layer_norm_eps: float = 1e-6 + qkv_bias: bool = True + attention_type: str = "divided_space_time" + drop_path_rate: int = 0 + + +__all__ = ["TimesformerConfig"]