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

基于Numpy与Scipy.special向量化序数回归PMF函数并移除循环的实现咨询

基于Numpy与Scipy.special向量化序数回归PMF函数并移除循环的实现咨询

你好!你的这个序数回归PMF函数完全可以通过向量化操作彻底移除for循环,而且逻辑会更简洁,运行效率也会更高(尤其是当类别数K很大的时候)。

核心思路:扩展Cutoff数组+向量化差分

我们可以利用Numpy的广播和差分操作,把原始循环中的分支逻辑统一成一套连续的计算:

原始代码中每个类别的概率本质上是相邻Sigmoid值的差:

  • 第0类:1 - expit(eta - c[0]) 等价于 expit(eta - (-inf)) - expit(eta - c[0])(因为expit(eta - (-inf)) = 1)
  • 中间类:expit(eta - c[k-1]) - expit(eta - c[k])
  • 最后一类:expit(eta - c[-1]) 等价于 expit(eta - c[-1]) - expit(eta + inf)(因为expit(eta + inf) = 0)

所以我们只需要给原始的c数组首尾分别添加负无穷和正无穷作为虚拟Cutoff点,然后计算所有Sigmoid值的相邻差分,就能直接得到所有类别的概率。

向量化后的完整代码

import numpy as np
import scipy.special as ss

def pmf(K: int, eta: np.ndarray, c: np.ndarray) -> np.array:
    """ 
    Example
    -------
    >>> K = 5
    >>> p = np.array([[0.1, 0.3, 0.2, 0.35, 0.05]])
    >>> cum_p = np.cumsum(p)
    >>> cum_logits = ss.logit(cum_p[:-1])
    >>> eta = np.zeros((1, 1))
    >>> p_K = pmf(K=K, eta=eta, c=cum_logits)
    >>> print(p_K)
    [[0.1  0.3  0.2  0.35 0.05]]
    """
    # 构造扩展的Cutoff数组:添加首尾虚拟 cutoff(负无穷、正无穷)
    c_extended = np.concatenate([[-np.inf], c, [np.inf]])
    
    # 利用广播计算所有 eta - c_extended 的Sigmoid值
    # eta[..., np.newaxis] 为 eta 添加一个维度,支持与 c_extended 广播
    expit_vals = ss.expit(eta[..., np.newaxis] - c_extended)
    
    # 计算相邻Sigmoid值的差,直接得到所有类别的概率
    p = np.diff(expit_vals, axis=-1)
    
    # 调整形状以匹配原始代码的输出(如果eta是(n_samples,1)的二维数组,挤压中间维度)
    if eta.ndim == 2 and eta.shape[1] == 1:
        p = p.squeeze(axis=1)
    
    return p

关键优势

  1. 彻底移除循环:用Numpy的向量化操作替代for循环,代码更简洁易读
  2. 支持批量样本:原始代码仅能处理单样本(eta形状为(1,1)),向量化后的代码可以处理任意数量的样本(比如eta形状为(100,1)时,输出会是(100,K)的数组,每行对应一个样本的类别概率)
  3. 性能提升:Numpy的底层C实现比Python循环快得多,尤其是当K或样本量很大时
  4. 逻辑统一:不再需要分支判断首尾类别,所有类别的计算逻辑完全一致

验证与原始代码一致性

运行你提供的示例代码,向量化后的函数会输出和原始代码完全相同的结果:

[[0.1  0.3  0.2  0.35 0.05]]

备注:内容来源于stack exchange,提问作者HJA24

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 18:24:35