请求将基于Theano的几何广告衰减代码转换为PyTensor版本
转换后的PyTensor实现代码
import pytensor.tensor as pt def adstock_geometric_pytensor(x, theta): x = pt.as_tensor_variable(x) def adstock_geometric_recurrence_pytensor(index, input_x, decay_x, theta): # 计算当前步的衰减值并更新到decay_x对应位置 updated_decay = pt.set_subtensor(decay_x[index], input_x + theta * decay_x[index - 1]) return updated_decay len_observed = x.shape[0] x_decayed = pt.zeros_like(x) # 初始化第一个时间步的衰减值 x_decayed = pt.set_subtensor(x_decayed[0], x[0]) output, _ = pt.scan( fn=adstock_geometric_recurrence_pytensor, sequences=[pt.arange(1, len_observed), x[1:len_observed]], outputs_info=x_decayed, non_sequences=theta, n_steps=len_observed - 1 ) return output[-1]
关键修改说明
- 替换Theano的
tensor模块为PyTensor的pytensor.tensor(导入后用pt别名) - 将
theano.scan替换为pt.scan,PyTensor的scan接口与Theano基本兼容,参数无需大幅调整 - 移除了原代码中多余的
tt.sum——因为input_x和decay_x[index-1]都是单时间步的标量,相加后本身就是标量,无需求和 - 递归函数名称调整为匹配PyTensor上下文,其余变量命名保持与原代码一致
内容的提问来源于stack exchange,提问作者Shuvayan Das
相关产品推荐
相关产品推荐

