CIFAR-100上Hierarchical Softmax实现异常:Loss振荡问题求助
层级Softmax(Hierarchical Softmax)实现异常问题
在作业中实现层级Softmax时遇到异常:Loss在前几轮迭代下降后,仅在两个数值间来回振荡。推测原因是梯度仅在前几轮非零,之后变为零,所有输入图像被归类到树的同一叶节点(通常是0)。
实现代码
1. 二叉树节点与构建逻辑
class Node: count_nodes = 0 def __init__(self, label=None, parent = None, grad = True, tree = None): Node.count_nodes += 1 self.label = label self.parent = parent self.left = None self.right = None self.probs = torch.FloatTensor(100, ).uniform_(0, 1) self.probs.requires_grad = grad if(label == None and tree != None): tree.append(self) def buildBinaryTree(labels, root, tree): if len(labels) == 1: return Node(labels[0], parent=root, grad=False) else: split_idx = len(labels) // 2 left_labels, right_labels = [], [] for i, label in enumerate(labels): if i < split_idx: left_labels.append(label) else: right_labels.append(label) node = Node(parent=root, tree=tree) node.left = buildBinaryTree(left_labels, root=node, tree=tree) node.right = buildBinaryTree(right_labels, root=node, tree=tree) return node # 基于CIFAR-100数据集构建二叉树 data = train_dataset.data labels = [i for i in range(100)] tree = [] root = buildBinaryTree(labels, None, tree)
2. 卷积网络模型
class HierarchicalConvNet(nn.Module): def __init__(self): super(HierarchicalConvNet, self).__init__() self.conv1 = nn.Conv2d(3, 6, 3, padding = 1) self.conv2 = nn.Conv2d(6, 8, 3, padding = 1) self.conv3 = nn.Conv2d(8, 10, 3, padding = 1) self.pool = nn.MaxPool2d(2, 2) self.fc1 = nn.Linear(4*4*10, 200) self.fc2 = nn.Linear(200, 100) def forward(self, x): # -> n, 3, 32, 32 x = self.pool(F.relu(self.conv1(x))) # -> n, 6, 16, 16 x = self.pool(F.relu(self.conv2(x))) # -> n, 8, 8, 8 x = self.pool(F.relu(self.conv3(x))) # -> n, 10, 4, 4 x = torch.flatten(x, start_dim = 1) # -> n, 160 x = F.relu(self.fc1(x)) # -> n, 200 xout = F.relu(self.fc2(x)) # -> n, 100 return xout H_model = HierarchicalConvNet()
3. 层级Softmax损失函数
class HierarchicalSoftmaxLoss(nn.Module): def __init__(self, root): super().__init__() self.root = root def forward(self, input, target): batch_size = input.shape[0] loss = [] for i in range(batch_size): x = input[i] label = target[i] node = label_node[label.item()] path = [] while node.parent != self.root: path.append(node.parent) path.append(node.parent) length = len(path) for j in range(length - 1, 0, -1): curr_node = path[j] w = curr_node.probs if curr_node.left == path[j - 1]: loss.append(F.binary_cross_entropy_with_logits(torch.reshape(torch.dot(x, w), (1,)), torch.tensor([1.0]))) else: loss.append(F.binary_cross_entropy_with_logits(torch.reshape(torch.dot(x, w), (1,)), torch.tensor([0.0]))) losses = sum(loss)/batch_size return losses
4. 优化器设置
H_criterion = HierarchicalSoftmaxLoss(root) H_optimizer = torch.optim.SGD(list(H_model.parameters()) + [node.probs for node in tree], lr=0.5)
训练结果
训练后Loss曲线显示:前5000轮每轮记录Loss,之后每5000轮记录一次,可见Loss仅在两个值间振荡。
内容的提问来源于stack exchange,提问作者karannb
相关产品推荐
相关产品推荐

