Flax模型init调用返回空参数及预测类型异常求助
问题排查与修正
你的代码存在两个核心问题,直接导致了参数为空和返回类型错误:
1. 未触发参数初始化逻辑
在MLP类的__call__方法中,你仅返回了nn.Sequential的实例,却没有调用它处理输入x。Flax的nn.compact装饰器只有当模块实际执行前向计算(传入输入并调用)时,才会自动初始化层参数。这里只是返回Sequential对象,完全没触发参数初始化,所以model.init得到的params是空字典。
2. 返回的是模块实例而非计算结果
调用model.apply时,返回的是你定义的Sequential对象本身,而不是它对输入x的计算输出,这就导致predictions的类型是flax.linen.combinators.Sequential而非预期的DeviceArray。
修正后的代码(两种写法)
写法一:调用Sequential处理输入
import jax import jax.numpy as jnp import flax.linen as nn import matplotlib.pyplot as plt class MLP(nn.Module): @nn.compact def __call__(self, x): # 关键修正:调用Sequential实例并传入输入x return nn.Sequential( [ nn.Dense(40), nn.relu, nn.Dense(40), nn.Dense(1), ] )(x) model = MLP() dummy_input = jnp.ones((40, 40, 1)) params = model.init(jax.random.PRNGKey(0), dummy_input) # 现在params会包含各Dense层的参数 print(jax.tree_util.tree_map(lambda x: x.shape, params)) n = 100 x_inputs = jnp.linspace(-10, 10, n).reshape(-1, 1) # 调整输入形状匹配Dense层要求(样本数×特征数) y_targets = jnp.sin(x_inputs) predictions = model.apply(params, x_inputs) plt.plot(x_inputs.reshape(-1), y_targets.reshape(-1)) plt.plot(x_inputs.reshape(-1), predictions.reshape(-1)) plt.show()
写法二:更符合Flax风格的层定义(推荐)
不依赖Sequential,直接在__call__里按顺序定义层,可读性和可调试性更好:
import jax import jax.numpy as jnp import flax.linen as nn import matplotlib.pyplot as plt class MLP(nn.Module): @nn.compact def __call__(self, x): x = nn.Dense(40)(x) x = nn.relu(x) x = nn.Dense(40)(x) x = nn.Dense(1)(x) return x model = MLP() dummy_input = jnp.ones((40, 40, 1)) params = model.init(jax.random.PRNGKey(0), dummy_input) print(jax.tree_util.tree_map(lambda x: x.shape, params)) n = 100 x_inputs = jnp.linspace(-10, 10, n).reshape(-1, 1) y_targets = jnp.sin(x_inputs) predictions = model.apply(params, x_inputs) plt.plot(x_inputs.reshape(-1), y_targets.reshape(-1)) plt.plot(x_inputs.reshape(-1), predictions.reshape(-1)) plt.show()
额外说明:修正代码中还调整了x_inputs的形状为(-1, 1),因为Flax的nn.Dense默认期望输入是(样本数, 特征数)的二维张量,原代码的(1, n)形状会导致第一层Dense的输入特征数不匹配,引发错误。
内容的提问来源于stack exchange,提问作者Bunny Rabbit
相关产品推荐
相关产品推荐

