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
相关产品推荐
相关产品推荐

