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

Python scoringrules库计算Energy Score时数组形状不兼容问题求助

问题:scoringrules库计算Energy Score时的形状不兼容错误

问题背景

使用Python的scoringrules库(功能类似R语言的scoringRules库)计算多变量预测的Energy Score和变异函数得分时,即便数组形状均为(1, 24),仍遇到数组形状不兼容问题。运行官方测试代码后同样出现形状兼容错误,期望得到形状为(100, 3)的Energy Score数组。

测试代码

import numpy as np
import pytest
from scoringrules._energy import energy_score
from scoringrules.backend import backends

# Define the backends you want to test
RUN_TESTS = ["numba", "jax"]
# Filter available backends based on RUN_TESTS
BACKENDS = [b for b in backends.available_backends if b in RUN_TESTS]
# Register the filtered backends
for backend in RUN_TESTS:
    backends.register_backend(backend)

ENSEMBLE_SIZE = 51
N = 100
N_VARS = 3

@pytest.mark.parametrize("backend", BACKENDS)
def test_energy_score(backend):
    obs = np.random.randn(N, N_VARS)
    fct = np.expand_dims(obs, axis=-2) + np.random.randn(N, ENSEMBLE_SIZE, N_VARS)
    res = energy_score(obs, fct, backend=backend)
    if backend in ["numpy", "numba"]:
        assert isinstance(res, np.ndarray)
    elif backend == "jax":
        assert isinstance(res, jax.Array)

x = test_energy_score(backend)
print(x)
print(x.shape)

报错信息

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
Cell In[12], line 32
     29     elif backend == "jax":
     30         assert isinstance(res, jax.Array)
---> 32 x = test_energy_score(backend)
     33 print(x)
     34 print(x.shape)

Cell In[12], line 25, in test_energy_score(backend)
     22 obs = np.random.randn(N, N_VARS)
     23 fct = np.expand_dims(obs, axis=-2) + np.random.randn(N, ENSEMBLE_SIZE, N_VARS)
---> 25 res = energy_score(obs, fct, backend=backend)
     27 if backend in ["numpy", "numba"]:
     28     assert isinstance(res, np.ndarray)

File ~\anaconda3\Lib\site-packages\scoringrules\_energy.py:52, in energy_score(forecasts, observations, m_axis, v_axis, backend)
     20 r"""Compute the Energy Score for a finite multivariate ensemble.
     21 
     22 The Energy Score is a multivariate scoring rule expressed as
   (...)
     49     The computed Energy Score.
     50 """
     51 backend = backend if backend is not None else backends._active
---> 52 forecasts, observations = multivariate_array_check(
     53     forecasts, observations, m_axis, v_axis, backend=backend
     54 )
     56 if backend == "numba":
     57     return energy._energy_score_gufunc(forecasts, observations)

File ~\anaconda3\Lib\site-packages\scoringrules\core\utils.py:31, in multivariate_array_check(fcts, obs, m_axis, v_axis, backend)
     29 m_axis = m_axis if m_axis >= 0 else fcts.ndim + m_axis
     30 v_axis = v_axis if v_axis >= 0 else fcts.ndim + v_axis
---> 31 _multivariate_shape_compatibility(fcts, obs, m_axis)
     32 return _multivariate_shape_permute(fcts, obs, m_axis, v_axis, backend=backend)

File ~\anaconda3\Lib\site-packages\scoringrules\core\utils.py:12, in _multivariate_shape_compatibility(fcts, obs, m_axis)
     10 o_shape_broadcast = o_shape[:m_axis] + (f_shape[m_axis],) + o_shape[m_axis:]
     11 if o_shape_broadcast != f_shape:
---> 12     raise ValueError(
     13         f"Forecasts shape {f_shape} and observations shape {o_shape} are not compatible for broadcasting!"
     14     )

ValueError: Forecasts shape (100, 3) and observations shape (100, 51, 3) are not compatible for broadcasting!

解决方案

1. 修正测试代码的调用逻辑

原代码直接调用test_energy_score(backend)存在两个错误:一是backend变量未定义,二是pytest装饰的测试函数不应手动直接调用。修改为如下可执行代码:

import numpy as np
from scoringrules._energy import energy_score
from scoringrules.backend import backends

RUN_TESTS = ["numba", "jax"]
BACKENDS = [b for b in backends.available_backends if b in RUN_TESTS]
for backend in RUN_TESTS:
    backends.register_backend(backend)

ENSEMBLE_SIZE = 51
N = 100
N_VARS = 3

def calculate_energy_score(backend):
    obs = np.random.randn(N, N_VARS)
    # 预测数组形状为(样本数, 集合大小, 变量数)
    fct = np.expand_dims(obs, axis=-2) + np.random.randn(N, ENSEMBLE_SIZE, N_VARS)
    # 明确指定集合维度(m_axis)和变量维度(v_axis)
    res = energy_score(fct, obs, m_axis=1, v_axis=2, backend=backend)
    print(f"Backend: {backend}, Result shape: {res.shape}")
    return res

# 遍历可用后端计算
for backend in BACKENDS:
    calculate_energy_score(backend)

2. 明确参数的维度要求

scoringrules的energy_score对输入数组形状有严格要求:

  • forecasts(预测)必须包含集合维度(m_axis)和变量维度(v_axis),示例形状为(样本数, 集合大小, 变量数)
  • observations(观测)的形状需与预测的非集合维度匹配,示例形状为(样本数, 变量数)
  • 调用时必须显式指定m_axis和v_axis,避免库自动识别维度出错

3. 针对(1,24)形状数据的处理

如果你的数据形状为(1,24),需根据实际含义调整:

  • 若24是变量数:预测需扩展集合维度,变为(1, 集合大小, 24),观测保持(1,24),调用时指定m_axis=1, v_axis=2
  • 若24是集合大小:预测形状应为(1,24, 变量数),观测为(1,变量数),对应调整m_axis和v_axis参数

通过以上调整,即可得到符合预期形状的Energy Score数组。


内容的提问来源于stack exchange,提问作者Alessandro Fabiani Bigaunah

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 00:55:55