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
相关产品推荐
相关产品推荐

