Task 1: Core Aggregator Component
Priority: Must be done first — all other tasks depend on this.
Create: Tree_Matching_Networks/GMN/pretrained_transformer_aggregator.py
Reference: Tree_Matching_Networks/GMN/transformer_tree_aggregator.py (existing aggregator to follow)
What This Does
Drop-in replacement for TransformerTreeAggregator that wraps a HuggingFace pre-trained transformer encoder. Must implement the exact same forward interface:
forward(node_states, graph_idx, n_graphs, from_idx=None, to_idx=None) → [n_graphs, graph_rep_dim]
Architecture
node_states [n_total_nodes, node_state_dim]
→ input_projection (Linear + LayerNorm): node_state_dim → hf_hidden_dim
→ + TreeShapePositionalEncoder (reuse existing, at hf_hidden_dim)
→ _group_and_pad (+ optional virtual CLS prepend)
→ HF encoder layers (pre-trained, from AutoModel.from_pretrained())
→ aggregation (mean pooling or CLS extraction)
→ output_norm (LayerNorm) + output_projection: hf_hidden_dim → graph_rep_dim
Key Implementation Details
-
Extract encoder from HF model — different families store it differently:
- BERT/RoBERTa/MiniLM:
model.encoder
- DistilBERT:
model.transformer
- Write a helper function
extract_encoder_from_hf_model() that tries both
-
HF attention mask format differs from PyTorch's nn.TransformerEncoder:
- PyTorch: boolean mask
- HF BERT-style: extended mask
[batch, 1, 1, seq_len] with large negative values for padding
- DistilBERT: simple float mask
[batch, seq_len]
- Handle both in forward()
-
Reuse grouping/padding logic from TransformerTreeAggregator:
_group_and_pad(): groups nodes by graph_idx, pads to max_nodes
_group_and_pad_with_cls(): same but prepends virtual CLS token
_aggregate_nodes(): mean pooling over non-padded positions
_extract_root_embeddings(): for root-as-CLS mode
- Copy these methods (they're short and well-tested). Refactor into shared base later if desired.
-
TreeShapePositionalEncoder — reuse existing class at hf_hidden_dim. Always newly initialized (trainable). This is the tree-structural information that distinguishes this from standard BERT.
-
Freeze/unfreeze — freeze_transformer() sets requires_grad=False on all encoder params. Input projection, pos encoder, output projection remain trainable.
-
get_parameter_groups(base_lr, pretrained_lr_scale) — returns list of param group dicts with differential learning rates. Encoder params get base_lr * pretrained_lr_scale, everything else gets base_lr. Frozen params (requires_grad=False) are excluded.
Constructor Signature
def __init__(self,
node_state_dim, # e.g. 1280 (from GNN) or 804 (raw features)
graph_rep_dim, # e.g. 2048
hf_model_name, # e.g. "sentence-transformers/all-MiniLM-L6-v2"
max_nodes=64,
positional_features=None,
positional_max_values=None,
use_cls_token=False,
cls_token_type="virtual",
freeze_transformer=False):
Verification
import torch
from Tree_Matching_Networks.GMN.pretrained_transformer_aggregator import PretrainedTransformerAggregator
agg = PretrainedTransformerAggregator(
node_state_dim=1280, graph_rep_dim=2048,
hf_model_name='sentence-transformers/all-MiniLM-L6-v2', max_nodes=64
)
node_states = torch.randn(33, 1280)
graph_idx = torch.cat([torch.zeros(10), torch.ones(15), torch.full((8,), 2)]).long()
from_idx = torch.randint(0, 33, (40,))
to_idx = torch.randint(0, 33, (40,))
out = agg(node_states, graph_idx, 3, from_idx=from_idx, to_idx=to_idx)
assert out.shape == (3, 2048)
# Test freeze
agg.freeze_transformer()
assert all(not p.requires_grad for p in agg.encoder.parameters())
assert all(p.requires_grad for p in agg.input_projection.parameters())
Where to Find Things
TransformerTreeAggregator: Tree_Matching_Networks/GMN/transformer_tree_aggregator.py — copy _group_and_pad, _group_and_pad_with_cls, _aggregate_nodes, _extract_root_embeddings, _compute_cls_positional_encoding
TreeShapePositionalEncoder: Tree_Matching_Networks/GMN/tree_shape_positional_encoder.py — instantiate with embed_dim=hf_hidden_dim
GraphEmbeddingNet: Tree_Matching_Networks/GMN/graphembeddingnetwork.py:586 — calls aggregator as self._aggregator(node_states, graph_idx, n_graphs, from_idx=from_idx, to_idx=to_idx)
See docs/plans/2026-03-09-pretrained-transformer-aggregation.md Task 1 for complete code.
Task 1: Core Aggregator Component
Priority: Must be done first — all other tasks depend on this.
Create:
Tree_Matching_Networks/GMN/pretrained_transformer_aggregator.pyReference:
Tree_Matching_Networks/GMN/transformer_tree_aggregator.py(existing aggregator to follow)What This Does
Drop-in replacement for
TransformerTreeAggregatorthat wraps a HuggingFace pre-trained transformer encoder. Must implement the exact same forward interface:Architecture
Key Implementation Details
Extract encoder from HF model — different families store it differently:
model.encodermodel.transformerextract_encoder_from_hf_model()that tries bothHF attention mask format differs from PyTorch's
nn.TransformerEncoder:[batch, 1, 1, seq_len]with large negative values for padding[batch, seq_len]Reuse grouping/padding logic from
TransformerTreeAggregator:_group_and_pad(): groups nodes by graph_idx, pads to max_nodes_group_and_pad_with_cls(): same but prepends virtual CLS token_aggregate_nodes(): mean pooling over non-padded positions_extract_root_embeddings(): for root-as-CLS modeTreeShapePositionalEncoder — reuse existing class at hf_hidden_dim. Always newly initialized (trainable). This is the tree-structural information that distinguishes this from standard BERT.
Freeze/unfreeze —
freeze_transformer()setsrequires_grad=Falseon all encoder params. Input projection, pos encoder, output projection remain trainable.get_parameter_groups(base_lr, pretrained_lr_scale)— returns list of param group dicts with differential learning rates. Encoder params getbase_lr * pretrained_lr_scale, everything else getsbase_lr. Frozen params (requires_grad=False) are excluded.Constructor Signature
Verification
Where to Find Things
TransformerTreeAggregator:Tree_Matching_Networks/GMN/transformer_tree_aggregator.py— copy_group_and_pad,_group_and_pad_with_cls,_aggregate_nodes,_extract_root_embeddings,_compute_cls_positional_encodingTreeShapePositionalEncoder:Tree_Matching_Networks/GMN/tree_shape_positional_encoder.py— instantiate withembed_dim=hf_hidden_dimGraphEmbeddingNet:Tree_Matching_Networks/GMN/graphembeddingnetwork.py:586— calls aggregator asself._aggregator(node_states, graph_idx, n_graphs, from_idx=from_idx, to_idx=to_idx)See
docs/plans/2026-03-09-pretrained-transformer-aggregation.mdTask 1 for complete code.