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

如何用Hypothesis生成带行列依赖的复杂Pandas DataFrame?

问题

有没有优雅的方法用hypothesis直接生成带有内部行、列依赖关系的复杂pandas DataFrame?例如需要包含以下列:

[longitude][latitude][some-text-meta][some-numeric-meta][numeric-data][some-junk][numeric-data][…

地理坐标可随机选取,但需来自同一大致区域(比如跨地球两端的点无法正常投影),这可以通过先选区域再生成对应坐标列实现。现有代码可生成关联的numpy数组:

@st.composite
def plaus_spamspam_arrs(
    draw,
    st_lonlat=plaus_lonlat_arr,
    st_values=plaus_val_arr,
    st_areas=plaus_area_arr,
    st_meta=plaus_meta_arr,
    bounds=ARR_LEN,
):
    """Returns plausible spamspamspam arrays"""
    size = draw(st.integers(*bounds))
    coords = draw(st_lonlat(size=size))
    values = draw(st_values(size=size))
    areas = draw(st_areas(size=size))
    meta = draw(st_meta(size=size))
    return PlausibleData(coords, values, areas, meta)

但示例中的数值列存在行级依赖(比如需要按行和动态归一化),如果单独生成各部分再用concat拼接会破坏draw的收缩功能,而且需要可扩展的优雅方案,而非临时hack。尝试过builds但没找到方法,希望得到示例启发。

更新:已能通过以下方式生成列,但需要更优雅的方式整合行级依赖:

@st.composite
def plaus_df_inputs(
    draw, *, nrows=None, ncols=None, nrow_bounds=ARR_LEN, ncol_bounds=COL_LEN
):
    """Returns …"""
    box_lon, box_lat = draw(plaus_box_geo())
    ncols_jnk = draw(st.integers(*ncol_bounds)) if ncols is None else ncols
    ncols_val = draw(st.integers(*ncol_bounds)) if ncols is None else ncols
    keys_val = draw(plaus_smp_key_elm(size=ncols_val))
    nrows = draw(st.integers(*nrow_bounds)) if nrows is None else nrows
    cols = (
        plaus_df_cols_lonlat(lons=plaus_lon(box_lon), lats=plaus_lat(box_lat))
        + plaus_df_cols_meta()
        + plaus_df_cols_value(keys=keys_val)
        + draw(plaus_df_cols_junk(size=ncols_jnk))
    )
    random.shuffle(cols)
    return draw(st_pd.data_frames(cols, index=plaus_df_idx(size=nrows)))

其中子策略示例:

@st.composite
def plaus_df_cols_junk(
    draw, *, size=1, names=plaus_meta(), dtypes=plaus_dtype(), unique=False
):
    """Returns strategy for list of columns of plausible junk data."""
    result = set()
    for _ in range(size):
        result.add(draw(names.filter(lambda name: name not in result)))
    return [
        st_pd.column(name=result.pop(), dtype=draw(dtypes), unique=unique)
        for _ in range(size)
    ]

现需更优雅的方案来处理行级依赖。


优雅处理行级依赖的方案

1. 先生成完整行数据,再组装DataFrame

放弃先定义列策略再生成DataFrame的方式,改用@st.composite直接生成每行的完整数据(包含所有列的关联值),再将行数据列表转换为DataFrame。这种方式能完美保留hypothesis的收缩功能,因为所有行级依赖都在单次draw流程内处理。

示例代码:

import pandas as pd
import hypothesis as st
import hypothesis.strategies as st
from hypothesis.extra.pandas import indices

# 定义单个行数据的策略
@st.composite
def plausible_row(draw, box_lon, box_lat, value_keys):
    # 生成当前行的地理坐标(绑定到指定区域)
    lon = draw(st.floats(min_value=box_lon[0], max_value=box_lon[1]))
    lat = draw(st.floats(min_value=box_lat[0], max_value=box_lat[1]))
    
    # 生成元数据
    text_meta = draw(st.text(alphabet=st.characters(whitelist_categories='L'), min_size=1))
    numeric_meta = draw(st.integers(min_value=0, max_value=1000))
    
    # 生成带行级依赖的数值列:按行求和归一化
    raw_values = draw(st.lists(st.floats(min_value=0, max_value=10), min_size=len(value_keys), max_size=len(value_keys)))
    row_sum = sum(raw_values) or 1  # 避免除以0
    normalized_values = [v / row_sum for v in raw_values]
    
    # 生成垃圾数据列(简化版)
    junk_name = draw(st.text(alphabet=st.characters(whitelist_categories='L'), min_size=5))
    junk_dtype = draw(st.sampled_from([int, float, str]))
    if junk_dtype == int:
        junk_val = draw(st.integers(min_value=-100, max_value=100))
    elif junk_dtype == float:
        junk_val = draw(st.floats(min_value=-100.0, max_value=100.0))
    else:
        junk_val = draw(st.text(min_size=1))
    junk_cols = {junk_name: junk_val}
    
    # 组装成单行字典
    row_data = {
        'longitude': lon,
        'latitude': lat,
        'some-text-meta': text_meta,
        'some-numeric-meta': numeric_meta,
        **dict(zip(value_keys, normalized_values)),
        **junk_cols
    }
    return row_data

# 生成完整DataFrame的复合策略
@st.composite
def plausible_dataframe(draw, nrow_bounds=(5, 20), ncol_bounds=(1, 3)):
    # 确定全局参数:地理区域、数值列数量/名称、行数
    box_lon, box_lat = draw(plaus_box_geo())  # 复用已有区域策略
    ncols_val = draw(st.integers(*ncol_bounds))
    value_keys = draw(st.lists(st.text(min_size=3, max_size=10), min_size=ncols_val, max_size=ncols_val, unique=True))
    nrows = draw(st.integers(*nrow_bounds))
    
    # 生成所有行数据
    rows = draw(st.lists(plausible_row(box_lon=box_lon, box_lat=box_lat, value_keys=value_keys), min_size=nrows, max_size=nrows))
    
    # 转换为DataFrame并随机打乱列顺序
    df = pd.DataFrame(rows)
    shuffled_cols = draw(st.permutations(df.columns))
    df = df[shuffled_cols]
    
    # 生成自定义索引
    df.index = draw(indices(min_size=nrows, max_size=nrows))
    return df

2. 处理行与行之间的关联依赖

如果需要行与行之间的依赖(比如前一行数值影响后一行),可以在生成行列表时逐步构建:

@st.composite
def plausible_dataframe_with_inter_row_deps(draw, nrow_bounds=(5, 20)):
    nrows = draw(st.integers(*nrow_bounds))
    box_lon, box_lat = draw(plaus_box_geo())
    value_keys = ['val1', 'val2']
    
    rows = []
    prev_val = draw(st.floats(min_value=0, max_value=10))
    for _ in range(nrows):
        # 当前行val1依赖前一行的val1
        curr_val1 = draw(st.floats(min_value=prev_val * 0.8, max_value=prev_val * 1.2))
        curr_val2 = draw(st.floats(min_value=0, max_value=10))
        row_sum = curr_val1 + curr_val2 or 1
        normalized_val1 = curr_val1 / row_sum
        normalized_val2 = curr_val2 / row_sum
        
        row = {
            'longitude': draw(st.floats(min_value=box_lon[0], max_value=box_lon[1])),
            'latitude': draw(st.floats(min_value=box_lat[0], max_value=box_lat[1])),
            'some-text-meta': draw(st.text(min_size=1)),
            'val1': normalized_val1,
            'val2': normalized_val2
        }
        rows.append(row)
        prev_val = curr_val1
    
    return pd.DataFrame(rows)

3. 复用已有列策略的混合方案

如果想保留之前的列策略定义,同时处理行级依赖,可以先生成基础DataFrame,再对需要行级依赖的列进行后处理——注意后处理要在draw流程内完成,避免破坏收缩功能:

@st.composite
def plausible_df_mixed(draw):
    # 用已有逻辑生成基础DataFrame
    base_df = draw(plaus_df_inputs())
    
    # 对以'val_'开头的数值列做行级归一化
    value_cols = [col for col in base_df.columns if col.startswith('val_')]
    if value_cols:
        row_sums = base_df[value_cols].sum(axis=1).replace(0, 1)
        base_df[value_cols] = base_df[value_cols].div(row_sums, axis=0)
    
    return base_df

这种方式适合依赖关系简单的场景,若需要完全保留收缩能力,优先选择第一种方案。

内容的提问来源于stack exchange,提问作者curlew77

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 20:05:13