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

MAML模型实现正确性排查:多模态6分类元学习问题咨询

多模态6分类MAML实现错误排查

我以图像与文本的CLIP嵌入作为输入,目标输出为0至5的6分类标签,尝试基于MAML(模型无关元学习)实现该多模态6分类元学习任务,但当前实现存在问题,烦请帮忙排查代码中的错误。

import numpy as np
import pandas as pd
from sklearn.preprocessing import LabelEncoder
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score

import warnings

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
device = "cuda" if torch.cuda.is_available() else "cpu"
print(device)

class CustomDataset(Dataset):
    def __init__(self, x, y):
        self.x = torch.tensor(x, dtype=torch.float32).to(device)
        self.y = torch.tensor(y, dtype=torch.long).to(device)
    
    def __len__(self):
        return len(self.x)
    
    def __getitem__(self, idx):
        return self.x[idx], self.y[idx]

class MAML(nn.Module):
    def __init__(self, input_dim, output_dim):
        super(MAML, self).__init__()
        self.input_dim = input_dim
        self.output_dim = output_dim
        self.num_samples = 10
        self.epochs = 20
        self.alpha = 0.001  # Adjusted learning rate
        self.beta = 0.001  # Adjusted meta learning rate
        self.theta = nn.Parameter(torch.randn(input_dim, output_dim).to(device))
        self.softmax = nn.Softmax(dim=1)

    def forward(self, x):
        a = torch.matmul(x, self.theta)
        return self.softmax(a)

    def sample_points(self, k, x, y):
        indices = np.random.choice(len(x), k)
        return x[indices], y[indices]

    def train(self, x_train, y_train, x_val, y_val):
        train_dataset = CustomDataset(x_train, y_train)
        train_loader = DataLoader(train_dataset, batch_size=self.num_samples, shuffle=True)

        optimizer = optim.Adam(self.parameters(), lr=self.alpha)

        for e in range(1, self.epochs + 1):
            self.theta_ = []
            for x_batch, y_batch in train_loader:
                x_batch = x_batch.to(device)
                y_batch = y_batch.to(device)

                y_hat = self.forward(x_batch)
                y_batch_encoded = torch.eye(self.output_dim, device=device)[y_batch]
                loss = -torch.mean(y_batch_encoded * torch.log(y_hat + 1e-7))

                optimizer.zero_grad()
                loss.backward()
                optimizer.step()

                self.theta_.append(self.theta.detach().clone())

            meta_gradient = torch.zeros_like(self.theta, dtype=torch.float32).to(device)
            for i in range(self.num_samples):
                x_test, y_test = self.sample_points(10, x_train, y_train)
                x_test = torch.tensor(x_test, dtype=torch.float32).to(device)
                y_pred = self.forward(x_test)
                y_test_encoded = torch.eye(self.output_dim)[y_test].to(device)
                meta_gradient += torch.matmul(x_test.T, (y_pred - y_test_encoded)) / self.num_samples

            self.theta.data -= self.beta * meta_gradient

            with warnings.catch_warnings():
                warnings.filterwarnings("ignore", category=UserWarning)
                x_val = torch.tensor(x_val, dtype=torch.float32).to(device).clone().detach().requires_grad_(True)
            y_val_pred = self.forward(x_val)
            val_loss = -torch.mean(torch.eye(self.output_dim, device=device)[y_val] * torch.log(y_val_pred + 1e-7))

    def predict(self, x):
        with torch.no_grad():
            x = torch.tensor(x, dtype=torch.float32).to(device)
            y_pred = self.forward(x)
            _, predictions = torch.max(y_pred, dim=1)
            return predictions.cpu().numpy()

# Load the dataset
data = pd.read_csv('data/text_image_embeddings.csv')
x_text = data['text_embedding'].str.split('\t', expand=True).astype(float).values
x_image = data['image_embedding'].str.split('\t', expand=True).astype(float).values
x = np.concatenate((x_text, x_image), axis=1)
label_encoder = LabelEncoder()
y = label_encoder.fit_transform(data['label'])
len(data)
num_labels = len(label_encoder.classes_)
print(num_labels)
models = []
accuracies = []
for i in range(num_labels):
    # Divide data into train and validation for the current label/task
    x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.8, stratify=y, random_state=i)
    
    # Create the CustomDataset for the current label/task
    train_dataset = CustomDataset(x_train, y_train)
    train_loader = DataLoader(train_dataset, batch_size=10, shuffle=True)
    
    # Create the MAML model for the current label/task
    model = MAML(input_dim=x.shape[1], output_dim=num_labels).to(device)
    models.append(model)
    
    # Train the model for the current label/task
    model.train(x_train, y_train, x_test, y_test)
    
    # Calculate accuracy on the validation dataset for the current label/task
    val_predictions = model.predict(x_test)
    accuracy = accuracy_score(y_test, val_predictions)
    accuracies.append(accuracy)

# Print the accuracies for each label/task
for label, accuracy in zip(label_encoder.classes_, accuracies):
    print(f"Label: {label}, Accuracy: {accuracy:.4f}")

错误点与修正说明

1. MAML核心逻辑完全偏离(最关键)

  • 原问题:把MAML写成了普通监督学习,用Adam直接更新元参数θ,没有实现MAML要求的「内循环小样本适配得到θ',再用θ'计算查询集损失来更新元参数」的核心逻辑。
  • 修正:内循环手动计算梯度得到适配后的θ',用θ'在查询集上计算损失,再通过反向传播更新原始元参数θ,严格遵循元学习的"先适配、再元更新"流程。

2. 任务划分逻辑错误

  • 原问题:循环6次训练独立模型,没有构建MAML所需的多任务元训练集(每个任务是N-way K-shot的小样本分类任务)。
  • 修正:生成多个6-way K-shot任务,每个任务包含小样本支持集和查询集,用这些任务来训练元模型,让模型学习跨任务的泛化能力。

3. 数据处理与设备兼容问题

  • 原问题:CustomDataset初始化时直接把数据放到GPU,会导致DataLoader多进程加载时报错;验证集无意义地设置requires_grad=True。
  • 修正:移除Dataset中的设备绑定,在使用张量时再移到对应设备;验证集不需要计算梯度,直接用torch.no_grad()包裹。

4. 损失计算与模型结构问题

  • 原问题:手动实现交叉熵容易出现数值不稳定,且单一线性层对CLIP复杂特征的拟合能力不足。
  • 修正:改用PyTorch官方nn.CrossEntropyLoss()(自带log_softmax,数值稳定性更强);可给模型增加隐藏层(比如nn.Sequential(nn.Linear(input_dim, 512), nn.ReLU(), nn.Linear(512, output_dim)))提升拟合能力。

5. 参数更新逻辑错误

  • 原问题:元梯度计算错误,直接用原始θ计算而不是适配后的θ';手动修改theta.data不符合PyTorch梯度流规范。
  • 修正:用θ'计算查询集损失,通过反向传播自动计算元梯度,用优化器完成元参数更新,保证梯度流的正确性。

内容的提问来源于stack exchange,提问作者varun80042

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 00:55:00