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

GPFlow v2.1.3多分类任务向量输入触发形状不匹配ValueError求助

GPFlow多分类SVGP维度报错解决方案

问题原因

报错触发的核心原因是核函数参数配置错误:

  • 基础核Matern32的lengthscales参数维度需要与输入特征维度(此处为10维)对齐,你传入了长度为5的列表,与输入维度不匹配,在核计算执行输入除以长度尺度的操作时触发维度对齐错误。
  • 多分类任务需要对应多个 latent GP 的多输出核,不能直接给基础单输出核传入长度等于类别数的参数列表,需要用专用的多输出核包装类处理。

修正后可运行代码

import gpflow
from gpflow.utilities import ops, print_summary, set_trainable
from gpflow.config import set_default_float, default_float, set_default_summary_fmt
from gpflow.ci_utils import ci_niter
import random
import numpy as np
import tensorflow as tf

np.random.seed(0)
tf.random.set_seed(123)

num_classes = 5
num_of_data_points = 1000
num_of_functions = num_classes
num_of_independent_vars = 10

data_gp_train = np.random.rand(num_of_data_points, num_of_independent_vars)
data_gp_train_target_hot = np.eye(num_classes)[np.array(random.choices(list(range(num_classes)), k=num_of_data_points))].astype(bool)
data_gp_train_target = np.apply_along_axis(np.argmax, 1, data_gp_train_target_hot)
data_gp_train_target = np.expand_dims(data_gp_train_target, axis=1)

data_gp = ( data_gp_train, data_gp_train_target )

# 修正:使用SharedIndependent包装基础核,适配多latent GP场景
base_kernel = gpflow.kernels.Matern32(
    variance=1.0, 
    lengthscales=[0.1]*num_of_independent_vars # 长度与输入特征维度匹配
)
kernel = gpflow.kernels.SharedIndependent(base_kernel, output_dim=num_of_functions)

# Robustmax Multiclass Likelihood
invlink = gpflow.likelihoods.RobustMax(num_of_functions)  # Robustmax inverse link function
likelihood = gpflow.likelihoods.MultiClass(num_of_functions, invlink=invlink)  # Multiclass likelihood

inducing_inputs = data_gp_train[::5].copy()  # inducing inputs (20% of obs are inducing)
   
m = gpflow.models.SVGP(
    kernel=kernel,
    likelihood=likelihood,
    inducing_variable=inducing_inputs,
    num_latent_gps=num_of_functions,
    whiten=True,
    q_diag=True,
)

set_trainable(m.inducing_variable, False)
print_summary(m)

opt = gpflow.optimizers.Scipy()
opt_logs = opt.minimize(
    m.training_loss_closure(data_gp), m.trainable_variables, options=dict(maxiter=ci_niter(1000))
)
print_summary(m, fmt="notebook")

关键修改说明

  • 基础核参数对齐输入维度:Matern32的lengthscales参数长度设置为输入特征维度10,匹配输入数据的特征维度
  • 多输出核包装:使用SharedIndependent类包装基础核,通过output_dim指定输出维度等于类别数5,适配多分类任务需要的多个 latent GP 要求,自动实现每个 latent GP 对应独立的核参数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 02:18:00