You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.23 08:45:28