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

PyTorch构建贝叶斯线性回归时RuntimeError错误解决求助

问题描述

运行以下贝叶斯线性回归代码时,遇到梯度计算相关的RuntimeError错误:

# Adstock Transformation Formula
def adstock_geometric(x, theta):
    x_len = data.shape[0]
    x_decay = torch.zeros_like(x)
    x_decay[0] = x[0]
    
    for i in range(1, x_len):
        x_decay[i]= x[i] + theta * x_decay[i - 1]
    return x_decay

# Hill Transformation Formula
def hill_transformation(x,alpha_sat,gamma_sat):
    x_len = data.shape[0]
    #x = torch.tensor(x)
    x_saturated = torch.zeros_like(x)
    x_max = torch.max(x)
    x_min = torch.min(x)
    a = torch.tensor([x_min,x_max], dtype=x.dtype)
    b = torch.tensor([1-gamma_sat,gamma_sat],dtype=x.dtype)
    inflexion = torch.dot(a,b)
    for i in range(x_len):
       x_saturated[i] = torch.pow(x[i],alpha_sat)/(torch.pow(x[i],alpha_sat)+torch.pow(inflexion,alpha_sat))
    return x_saturated



def linear_regression(x, y):
    slope = pyro.sample("slope", dist.HalfNormal(2))
    intercept = pyro.sample("intercept", dist.HalfNormal(2))
    theta = pyro.sample("theta", dist.Beta(1,3))
    alpha_sat = pyro.sample("alpha_sat", dist.Gamma(3,1))
    gamma_sat = pyro.sample("gamma_sat", dist.Beta(2,2))
    with pyro.plate("data", len(y)):
        x_adstocked = adstock_geometric(x, theta)
        x_transformed = hill_transformation(x_adstocked, alpha_sat, gamma_sat)
        y_pred = slope * x_transformed + intercept
        pyro.sample("obs", dist.Normal(y_pred, 1), obs=y)
x=data['spends']
y=data['Sales']


x = torch.tensor(x.values, dtype=torch.float32)
y = torch.tensor(y.values, dtype=torch.float32)

    
nuts_kernel = NUTS(linear_regression, jit_compile=True)
mcmc_run = MCMC(nuts_kernel, num_samples=1000, warmup_steps=200)
mcmc_run.run(x, y)

错误信息

RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation: [torch.FloatTensor []], which is output 0 of AsStridedBackward0, is at version 92; expected version 91 instead. Hint: the backtrace further above shows the operation that failed to compute its gradient. The variable in question was changed in there or anywhere later. Good luck!

错误原因

代码中的adstock_geometric和hill_transformation函数使用了inplace赋值操作(如x_decay[i] = ...、x_saturated[i] = ...),这类操作会直接修改张量内存内容,破坏PyTorch的计算图追踪机制。而NUTS采样依赖完整计算图计算梯度,因此触发张量版本不匹配的错误。

修复方案

将inplace赋值替换为非inplace的向量化操作,同时移除对全局变量data的依赖:

修复后的代码

import torch
import pyro
import pyro.distributions as dist
from pyro.infer import MCMC, NUTS

# 修复Adstock变换:用向量化累积计算替代循环inplace赋值
def adstock_geometric(x, theta):
    seq_len = len(x)
    # 生成几何衰减权重序列
    weights = theta ** torch.arange(seq_len, dtype=x.dtype, device=x.device)
    # 构造移位补零的x矩阵,按权重加权后求和
    padded_x = torch.nn.functional.pad(x.unsqueeze(0), (0, seq_len-1)).unfold(1, seq_len, 1)
    x_decay = torch.sum(padded_x * weights.unsqueeze(0), dim=0)
    return x_decay

# 修复Hill变换:用全张量运算替代循环inplace赋值
def hill_transformation(x, alpha_sat, gamma_sat):
    x_max = torch.max(x)
    x_min = torch.min(x)
    inflexion = x_min * (1 - gamma_sat) + x_max * gamma_sat
    # 直接对整个张量执行运算,无需循环赋值
    numerator = torch.pow(x, alpha_sat)
    denominator = numerator + torch.pow(inflexion, alpha_sat)
    x_saturated = numerator / denominator
    return x_saturated

def linear_regression(x, y):
    slope = pyro.sample("slope", dist.HalfNormal(2))
    intercept = pyro.sample("intercept", dist.HalfNormal(2))
    theta = pyro.sample("theta", dist.Beta(1,3))
    alpha_sat = pyro.sample("alpha_sat", dist.Gamma(3,1))
    gamma_sat = pyro.sample("gamma_sat", dist.Beta(2,2))
    with pyro.plate("data", len(y)):
        x_adstocked = adstock_geometric(x, theta)
        x_transformed = hill_transformation(x_adstocked, alpha_sat, gamma_sat)
        y_pred = slope * x_transformed + intercept
        pyro.sample("obs", dist.Normal(y_pred, 1), obs=y)

# 假设data是已加载的DataFrame
x = torch.tensor(data['spends'].values, dtype=torch.float32)
y = torch.tensor(data['Sales'].values, dtype=torch.float32)

nuts_kernel = NUTS(linear_regression, jit_compile=True)
mcmc_run = MCMC(nuts_kernel, num_samples=1000, warmup_steps=200)
mcmc_run.run(x, y)

优化说明

  1. 移除了对全局变量data的依赖,改用输入张量的长度计算,增强函数独立性
  2. 用向量化操作替代循环,彻底避免inplace操作对计算图的破坏,同时提升运行效率

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 02:35:01