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

如何确保Hypothesis生成的numpy数组包含全部允许元素?

如何确保Hypothesis生成的numpy数组包含全部允许元素?

嘿,这个需求我之前做测试的时候也遇到过,给你分享两个靠谱的解决办法,既能保证数组里同时有0和1,又能严格符合你要求的np.int8类型和固定形状:

方法一:用过滤策略筛选符合条件的数组

这是最直观的方式——先按你原来的逻辑生成数组,再过滤掉那些只包含0或者只包含1的无效样本:

import numpy as np
import hypothesis.strategies as st
from hypothesis.extra import numpy as st_np

# 定义基础的数组生成策略,和你原来的写法一致
base_arr_strategy = st_np.arrays(
    dtype=np.int8,
    shape=10,
    elements=st.integers(0, 1)
)

# 添加过滤条件,确保数组同时包含0和1
valid_arr_strategy = base_arr_strategy.filter(
    lambda arr: (0 in arr) and (1 in arr)
)

# 可以封装成复合策略方便调用
@st.composite
def valid_array(draw):
    return draw(valid_arr_strategy)

这个方法的好处是代码改动小,容易理解。不过要注意:如果你的数组形状很小(比如shape=1),过滤会导致Hypothesis很难找到有效样本,但你的shape是10,这种情况概率极低,完全不用担心效率问题。

方法二:自定义复合策略直接生成有效数组

如果想避免过滤带来的潜在开销,你可以直接构造满足条件的数组,从根源上保证有效性:

import numpy as np
import hypothesis.strategies as st

@st.composite
def guaranteed_valid_array(draw):
    # 随机选择至少1个、最多9个位置用来放0(剩下的位置自然放1,保证至少有1个1)
    zero_positions = draw(
        st.sets(st.integers(0, 9), min_size=1, max_size=9)
    )
    # 初始化数组为全0,然后把非0位置设为1
    arr = np.zeros(10, dtype=np.int8)
    arr[list(set(range(10)) - zero_positions)] = 1
    return arr

这种方式的优势是效率更高,因为我们直接生成符合要求的数组,不需要Hypothesis反复尝试过滤无效样本。而且完全能保证数组里同时存在0和1,类型和形状也严格符合你的要求。

不管用哪种方法,最终生成的数组都是np.int8类型、形状为10的numpy数组,完美匹配你的需求~

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.15 10:48:04