ResNet图像对比模型无法区分同色单调图像的问题及解决
图像对比模型收敛异常问题及解决
训练图像对比模型时遇到问题:将任务简化为输入3×128×128的纯黑或纯白图像对,模型通过两个独立ResNet提取特征后拼接,经全连接层输出1.0(同色)或0.0(异色)。但模型始终收敛至约0.5,尽管任务看似简单。
初始模型
class TemplateEvaluator(nn.Module): def __init__(self, q_encoder=resnet18(), t_encoder=resnet18()): super(TemplateEvaluator, self).__init__() self.q_encoder = q_encoder self.t_encoder = t_encoder # 设置ResNet参数可训练 for param in self.q_encoder.parameters(): param.requires_grad = True for param in self.t_encoder.parameters(): param.requires_grad = True self.fc = nn.Sequential( nn.Linear(2000, 1), nn.Sigmoid() ) def forward(self, data): q = data[0] t = data[1] # 处理单张图像 if q.ndim == 3: q = q.unsqueeze(0) if t.ndim == 3: t = t.unsqueeze(0) q = self.q_encoder(q) t = self.t_encoder(t) res = self.fc(torch.cat([q,t],-1)).flatten() return res
数据加载器
class BlackOrWhiteDataset(Dataset): def __init__(self): self.tf = transforms.ToTensor() def __getitem__(self, i): black = (255,255,255) white = (0,0,0) x1_col = black if (np.random.random() > 0.5) else white x2_col = black if (np.random.random() > 0.5) else white y = torch.tensor(x1_col == x2_col, dtype=torch.float) x1 = Image.new('RGB', (img_width,img_width), x1_col) x2 = Image.new('RGB', (img_width,img_width), x2_col) return self.tf(x1), self.tf(x2), y def __len__(self): return 100 def create_data_loader(dataset, batch_size, verbose=True): dl = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=True, collate_fn=lambda x: tuple(x_.to(device) for x_ in default_collate(x))) return dl
训练代码
t_eval = TemplateEvaluator().to(device) opt = optim.SGD(t_eval.parameters(), lr=0.001, momentum=0.01) epochs = 10 losses = [] for epoch in tqdm(range(epochs)): t_eval.train() for X1, X2, Y in dl: Y_pred = t_eval(torch.stack([X1,X2])) loss = F.mse_loss(Y_pred,Y) opt.zero_grad() loss.backward() opt.step() sys.stdout.write('\r') sys.stdout.write("loss: %f" % loss.item()) sys.stdout.flush() losses.append(loss.item()) plt.plot(losses) plt.ylim(0,1)
训练结果
训练后loss稳定在0.25左右,模型输出始终约0.5。测试发现,单独判断单张图像是否为黑色时模型可正常收敛,仅在融合双图像特征时失效。
问题解决
经分析,单层感知器无法解决等价判断(NXOR/XOR问题)。将特征拼接改为元素相乘后,模型成功收敛:
class TemplateEvaluator(nn.Module): def __init__(self, q_encoder=resnet18(), t_encoder=resnet18()): super(TemplateEvaluator, self).__init__() self.q_encoder = q_encoder self.t_encoder = t_encoder self.fc = nn.Sequential( nn.Linear(1000, 1), nn.Sigmoid() ) def forward(self, data): q = data[0] t = data[1] if q.ndim == 3: q = q.unsqueeze(0) if t.ndim == 3: t = t.unsqueeze(0) q_features = self.q_encoder(q) t_features = self.t_encoder(t) combined_features = q_features * t_features res = self.fc(combined_features).flatten() return res
新模型训练后loss快速趋近于0,可准确判断图像对是否同色。
内容的提问来源于stack exchange,提问作者Th F
相关产品推荐
相关产品推荐

