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

如何实现scipy中make_smoothing_spline的批量处理?

实现make_smoothing_spline()的批量处理功能

核心思路

由于scipy.interpolate.make_smoothing_spline()仅支持一维y输入,要实现和make_interp_spline()一致的批量处理(支持y为(m, ...)形状),最简洁的方式是封装循环逻辑,生成一个能处理批量输入的复合样条函数——numpy.vectorize()无法满足需求,因为它仅能做元素级批量操作,无法生成统一的批量计算函数。

实现方案

可以写一个包装函数,遍历y的额外维度,为每个维度单独创建平滑样条,然后返回一个新函数,该函数能对输入的x_new批量计算所有维度的结果:

import numpy as np
from scipy.interpolate import make_smoothing_spline

def make_batch_smoothing_spline(x, y, **kwargs):
    # 获取y的额外批量维度(除了与x匹配的第一维)
    batch_dims = y.shape[1:]
    # 将y重塑为(m, N),N是批量元素的总数量
    y_reshaped = y.reshape(y.shape[0], -1)
    
    # 为每个批量元素生成平滑样条
    splines = [make_smoothing_spline(x, y_reshaped[:, i], **kwargs) 
               for i in range(y_reshaped.shape[1])]
    
    # 定义批量计算的封装函数
    def batch_spline(x_new):
        # 计算所有样条的结果,再重塑回原批量形状
        results = np.array([s(x_new) for s in splines])
        return results.T.reshape(x_new.shape[0], *batch_dims)
    
    return batch_spline

使用示例

# 构造测试数据
x = np.linspace(0, 10, 20)
# y为(20, 3, 2)的批量形状
y = np.sin(x[:, None, None] + np.linspace(0, np.pi, 3)[None, :, None]) + np.random.normal(0, 0.1, (20,3,2))

# 创建批量平滑样条
batch_spline = make_batch_smoothing_spline(x, y, lam=0.1)

# 测试批量计算
x_new = np.linspace(0,10,100)
y_new = batch_spline(x_new)
print(y_new.shape)  # 输出(100, 3, 2),符合批量处理预期

补充说明

  • 该方案通过循环生成单个样条后统一封装,逻辑清晰且易于维护;
  • 官方已关注make_smoothing_spline()的批量支持需求,后续版本可能会原生提供该功能。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 11:08:10