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)
问题排查与修复建议
- 核心参数传递错误:在
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)
内循环超参数调整:正弦任务比线性回归更复杂,当前内循环学习率
0.1可能过大,建议降低到0.01或0.001;同时增加内循环步数到3-5步,让模型在每个任务上有足够的适配空间。输入数据归一化:当前生成的xs范围是
(-2π, 2π),可以尝试对输入xs做归一化(比如除以2π),缩小输入数据范围,帮助模型更稳定地学习任务间的共性。元学习率调整:当前Adam的学习率
1e-3对于MAML来说可能偏大,建议尝试1e-4,防止元更新幅度过大导致训练不稳定。
内容的提问来源于stack exchange,提问作者Sefton de Pledge
相关产品推荐
相关产品推荐

