GraphTokenLM / glm.py
naos-ku's picture
Export GraphTokenLM to Hugging Face Hub
a342f8b
Raw History Blame Contribute Delete
21.8 kB
from typing import Literal
import torch
import torch.nn as nn
from torch_geometric.nn import (
GATConv,
GCNConv,
GINConv,
GraphSAGE,
TransformerConv,
global_add_pool,
global_max_pool,
global_mean_pool,
)
from torch_geometric.utils import to_dense_batch
from transformers import (
AutoConfig,
AutoModelForCausalLM,
PretrainedConfig,
PreTrainedModel,
)
from transformers.generation.utils import GenerationMixin
from transformers.modeling_outputs import CausalLMOutputWithPast
VALID_GRAPH_POOLING = ["mean", "sum", "max"]
class GraphTokenLMConfig(PretrainedConfig):
model_type = "graph_token_lm"
def __init__(
self,
base_model="Qwen/Qwen3-4B-Base",
gnn_type: Literal["GCN", "GAT", "GIN", "GraphSAGE", "GraphTransformer"] = "GCN",
node_feat_dim=8,
lpe_dim: int | None = None,
use_degree_emb: bool = False,
pos_emb_dim=8,
gnn_hidden_dim=256,
gnn_out_dim=512,
num_gnn_layers=2,
graph_pooling: list[Literal["mean", "sum", "max"]] = ["mean"],
num_proj_layers=1,
num_graph_tokens=4,
num_max_nodes=20, # maximum number of nodes per batch
freeze_llm=True,
enable_lora: bool = False,
tie_word_embeddings=True,
**kwargs,
):
"""Initialize a GraphToken language model configuration.
Parameters
----------
base_model : str, default="Qwen/Qwen3-4B-Base"
Hugging Face model name or path for the underlying LLM.
gnn_type : {"GCN", "GAT", "GIN", "GraphSAGE", "GraphTransformer"}, default="GCN"
Type of GNN layer to use for encoding graph nodes.
node_feat_dim : int, default=8
Dimensionality of the raw node features.
lpe_dim : int or None, default=None
Dimensionality for Laplacian positional encodings; defaults to ``node_feat_dim``.
use_degree_emb : bool, default=False
Whether to add degree-based embeddings to node features.
pos_emb_dim : int, default=8
Dimensionality of learned positional embeddings for nodes.
gnn_hidden_dim : int, default=256
Hidden dimensionality for the GNN stack.
gnn_out_dim : int, default=512
Output dimensionality of the GNN encoder.
num_gnn_layers : int, default=2
Number of GNN layers to apply.
num_proj_layers : int, default=1
Number of projection layers mapping graph reps to tokens.
num_graph_tokens : int, default=4
Number of graph tokens to prepend to the LLM.
num_max_nodes : int, default=20
Maximum number of nodes per graph in a batch.
graph_pooling : {"mean", "sum", "max"} list, default=["mean"]
Pooling strategy for graph-level aggregation. When multiple values are
provided, pooled vectors are concatenated.
freeze_llm : bool, default=True
Whether to freeze the underlying LLM parameters.
enable_lora : bool, default=False
Whether LoRA adapters are expected to be active, in which case the
base LLM should remain in training mode unless explicitly set
elsewhere.
tie_word_embeddings : bool, default=True
Whether to tie input/output embeddings in the LLM config.
**kwargs
Additional arguments forwarded to ``PretrainedConfig``.
"""
self.base_model = base_model
self.llm_name = base_model # backward compatibility
self.gnn_type = gnn_type
self.node_feat_dim = node_feat_dim
self.lpe_dim = lpe_dim if lpe_dim is not None else node_feat_dim
self.use_degree_emb = bool(use_degree_emb)
self.pos_emb_dim = pos_emb_dim
self.node_pos_emb_dim = pos_emb_dim # backward compatibility
self.gnn_hidden_dim = gnn_hidden_dim
self.gnn_hidden = gnn_hidden_dim # backward compatibility
self.gnn_out_dim = gnn_out_dim
self.gnn_out = gnn_out_dim # backward compatibility
self.num_gnn_layers = num_gnn_layers
self.num_proj_layers = num_proj_layers
self.num_graph_tokens = num_graph_tokens
self.num_max_nodes = num_max_nodes
self.graph_pooling = graph_pooling
self.freeze_llm = freeze_llm
self.enable_lora = enable_lora
# Keep generation-related fields for compatibility (updated later).
self.vocab_size = kwargs.get("vocab_size", None)
self.pad_token_id = kwargs.get("pad_token_id", None)
self.bos_token_id = kwargs.get("bos_token_id", None)
self.eos_token_id = kwargs.get("eos_token_id", None)
super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs)
# Maintain compatibility with transformers.GenerationConfig.
def get_text_config(self, decoder: bool | None = None, **kwargs):
return self
class GNNEncoder(nn.Module):
def __init__(
self,
in_dim: int,
hid_dim: int,
out_dim: int,
max_nodes: int,
num_layers: int = 2,
node_pos_emb_dim: int = 8,
dropout: float = 0.1,
gnn_type: Literal["GCN", "GAT", "GIN", "GraphSAGE", "GraphTransformer"] = "GCN",
):
"""Encode graph node features with a configurable GNN stack.
Parameters
----------
in_dim : int
Dimensionality of the input node features.
hid_dim : int
Hidden dimensionality used for intermediate layers.
out_dim : int
Dimensionality of the output node representations.
max_nodes : int
Maximum number of nodes per graph in a batch.
num_layers : int, default=2
Number of graph convolution layers.
node_pos_emb_dim : int, default=8
Dimensionality of the optional learned positional embeddings.
dropout : float, default=0.1
Dropout probability applied between hidden layers.
gnn_type : {"GCN", "GAT", "GIN", "GraphSAGE", "GraphTransformer"}, default="GCN"
Type of graph convolution layer to build.
"""
super().__init__()
self.max_nodes = max_nodes
self.pos_emb = nn.Embedding(max_nodes, node_pos_emb_dim) if node_pos_emb_dim > 0 else None
in_channels = in_dim + (node_pos_emb_dim if node_pos_emb_dim > 0 else 0)
if in_channels <= 0:
raise ValueError("GNNEncoder requires a positive input feature dimension.")
hidden_dims = [hid_dim] * max(num_layers - 1, 0)
dims = [in_channels, *hidden_dims, out_dim]
self.convs = nn.ModuleList()
for i in range(len(dims) - 1):
match gnn_type:
case "GCN":
self.convs.append(GCNConv(dims[i], dims[i + 1]))
case "GAT":
self.convs.append(GATConv(dims[i], dims[i + 1]))
case "GIN":
self.convs.append(GINConv(nn.Linear(dims[i], dims[i + 1])))
case "GraphSAGE":
self.convs.append(GraphSAGE(dims[i], dims[i + 1], 1))
case "GraphTransformer":
# Multi-head attention with concat disabled to keep the output dim aligned.
self.convs.append(TransformerConv(dims[i], dims[i + 1], heads=4, concat=False, dropout=dropout))
case _:
raise ValueError(f"Unsupported gnn_type: {gnn_type}")
self.act = nn.ReLU()
self.dropout = nn.Dropout(dropout)
def forward(self, x, edge_index, batch):
if self.pos_emb is not None:
# Assign positional indices per graph starting from zero within the batch.
_, mask = to_dense_batch(x, batch, max_num_nodes=self.max_nodes)
pos_idx = torch.arange(self.max_nodes, device=x.device).unsqueeze(0).expand(mask.size(0), -1)
pos_idx = pos_idx[mask]
x = torch.cat([x, self.pos_emb(pos_idx)], dim=-1)
for i, conv in enumerate(self.convs):
x = conv(x, edge_index)
if i < len(self.convs) - 1:
x = self.act(x)
x = self.dropout(x)
return x # [num_nodes, out_dim]
def _normalize_graph_pooling(graph_pooling: list[str]) -> list[str]:
"""Normalize and validate graph pooling values into a canonical list."""
poolings = list(graph_pooling)
if not (1 <= len(poolings) <= 3):
raise ValueError(f"Unsupported graph_pooling length: {len(poolings)}")
if len(set(poolings)) != len(poolings):
raise ValueError(f"Duplicate graph_pooling values: {poolings}")
if not set(poolings).issubset(set(VALID_GRAPH_POOLING)):
raise ValueError(f"Unsupported graph_pooling values: {poolings}")
return poolings
class DomainProjector(nn.Module):
def __init__(
self,
gnn_out_dim,
llm_hidden_size,
num_graph_tokens=4,
num_layers=1,
graph_pooling: list[str] = ["mean"],
):
"""Project graph-level representations into graph tokens.
The projector first pools node embeddings into a graph representation and
then maps it into ``k`` graph tokens that match the language model's hidden
dimension.
Parameters
----------
gnn_out_dim : int
Dimensionality of the encoder output to project from.
llm_hidden_size : int
Target dimensionality matching the language model embeddings.
num_graph_tokens : int, default=4
Number of graph tokens to produce.
graph_pooling : {"mean", "sum", "max"} list, default=["mean"]
Pooling strategy used to aggregate node embeddings. When multiple values are
provided, pooled vectors are concatenated.
num_layers : int, default=1
Number of linear/GELU projection layers.
"""
super().__init__()
if num_layers < 1:
raise ValueError("DomainProjector requires at least one projection layer.")
poolings = _normalize_graph_pooling(graph_pooling)
self.num_graph_tokens = num_graph_tokens
self.graph_pooling = poolings
layers = []
in_dim = gnn_out_dim * len(poolings)
final_dim = llm_hidden_size * num_graph_tokens
if num_layers == 1:
layer_dims = [final_dim]
else:
ratio = final_dim / in_dim
layer_dims = []
prev_dim = in_dim
for layer_idx in range(num_layers):
t = (layer_idx + 1) / num_layers
dim = max(1, int(round(in_dim * (ratio ** t))))
if dim % 2 != 0:
dim += 1
dim = max(dim, prev_dim)
layer_dims.append(dim)
prev_dim = dim
layer_dims[-1] = final_dim
for layer_idx, out_dim in enumerate(layer_dims):
layers.append(nn.Linear(in_dim, out_dim))
if layer_idx < num_layers - 1:
layers.append(nn.GELU())
in_dim = out_dim
self.project = nn.Sequential(*layers)
# Optional learned positional embeddings for graph tokens.
self.graph_pos = nn.Embedding(num_graph_tokens, llm_hidden_size)
def forward(self, node_repr, batch_index):
"""Aggregate node embeddings into graph tokens.
Parameters
----------
node_repr : torch.Tensor
Node representations of shape ``(num_nodes_total, gnn_out_dim)``.
batch_index : torch.Tensor
Batch indices identifying the graph for each node. Shape
``(num_nodes_total,)``.
Returns
-------
torch.Tensor
Graph token tensor of shape ``(batch_size, num_graph_tokens, hidden)``.
"""
# Pool node representations into a graph-level vector.
pooled_list = []
if "mean" in self.graph_pooling:
pooled_list.append(global_mean_pool(node_repr, batch_index)) # [B, gnn_out_dim]
if "sum" in self.graph_pooling:
pooled_list.append(global_add_pool(node_repr, batch_index)) # [B, gnn_out_dim]
if "max" in self.graph_pooling:
pooled_list.append(global_max_pool(node_repr, batch_index)) # [B, gnn_out_dim]
pooled = (
torch.cat(pooled_list, dim=-1) if len(pooled_list) > 1 else pooled_list[0]
) # [B, gnn_out_dim * num_poolings]
B = pooled.size(0)
# Expand into k tokens via the projection stack.
tokens = self.project(pooled) # [B, k * hidden]
Hk = tokens.view(B, self.num_graph_tokens, -1) # [B, k, hidden]
# Add learned positional embeddings.
pos = self.graph_pos.weight.unsqueeze(0).expand(B, -1, -1) # [B, k, hidden]
Hk = Hk + pos
return Hk # [B, k, hidden]
class GraphTokenLM(PreTrainedModel, GenerationMixin):
"""Language model that prepends graph tokens to textual inputs.
A graph neural network encodes node features, pools them, and projects the
result into learned graph tokens that are concatenated with language model
embeddings before decoding.
Parameters
----------
config : GraphTokenLMConfig
Model configuration describing the graph encoder and base LLM.
load_llm_weights : bool, default=True
Whether to load pretrained weights for the base language model.
"""
_tied_weights_keys = ["llm.lm_head.weight"]
_keys_to_ignore_on_load_missing = [r"^llm\.lm_head\.weight$"]
config_class = GraphTokenLMConfig
base_model_prefix = "llm"
def __init__(self, config: GraphTokenLMConfig, load_llm_weights: bool = True):
super().__init__(config)
# LLM
if load_llm_weights:
self.llm = AutoModelForCausalLM.from_pretrained(
config.base_model, trust_remote_code=True, tie_word_embeddings=True
)
else:
llm_cfg = AutoConfig.from_pretrained(config.base_model, dtype=torch.float32)
self.llm = AutoModelForCausalLM.from_config(llm_cfg)
self.num_graph_tokens = config.num_graph_tokens
# GNN + Domain Projector
self.gnn = GNNEncoder(
gnn_type=config.gnn_type,
node_pos_emb_dim=config.pos_emb_dim,
in_dim=config.node_feat_dim,
hid_dim=config.gnn_hidden_dim,
out_dim=config.gnn_out_dim,
num_layers=config.num_gnn_layers,
max_nodes=config.num_max_nodes,
)
self.tokenizer_head = DomainProjector(
gnn_out_dim=config.gnn_out_dim,
llm_hidden_size=self.llm.config.hidden_size,
num_graph_tokens=config.num_graph_tokens,
num_layers=config.num_proj_layers,
graph_pooling=config.graph_pooling,
)
if config.freeze_llm:
for p in self.llm.parameters():
p.requires_grad = False
if not config.enable_lora:
self.llm.eval()
# --- sync basic generation fields so GenerationMixin works cleanly ---
mirror_keys = [
"vocab_size",
"pad_token_id",
"bos_token_id",
"eos_token_id",
"hidden_size",
"num_hidden_layers",
"num_attention_heads",
]
for k in mirror_keys:
if hasattr(self.llm.config, k):
setattr(self.config, k, getattr(self.llm.config, k))
# make sure tying is done once at init (harmless if already tied)
if getattr(self.config, "tie_word_embeddings", False):
self.tie_weights()
@property
def device(self):
return next(self.parameters()).device
def _concat_graph_tokens(
self,
input_ids=None,
attention_mask=None,
labels=None,
inputs_embeds=None,
graph=None,
):
"""Prepend graph tokens to language model embeddings.
Parameters
----------
input_ids : torch.Tensor, optional
Token IDs used to derive embeddings if ``inputs_embeds`` is not provided.
attention_mask : torch.Tensor, optional
Attention mask aligned with ``input_ids``.
labels : torch.Tensor, optional
Label tensor passed through unchanged.
inputs_embeds : torch.Tensor, optional
Precomputed language model embeddings.
graph : Mapping[str, torch.Tensor], optional
Graph structure containing ``x``, ``edge_index``, and ``batch`` as
produced by the collator.
Returns
-------
tuple of torch.Tensor
Tuple ``(inputs_embeds, attention_mask, labels)`` with graph tokens
concatenated at the front of the sequence.
"""
# Obtain embeddings from the base language model if necessary.
if inputs_embeds is None:
inputs_embeds = self.llm.get_input_embeddings()(input_ids)
B, T, H = inputs_embeds.size()
# ---- Graph to tokens ----
graph_device = next(self.gnn.parameters()).device
if hasattr(graph, "to"):
graph = graph.to(graph_device)
x = graph["x"] # [N_nodes, node_feat_dim]
edge_index = graph["edge_index"] # [2, N_edges]
batch = graph["batch"] # [N_nodes]
node_repr = self.gnn(x, edge_index, batch) # [N_nodes, gnn_out]
graph_tokens = self.tokenizer_head(node_repr, batch) # [B, k, H]
if inputs_embeds is not None:
graph_tokens = graph_tokens.to(inputs_embeds.device)
# ---- Concatenate by prepending graph tokens ----
new_inputs = torch.cat([graph_tokens, inputs_embeds], dim=1) # [B, k+T, H]
# Prepend ones to the attention mask for the graph tokens.
if attention_mask is None:
attention_mask = input_ids.ne(self.llm.config.pad_token_id).long()
new_attention = torch.cat(
[torch.ones((B, self.num_graph_tokens), dtype=attention_mask.dtype, device=self.device), attention_mask],
dim=1,
)
return new_inputs, new_attention, labels
def forward(
self,
input_ids=None,
attention_mask=None,
inputs_embeds=None,
graph=None,
labels=None,
**generate_kwargs,
) -> CausalLMOutputWithPast:
if (input_ids is None) and (inputs_embeds is None):
raise ValueError("Either input_ids or inputs_embeds must be provided")
if graph is not None:
inputs_embeds, attention_mask, labels = self._concat_graph_tokens(
input_ids=input_ids,
attention_mask=attention_mask,
labels=labels,
inputs_embeds=inputs_embeds,
graph=graph,
)
else:
# For generation, inputs may already include concatenated graph tokens.
if inputs_embeds is None:
inputs_embeds = self.llm.get_input_embeddings()(input_ids)
if attention_mask is None:
if input_ids is None:
raise ValueError("When attention_mask is not provided, input_ids must also be provided")
attention_mask = input_ids.ne(self.llm.config.pad_token_id).long()
# Exclude auxiliary keys that may interfere with Qwen3 loss computation.
blocked = {
"num_items_in_batch",
"label_smoothing", # Added by TRL/transformers in some setups.
"labels_shifted", # Same as above.
}
passdown = {
k: v
for k, v in generate_kwargs.items()
if k not in blocked and k in {"use_cache", "output_attentions", "output_hidden_states", "past_key_values"}
}
out = self.llm(
inputs_embeds=inputs_embeds,
attention_mask=attention_mask,
labels=labels,
**passdown,
)
return out
def prepare_inputs_for_generation(
self, input_ids=None, inputs_embeds=None, attention_mask=None, graph=None, **kwargs
):
if input_ids is None and inputs_embeds is None:
raise ValueError("Either input_ids or inputs_embeds must be provided")
if input_ids is not None and inputs_embeds is not None:
raise ValueError("Both input_ids and inputs_embeds cannot be provided at the same time")
if inputs_embeds is None:
inputs_embeds = self.llm.get_input_embeddings()(input_ids)
if graph is not None:
inputs_embeds, attention_mask, _ = self._concat_graph_tokens(
input_ids=None,
attention_mask=attention_mask,
labels=None,
inputs_embeds=inputs_embeds,
graph=graph,
)
return {"inputs_embeds": inputs_embeds, "attention_mask": attention_mask, "graph": None}
# delegate embeddings to inner LLM so HF can tie weights correctly
def get_input_embeddings(self):
return self.llm.get_input_embeddings()
def set_input_embeddings(self, new_embeddings):
self.llm.set_input_embeddings(new_embeddings)
def get_output_embeddings(self):
return self.llm.get_output_embeddings()
def set_output_embeddings(self, new_embeddings):
self.llm.set_output_embeddings(new_embeddings)
def tie_weights(self):
# honor config.tie_word_embeddings and delegate
if getattr(self.config, "tie_word_embeddings", False):
# inner LLM handles actual tying (lm_head <-> embeddings)
self.llm.tie_weights()
# keep parent behavior (no-op for most models)
return super().tie_weights()