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

GraphSAGE聚合层权重未训练问题:仅边预测层权重可更新

问题

我构建了一个包含两个带独立权重与偏置的聚合层的GraphSAGE模型,同时设置边预测层用于预测节点间是否存在边。训练过程中仅边预测层的权重被学习,聚合层的权重完全不更新。模型逻辑是先聚合邻居节点表征得到目标节点表征,再通过边预测层计算损失。以下是相关代码片段:

训练模型中获取损失的代码片段

...
with tf.GradientTape() as tape:
    
    # Generate inputs to be fed into model
    # Note: next_minibatch_feed_dict() feeds in the next batch of training edges
    train_feed_dict = minibatch.next_minibatch_feed_dict()
    self.inputs1 = train_feed_dict['batch1']
    self.inputs2 = train_feed_dict['batch2']
    self.batch_size = train_feed_dict['batch_size']

    # Start taking note of training time
    start_time = time.time()

    # Pass in inputs to model and get outputs
    input1_outputs, input2_outputs, neg_outputs = self.get_outputs()

    # Run model on train inputs to get train_loss, train_mrr 
    train_loss = self.calc_loss((input1_outputs, input2_outputs, neg_outputs))
    train_loss = train_loss / tf.cast(self.batch_size, tf.float32)
...

计算损失的函数

def calc_loss(self, outputs):
    
    loss = 0.
    input1_outputs, input2_outputs, neg_outputs = outputs
    
    # Get aggregators
    aggregators = [self.layer0_aggregator, self.layer1_aggregator]
    
    # For each aggregator,
    for aggregator in aggregators:

        # Add l2 loss of aggregator's weights
        loss += self.params_decay * tf.nn.l2_loss(aggregator.neigh_weights)
        loss += self.params_decay * tf.nn.l2_loss(aggregator.self_weights)
        
        if aggregator.bias:
            loss += self.params_decay * tf.nn.l2_loss(aggregator.bias_vals)
        
    # Add loss from link prediction
    loss += self.link_predictor.loss(input1_outputs, input2_outputs, neg_outputs)
    
    tf.summary.scalar('loss', loss)
    return loss

获取输出的函数

def get_outputs(self):
    
    # Prepare neg_nodes
    labels = tf.reshape(tf.cast(self.inputs2, dtype = tf.int64), [self.batch_size, 1])
    
    # Convert numpy array to list and flatten 2D list to a 1D list
    unigrams = self.degrees.flatten()
    unigrams = list(unigrams)
    
    self.neg_samples, _, _ = tf.nn.fixed_unigram_candidate_sampler(
        true_classes=labels,
        num_true=1,
        num_sampled=self.neg_sample_size,
        unique=False,
        range_max=len(self.degrees),
        distortion=0.75,
        unigrams=unigrams)
     
    # Get the samples and support_size for batch1, batch2 and neg_nodes in an array
    samples_input1, support_sizes_input1 = self.sample(self.inputs1, self.layer_infos)
    samples_input2, support_sizes_input2 = self.sample(self.inputs2, self.layer_infos)
    neg_samples, neg_support_sizes = self.sample(self.neg_samples, self.layer_infos, self.neg_sample_size)
    
    # Get number of samples
    num_samples = [layer_info.num_samples for layer_info in self.layer_infos]
    
    # Aggregate the samples and get output
    input1_outputs = self.aggregate(samples_input1, [self.features], self.dims, num_samples,
                              support_sizes_input1, concat = self.concat, model_size = self.model_size)
    input2_outputs = self.aggregate(samples_input2, [self.features], self.dims, num_samples,
                              support_sizes_input2, concat = self.concat, model_size = self.model_size)
    neg_outputs = self.aggregate(neg_samples, [self.features], self.dims, num_samples, neg_support_sizes,
                                 batch_size = self.neg_sample_size, concat = self.concat, model_size = self.model_size)
    
    # return feature representation for batch 1 nodes, " for batch 2 nodes, " for negative output nodes
    return (input1_outputs, input2_outputs, neg_outputs)

聚合函数(调用自定义聚合层)

def aggregate(self, samples, input_features, dims, num_samples, support_sizes, batch_size = None,
             name = None, concat = False, model_size = 'small'):
       
    # Determine batch size
    if batch_size is None:
        batch_size = self.batch_size
        
    # hidden is a list that contains the embeddings or hidden representations of the node samples from each layer.
    # Each element in the hidden list corresponds to a layer, and the shape of each element is determined by the number 
    # of node samples in that layer.
    # hidden: [base nodes to get representation for, nodes 2 layers away (Layer 2), nodes 1 layer away (Layer 1)]
    hidden = [tf.nn.embedding_lookup(input_features, np.asarray(node_samples, dtype = np.int32)) for node_samples in samples]
           
    # Start aggregating by layers
    for layer in range(len(num_samples)):
                                
        # After determining aggregators, start from hops furthest away until reach back to base parent node
        next_hidden = []
                    
        for hop in range(len(num_samples) - layer):
            
            # Determine dimension multiplier
            dim_mult = 2 if concat and layer != 0 else 1
            
            # Convert batch size from tensor to numpy if needed
            if tf.is_tensor(batch_size):
                batch_size = batch_size.numpy()
                
            # Shape: [dims for current layer, dims for neighbour layer, dims of current layer]
            neigh_dims = [batch_size*support_sizes[hop], num_samples[len(num_samples) - hop - 1], dim_mult*dims[layer]]

            # Doing aggregator(input_data) calls .call() and .build() implictly
            if hop == len(num_samples) - 1:
                aggregated_res = self.layer1_aggregator((hidden[hop], tf.reshape(hidden[hop + 1], neigh_dims)))
            else:
                aggregated_res = tf.nn.relu(self.layer0_aggregator((hidden[hop], tf.reshape(hidden[hop + 1], neigh_dims))))
            
            next_hidden.append(aggregated_res)
            
    # Re-assign hidden as new aggregated output
    hidden = next_hidden

    # Return node representation of base node
    return hidden[0]

聚合层的_call()函数(接收目标节点批次及其邻居节点)

def _call(self, inputs):
    
    # Feature representation for base nodes, feature representation for neighbour nodes
    self_vecs, neigh_vecs = inputs
    input_batchsize = self_vecs.shape[0] # 50
    output_dim = self_vecs.shape[1] # 14
    
    # Perform dropout
    # neigh_vecs = tf.nn.dropout(neigh_vecs, self.dropout)
    # self_vecs = tf.nn.dropout(self_vecs, self.dropout)
    
    # Array to store output
    output_np = np.zeros(shape = (input_batchsize, output_dim), dtype = np.float32)

    # Aggregate for each base node in batch
    for i in range(len(self_vecs)):
        
        # shape: (14,)
        cur_node_features = self_vecs[i]
        # shape: (5.14)
        neigh_nodes_features = neigh_vecs[i]
        
        # Aggregate features of neighbours by taking the mean
        # shape: (14,)
        neigh_nodes_mean = tf.reduce_mean(neigh_nodes_features, axis = 0)
        
        # Expand dims to perform matmul operation
        # shape: (14, 1)
        neigh_nodes_mean = tf.expand_dims(neigh_nodes_mean, axis = 1)
        cur_node_features = tf.expand_dims(cur_node_features, axis = 1)
        
        # shape: (14, 1)
        from_neighs = tf.matmul(self.neigh_weights, neigh_nodes_mean)
        from_self = tf.matmul(self.self_weights, cur_node_features)
        
        # Element wise addition if not concatenating else concate
        if not self.concat:
            # output shape: (14,1)
            output = tf.add_n([from_neighs, from_self])
        else:
            output = tf.concat([from_neighs, from_self], axis = 1)
        
        # Reduce dimensions to add with bias
        # shape: (14,)
        output = tf.squeeze(output)

        # Add bias, depending on self.bias
        if self.bias:
            output += self.bias_vals
        
        # output_np[i] shape: (14,)
        # output shape: (14,)
        output_np[i] = output
    
    output = tf.convert_to_tensor(output_np, np.float32)
    return output
原因分析与解决方案

1. NumPy数组操作切断梯度传播链

在聚合层的_call函数中,你使用NumPy数组output_np存储中间结果,再转换为Tensor返回。NumPy操作不属于TensorFlow计算图范畴,会直接切断梯度从输出回溯到聚合层权重的路径,导致梯度无法更新这些权重。

修复方法:全程使用TensorFlow张量操作替代NumPy数组,同时批量处理样本(避免循环单个样本,提升效率的同时保证梯度追踪):

def _call(self, inputs):
    self_vecs, neigh_vecs = inputs
    input_batchsize = tf.shape(self_vecs)[0]
    output_dim = tf.shape(self_vecs)[1]

    # 批量计算邻居节点均值,无需循环单个样本
    neigh_nodes_mean = tf.reduce_mean(neigh_vecs, axis=1)  # shape: (batch_size, output_dim)
    
    # 扩展维度以支持批量矩阵乘法
    neigh_nodes_mean = tf.expand_dims(neigh_nodes_mean, axis=-1)  # shape: (batch_size, output_dim, 1)
    self_vecs_expanded = tf.expand_dims(self_vecs, axis=-1)  # shape: (batch_size, output_dim, 1)
    
    # 批量矩阵乘法,替代循环计算
    from_neighs = tf.matmul(self.neigh_weights[None, ...], neigh_nodes_mean)  # shape: (batch_size, output_dim, 1)
    from_self = tf.matmul(self.self_weights[None, ...], self_vecs_expanded)  # shape: (batch_size, output_dim, 1)
    
    if not self.concat:
        output = tf.add_n([from_neighs, from_self])
    else:
        output = tf.concat([from_neighs, from_self], axis=1)
    
    output = tf.squeeze(output, axis=-1)  # 压缩维度至(batch_size, output_dim)
    
    if self.bias:
        output += self.bias_vals
    
    return output

2. Tensor类型batch_size转NumPy破坏梯度追踪

在aggregate函数中,你将Tensor类型的batch_size转换为NumPy值:

if tf.is_tensor(batch_size):
    batch_size = batch_size.numpy()

这会导致后续依赖batch_size的形状计算脱离GradientTape追踪,进而影响聚合层操作的梯度传播。

修复方法:保留batch_size的Tensor类型,使用TensorFlow原生操作处理形状:

# 替换原batch_size转换代码
batch_size = tf.cast(batch_size, tf.int32)
# 用TensorFlow的stack操作构建形状,避免依赖NumPy数值
neigh_dims = tf.stack([
    batch_size * support_sizes[hop], 
    num_samples[len(num_samples) - hop - 1], 
    dim_mult * dims[layer]
])
# 或者直接使用动态形状重塑张量
neigh_vecs_reshaped = tf.reshape(
    hidden[hop + 1], 
    (-1, num_samples[len(num_samples) - hop - 1], dim_mult * dims[layer])
)

3. 确保聚合层权重是可训练的TensorFlow变量

检查聚合层的初始化代码,确认neigh_weights、self_weights、bias_vals是用tf.Variable(或Layer.add_weight)创建的可训练变量,而非普通tf.Tensor。只有tf.Variable才会被GradientTape追踪并更新梯度。

示例初始化代码:

class Aggregator(tf.keras.layers.Layer):
    def __init__(self, output_dim, concat=True, bias=True, **kwargs):
        super().__init__(**kwargs)
        self.output_dim = output_dim
        self.concat = concat
        self.bias = bias
    
    def build(self, input_shape):
        # 输入形状为(self_vecs_shape, neigh_vecs_shape)
        input_dim = input_shape[0][-1]
        # 创建可训练的邻居权重
        self.neigh_weights = self.add_weight(
            name='neigh_weights',
            shape=(input_dim, self.output_dim),
            initializer='glorot_uniform',
            trainable=True
        )
        # 创建可训练的自身节点权重
        self.self_weights = self.add_weight(
            name='self_weights',
            shape=(input_dim, self.output_dim),
            initializer='glorot_uniform',
            trainable=True
        )
        # 创建可训练偏置(如果需要)
        if self.bias:
            self.bias_vals = self.add_weight(
                name='bias',
                shape=(self.output_dim,),
                initializer='zeros',
                trainable=True
            )
        super().build(input_shape)

内容的提问来源于stack exchange,提问作者pjyamas

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 18:15:55