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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 23:03:13