【问题标题】:Pytorch and numpy least squares with an intercept: performance complications带有截距的 Pytorch 和 numpy 最小二乘:性能并发症
【发布时间】:2022-10-19 15:15:34
【问题描述】:

我正在对一些相当大的向量进行回归分析(现在,如果我让计算机在一夜之间工作,使用 numpy 和其他科学工具是可以的)但它们最终会增长几个因素,所以我希望提高性能,移动pytorch 的实现。

回归相当简单。我有 2 个向量,predictionsbetas,尺寸分别为 (750, 6340) 和 (750, 4313)。我正在寻找的最小二乘解决方案是predictions * x = betas,其中 x 的尺寸为 (6340, 4313),但我必须考虑回归中的截距。使用 numpy,我通过迭代 predictions 中的第二维来解决这个问题,创建一个包含每列 + 一列的向量,并将其作为第一个参数传递

for candidate in range(0, predictions.shape[1])): #each column is a candidate
    prediction = predictions[:, candidate]
    #allow for an intercept by adding a column with ones
    prediction = np.vstack([prediction, np.ones(prediction.shape[0])]).T
    sol = np.linalg.lstsq(prediction, betas, rcond=-1)

第一个问题是:有没有办法避免迭代每个候选者,以允许最小二乘计算来解释截距?这将大大缩短计算时间。

我尝试使用statsmodels.regression.linear_model.ols,默认情况下允许这样做(如果你想删除它,你可以在公式中添加-1),但使用这种方法要么迫使我遍历每个候选者(使用apply 很有吸引力,但并没有真正显着改善计算时间)或者我缺少一些东西。那么问题 1.5:我能以这种方式使用这个工具吗?

同样在pytorch我会做

t_predictions = torch.tensor(predictions, dtype=torch.float)
t_betas_roi = torch.tensor(betas, dtype=torch.float)
t_sol = torch.linalg.lstsq(t_predictions, t_betas_roi)

它确实很快,但我错过了这里的拦截。我认为如果我使用 numpy 而不是像我那样进行迭代,它也会更快,但无论哪种方式,如果问题 1 有一个解决方案,我想它可以类似地应用在这里,对吧?

【问题讨论】:

    标签: python numpy pytorch least-squares


    【解决方案1】:

    正如您所提到的,torch.linalg.lstsq 速度很快,但假定截距为零。所以我们可以抵消数据以确保它们的均值为零!

    这里我用X 表示你的predictionsY 表示你的betas

    X_mean = X.mean(dim=0)
    X_centered = X - X_mean 
    Y_mean = Y.mean(dim=0)
    Y_centered = Y - Y_mean 
    
    solution_centered, _, _, _ = lstsq(X_centered, Y_centered)
    
    coef = solution_centered
    intercept = Y_mean - X_mean @ solution_centered
    

    这是有效的,因为计算平均值和偏移由 pyTorch 处理。

    数学

    请注意,这里我使用数学首选的约定(与 pyTorch 不同):

    • 批次由多行数据点组成
    • 操作符在左边
    solution @ (X - X_mean) =fits= Y - Y_mean  
    solution @ X - solution @ X_mean + Y_mean =fits= Y
    

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2021-07-28
      • 2015-04-04
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2018-05-16
      • 2023-03-23
      • 2021-08-14
      相关资源
      最近更新 更多