From ce7f12fc674b289cff89d938ba85b1aac13907c7 Mon Sep 17 00:00:00 2001 From: Andrew White Date: Mon, 27 Jul 2026 07:38:40 -0500 Subject: [PATCH] fix(network): reset_parameters checks nonexistent (+2 more) Signed-off-by: Julius Berner --- fastgen/networks/cosmos_predict2/network.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/fastgen/networks/cosmos_predict2/network.py b/fastgen/networks/cosmos_predict2/network.py index b6215f7..6e9aa3a 100644 --- a/fastgen/networks/cosmos_predict2/network.py +++ b/fastgen/networks/cosmos_predict2/network.py @@ -431,6 +431,11 @@ def forward( - [output, features]: If feature_indices is non-empty - (output, logvar): If return_logvar=True """ + if feature_indices is None: + feature_indices = set() + if return_features_early and len(feature_indices) == 0: + return [] + x_B_T_H_W_D, rope_emb_L_1_1_D, extra_pos_emb = self.prepare_embedded_sequence( x_B_C_T_H_W, fps=fps, @@ -1007,8 +1012,8 @@ def reset_parameters(self): nn.init.zeros_(m.bias) # Reinitialize RoPE buffers (non-persistent buffers must be recomputed) - if hasattr(self.transformer, "rope_embedder") and self.transformer.rope_embedder is not None: - rope = self.transformer.rope_embedder + if hasattr(self.transformer, "pos_embedder") and self.transformer.pos_embedder is not None: + rope = self.transformer.pos_embedder device = next(rope.buffers()).device # Recompute the sequence buffer @@ -1327,7 +1332,7 @@ def forward( - ([output, features], logvar): If both features and logvar requested """ if feature_indices is None: - feature_indices = {} + feature_indices = set() if return_features_early and len(feature_indices) == 0: # Exit immediately if user requested this. return []