diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 82000f6..886f800 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -34,13 +34,6 @@ type EdgeAttr = GVPTuple """Edge (scalar, vector) attributes.""" -# Forward return type: either pooled residue embeddings or full GVP features -type ForwardOutput = tuple[NodeFeatures | torch.Tensor, EdgeAttr | None] -"""Return type of forward(): -- When pool_residue=True: (residue_embed: Tensor, None) -- When pool_residue=False: (node_features: GVPTuple, edge_attr: GVPTuple | None) -""" - def edge_vectors( pos: torch.Tensor, edge_index: torch.Tensor @@ -327,7 +320,7 @@ def _compute_edge_attr(self, data: Batch): V_edge = u.unsqueeze(1) return (s_edge, V_edge), s_edge_raw - def forward(self, data: Batch) -> ForwardOutput: + def forward(self, data: Batch) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Forward pass through the GVP encoder. @@ -341,8 +334,9 @@ def forward(self, data: Batch) -> ForwardOutput: - edge_unit_vectors: (E, 3) unit edge vectors Returns: - x: tuple (s, V) of node scalar and vector features - edge_attr: tuple (s_edge, V_edge) of edge features, or None if use_edge_update=False + s: (N, scalar_dim) scalar node features + V: (N, vector_dim, 3) vector node features (empty tensor if pooling) + edge_attr: tuple (s_edge, V_edge) of edge features, or None if pooling or edge updates disabled """ x_scalar = self.input_scalar_encoder(data.x) @@ -376,11 +370,12 @@ def forward(self, data: Batch) -> ForwardOutput: res_embed = self._pool_by_residue( atom_embed, data.residue_index, int(data.num_residues) ) - return res_embed, None # No edge features when pooling + # Return empty V tensor (N, 0, 3) to match 3-tuple signature + return res_embed, res_embed.new_empty(res_embed.size(0), 0, 3), None # Return edge_attr only if edge_update was used final_edge_attr = edge_attr if self.edge_update is not None else None - return x, final_edge_attr + return x[0], x[1], final_edge_attr def load_encoder_from_checkpoint( @@ -527,7 +522,7 @@ def forward( enc_data = make_gvp_encoder_data(data) with torch.set_grad_enabled(not self._freeze): - (s, V), edge_attr = self.encoder(enc_data) + s, V, edge_attr = self.encoder(enc_data) return s, V, edge_attr diff --git a/tests/test_encoder.py b/tests/test_encoder.py index 4925384..1df6276 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -257,17 +257,19 @@ def test_encoder_initialization(self, simple_encoder): def test_encoder_forward_with_pooling( self, simple_encoder, sample_homogeneous_data ): - """Test forward pass with residue pooling.""" - output, edge_attr = simple_encoder(sample_homogeneous_data) - assert output.shape == ( + """Test forward pass with residue pooling returns (s, V, edge_attr).""" + s, V, edge_attr = simple_encoder(sample_homogeneous_data) + assert s.shape == ( sample_homogeneous_data.num_residues, simple_encoder.pooled_dim, ) + # Pooling mode returns empty V tensor (N, 0, 3) to match flat 3-tuple signature + assert V.shape == (sample_homogeneous_data.num_residues, 0, 3) # Pooling mode returns None for edge features assert edge_attr is None def test_encoder_forward_no_pooling(self, sample_homogeneous_data): - """Test encoder without residue pooling returns ((s, V), edge_attr) tuple.""" + """Test encoder without residue pooling returns (s, V, edge_attr) tuple.""" encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), @@ -277,7 +279,7 @@ def test_encoder_forward_no_pooling(self, sample_homogeneous_data): num_edge_rbf=16, ) - (s, v), edge_attr = encoder(sample_homogeneous_data) + s, v, edge_attr = encoder(sample_homogeneous_data) assert s.shape == (sample_homogeneous_data.num_nodes, 64) assert v.shape == (sample_homogeneous_data.num_nodes, 16, 3) @@ -302,8 +304,10 @@ def test_encoder_empty_graph(self, simple_encoder): num_residues=0, ) - output, edge_attr = simple_encoder(data) - assert output.shape == (0, simple_encoder.pooled_dim) + s, V, edge_attr = simple_encoder(data) + assert s.shape == (0, simple_encoder.pooled_dim) + # Pooling mode returns empty V tensor (N, 0, 3) to match flat 3-tuple signature + assert V.shape == (0, 0, 3) # Pooling mode returns None for edge features assert edge_attr is None