Skip to content

Pre-trained HF Aggregation: PretrainedTransformerAggregator core component #3

Description

@jlunder00

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

  1. 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
  2. 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()
  3. 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.
  4. TreeShapePositionalEncoder — reuse existing class at hf_hidden_dim. Always newly initialized (trainable). This is the tree-structural information that distinguishes this from standard BERT.

  5. Freeze/unfreezefreeze_transformer() sets requires_grad=False on all encoder params. Input projection, pos encoder, output projection remain trainable.

  6. 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.

Metadata

Metadata

Assignees

No one assigned

    Labels

    enhancementNew feature or request

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions