如何基于Hugging Face Transformer自定义融合边信息的位置嵌入?
问题描述
我正在用Hugging Face的Transformer模型做机器翻译,输入数据包含token间的关联信息,想构建图结构:把句子里的每个token作为节点,token之间嵌入边信息。
标准Transformer里,token会转成token embedding,每个位置对应positional embedding。我想给边信息做类似的嵌入处理——理论上可以结合边类型与节点位置编码生成嵌入,再把所有边嵌入加到对应节点的位置嵌入中,但不知道怎么修改Longformer的位置嵌入代码:
self.position_embeddings = nn.Embedding( config.max_position_embeddings, config.hidden_size, padding_idx=self.padding_idx )
解决方案
1. 扩展模型配置
先给LongformerConfig新增边类型数量的参数,用来定义边嵌入层的规模:
from transformers import LongformerConfig # 自定义配置,新增num_edge_types字段 class GraphLongformerConfig(LongformerConfig): def __init__(self, num_edge_types=4, **kwargs): super().__init__(**kwargs) self.num_edge_types = num_edge_types
2. 修改嵌入层结构
在LongformerEmbeddings类中,新增边类型嵌入层,用于将边类型转为向量表示:
from transformers.models.longformer.modeling_longformer import LongformerEmbeddings import torch.nn as nn import torch class GraphLongformerEmbeddings(LongformerEmbeddings): def __init__(self, config): super().__init__(config) # 新增边类型嵌入层 self.edge_type_embedding = nn.Embedding(config.num_edge_types, config.hidden_size) def forward( self, input_ids=None, token_type_ids=None, position_ids=None, inputs_embeds=None, past_key_values_length=0, edge_info=None, # 新增:传入边信息,形状为(batch_size, num_edges, 3),每个元素是[src_idx, tgt_idx, edge_type] ): # 执行原位置嵌入逻辑 if position_ids is None: position_ids = self.create_position_ids_from_input_ids(input_ids) device = inputs_embeds.device if inputs_embeds is not None else input_ids.device position_ids = position_ids.to(device) position_embeddings = self.position_embeddings(position_ids) # 处理边嵌入并聚合到对应节点 if edge_info is not None: src_indices = edge_info[:, :, 0] # 边的起始节点索引 tgt_indices = edge_info[:, :, 1] # 边的目标节点索引 edge_types = edge_info[:, :, 2] # 边的类型 # 获取边类型对应的嵌入向量 edge_embs = self.edge_type_embedding(edge_types) # shape: (batch_size, num_edges, hidden_size) # 初始化边嵌入聚合张量 edge_agg_embeddings = torch.zeros_like(position_embeddings) # 用scatter_add高效将边嵌入加到对应节点的位置嵌入上 batch_indices = torch.arange(edge_embs.size(0)).unsqueeze(1).expand(-1, edge_embs.size(1)).to(edge_embs.device) # 给起始节点加边嵌入 edge_agg_embeddings.scatter_add_( dim=1, index=src_indices.unsqueeze(-1).expand(-1, -1, edge_embs.size(2)), src=edge_embs ) # 给目标节点加边嵌入 edge_agg_embeddings.scatter_add_( dim=1, index=tgt_indices.unsqueeze(-1).expand(-1, -1, edge_embs.size(2)), src=edge_embs ) # 将边聚合嵌入加到位置嵌入中 position_embeddings += edge_agg_embeddings # 继续执行原嵌入层的后续逻辑 embeddings = self.token_embeddings(input_ids) if input_ids is not None else inputs_embeds if token_type_ids is not None: embeddings += self.token_type_embeddings(token_type_ids) embeddings += position_embeddings embeddings = self.LayerNorm(embeddings) embeddings = self.dropout(embeddings) return embeddings
3. 替换模型的嵌入层
初始化Longformer模型时,把默认的embeddings换成自定义的GraphLongformerEmbeddings:
from transformers import LongformerModel class GraphLongformerModel(LongformerModel): def __init__(self, config): super().__init__(config) # 替换嵌入层 self.embeddings = GraphLongformerEmbeddings(config) # 重新初始化权重 self.init_weights()
4. 模型使用示例
# 初始化自定义配置与模型 config = GraphLongformerConfig(num_edge_types=4, max_position_embeddings=512, hidden_size=768) model = GraphLongformerModel(config) # 准备输入数据 input_ids = torch.randint(0, config.vocab_size, (2, 10)) # batch_size=2, seq_len=10 # 边信息:每个batch有3条边,每条边格式为[src_idx, tgt_idx, edge_type] edge_info = torch.tensor([ [[0, 1, 1], [2, 3, 2], [5, 7, 1]], [[1, 4, 3], [3, 6, 2], [0, 8, 1]] ]) # 前向传播,传入edge_info outputs = model(input_ids=input_ids, edge_info=edge_info)
关键说明
- 边信息的聚合方式这里用了求和,你也可以根据需求换成平均、最大池化等操作
- 如果需要结合节点的位置信息生成边嵌入,可以在边嵌入层中加入相对位置编码(比如用
nn.Embedding(config.max_position_embeddings*2, config.hidden_size)表示相对位置,再和边类型嵌入拼接/相加) - 机器翻译场景如果用Encoder-Decoder结构,需要对Decoder的嵌入层做类似修改,或者只在Encoder中加入边信息处理
内容的提问来源于stack exchange,提问作者Exploring
相关产品推荐
相关产品推荐

