diff --git a/dev/trainer_rank_recompute_memory.csv b/dev/trainer_rank_recompute_memory.csv new file mode 100644 index 000000000..e9bc8dc5c --- /dev/null +++ b/dev/trainer_rank_recompute_memory.csv @@ -0,0 +1,74 @@ +phase,source_sha,driver_sha256,model,layers,gdn_layers,tp,ep,etp,sequence_parallel,compiled,concentrated_routing,mode,modules,lengths,shared_prefix,logical_tokens,packed_tokens,grad_segment_count,selected_max_depth,memory_minimal,estimate_bytes,available_min_bytes,measured_rank_samples,refused_rank_samples,forward_peak_max_bytes,first_forward_peak_max_bytes,repeat_forward_peak_max_bytes,retained_max_bytes,forward_backward_peak_max_bytes,baseline_min_bytes,baseline_max_bytes,finite_loss_all,finite_gradients_all,estimate_covers_forward_all,min_estimate_to_forward_ratio +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,12288+12288,3686,24576,20892,,2,False,124038771507,128845011620,8,0,112387296256,112387296256,112387296256,112146179072,113656130560,13966070784,13966070784,True,True,True,1.103672529183902 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,19221+19222,5733,38443,32712,,2,True,194173605683,128845011620,0,4,,,,,,,,,,, +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,2048+2048,614,4096,4096,,1,False,24369745100,128845011620,8,0,22027530752,22027530752,22027530752,21972978176,22224638464,13966070784,13966070784,True,True,True,1.106331225881382 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,256+256,76,512,512,,1,False,3110810419,128847108772,8,0,3042869760,3042869760,2975759872,3036048896,3114492928,13865406464,13966070784,True,True,True,1.022327823521438 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,4096+4096,1228,8192,6964,,2,False,41395470336,128845011620,8,0,37665950208,37665950208,37665950208,37586216960,38089535488,13966070784,13966070784,True,True,True,1.0990156920880725 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,none,,8192+8192,2457,16384,13928,,2,False,82717120921,128845011620,8,0,74952197632,74952197632,74952197632,74790640128,75797275136,13966070784,13966070784,True,True,True,1.1035983404666023 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,16384+16384,4915,32768,27856,,2,False,115282732646,135698216653,16,0,104854994944,104854994944,104854994944,104569116160,106582384128,7119606784,7119606784,True,True,True,1.0994491269354325 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,19221+19222,5733,38443,32712,,2,False,135366117990,135698207949,16,0,123127285760,123127285760,123127277056,122791571968,125156542976,7119606784,7119615488,True,True,True,1.0993998377732126 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,2048+2048,614,4096,4096,,1,False,17006224998,135700313805,16,0,15596174336,15596174336,15531161600,15554143744,15805804032,7018942464,7119606784,True,True,True,1.0904100346419723 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,4096+4096,1228,8192,6968,,2,False,28892538470,135698216653,16,0,26428632064,26428632064,26428632064,26357117440,26860435968,7119606784,7119606784,True,True,True,1.093228677142024 +before_segment_states,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,none,,8192+8192,2457,16384,13928,,2,False,57678276198,135698216653,16,0,52559443456,52559443456,52559443456,52416501248,53423136256,7119606784,7119606784,True,True,True,1.0973913041199763 +before_segment_states,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,False,selective,core_attn,10000+10000,6000,20000,14000,,2,False,53618370150,133630881485,4,0,48752947200,48752947200,48752947200,48458028032,49441070080,9191136256,9191136256,True,True,True,1.0997975143951912 +before_segment_states,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.5-9B,32,24,2,1,1,True,True,False,selective,core_attn,5000+5000,3000,10000,7000,,2,False,26846094950,133632978637,4,0,24515836928,24515836928,24460253184,24369065984,24860588032,9090471936,9191136256,True,True,True,1.0950511307790014 +cold_workspace_negative_control,71055f4639b2d9f247d7f3a3421a7f79cecc2c61,9bfae96d3b04221b7ccac18086bce2bcc13b3544055be38c83db4c4505c3f370,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,256+256,76,512,436,,2,False,813621248,140941383373,4,0,829413888,829413888,768536064,827618816,866536960,1781843968,1882731520,True,True,False,0.9809592771130449 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,1024+1024,307,2048,1742,,2,False,3324583116,140940323533,4,0,2976541696,2976541696,2975649792,2969369600,3019703296,1882899456,1883791360,True,True,True,1.1169281184495794 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,128+128,38,256,218,,2,False,480630374,140941438669,4,0,379880448,379880448,379768832,378982400,385275904,1882564608,1882676224,True,True,True,1.2652148235857614 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,16384+16384,4915,32768,27854,,2,False,52052538982,140914528461,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894488064,1909586432,True,True,True,1.0943870746538502 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,256+256,76,512,436,,2,False,887440998,140941215437,4,0,768759296,768759296,768536064,766964224,779549184,1882676224,1882899456,True,True,True,1.154380834960336 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,4096+4096,1228,8192,6964,,2,False,13069429964,140936757965,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883791360,1887356928,True,True,True,1.0944747232661334 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,64+64,19,128,110,,2,False,279085875,140941550285,4,0,256186368,256186368,189020160,255733248,290649600,1781843968,1882564608,True,True,True,1.0893861261189355 +eager_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,8192+8192,2457,16384,13928,,2,False,26065040179,140929626829,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1887356928,1894488064,True,True,True,1.0923247616155016 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,False,selective,core_attn,4096+4096,1228,8192,6964,3,2,False,13069429964,140938041037,4,0,9483472896,9483472896,9414886400,9402919936,9604248576,1781843968,1886073856,True,True,True,1.3781269907475042 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3-1.7B,28,0,2,1,1,True,True,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,26065040179,140930909901,4,0,18806726144,18806726144,18799595008,18645209600,19047864832,1886073856,1893204992,True,True,True,1.385942453748956 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn,1500+1500,450,3000,3000,2,1,False,19630804582,124212695757,8,0,14968917504,14968917504,14561442304,14694942208,14968917504,17555830784,17994856448,True,True,True,1.311437822858884 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn,3000+3000,900,6000,5100,3,2,False,33310583398,123993616077,8,0,24725110272,24725110272,24549762048,24262462976,24725110272,17991170048,18165701632,True,True,True,1.3472370004239227 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn,6000+6000,1800,12000,10200,3,2,False,66441104998,123551084237,8,0,49004391936,49004391936,48701675008,48079419904,49004391936,18149577728,18494987264,True,True,True,1.3558193944080041 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn+moe,1500+1500,450,3000,3000,2,1,False,4606881382,123971948237,8,0,4355604992,4355604992,3704471040,3923585536,4501488128,17555830784,18235603968,True,True,True,1.0576903531108819 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn+moe,3000+3000,900,6000,5100,3,2,False,7769913958,123571519693,8,0,6649523200,6649523200,6301277184,5920428544,6840852992,18233154560,18589895168,True,True,True,1.1684918939751952 +final,7760536e7bf10e4054d708383a1333e7b9f0e120,259914aa7a1818ac7d2b6cf68b419e409f5c741f8bcbd9af7d459cf84123fbc1,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,False,True,selective,core_attn+moe,6000+6000,1800,12000,10200,3,2,False,15359766118,122728326861,8,0,13253387776,13253387776,12513351168,11752620032,13627614208,18589530112,19319841792,True,True,True,1.1589313145891915 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn,1500+1500,450,3000,3000,2,1,False,19630804582,124206687908,8,0,6961144320,6961144320,6634839040,6897091072,6971726336,17555830784,17994123264,True,True,True,2.820054243897647 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn,3000+3000,900,6000,5100,3,2,False,33310583398,123984802980,8,0,11145096192,11145096192,11101520896,11016890880,11164938752,17993316352,18169870848,True,True,True,2.9888107580363896 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn,6000+6000,1800,12000,10200,3,2,False,66441104998,123536855204,8,0,21793972736,21793972736,21733763072,21565764096,21860678144,18163617792,18502475264,True,True,True,3.0485999869243847 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn+moe,1500+1500,450,3000,3000,2,1,False,4606881382,123961172644,8,0,4025781248,4025781248,3373262336,3903570944,4122183680,17555830784,18239638528,True,True,True,1.1443446869570197 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn+moe,3000+3000,900,6000,5100,3,2,False,7769913958,123548129444,8,0,6099939840,6099939840,5727812096,5894723072,6208201728,18236235264,18606544384,True,True,True,1.2737689488426167 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-35B-A3B,40,30,4,4,1,True,True,False,selective,core_attn+moe,6000+6000,1800,12000,10200,3,2,False,15359766118,122723152548,8,0,12123608576,12123608576,11390174208,11691548160,12330070016,18591955456,19316177920,True,True,True,1.2669302231025774 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,10240+10240,5120,20480,15360,3,2,False,28776870707,138374372045,4,0,25853539840,25853539840,25853539840,25068629504,26035840512,4447645696,4447645696,True,True,True,1.1130727507757792 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,128+128,64,256,256,2,1,False,662215065,138374372045,4,0,498778624,498778624,498778624,484490240,501137408,4447645696,4447645696,True,True,True,1.3276733066251052 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,2048+2048,1024,4096,4096,2,1,False,7788272025,138374372045,4,0,6917972480,6917972480,6917972480,6701941248,6959876608,4447645696,4447645696,True,True,True,1.1258026896632003 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,5120+5120,2560,10240,10240,2,1,False,19189963161,138374372045,4,0,17218811392,17218811392,17218811392,16678733312,17323568640,4447645696,4447645696,True,True,True,1.1144766455782082 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,64+64,32,128,128,2,1,False,424679833,138376469197,4,0,336145408,336145408,269035520,328608256,432695808,4346981376,4447645696,True,True,True,1.2633813310934772 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,2,1,1,True,False,False,selective,core_attn,256+256+256+256+256+256+256+256,0,2048,2048,8,1,False,22748594176,115395819213,4,0,20101526016,20101526016,20101526016,20025496064,20101526016,27428295680,27428295680,True,True,True,1.1316849356557825 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,2,1,1,True,False,False,selective,core_attn,512+512+512+512+512+512+512+512,0,4096,4096,8,1,False,44068660838,115395819213,4,0,39580185088,39580185088,39580185088,39428125184,39580185088,27428295680,27428295680,True,True,True,1.1134020909710405 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,2,1,1,True,False,False,selective,core_attn,64+64+64+64+64+64+64+64,0,512,512,8,1,False,6758544179,115395819213,4,0,5649679872,5649679872,5590548992,5628575232,5864041472,27327631360,27428295680,True,True,True,1.1962702900204254 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,12288+12288,3686,24576,20892,3,2,False,124292779212,128845011620,8,0,112506383872,112506383872,112506383872,112259961856,113769913344,13966070784,13966070784,True,True,True,1.1047620138018972 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,2048+2048,614,4096,4096,2,1,False,24539083571,128847108772,8,0,22065410560,22065410560,22065410560,22009815552,22261475840,13966070784,13966070784,True,True,True,1.1121063668529356 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,4096+4096,1228,8192,6964,3,2,False,41649478041,128845011620,8,0,37657602560,37657602560,37657602560,37575004672,38078323200,13966070784,13966070784,True,True,True,1.1060045039946378 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,64+64,19,128,128,2,1,False,1002405888,128847108772,8,0,903788032,903788032,836678144,902051328,938049536,13865406464,13966070784,True,True,True,1.1091161339919138 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,82971128627,128845011620,8,0,75082050048,75082050048,75082050048,74916958208,75923593216,13966070784,13966070784,True,True,True,1.1050727647148222 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,12288+12288,3686,24576,20892,3,2,False,124292779212,128845011620,8,0,112379268096,112379268096,112379268096,112138194432,113648145920,13966070784,13966070784,True,True,True,1.1060116453670341 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,19221+19222,5733,38443,32712,3,2,True,194427613388,128845011620,0,4,,,,,,,,,,, +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,2048+2048,614,4096,4096,2,1,False,24539083571,128845011620,8,0,22025957888,22025957888,22025957888,21971411456,22223071744,13966070784,13966070784,True,True,True,1.1140983604789865 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,256+256,76,512,512,2,1,False,3280148889,128845011620,8,0,2975563264,2975563264,2975563264,2968744960,3000204288,13966070784,13966070784,True,True,True,1.1023623421773767 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,4096+4096,1228,8192,6964,3,2,False,41649478041,128845011620,8,0,37663271424,37663271424,37663271424,37583555072,38086873600,13966070784,13966070784,True,True,True,1.10583803441089 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,64+64,19,128,128,2,1,False,1002405888,128847108772,8,0,900511232,900511232,833401344,898807296,997523456,13865406464,13966070784,True,True,True,1.1131520100795367 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,82971128627,128845011620,8,0,74946848256,74946848256,74946848256,74785321984,75791956992,13966070784,13966070784,True,True,True,1.1070662817413086 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,16384+16384,4915,32768,27856,3,2,False,90497895628,128845011620,8,0,81347673600,81347673600,81347673600,80013564928,82026832896,13966070784,13966070784,True,True,True,1.1124828974580534 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,2048+2048,614,4096,4096,2,1,False,13493803417,128845011620,8,0,11972481536,11972481536,11972481536,11769029120,12035353088,13966070784,13966070784,True,True,True,1.1270682169294264 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,4096+4096,1228,8192,6964,3,2,False,22870344499,128845011620,8,0,20407347712,20407347712,20407347712,20073009664,20576328192,13966070784,13966070784,True,True,True,1.1206916656568604 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,64+64,19,128,128,2,1,False,657240883,128847108772,8,0,548321280,548321280,481211392,539703808,583254016,13865406464,13966070784,True,True,True,1.1986419403602209 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,8192+8192,2457,16384,13928,3,2,False,45412861542,128845011620,8,0,40674358272,40674358272,40674358272,40005276672,41011911680,13966070784,13966070784,True,True,True,1.1164985379317456 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,16384+16384,4915,32768,27856,3,2,False,115409736499,135698216653,16,0,104849291776,104849291776,104849291776,104563468800,106576736768,7119606784,7119606784,True,True,True,1.100720229427599 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,19221+19222,5733,38443,32712,3,2,False,135493121843,135698207949,16,0,123120420864,123120420864,123120412160,122784771584,125149742592,7119606784,7119615488,True,True,True,1.100492679379865 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,2048+2048,614,4096,4096,2,1,False,17090894233,135698216653,16,0,15528278528,15528278528,15526181376,15486253568,15737913856,7119606784,7119606784,True,True,True,1.100630324358386 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,4096+4096,1228,8192,6968,3,2,False,29019542323,135698216653,16,0,26427723264,26427723264,26427723264,26356226560,26859545088,7119606784,7119606784,True,True,True,1.09807197665531 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,64+64,19,128,128,2,1,False,687626649,135700313805,16,0,652222464,652222464,584981504,650909184,748507136,7018942464,7119606784,True,True,True,1.0542823759593782 +final,850e20a2c4c60247430b3d2f4e463eaf860f92a7,defe4bf66ace7e2745c1ae7f029f4a0e1ccb95ee8b268d64bfd536cb124bb4d9,Qwen/Qwen3.8-27B,64,48,8,1,1,True,True,False,selective,core_attn,8192+8192,2457,16384,13928,3,2,False,57805280051,135698216653,16,0,52556722176,52556722176,52556722176,52413810688,53420445696,7119606784,7119606784,True,True,True,1.0998646349637982 +segment_state_negative_control,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,False,False,selective,core_attn,64+64,19,128,128,,1,False,833067417,128847108772,8,0,892974592,892974592,836678144,890582528,989954048,13865406464,13966070784,True,True,False,0.9329127888556991 +segment_state_negative_control,45297a4af,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.8-27B,64,48,4,1,1,True,True,False,selective,core_attn+mlp,64+64,19,128,128,,1,False,487902412,128847108772,8,0,548321280,548321280,481211392,539703808,639605248,13865406464,13966070784,True,True,False,0.8898111924454217 +segment_state_negative_control,45297a4af6f1a17211a5a7c1ab02f19dc423f4a8,c6dd62dc57644a785a49114c5b489c21b44f1cfff3d21152b9d87d59d61e1edd,Qwen/Qwen3.5-4B,32,24,2,1,1,True,True,False,selective,core_attn+mlp,128+128,64,256,256,,1,False,548890214,138376469197,4,0,563791360,563791360,498778624,549502976,660929536,4346981376,4447645696,True,True,False,0.973569751051169 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,1024+1024,307,2048,1742,,2,False,2756291788,140940714701,4,0,3043651584,3043651584,2975649792,3036479488,3091428352,1781843968,1883400192,True,True,False,0.9055871580339203 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,16384+16384,4915,32768,27854,,2,False,44072283340,140914919629,4,0,47563188736,47563188736,47548090368,47448521216,48253829632,1894096896,1909195264,True,True,False,0.9266048915396251 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,4096+4096,1228,8192,6964,,2,False,11018859315,140937149133,4,0,11941280768,11941280768,11937715200,11912611840,12113940480,1883400192,1886965760,True,True,False,0.9227535579372788 +unfused_negative_control,d31f9423edf7c1fbdd3e125a08d254ea662377de,42e9898ce34d566da5b9ebd2d3ef608317ae61b2d51f3ede058216d2ee61e788,Qwen/Qwen3-1.7B,28,0,2,1,1,True,False,False,selective,core_attn,8192+8192,2457,16384,13928,,2,False,22037718630,140930017997,4,0,23861987840,23861987840,23854856704,23804649984,24207305216,1886965760,1894096896,True,True,False,0.9235491518044459 diff --git a/dev/trainer_rank_recompute_memory.md b/dev/trainer_rank_recompute_memory.md new file mode 100644 index 000000000..4f6ffe830 --- /dev/null +++ b/dev/trainer_rank_recompute_memory.md @@ -0,0 +1,181 @@ +# Recompute memory calibration + +The revised estimate is **10–13% above measured 27B peaks** for the paired +2k-and-larger workloads, including MLP recompute. The previous TP4 estimate was +about 60% high. The new floor admits a 12k TP4 pair and, at TP8, the original +19,221 + 19,222-token pair from #913. Both complete forward and backward. + +All numbers below are **incremental allocated GiB above the pre-forward +baseline**, not total GPU use or reserved memory. The estimate includes the +existing 10% safety margin. Measurements use native H200 execution with bf16, +rank-1 LoRA, random model/adaptor weights, SP enabled, and DP/CP/PP=1. + +## Requested 27B series and admission boundary + +Qwen3.8-27B has 48 GDN and 16 full-attention layers. Pairs share 30% of their +prefix; the planner chooses the layout, including TP padding. + +| Tokens per sequence | TP | Packed tokens | Previous estimate | New estimate | Compiled forward | Eager forward | +| --: | --: | --: | --: | --: | --: | --: | +| 2,048 | 4 | 4,096 | 32.974 | 22.854 | 20.513 | 20.550 | +| 4,096 | 4 | 6,964 | 56.075 | 38.789 | 35.077 | 35.071 | +| 8,192 | 4 | 13,928 | 112.151 | 77.273 | 69.800 | 69.926 | +| 12,288 | 4 | 20,892 | — | 115.757 | 104.661 | 104.780 | +| 2,048 | 8 | 4,096 | 19.259 | 15.917 | 14.462 | — | +| 4,096 | 8 | 6,968 | 32.775 | 27.027 | 24.613 | — | +| 8,192 | 8 | 13,928 | 65.513 | 53.835 | 48.947 | — | +| 16,384 | 8 | 27,856 | 131.025 | 107.484 | 97.649 | — | +| 19,221 + 19,222 | 8 | 32,712 | 153.866 | 126.188 | 114.665 | — | + +The 12k TP4 and original long TP8 shapes were not used to fit the component +coefficients. The latter was previously refused, so its true selective peak +was unknown. TP4 still refuses the original pair, now at 181.075 GiB. The new +8k TP4 estimate fits both the measured budget and #913's 119.289 GiB budget. +The issue's error labels say GB but divide bytes by 1024³. + +Adding `mlp` to selective recomputation at TP4 gives: + +| Tokens per sequence | Estimate | Forward peak | Forward + backward peak | +| --: | --: | --: | --: | +| 2,048 | 12.567 | 11.150 | 11.209 | +| 4,096 | 21.300 | 19.006 | 19.163 | +| 8,192 | 42.294 | 37.881 | 38.195 | +| 16,384 | 84.283 | 75.761 | 76.393 | + +## What is being priced + +Component hooks on eager native layers separated MLP, attention, and GDN +retention. They exposed native activation fusion as the main MLP distinction: +Qwen3.8 uses fused SwiGLU even in eager execution, while Qwen3-1.7B uses the +unfused path. Compilation alone is not a reliable discount because it can fall +back to eager. + +Let H be hidden width, F dense FFN width, Q the full query width, KV the full +key/value width, and K/V the GDN key/value widths. SP is TP when sequence +parallel is enabled, otherwise 1. The common per-component term is +`2H/SP + gathered`, with `gathered = H` when SP>1 and zero otherwise: without +sequence sharding, the LoRA input aliases norm output. + +- Dense MLP storage is `common + cF/TP`: c=3 for native fused activations, + c=5 for unfused SwiGLU, with an additional allowance for unfused clamping. + These count gate/up, activation output, and the unfused SiLU/offset tensors. +- Full attention costs `common + (5Q + 3KV)/TP`, with another `2Q/TP` for + output gating. If KV groups are fewer than TP ranks, ART's replicated-QKV + LoRA path adds the global QKV storage that survives slicing. Omitting that + term underprices TP8. +- GDN costs `common + (4K + 8V)/TP`. Attention and GDN terms are multiplied by + their actual layer counts. These are calibrated envelopes of the measured + native paths, not an exact liveness proof for every kernel. +- Checkpointed dense MLPs retain their inputs. Native MoE checkpoint flags + determine the discounted MoE layer count; its external norm is also priced. + One full MLP workspace remains charged, including worst-case expert dispatch. +- Each gradient-enabled prefix segment gets an allowance for initial/final GDN + recurrent states (fp32) and convolution history. Exact plans use actual + segment counts; cheap admission uses the radix-tree bound, and optimistic + pruning omits this nonnegative term. A separate 64 MiB allowance covers the + observed roughly 58 MiB cold kernel setup cost. + +The existing static heuristic and routed FC2 bound remain floors. Profiles can +only increase the estimate, and output storage and the 10% safety factor are +applied afterward. Full-recompute and no-grad requests keep their existing path. + +## Cross-checks and remaining conservatism + +The [CSV](trainer_rank_recompute_memory.csv) includes short-input failures that +motivated the fixed costs. Without segment states, a 64-token 27B TP4 pair was +estimated at 0.776 GiB against 0.832 GiB observed. The corrected estimate is +0.934 GiB against a fresh eager peak of 0.842 GiB. Eight unshared 64-token +sequences at TP2 also pass: 6.294 estimated versus 5.262 observed. Larger +8-sequence batches and 4B MLP-checkpointed pairs provide additional checks. + +Qwen3.5-4B TP2 with MLP recompute, pairs of 2,048 / 5,120 / 10,240 tokens and a +50% common prefix, yields estimates 7.253 / 17.872 / 26.801 against peaks +6.443 / 16.036 / 24.078 GiB. A 9B TP2 check used previously unmeasured 5k / 10k +lengths with a 60% prefix: 25.002 / 49.936 estimated against 22.832 / 45.405 GiB, +before adding the nonnegative segment-state term. + +The unfused attention-only eager control is covered from 64 through 16,384 +tokens per sequence. Its larger inputs have about 9% headroom. Compiled-only +Qwen3-1.7B peaks are lower: 12.172 / 24.275 estimated against 8.832 / 17.515 GiB +for 4k / 8k pairs. That remaining 38–39% overestimate preserves the eager +fallback allowance; it does not apply to the natively fused 27B MLP. + +MoE cannot assume balanced expert dispatch. Qwen3.5-35B-A3B, TP4/EP4/ETP1: + +| Tokens per sequence | Default estimate | Balanced peak | Concentrated peak | With `moe`: estimate | Balanced peak | Concentrated peak | +| --: | --: | --: | --: | --: | --: | --: | +| 1,500 | 18.283 | 6.483 | 13.941 | 4.290 | 3.749 | 4.056 | +| 3,000 | 31.023 | 10.380 | 23.027 | 7.236 | 5.681 | 6.193 | +| 6,000 | 61.878 | 20.297 | 45.639 | 14.305 | 11.291 | 12.343 | + +The concentrated eager run routes about 31/32 of tokens to the first expert +rank; the balanced run uses the native random-weight router with compilation. +Their difference includes both routing and execution mode. The routed term +uses actual top-k expert plus shared FFN widths, not the unrelated dense FFN +fallback. It still makes no EP/ETP discount and is deliberately loose for +balanced routing. MoE recomputation removes most of that retained storage. + +A fully concentrated stress attempt segfaulted in Transformer Engine's grouped +GEMM on peers receiving zero tokens, before a complete measurement. The revised +stress keeps those peers nonempty. This kernel limitation is not fixed by the +memory estimate and is not counted as a successful calibration cell. + +## Evidence and reproduction + +The final campaign contains **45 cells, 360 measured rank-samples, and 4 refused +rank-samples**. Every measured forward is covered; every backward completes; +all measured losses and adapter gradients are finite. The smallest observed +forward headroom is 5.4%. The CSV also preserves earlier controls and negative +controls as separate phases, with the actual source/driver hash for each row. +It aggregates maximum peaks across every rank and both repetitions, and minimum +available memory. Refusals have no invented observed peak. + +Estimator source: `850e20a2c4c60247430b3d2f4e463eaf860f92a7`. Stress-driver-only +revision: `7760536e7bf10e4054d708383a1333e7b9f0e120`. A later type annotation does +not change the calculation. Native execution uses PyTorch 2.11.0+cu128, CUDA +12.8, and H200s. Component instrumentation was used only for diagnosis; final +measurements have no component hooks. + +Every sample clears learned profiles, gradients, and unused cached allocations. +The first workload in each process includes cold kernel setup; later lengths +can reuse initialized kernels. CSV `first_` and `repeat_` columns refer to the +two executions of each shape, not independent fresh processes. The driver tries +the minimum-memory layout before refusing, executes no split or budget bypass, +and backpropagates a mean-square hidden-state loss. Peak collection precedes +the finite-gradient checks. No optimizer step is included. + +After `INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh`: + +```sh +ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE=4 \ +ART_MEGATRON_CONTEXT_PARALLEL_SIZE=1 \ +ART_MEGATRON_DATA_PARALLEL_SIZE=1 \ +ART_MEGATRON_PIPELINE_MODEL_PARALLEL_SIZE=1 \ +uv run --project megatron_runtime --no-sync python -m torch.distributed.run \ + --standalone --nproc-per-node=4 dev/trainer_rank_recompute_memory.py \ + --mode selective --pairs --tokens 64 256 2048 4096 8192 12288 --reported-pair \ + --evidence scratch/recompute-memory-calibrated/tp4-selective.jsonl +``` + +Use `ART_DISABLE_MEGATRON_COMPILE=1` for eager, `--modules core_attn mlp` for +dense MLP checkpoints, or `--modules core_attn moe` for MoE checkpoints. +`--sequences 8 --prefix-fraction 0` checks many short unshared sequences. +`--concentrate-routing` enables the expert-imbalance stress. When running four +processes on an eight-GPU host, also set +`ART_MEGATRON_EXPERT_MODEL_PARALLEL_SIZE=4` and +`ART_MEGATRON_EXPERT_TENSOR_PARALLEL_SIZE=1` for the MoE case. + +The [SkyPilot task](trainer_rank_recompute_memory.sky.yaml) uses free Kubernetes. +Final jobs: `art-915-tight-tp4-0917:6,7,8,9` and +`art-915-tight-tp8-0917:5,7`; TP2 controls ran locally. Both clusters had autodown +and two-hour pod deadlines and were explicitly torn down after evidence +retrieval. No paid clusters were used. Raw JSONL is retained locally under +`scratch/recompute-memory-calibrated/`, with diagnostic and negative-control +runs in `scratch/recompute-memory-components/` and `scratch/recompute-memory-tight/`. + +The [previous report](https://github.com/OpenPipe/ART/blob/57f9de2f9/dev/trainer_rank_recompute_memory.md) +records the looser estimate and the earlier full-recompute underestimation. +This work does not validate that legacy path, larger LoRA ranks, CP, pretrained +routing distributions, arbitrary deep prefix trees, or other hardware/kernels. +These measurements support the calibrated native paths, not a universal memory +bound. diff --git a/dev/trainer_rank_recompute_memory.py b/dev/trainer_rank_recompute_memory.py new file mode 100644 index 000000000..c888b3766 --- /dev/null +++ b/dev/trainer_rank_recompute_memory.py @@ -0,0 +1,315 @@ +"""Measure cold unsplit admission and native CUDA peaks for one recompute mode. + +Run with torchrun; use a fresh process per mode/topology. Random weights keep +the native model geometry and kernels without downloading a checkpoint. This +measures memory, not pretrained-model correctness. Like public admission, try +the minimum-memory layout before refusing an unsplit request; never split or +bypass the budget. Each repetition clears the learned memory profile, while +sample 0 includes cold compilation/autotuning. Output is one JSONL row per rank. +""" + +import argparse +from dataclasses import asdict +import gc +import hashlib +import json +import os +from pathlib import Path +import subprocess +import time + +from dotenv import load_dotenv +import torch +import torch.distributed as dist + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--model", default="Qwen/Qwen3.8-27B") + parser.add_argument("--mode", choices=("full", "selective", "none"), required=True) + parser.add_argument("--layers", type=int, default=0) + parser.add_argument("--tokens", type=int, nargs="+", default=[1024, 2048, 4096]) + parser.add_argument("--repeat", type=int, default=2) + parser.add_argument("--reported-pair", action="store_true") + parser.add_argument( + "--pairs", action="store_true", help="Two sequences at each token length" + ) + parser.add_argument( + "--sequences", type=int, default=0, help="Override sequence count per batch" + ) + parser.add_argument("--prefix-fraction", type=float, default=0.3) + parser.add_argument("--modules", nargs="+", default=["core_attn"]) + parser.add_argument( + "--components", action="store_true", help="Attribute eager layer allocations" + ) + parser.add_argument( + "--concentrate-routing", + action="store_true", + help="Concentrate 31/32 of tokens on the first expert rank", + ) + parser.add_argument("--evidence", type=Path, required=True) + args = parser.parse_args() + if args.repeat < 1 or args.layers < 0 or any(n < 1 for n in args.tokens): + parser.error("repeat/tokens must be positive and layers nonnegative") + if args.sequences < 0: + parser.error("sequences must be nonnegative") + if not 0 <= args.prefix_fraction < 1: + parser.error("prefix-fraction must be in [0, 1)") + load_dotenv(".env") + from trainer_rank_support import load_random_checkpoints + + from art.megatron import train + from art.trainer_rank import ForwardInput, TrainerRank + + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group("nccl") + + def configure(provider): + if args.layers: + provider.num_layers = args.layers + provider.recompute_granularity = None if args.mode == "none" else args.mode + provider.recompute_method = "uniform" if args.mode == "full" else None + provider.recompute_num_layers = 1 if args.mode == "full" else None + provider.recompute_modules = args.modules if args.mode == "selective" else [] + + def emit(row): + gathered = [None] * dist.get_world_size() + dist.all_gather_object(gathered, {**facts, **row, "rank": dist.get_rank()}) + if dist.get_rank() == 0: + with args.evidence.open("a") as stream: + for item in gathered: + stream.write(json.dumps(item, default=str, sort_keys=True) + "\n") + print("MEMORY_CALIBRATION " + json.dumps(gathered, default=str), flush=True) + + try: + torch.manual_seed(913) + runtime = train.build_training_runtime( + model_identifier=args.model, + model_initialization="random", + provider_configure=configure, + print_env=False, + ) + if args.concentrate_routing: + # Stress dispatch imbalance without replacing native routing/experts. + for module in runtime.model[0].modules(): + if type(module).__name__ == "TopKRouter": + + def gating(inputs, original=module.gating): + logits = original(inputs) + experts = logits.shape[-1] + ep = runtime.provider.expert_model_parallel_size + rows = torch.arange( + logits.numel() // experts, device=logits.device + ).reshape(*logits.shape[:-1], 1) + # Fully empty EP peers crash TE's grouped GEMM. Leave + # 1/32 of tokens on peers while stressing near-max load. + owner = ( + torch.where( + rows % 32 == 0, 1 + (rows // 32) % max(1, ep - 1), 0 + ) + if ep > 1 + else torch.zeros_like(rows) + ) + expert = torch.arange(experts, device=logits.device) + selected = (expert >= owner * (experts // ep)) & ( + expert + < owner * (experts // ep) + runtime.provider.moe_router_topk + ) + return logits + selected * 10000 + + module.gating = gating + for chunk in runtime.model: + chunk.train() + if args.components and runtime.transformer_layers_compiled: + parser.error("--components requires ART_DISABLE_MEGATRON_COMPILE=1") + rank = TrainerRank(runtime) + [slot] = load_random_checkpoints( + runtime, rank, 1, base_model=args.model, lora_rank=1 + ) + facts = { + "schema": "art.dev.recompute_memory.v1", + "driver_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(), + "source_sha": os.environ.get("ART_CALIBRATION_SOURCE_SHA") + or subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(), + "model": args.model, + "initialization": "random", + "concentrated_routing": args.concentrate_routing, + "lora_rank": 1, + "mode": rank._recompute_granularity, + "recompute_method": runtime.provider.recompute_method, + "recompute_num_layers": runtime.provider.recompute_num_layers, + "recompute_modules": runtime.provider.recompute_modules, + "geometry": rank._geometry.as_dict(), + "layers": rank._num_layers, + "gdn_layers": rank._gdn_layers, + "sequence_parallel": rank._sequence_parallel, + "topology_dp_tp_cp_pp": rank._topology_key(), + "parallel_shape": asdict(rank._parallel_shape), + "dtype": str(next(runtime.model[0].parameters()).dtype), + "device": torch.cuda.get_device_name(), + "device_total_bytes": torch.cuda.get_device_properties(0).total_memory, + "torch": torch.__version__, + "cuda": torch.version.cuda, + "transformer_layers_compiled": runtime.transformer_layers_compiled, + "activation_config": { + name: str(getattr(runtime.provider, name, None)) + for name in ( + "bias_activation_fusion", + "use_te_activation_func", + "gated_linear_unit", + "activation_func", + "attention_output_gate", + "qk_layernorm", + ) + }, + } + args.evidence.parent.mkdir(parents=True, exist_ok=True) + count = args.sequences or (2 if args.pairs else 1) + workloads = [ + ([length] * count, int(length * args.prefix_fraction) if count > 1 else 0) + for length in args.tokens + ] + if args.reported_pair: + # Same logical/fully-shared packed counts as #913; synthetic IDs. + workloads.append(([19221, 19222], 5733)) + generator = torch.Generator().manual_seed(913) + for lengths, prefix in workloads: + tokens = [ + torch.randint(100, 10000, (n,), generator=generator) for n in lengths + ] + for item in tokens[1:]: + item[:prefix] = tokens[0][:prefix] + requests = [ + ForwardInput(input_tokens=item, hidden_states=True) for item in tokens + ] + plan = rank._plan_flat_forward(requests, checkpoint=slot) + memory_minimal = False + for sample in range(args.repeat): + rank.zero_grad() + rank._memory_profiles.clear() + gc.collect() + torch.cuda.empty_cache() + dist.barrier() + torch.cuda.synchronize() + check = rank._memory_check(plan) + if not check.fits and not memory_minimal: + plan = rank._plan_flat_forward( + requests, checkpoint=slot, memory_minimal=True + ) + memory_minimal = True + check = rank._memory_check(plan) + row = { + "lengths": lengths, + "shared_prefix": prefix, + "logical_tokens": plan.logical_tokens, + "packed_tokens": plan.packed_tokens, + "grad_segment_count": plan.grad_segment_count, + "output_bytes": plan.output_bytes, + "selected_max_depth": plan.selected_max_depth, + "memory_minimal": memory_minimal, + "sample": sample, + "admission": asdict(check), + } + if not check.fits: + emit({**row, "status": "refused"}) + break + baseline = torch.cuda.memory_allocated() + reserved = torch.cuda.memory_reserved() + torch.cuda.reset_peak_memory_stats() + started = time.monotonic() + components, handles, entries = [], [], {} + if args.components: + from megatron.core.transformer.transformer_layer import ( + TransformerLayer, + ) + + def enter(module, inputs): + entries[id(module)] = torch.cuda.memory_allocated() + + def leave(name): + def record(module, inputs, output): + components.append( + { + "name": name, + "type": type(module).__name__, + "input_shape": list(inputs[0].shape) + if inputs and isinstance(inputs[0], torch.Tensor) + else None, + "retained_delta_bytes": torch.cuda.memory_allocated() + - entries[id(module)], + } + ) + + return record + + for name, layer in runtime.model[0].named_modules(): + if isinstance(layer, TransformerLayer): + for part, module in ( + ("layer", layer), + ("attention", layer.self_attention), + ("mlp", layer.mlp), + ): + handles.extend( + ( + module.register_forward_pre_hook(enter), + module.register_forward_hook( + leave(f"{name}.{part}") + ), + ) + ) + # Direct execution isolates the selected unsplit plan. It uses + # the same native forward as public admission, with no split. + outputs = rank._execute_flat_plan(plan) + for handle in handles: + handle.remove() + torch.cuda.synchronize() + forward_peak = torch.cuda.max_memory_allocated() + retained = torch.cuda.memory_allocated() + forward_seconds = time.monotonic() - started + terms = [ + output.hidden_states.float().square().mean() + for output in outputs + if output.hidden_states is not None + ] + assert len(terms) == len(requests) + loss = torch.stack(terms).sum() + loss.backward() + torch.cuda.synchronize() + backward_peak = torch.cuda.max_memory_allocated() + backward_seconds = time.monotonic() - started + gradients = [ + p.grad + for p in rank._checkpoint_slots[slot].params + if p.grad is not None + ] + emit( + { + **row, + "status": "measured", + "components": components, + "baseline_allocated_bytes": baseline, + "baseline_reserved_bytes": reserved, + "forward_peak_delta_bytes": forward_peak - baseline, + "retained_delta_bytes": retained - baseline, + "forward_backward_peak_delta_bytes": backward_peak - baseline, + "peak_reserved_bytes": torch.cuda.max_memory_reserved(), + "forward_seconds": forward_seconds, + "forward_backward_seconds": backward_seconds, + "finite_loss": bool(torch.isfinite(loss).item()), + "gradient_tensors": len(gradients), + "finite_gradients": bool(gradients) + and all( + bool(torch.isfinite(g).all().item()) for g in gradients + ), + "estimate_covers_forward": check.estimated_required_bytes + >= forward_peak - baseline, + } + ) + del outputs, loss, terms, gradients + rank.zero_grad() + finally: + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/dev/trainer_rank_recompute_memory.sky.yaml b/dev/trainer_rank_recompute_memory.sky.yaml new file mode 100644 index 000000000..ad5b95384 --- /dev/null +++ b/dev/trainer_rank_recompute_memory.sky.yaml @@ -0,0 +1,56 @@ +# Launch with --infra k8s/cks-wb3 --idle-minutes-to-autostop 10 --down. +# Set ART_CALIBRATION_SOURCE_SHA to the commit synced in workdir. +name: trainer-rank-recompute-memory +workdir: . + +resources: + infra: k8s/cks-wb3 + accelerators: H200:4 + cpus: 32+ + memory: 256+ + image_id: docker:docker.io/bradhiltonnw/art-gpu:latest + +envs: + ART_CALIBRATION_SOURCE_SHA: unset + ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE: "4" + ART_MEGATRON_CONTEXT_PARALLEL_SIZE: "1" + ART_MEGATRON_DATA_PARALLEL_SIZE: "1" + ART_MEGATRON_PIPELINE_MODEL_PARALLEL_SIZE: "1" + OMP_NUM_THREADS: "1" + OPENBLAS_NUM_THREADS: "1" + MKL_NUM_THREADS: "1" + PYTHONUNBUFFERED: "1" + TOKENIZERS_PARALLELISM: "false" + MODEL: Qwen/Qwen3.8-27B + MODES: selective none + MODULES: core_attn + CALIBRATION_ARGS: --pairs --tokens 2048 4096 8192 --reported-pair + EVIDENCE_DIR: scratch/recompute-memory-sharded + +setup: | + INSTALL_VLLM_RUNTIME=false bash src/art/megatron/setup.sh + +run: | + set -euo pipefail + export PYTHONPATH="$PWD/src:$PWD" + for mode in ${MODES}; do + timeout --signal=TERM --kill-after=30s 20m \ + megatron_runtime/.venv/bin/python -m torch.distributed.run \ + --standalone --nproc-per-node="$ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE" \ + dev/trainer_rank_recompute_memory.py --model "$MODEL" --mode "$mode" \ + --modules ${MODULES} ${CALIBRATION_ARGS} \ + --evidence "$EVIDENCE_DIR/tp$ART_MEGATRON_TENSOR_MODEL_PARALLEL_SIZE-$mode.jsonl" + done + +config: + kubernetes: + pod_config: + spec: + schedulerName: binpack-scheduler + activeDeadlineSeconds: 7200 + containers: + - name: ray-node + imagePullPolicy: Always + env: + - name: UV_LINK_MODE + value: copy diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 8a0e9f038..8b0603348 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -941,6 +941,12 @@ def active_logical_tokens(self) -> int: # Keep total-input telemetry while pricing only executed requests. return self.logical_tokens - self.inactive_logical_tokens + @property + def grad_segment_count(self) -> int: + return sum( + len(group.packed.segments) for group in self.groups if group.grad_enabled + ) + @property def subforward_count(self) -> int: return 1 @@ -1357,21 +1363,13 @@ def __init__(self, runtime: TrainingRuntime) -> None: "therefore requires PP=1 with exactly one local model chunk; " f"got pp={pp_size}, chunks={len(runtime.model)}" ) - if getattr(runtime.provider, "recompute_granularity", None) == "selective": - raise TrainerRankRuntimeSupportError( - "TrainerRank memory planning does not support selective recompute; " - "its activation estimate assumes full recompute. Use " - "ART_MEGATRON_RECOMPUTE_GRANULARITY=full with " - "ART_MEGATRON_RECOMPUTE_METHOD=uniform and " - "ART_MEGATRON_RECOMPUTE_NUM_LAYERS=1." - ) # Tensor parallelism is admitted: the vocab-parallel head, sequence- # parallel gather, TP padding of packed batches and sharded LoRA # gradient reduction pre-date the planner, memory checks all-reduce # within the TP x CP group, the memory profile is keyed by topology so # TP calibrates itself online, and the fitted layout cost model prices - # TP explicitly. Known limitation: the cold static memory estimate - # ignores sharding (conservative). + # TP explicitly. The cold retained-activation floor also distinguishes + # tensor/sequence-parallel storage from gathered LoRA inputs. self.runtime: TrainingRuntime = runtime self.device: torch.device = next(runtime.model[0].parameters()).device self._param_dtype_size = _dtype_size(next(runtime.model[0].parameters()).dtype) @@ -1388,6 +1386,28 @@ def __init__(self, runtime: TrainingRuntime) -> None: or getattr(runtime.provider, "num_layers", 1) or 1 ) + memory_config = getattr(metadata_model, "config", None) or runtime.provider + + def memory_field(name: str, default: Any = None) -> Any: + return getattr( + memory_config, name, getattr(runtime.provider, name, default) + ) + + self._recompute_granularity = memory_field("recompute_granularity", None) + self._recompute_modules: frozenset[str] = frozenset( + memory_field("recompute_modules", ()) or () + ) + self._sequence_parallel = bool(memory_field("sequence_parallel", False)) + self._attention_output_gate = bool(memory_field("attention_output_gate", False)) + # Native fused SwiGLU retains gate/up and the output (3F). Eager + # unfused SwiGLU also retains SiLU and offset tensors (5F). Compilation + # may fall back, so only the native fusion setting earns this discount. + self._mlp_activation_factor = ( + 3 + if memory_field("bias_activation_fusion", False) + and not memory_field("use_te_activation_func", False) + else 5 + 2 * (memory_field("activation_func_clamp_value", None) is not None) + ) # Layers that run the gated-delta-net path (Qwen3.5-4B: 24 of 32); the # cost model prices GDN state hand-offs per GDN layer, not per layer. self._gdn_layers = _gdn_layer_count(runtime.model[0]) @@ -1408,6 +1428,10 @@ def __init__(self, runtime: TrainingRuntime) -> None: ) spec = getattr(runtime, "model_support_spec", None) self._moe_layers = _moe_layer_count(runtime.model[0]) + self._checkpointed_moe_layers = sum( + getattr(module, "moe_layer_recompute", False) is True + for module in runtime.model[0].modules() + ) is_moe = bool( self._moe_layers or getattr(spec, "is_moe", False) @@ -2909,6 +2933,7 @@ def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: output_bytes=plan.output_bytes, signature=plan.signature, logical_tokens=plan.active_logical_tokens, + gdn_segments=plan.grad_segment_count, ) def _subforward_cost( @@ -2918,12 +2943,14 @@ def _subforward_cost( output_bytes: int, signature: _MemorySignature, logical_tokens: int, + gdn_segments: int = 0, ) -> _SubforwardCost: required = self._estimate_required_memory_bytes_from_values( packed_tokens=packed_tokens, output_bytes=output_bytes, signature=signature, logical_tokens=logical_tokens, + gdn_segments=gdn_segments, ) retained = self._retained_memory_bytes( signature, @@ -3921,6 +3948,12 @@ def priced( output_bytes=output_bytes, signature=signature, logical_tokens=logical_tokens, + # A radix tree has fewer than twice as many segments as + # active requests; the exact plan uses its actual count. + gdn_segments=2 + * sum( + _request_mix_key(r) != "inactive" for r in local_requests + ), ) return ( self._memory_check_required(required, sync_across_dp=True), @@ -5062,6 +5095,7 @@ def _memory_check( output_bytes=forward.output_bytes, signature=forward.signature, logical_tokens=forward.active_logical_tokens, + gdn_segments=forward.grad_segment_count, ) return self._memory_check_required(required, sync_across_dp=sync_across_dp) @@ -5123,6 +5157,7 @@ def _estimate_required_memory_bytes_from_values( output_bytes: int, signature: _MemorySignature, logical_tokens: int | None = None, + gdn_segments: int = 0, ) -> int: if packed_tokens <= 0: return output_bytes @@ -5134,6 +5169,95 @@ def _estimate_required_memory_bytes_from_values( * self._param_dtype_size * activation_factor ) + if signature.grad_enabled and self._recompute_granularity != "full": + geometry = self._geometry + hidden = self._hidden_size + tp = max(1, self._topology_key()[1]) + sp = tp if self._sequence_parallel else 1 + # Gathered LoRA inputs alias norm output without sequence sharding. + gathered = hidden if sp > 1 else 0 + common = 2 * hidden / sp + gathered + attention_width = ( + geometry.num_attention_heads * geometry.kv_channels or hidden + ) + kv_width = geometry.num_query_groups * geometry.kv_channels or hidden + gated = self._attention_output_gate + attention = ( + common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp + ) + if 0 < geometry.num_query_groups < tp: + # SelfAttentionLinearQKVLoRA constructs global QKV before + # slicing it when KV groups cannot be partitioned across TP. + attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( + 1 - 1 / tp + ) + gdn = ( + common + + ( + 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim + ) + / tp + ) + ffn_width = geometry.ffn_hidden_size or 4 * hidden + mlp = common + self._mlp_activation_factor * ffn_width / tp + if geometry.moe_experts: + # Keep the worst-case dispatch envelope: random-weight runs + # cannot establish an EP discount for pretrained routing. + ffn_width = ( + geometry.moe_topk * geometry.moe_ffn_hidden_size + + geometry.moe_shared_expert_ffn + ) or ffn_width + mlp = ( + common + 6 * ffn_width + 2 * hidden * max(0, geometry.moe_topk - 1) + ) + gdn_layers = min(self._num_layers, self._gdn_layers) + retained_features = ( + (self._num_layers - gdn_layers) * attention + + gdn_layers * gdn + + self._num_layers * mlp + ) + if self._recompute_granularity == "selective": + checkpointed = ( + self._checkpointed_moe_layers + if geometry.moe_experts + else self._num_layers + if "mlp" in self._recompute_modules + else 0 + ) + # Checkpoints keep their input (and MoE's external norm); one + # live MLP still needs workspace, including worst-case dispatch. + checkpoint_input = (2 if geometry.moe_experts else 1) * hidden / sp + retained_features -= max(0, checkpointed - 1) * (mlp - checkpoint_input) + # Each GDN segment can retain an initial and a final recurrent + # state (fp32), plus convolution history. Unlike token activations, + # these do not shrink with segment length. + gdn_state_bytes = ( + 2 + * gdn_segments + * gdn_layers + / tp + * ( + 4 + * geometry.gdn_value_heads + * geometry.gdn_key_head_dim + * geometry.gdn_value_head_dim + + self._param_dtype_size + * ( + 2 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + geometry.gdn_value_heads * geometry.gdn_value_head_dim + ) + * max(0, geometry.gdn_conv_kernel - 1) + ) + ) + static_compute = max( + static_compute, + # Cold eager runs allocate ~58 MiB beyond warm retention for + # native kernel initialization; a slope alone misses short inputs. + 64 * 2**20 + + gdn_state_bytes + + packed_tokens * self._param_dtype_size * retained_features, + ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. static_compute = max( diff --git a/tests/unit/test_trainer_rank_active_memory.py b/tests/unit/test_trainer_rank_active_memory.py index 74ced9d63..d9c1d710a 100644 --- a/tests/unit/test_trainer_rank_active_memory.py +++ b/tests/unit/test_trainer_rank_active_memory.py @@ -34,7 +34,9 @@ def _rank(): SimpleNamespace( model=[_Model()], optimizer=None, - provider=SimpleNamespace(hidden_size=8, num_layers=4), + provider=SimpleNamespace( + hidden_size=8, num_layers=4, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ), ) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index e0dabad67..ccb1ed6a2 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -75,7 +75,9 @@ def _rank(layer=None): SimpleNamespace( model=[model], optimizer=None, - provider=SimpleNamespace(hidden_size=2048, num_layers=40), + provider=SimpleNamespace( + hidden_size=2048, num_layers=40, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ), ) diff --git a/tests/unit/test_trainer_rank_recompute_memory.py b/tests/unit/test_trainer_rank_recompute_memory.py new file mode 100644 index 000000000..56f03dd95 --- /dev/null +++ b/tests/unit/test_trainer_rank_recompute_memory.py @@ -0,0 +1,370 @@ +"""Cold admission checks; these CPU estimates are not native GPU measurements.""" + +from dataclasses import replace +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch +from torch.utils.checkpoint import checkpoint + +from art.trainer_rank import ForwardInput, TrainerRank, TrainerRankMemoryError +from art.trainer_rank._impl import _MemoryProfile + + +def _rank( + granularity: str | None = "full", + *, + dtype: torch.dtype = torch.bfloat16, + **geometry: Any, +) -> TrainerRank: + provider = SimpleNamespace( + **{ + "hidden_size": 5120, + "ffn_hidden_size": 17408, + "num_layers": 64, + "num_attention_heads": 24, + "num_query_groups": 4, + "kv_channels": 256, + "recompute_granularity": granularity, + **geometry, + } + ) + return TrainerRank( + cast( + Any, + SimpleNamespace( + model=[torch.nn.Linear(1, 1, dtype=dtype)], + provider=provider, + optimizer=None, + model_support_handler=SimpleNamespace( + build_gdn_execution_spec=bool( + geometry.get("linear_num_value_heads") + ), + is_moe=bool(geometry.get("num_moe_experts")), + ), + ), + ) + ) + + +def _plan(rank: TrainerRank, tokens: int = 32710, *, no_grad: bool = False): + request = ForwardInput( + input_tokens=torch.tensor([1, 2]), hidden_states=True, no_grad=no_grad + ) + return replace( + rank._plan_flat_forward([request]), + packed_tokens=tokens, + logical_tokens=tokens, + output_bytes=tokens * rank.hidden_size * rank._param_dtype_size, + ) + + +@pytest.mark.parametrize("tp", (2, 4)) +@pytest.mark.parametrize("granularity", (None, "selective")) +def test_reported_cold_request_is_refused_before_execution( + monkeypatch: pytest.MonkeyPatch, tp: int, granularity: str | None +) -> None: + rank = _rank(granularity) + monkeypatch.setattr(rank, "_topology_key", lambda: (1, tp, 1, 1)) + monkeypatch.setattr(rank, "_dp_rank_and_size", lambda: (0, 1)) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: int(119.289 * 2**30)) + monkeypatch.setattr( + rank, "_execute_flat_plan", lambda _: pytest.fail("unsafe forward admitted") + ) + assert not rank._memory_check(_plan(rank)).fits + assert rank._memory_check(_plan(rank, tokens=1024)).fits + with pytest.raises(TrainerRankMemoryError): + rank.dp_rank_forward( + [ForwardInput(input_tokens=torch.arange(32710), hidden_states=True)] + ) + + +def test_full_recompute_keeps_existing_estimate() -> None: + rank = _rank() + plan = _plan(rank) + assert rank._memory_check(plan).estimated_required_bytes == int( + (plan.output_bytes + 32710 * 5120 * 2 * 16) * 1.1 + ) + + +@pytest.mark.parametrize("granularity", (None, "selective")) +def test_no_grad_does_not_pay_for_retained_layers(granularity: str | None) -> None: + rank, full = _rank(granularity), _rank() + assert rank._memory_check(_plan(rank, no_grad=True)) == full._memory_check( + _plan(full, no_grad=True) + ) + assert rank._memory_check(_plan(rank)).estimated_required_bytes > ( + rank._memory_check(_plan(rank, no_grad=True)).estimated_required_bytes + ) + + +@pytest.mark.parametrize( + "geometry", + [ + {"num_layers": 128}, + {"ffn_hidden_size": 34816}, + {"num_attention_heads": 48}, + {"ffn_hidden_size": 0}, + { + "linear_num_key_heads": 16, + "linear_key_head_dim": 128, + "linear_num_value_heads": 48, + "linear_value_head_dim": 128, + }, + {"num_moe_experts": 64, "moe_router_topk": 8, "moe_ffn_hidden_size": 8192}, + ], +) +def test_retained_estimate_tracks_model_geometry(geometry: dict[str, Any]) -> None: + rank, base = _rank("selective", **geometry), _rank("selective") + assert rank._memory_check(_plan(rank)).estimated_required_bytes > ( + base._memory_check(_plan(base)).estimated_required_bytes + ) + + +def test_profile_cannot_erase_recompute_floor() -> None: + rank = _rank("selective") + plan = _plan(rank) + cold = rank._memory_check(plan).estimated_required_bytes + rank._memory_profiles[plan.signature] = _MemoryProfile(0.0, plan.packed_tokens) + assert rank._memory_check(plan).estimated_required_bytes == cold + rank._memory_profiles[plan.signature] = _MemoryProfile( + 2 * cold / plan.packed_tokens, plan.packed_tokens + ) + assert rank._memory_check(plan).estimated_required_bytes > cold + + +def _hybrid_rank(monkeypatch: pytest.MonkeyPatch, tp: int) -> TrainerRank: + rank = _rank( + "selective", + sequence_parallel=True, + bias_activation_fusion=True, + attention_output_gate=True, + linear_num_key_heads=16, + linear_key_head_dim=128, + linear_num_value_heads=48, + linear_value_head_dim=128, + ) + rank._gdn_layers = 48 + monkeypatch.setattr(rank, "_topology_key", lambda: (1, tp, 1, 1)) + return rank + + +@pytest.mark.parametrize( + "tp,peak", [(1, 33758228992), (2, 19541553664), (4, 11132523008)] +) +def test_sharded_floor_covers_recorded_native_gdn_peaks(monkeypatch, tp, peak): + # H200, 64-layer Qwen3.8-27B, LoRA r1, SP, cold/warm max, 2,048 tokens. + # First unsharded campaign, linked from dev/trainer_rank_recompute_memory.md. + rank = _hybrid_rank(monkeypatch, tp) + estimate = rank._memory_check(_plan(rank, tokens=2048)).estimated_required_bytes + assert peak <= estimate <= 1.2 * peak + + +@pytest.mark.parametrize( + "tp,mlp,packed,logical,peak_gib", + [ + (4, False, 4096, 4096, 20.613), + (4, False, 6964, 8192, 35.090), + (4, False, 13928, 16384, 69.924), + (8, False, 4096, 4096, 14.529), + (8, False, 6968, 8192, 24.654), + (8, False, 13928, 16384, 48.950), + (4, True, 4096, 4096, 11.213), + (4, True, 6964, 8192, 19.006), + (4, True, 13928, 16384, 37.881), + ], +) +def test_hybrid_estimate_is_close_to_recorded_peaks( + monkeypatch, tp, mlp, packed, logical, peak_gib +): + # Native H200 cold/warm witnesses in the calibration report. Coverage alone + # would allow the former 60%-high estimate; also check useful admission. + rank = _hybrid_rank(monkeypatch, tp) + rank._recompute_modules = frozenset(("core_attn", "mlp") if mlp else ("core_attn",)) + plan = replace(_plan(rank, tokens=packed), output_bytes=logical * 5120 * 2) + estimate = rank._memory_check(plan).estimated_required_bytes / 2**30 + assert peak_gib <= estimate <= 1.15 * peak_gib + + +def test_native_fusion_discount_is_independent_of_compilation(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + plan = _plan(rank, tokens=4096) + fused = rank._memory_check(plan).estimated_required_bytes + rank.runtime.transformer_layers_compiled = False + assert rank._memory_check(plan).estimated_required_bytes == fused + rank._mlp_activation_factor = 5 + assert rank._memory_check(plan).estimated_required_bytes > fused + + +def test_mlp_discount_requires_selective_and_keeps_one_live_workspace(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + rank._recompute_granularity = None + plan = _plan(rank) + unrecomputed = rank._memory_check(plan).estimated_required_bytes + rank._recompute_modules = frozenset(("mlp",)) + assert rank._memory_check(plan).estimated_required_bytes == unrecomputed + rank._recompute_granularity = "selective" + assert rank._memory_check(plan).estimated_required_bytes < unrecomputed + rank._num_layers = 1 + checkpointed = rank._memory_check(plan).estimated_required_bytes + rank._recompute_modules = frozenset() + assert rank._memory_check(plan).estimated_required_bytes == checkpointed + + +def test_tp4_admits_eight_k_sibling_pair_and_prices_actual_layer_mix(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + monkeypatch.setattr(rank, "_available_memory_bytes", lambda: 119 * 2**30) + plan = replace( + _plan(rank, tokens=13928), logical_tokens=16384, output_bytes=16384 * 5120 * 2 + ) + check = rank._memory_check(plan) + assert check.fits + rank._gdn_layers = 64 + assert ( + rank._memory_check(plan).estimated_required_bytes + > check.estimated_required_bytes + ) + rank._gdn_layers = 0 + assert ( + rank._memory_check(plan).estimated_required_bytes + < check.estimated_required_bytes + ) + + +def test_sp_discount_excludes_gathered_lora_inputs(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + plan = _plan(rank, tokens=2048) + tp4 = rank._memory_check(plan).estimated_required_bytes + rank._sequence_parallel = False + assert rank._memory_check(plan).estimated_required_bytes > tp4 + monkeypatch.setattr(rank, "_topology_key", lambda: (1, 1, 1, 1)) + assert tp4 > rank._memory_check(plan).estimated_required_bytes / 4 + + +@pytest.mark.parametrize( + "packed,logical,peak", + [ + (436, 512, 829413888), # Cold 256-token pair before adding workspace. + (1742, 2048, 3043651584), + (6964, 8192, 11941280768), + (13928, 16384, 23861987840), + (27854, 32768, 47563188736), + ], +) +def test_gathered_inputs_cover_attention_only_cold_peak( + monkeypatch, packed, logical, peak +): + # Native eager Qwen3-1.7B TP2 paired requests at d31f9423, max across ranks + # and cold/warm repetitions. Four FFN widths missed all four peaks. + rank = _rank( + "selective", + hidden_size=2048, + ffn_hidden_size=6144, + num_layers=28, + num_attention_heads=16, + num_query_groups=8, + kv_channels=128, + sequence_parallel=True, + ) + monkeypatch.setattr(rank, "_topology_key", lambda: (1, 2, 1, 1)) + plan = replace( + _plan(rank, tokens=packed), + logical_tokens=logical, + output_bytes=logical * 2048 * 2, + ) + assert rank._memory_check(plan).estimated_required_bytes >= peak + + +def test_non_full_estimate_covers_retained_gated_mlp_tensors() -> None: + # Selective core-attention recompute leaves this MLP graph live. Count + # distinct saved activation storage, excluding model parameters/views. + layers = [ + torch.nn.Sequential( + torch.nn.LayerNorm(16), + torch.nn.Linear(16, 256), + torch.nn.GLU(), + torch.nn.Linear(128, 16), + ) + for _ in range(8) + ] + parameters = { + parameter.untyped_storage().data_ptr() + for layer in layers + for parameter in layer.parameters() + } + + def retained(full: bool) -> int: + saved: dict[int, int] = {} + + def pack(tensor: torch.Tensor) -> torch.Tensor: + storage = tensor.untyped_storage() + if storage.data_ptr() not in parameters: + saved[storage.data_ptr()] = storage.nbytes() + return tensor + + with torch.autograd.graph.saved_tensors_hooks(pack, lambda tensor: tensor): + value = torch.zeros(32, 16, requires_grad=True) + for layer in layers: + value = ( + checkpoint(layer, value, use_reentrant=False) + if full + else layer(value) + ) + return sum(saved.values()) + + rank = _rank( + "selective", + dtype=torch.float32, + hidden_size=16, + ffn_hidden_size=128, + num_layers=8, + num_attention_heads=2, + kv_channels=8, + ) + assert ( + retained(True) + < retained(False) + <= rank._memory_check(_plan(rank, tokens=32)).estimated_required_bytes + ) + + +def test_moe_discount_uses_effective_native_checkpoint_count(monkeypatch): + rank = _rank( + "selective", + sequence_parallel=True, + num_moe_experts=256, + moe_router_topk=8, + moe_ffn_hidden_size=512, + moe_shared_expert_intermediate_size=512, + recompute_modules=["moe"], + ) + monkeypatch.setattr(rank, "_topology_key", lambda: (1, 4, 1, 1)) + plan = _plan(rank, tokens=4096) + # A module list alone cannot guarantee the native MoE checkpoint is active. + undiscounted = rank._memory_check(plan).estimated_required_bytes + rank._checkpointed_moe_layers = rank._num_layers + assert rank._memory_check(plan).estimated_required_bytes < undiscounted + rank._recompute_granularity = None + assert rank._memory_check(plan).estimated_required_bytes == undiscounted + + +def test_short_hybrid_pairs_pay_for_recurrent_states(monkeypatch): + rank = _hybrid_rank(monkeypatch, 4) + requests = [ + ForwardInput(input_tokens=torch.arange(64) + offset, hidden_states=True) + for offset in (0, 100) + ] + plan = rank._plan_flat_forward(requests) + assert plan.grad_segment_count == 2 + estimate = rank._memory_check(plan).estimated_required_bytes + # Cold eager TP4 at 45297a4af missed this peak without segment states. + assert 892974592 <= estimate <= 1.2 * 892974592 + assert rank._plan_cost(plan).required == estimate + rank._memory_profiles[plan.signature] = _MemoryProfile(0, plan.packed_tokens) + assert rank._memory_check(plan).estimated_required_bytes == estimate + inactive = rank._plan_flat_forward( + requests + [ForwardInput(input_tokens=torch.arange(1000))] + ) + assert inactive.grad_segment_count == 2 + assert rank._memory_check(inactive).estimated_required_bytes == estimate diff --git a/tests/unit/test_trainer_rank_split.py b/tests/unit/test_trainer_rank_split.py index ce2ad6cdb..7b8bb812f 100644 --- a/tests/unit/test_trainer_rank_split.py +++ b/tests/unit/test_trainer_rank_split.py @@ -83,7 +83,9 @@ def _runtime() -> "TrainingRuntime": return SimpleNamespace( model=[_FakeGPT()], optimizer=None, - provider=SimpleNamespace(hidden_size=8, num_layers=4), + provider=SimpleNamespace( + hidden_size=8, num_layers=4, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ) # type: ignore diff --git a/tests/unit/test_trainer_rank_split_peak.py b/tests/unit/test_trainer_rank_split_peak.py index 7eb6cab04..54473e12e 100644 --- a/tests/unit/test_trainer_rank_split_peak.py +++ b/tests/unit/test_trainer_rank_split_peak.py @@ -31,7 +31,9 @@ def _rank(): SimpleNamespace( model=[_Model()], optimizer=None, - provider=SimpleNamespace(hidden_size=8, num_layers=4), + provider=SimpleNamespace( + hidden_size=8, num_layers=4, recompute_granularity="full" + ), model_support_handler=SimpleNamespace(build_gdn_execution_spec=False), ), ) diff --git a/tests/unit/test_trainer_rank_topology.py b/tests/unit/test_trainer_rank_topology.py index e706b36a7..45ec291dd 100644 --- a/tests/unit/test_trainer_rank_topology.py +++ b/tests/unit/test_trainer_rank_topology.py @@ -65,14 +65,7 @@ def test_trainer_rank_recompute_support( ) -> None: runtime = _runtime(tp=tp) runtime.provider.recompute_granularity = granularity - if granularity == "selective": - with pytest.raises( - TrainerRankRuntimeSupportError, - match="selective recompute.*ART_MEGATRON_RECOMPUTE_GRANULARITY=full", - ): - TrainerRank(runtime) - else: - TrainerRank(runtime) + TrainerRank(runtime) def test_trainer_rank_still_refuses_pipeline_parallel_runtimes() -> None: