11from __future__ import annotations
22
3+ import inspect
34import sys
45import types
6+ from collections .abc import Callable
57from typing import Any
68
79from nemotron .recipes .super3 .stage1_sft import packed_compat_step
@@ -47,17 +49,11 @@ def forward(
4749 return tokens , packed_seq_params
4850
4951
50- def _install_stub_gpt_step () -> None :
52+ def _install_stub_gpt_step (forward_step : Callable [..., Any ] ) -> None :
5153 for parent in ("megatron" , "megatron.bridge" , "megatron.bridge.training" ):
5254 sys .modules .setdefault (parent , types .ModuleType (parent ))
5355
5456 gpt_step = types .ModuleType ("megatron.bridge.training.gpt_step" )
55-
56- def forward_step (data_iterator : Any , model : Any ) -> tuple [Any , Any ]:
57- del data_iterator
58- output = model (tokens = "batch" , packed_seq_params = {"packed" : True })
59- return output , lambda output_tensor : output_tensor
60-
6157 gpt_step .forward_step = forward_step
6258 sys .modules ["megatron.bridge.training.gpt_step" ] = gpt_step
6359
@@ -72,12 +68,31 @@ def _remove_stub_gpt_step() -> None:
7268 sys .modules .pop (module_name , None )
7369
7470
75- def test_packed_compat_step_drops_packed_seq_params_for_mamba_like_leaf () -> None :
71+ def _two_arg_upstream (data_iterator : Any , model : Any ) -> tuple [Any , Any ]:
72+ del data_iterator
73+ output = model (tokens = "batch" , packed_seq_params = {"packed" : True })
74+ return output , lambda output_tensor : output_tensor
75+
76+
77+ def test_packed_compat_step_signature_supports_state_aware_bridge_arity () -> None :
78+ signature = inspect .signature (packed_compat_step .forward_step )
79+
80+ assert list (signature .parameters ) == [
81+ "state_or_data_iterator" ,
82+ "data_iterator_or_model" ,
83+ "model" ,
84+ "return_schedule_plan" ,
85+ ]
86+ assert signature .parameters ["model" ].default is None
87+ assert signature .parameters ["return_schedule_plan" ].default is False
88+
89+
90+ def test_packed_compat_step_keeps_existing_two_arg_stub_behavior () -> None :
7691 leaf = _MambaLikeLeaf ()
7792 model = _ForwardingWrapper (_ForwardingWrapper (leaf ))
7893 assert not packed_compat_step ._model_forward_accepts_kwarg (model )
7994
80- _install_stub_gpt_step ()
95+ _install_stub_gpt_step (_two_arg_upstream )
8196 try :
8297 output , _loss = packed_compat_step .forward_step (iter (()), model )
8398 finally :
@@ -87,16 +102,79 @@ def test_packed_compat_step_drops_packed_seq_params_for_mamba_like_leaf() -> Non
87102 assert leaf .calls == [{"tokens" : "batch" , "position_ids" : None }]
88103
89104
90- def test_packed_compat_step_preserves_packed_seq_params_for_supported_leaf () -> None :
105+ def test_packed_compat_step_drops_packed_seq_params_for_state_aware_mamba_like_leaf () -> None :
106+ upstream_calls : list [dict [str , Any ]] = []
107+
108+ def state_aware_upstream (
109+ state : Any ,
110+ data_iterator : Any ,
111+ model : Any ,
112+ return_schedule_plan : bool = False ,
113+ ) -> tuple [Any , Any ]:
114+ upstream_calls .append (
115+ {
116+ "state" : state ,
117+ "data_iterator" : data_iterator ,
118+ "return_schedule_plan" : return_schedule_plan ,
119+ }
120+ )
121+ output = model (tokens = "batch" , packed_seq_params = {"packed" : True })
122+ return output , lambda output_tensor : output_tensor
123+
124+ leaf = _MambaLikeLeaf ()
125+ model = _ForwardingWrapper (_ForwardingWrapper (leaf ))
126+ data_iterator = iter (())
127+
128+ _install_stub_gpt_step (state_aware_upstream )
129+ try :
130+ output , _loss = packed_compat_step .forward_step ("state" , data_iterator , model )
131+ finally :
132+ _remove_stub_gpt_step ()
133+
134+ assert output == ("batch" , None )
135+ assert leaf .calls == [{"tokens" : "batch" , "position_ids" : None }]
136+ assert upstream_calls == [
137+ {"state" : "state" , "data_iterator" : data_iterator , "return_schedule_plan" : False }
138+ ]
139+
140+
141+ def test_packed_compat_step_preserves_packed_seq_params_for_state_aware_supported_leaf () -> None :
142+ upstream_calls : list [dict [str , Any ]] = []
143+
144+ def state_aware_upstream (
145+ state : Any ,
146+ data_iterator : Any ,
147+ model : Any ,
148+ return_schedule_plan : bool = False ,
149+ ) -> tuple [Any , Any ]:
150+ upstream_calls .append (
151+ {
152+ "state" : state ,
153+ "data_iterator" : data_iterator ,
154+ "return_schedule_plan" : return_schedule_plan ,
155+ }
156+ )
157+ output = model (tokens = "batch" , packed_seq_params = {"packed" : True })
158+ return output , lambda output_tensor : output_tensor
159+
91160 leaf = _PackedAwareLeaf ()
92161 model = _ForwardingWrapper (_ForwardingWrapper (leaf ))
93162 assert packed_compat_step ._model_forward_accepts_kwarg (model )
163+ data_iterator = iter (())
94164
95- _install_stub_gpt_step ()
165+ _install_stub_gpt_step (state_aware_upstream )
96166 try :
97- output , _loss = packed_compat_step .forward_step (iter (()), model )
167+ output , _loss = packed_compat_step .forward_step (
168+ "state" ,
169+ data_iterator ,
170+ model ,
171+ return_schedule_plan = True ,
172+ )
98173 finally :
99174 _remove_stub_gpt_step ()
100175
101176 assert output == ("batch" , {"packed" : True })
102177 assert leaf .calls == [{"tokens" : "batch" , "packed_seq_params" : {"packed" : True }}]
178+ assert upstream_calls == [
179+ {"state" : "state" , "data_iterator" : data_iterator , "return_schedule_plan" : True }
180+ ]
0 commit comments