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

JAX实现MAML时参数归零问题(元学习)

MAML在正弦任务分布训练时权重收敛到0的问题

我正在用JAX实现MAML(参考论文:Model-Agnostic Meta-Learning for Fast Adaptation of Deep Networks),在简单线性回归任务分布上训练时模型表现正常(收敛较慢但最终有效);但在形如A*sin(B+X)(A、B为随机变量)的任务分布上训练时,网络所有权重均收敛到0,训练后模型几乎没有预测能力。

精简代码如下:

任务生成代码

class MAMLDataLoader:
    def __init__(self, sample_task_fn, num_tasks, batch_size):
        self.sample_task_fn = sample_task_fn
        self.num_tasks = num_tasks
        self.batch_size = batch_size

    def sample_tasks(self, key):
        XS = jnp.empty((self.num_tasks, 2 * self.batch_size, 1))
        YS = jnp.empty((self.num_tasks, 2 * self.batch_size, 1))

        for i in range(self.num_tasks):
            key, subkey = random.split(key)
            xs, ys = self.sample_task_fn(self.batch_size * 2, subkey)
            XS = XS.at[i].set(xs)
            YS = YS.at[i].set(ys)

        x_train, x_test = XS[:, :self.batch_size], XS[:, self.batch_size:]
        y_train, y_test = YS[:, :self.batch_size], YS[:, self.batch_size:]

        return x_train, y_train, x_test, y_test

    def dummy_input(self):
        key = random.PRNGKey(0)
        x = self.sample_task_fn(1, key)[0][0]
        return x

def sample_sinusoidal_task(samples, key):
    # y = a * sin(b + x)
    xs_key, amplitude_key, phase_key = random.split(key, num=3)
    amplitude = random.uniform(amplitude_key, (1, 1))
    phase = random.uniform(phase_key, (1, 1)) * jnp.pi * 2

    xs = (random.uniform(xs_key, (samples, 1)) * 4 - 2) * jnp.pi
    ys = amplitude * jnp.sin(xs + phase)
    return xs, ys

MAML核心代码

class MAMLTrainer:
    def __init__(self, model, alpha, optimiser, inner_steps=1):
        self.model = model
        self.alpha = alpha
        self.optimiser = optimiser
        self.inner_steps = inner_steps

        self.jit_step = jit(self.step)

    def loss(self, params, x, y):
        preds = self.model.apply(params, x)
        return jnp.mean(jnp.inner(y - preds, y - preds) / 2.0)

    def update(self, params, x, y, inner_steps=None):
        if inner_steps is None:
            inner_steps = self.inner_steps

        loss_grad = grad(self.loss)

        def _update(i, params):
            grads = loss_grad(params, x, y)
            new_params = tree_map(lambda p, g: p - self.alpha * g, params, grads)
            return new_params

        return lax.fori_loop(0, inner_steps, _update, params)

    def meta_loss(self, params, x1, y1, x2, y2):
        return self.loss(self.update(params, x1, x2), x2, y2) 

    def batch_meta_loss(self, params, x1, y1, x2, y2):
        return jnp.mean(vmap(partial(self.meta_loss, params))(x1, y1, x2, y2))

    def step(self, params, optimiser, x1, y1, x2, y2):
        loss, grads = value_and_grad(self.batch_meta_loss)(params, x1, y1, x2, y2)

        updates, opt_state = self.optimiser.update(grads, optimiser, params)
        params = optax.apply_updates(params, updates)

        return params, loss


    def train(self, dataloader, steps, key, params=None):
        if params is None:
            key, subkey = random.split(key)
            params = self.model.init(subkey, dataloader.dummy_input())

        optimiser = self.optimiser.init(params)

        pbar, losses = tqdm(range(steps), desc='Training'), []

        for epoch in pbar:
            key, subkey = random.split(key)
            params, loss = self.jit_step(params, optimiser, *dataloader.sample_tasks(subkey))
            losses.append(loss)

            if epoch % 100 == 0:
                avg_loss = jnp.mean(jnp.array(losses[-100:]))

            pbar.set_postfix_str(f'current_loss: {loss:.3f}, running_loss_100_epochs: {avg_loss:.3f}')

        return params, jnp.array(losses)

    def n_shot_learn(self, x_train, y_train, params, n):
        return self.update(params, x_train, y_train, n)

训练代码

class SimpleMLP(nn.Module):
    features: Sequence[int]

    @nn.compact
    def __call__(self, inputs):
        x = inputs
        for i, feat in enumerate(self.features[:-1]):
            x = nn.Dense(feat)(x)
            x = nn.relu(x)

        return nn.Dense(self.features[-1])(x)

model = SimpleMLP([64, 64, 1])
optimiser = optax.adam(1e-3)
trainer = MAMLTrainer(model, 0.1, optimiser, 1)
dataloader = MAMLDataLoader(sample_sinusoidal_task, 2, 100)
key = random.PRNGKey(0)
params, losses = trainer.train(dataloader, 10000, key)

问题排查与修复建议

  1. 核心参数传递错误:在meta_loss方法中,调用内循环更新时错误传入了x1, x2,正确应该传入任务的训练数据x1, y1。这个错误会导致内循环的损失计算完全偏离任务目标,元梯度完全失效,最终模型权重被优化到0来“最小化”无意义的损失。修正代码:
def meta_loss(self, params, x1, y1, x2, y2):
    # 修正内循环更新的参数
    return self.loss(self.update(params, x1, y1), x2, y2) 
  1. 内循环超参数调整:正弦任务比线性回归更复杂,当前内循环学习率0.1可能过大,建议降低到0.01或0.001;同时增加内循环步数到3-5步,让模型在每个任务上有足够的适配空间。

  2. 输入数据归一化:当前生成的xs范围是(-2π, 2π),可以尝试对输入xs做归一化(比如除以2π),缩小输入数据范围,帮助模型更稳定地学习任务间的共性。

  3. 元学习率调整:当前Adam的学习率1e-3对于MAML来说可能偏大,建议尝试1e-4,防止元更新幅度过大导致训练不稳定。

内容的提问来源于stack exchange,提问作者Sefton de Pledge

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 05:30:59