@@ -918,8 +918,11 @@ def _ported_forward_pass(
918918 logits = self .unembed (residual )
919919 loss = self ._calculate_loss (logits , tokens , loss_per_token )
920920 return logits , loss
921+ elif return_type is None :
922+ # Return None when explicitly requested
923+ return None
921924 else :
922- # Return final residual
925+ # Return final residual for any other return_type
923926 return residual
924927
925928 def _calculate_loss (self , logits , tokens , loss_per_token = False ):
@@ -1678,6 +1681,9 @@ def _handle_return_type(self, logits, input_ids, return_type, loss_per_token):
16781681 shift_logits = logits [:, :- 1 , :].contiguous ()
16791682 loss = F .cross_entropy (shift_logits .view (- 1 , shift_logits .size (- 1 )), labels .view (- 1 ))
16801683 return loss , logits
1684+ elif return_type is None :
1685+ # Return None when explicitly requested
1686+ return None
16811687 else :
16821688 return logits
16831689
@@ -2237,6 +2243,9 @@ def _true_hf_format_forward_pass(
22372243 shift_logits .view (- 1 , shift_logits .size (- 1 )), targets .view (- 1 ), reduction = "mean"
22382244 )
22392245 return (loss , logits )
2246+ elif return_type is None :
2247+ # Return None when explicitly requested
2248+ return None
22402249 else :
22412250 return logits
22422251
@@ -2421,6 +2430,9 @@ def _hf_format_forward_pass(
24212430 shift_logits .view (- 1 , shift_logits .size (- 1 )), targets .view (- 1 ), reduction = "mean"
24222431 )
24232432 return (loss , logits )
2433+ elif return_type is None :
2434+ # Return None when explicitly requested
2435+ return None
24242436 else :
24252437 return logits
24262438
@@ -3719,7 +3731,8 @@ def forward(
37193731 loss = self .loss_fn (logits , input_ids , per_token = loss_per_token )
37203732 return logits , loss
37213733 elif return_type is None :
3722- return output
3734+ # Return None when explicitly requested (don't return output/logits)
3735+ return None
37233736 else :
37243737 raise ValueError (f"Invalid return_type: { return_type } " )
37253738
0 commit comments