Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 8 additions & 13 deletions src/gvp_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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)

Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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

Expand Down
18 changes: 11 additions & 7 deletions tests/test_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand All @@ -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)
Expand All @@ -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

Expand Down
Loading