Julia Knet结合预训练编码器与MLP分类的维度不匹配问题
问题根因
维度不匹配的核心原因不是自编码器结构的固有约束,是两个任务逻辑混用导致的:
- 构造小批次数据集时沿用了自编码器训练的逻辑,把原始输入
x_pN当标签,标签维度为22283,但分类头输出维度为3,计算损失时两个张量维度无法对齐 - 损失函数沿用了自编码器的平方差重构损失,不符合3分类任务的训练要求,同时当前代码没有加入预训练编码器的权重加载、冻结逻辑,没有实现「复用预训练编码器」的设计目标
分步修复方案
1. 构造分类任务专属数据集
分类任务的标签是疾病分期标签,不是原始输入。测试阶段可以先生成随机模拟标签验证流程,正式训练替换成真实标注即可:
# 测试用:生成440个样本对应的3分类one-hot标签,维度为(3, 440) y_stage = Flux.onehotbatch(rand(1:3, 440), 1:3) # 构造分类任务minibatch:输入为原始特征,标签为分期标注 trn_pN_mb_clf = minibatch(x_pN, y_stage, 16, shuffle=false, partial=true)
2. 替换分类任务适配的损失函数
删除原有的平方差损失定义,替换为分类任务通用的交叉熵损失,注意最后一层分类头用identity激活输出原始logits,正好匹配Knet交叉熵损失的输入要求:
using Knet, NNHelferlein # 单批次损失:负对数似然(即交叉熵) (ae::EncoderMLP)(x, y) = nll(ae(x), y) # 全数据集平均损失 (ae::EncoderMLP)(d::Knet.Data) = mean(ae(x,y) for (x,y) in d)
如果要复用预训练好的自编码器权重,在实例化mlp_encoder_diseaseStage后,先把训好的自编码器encoder部分的参数加载到encoder_mlp_diseaseStage中;训练初期可以先冻结编码器参数,只更新分类头参数,训练3-5轮分类头收敛后再解冻编码器做小学习率微调,效果会更稳定。
3. 修正训练调用逻辑
把训练时传入的数据集换成分类任务的minibatch,不要传自编码器用的重构数据集。正式训练前先跑一遍维度校验,避免再出现维度不匹配问题:
# 维度校验:输入16个样本的随机batch,检查输出维度是否符合预期 x_test = rand(Float32, 22283, 16) @assert size(mlp_encoder_diseaseStage(x_test)) == (3, 16) "输出维度不匹配,请检查网络结构" # 启动训练 mlp_encoder_diseaseStage = tb_train!( mlp_encoder_diseaseStage, Adam, trn_pN_mb_clf, epochs=50, lr=0.0002, l2=0.0001, l1=0.00002, tb_name="AE_MLP_Disease_Stage", tb_text="AE_MLP" )
内容的提问来源于stack exchange,提问作者Clara
相关产品推荐
相关产品推荐

