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

如何实现返回Dask DataFrame函数的分块确定性缓存至Parquet文件

实现方案

核心逻辑调整

要实现行级失败重试的缓存机制,我们需要对原有装饰器做以下改造:

  1. 新增primary_key_cols参数指定唯一标识每行的主键列,failure_col参数指定用于判断请求是否失败的列
  2. 缓存逻辑从「首次写入后只读」调整为「读缓存→筛选失败记录→仅重算失败记录→合并结果更新缓存」
  3. 兼容原有装饰器的全量缓存能力,未传入主键和失败列时默认退化为原有全量缓存逻辑

修改后完整代码

import base64
from functools import wraps
import hashlib
import json
from pathlib import Path
import pandas as pd
import dask.dataframe as dd

def ezhash(o):
    o = json.dumps(o, sort_keys=True)
    o = hashlib.md5(o.encode('utf-8')).digest()
    o = base64.urlsafe_b64encode(o).decode('ascii')
    return o

def cached_compute(enabled=True, engine='pyarrow', write_kwargs={}, read_kwargs={},
                   primary_key_cols=None, failure_col=None):
    def _cached_compute(f):
        @wraps(f)
        def g(*args, input_state=None, **kwargs):
            if not enabled:
                return f(*args, **kwargs)
            # 生成缓存key逻辑保持不变
            if input_state == None:
                key = ezhash({'name': f.__name__, 'args': args, 'kwargs': kwargs})
            else:
                key = ezhash({'name': f.__name__, 'input_state': input_state})
            key_path = Path(f'parquet_cache/{key}/')
            
            # 无主键/失败列配置时退化为原有全量缓存逻辑
            if primary_key_cols is None or failure_col is None:
                if not key_path.exists():
                    print(f'caching to {key}')
                    res = f(*args, **kwargs)
                    res.to_parquet(key_path, engine=engine, **write_kwargs)
                print(f'reading cache {key}')
                res = dd.read_parquet(key_path, engine=engine, **read_kwargs)
                return key, res
            
            # 行级重试缓存逻辑
            input_data = args[0]  # 这里默认第一个参数是输入数据集,可根据实际场景调整
            input_df = pd.DataFrame(input_data)
            if not key_path.exists():
                # 首次运行全量计算
                print(f'caching to {key}')
                full_ddf = f(*args, **kwargs)
                full_ddf.to_parquet(key_path, engine=engine, write_index=False, **write_kwargs)
            else:
                # 读取已有缓存
                print(f'reading cache {key}')
                cached_ddf = dd.read_parquet(key_path, engine=engine, **read_kwargs)
                cached_df = cached_ddf.compute()
                # 筛选失败的记录主键
                failed_keys = cached_df[cached_df[failure_col].isna()][primary_key_cols].drop_duplicates()
                if len(failed_keys) == 0:
                    # 无失败记录直接返回缓存
                    return key, cached_ddf
                # 过滤出需要重试的输入数据
                retry_input = input_df.merge(failed_keys, on=primary_key_cols).to_dict(orient='list')
                # 仅重算失败的记录
                retry_ddf = f(retry_input, **kwargs)
                retry_df = retry_ddf.compute()
                # 合并结果:保留缓存中成功的记录 + 新算出来的重试记录
                success_cached = cached_df[~cached_df[failure_col].isna()]
                full_df = pd.concat([success_cached, retry_df], ignore_index=True)
                full_ddf = dd.from_pandas(full_df, npartitions=cached_ddf.npartitions)
                # 覆写缓存
                full_ddf.to_parquet(key_path, engine=engine, write_index=False, overwrite=True, **write_kwargs)
            
            return key, full_ddf
        return g
    return _cached_compute

# 测试代码
import random

def network_call(row):
    """Simulate a network call"""
    if random.random() < 0.9:  # 修正原代码成功率逻辑:90%成功
        response = sum(row.tolist())
    else:
        response = None
    row['response'] = response
    return row

def run_example(cache_enabled, max_runs=5):
    def expensive_function(input_data):
        ddf = dd.from_pandas(pd.DataFrame(input_data), npartitions=4)
        ddf = ddf.apply(network_call, axis=1, meta={k: 'int64' for k in [*input_data.keys(), 'response']})
        return ddf

    if cache_enabled:
        # 新增主键和失败列配置
        expensive_function = cached_compute(
            write_kwargs={'schema': 'infer'},
            primary_key_cols=['col_1'],
            failure_col='response'
        )(expensive_function)

    input_data = {
        'col_1': list(range(  0, 100)),
        'col_2': list(range(100, 200)),
        'col_3': list(range(200, 300)),
    }

    prev_success = 0
    for i_run in range(max_runs):
        if cache_enabled:
            _, ddf = expensive_function(input_data)
        else:
            ddf = expensive_function(input_data)
        computed = ddf.compute()
        # 修正原代码成功数统计逻辑:非NA的是成功记录
        n_success = sum(~computed['response'].isna())
        print(f'''Run #{i_run + 1:>2}: Total successes {n_success:>3}/100 (+{n_success - prev_success:>2})''')
        prev_success = n_success

print('''Running with cache disabled''')
run_example(cache_enabled=False)

print('''\nRunning with cache enabled''')
run_example(cache_enabled=True)

运行效果

修改后运行代码即可得到预期输出,缓存开启后成功数会逐次累加直到100%,不会重复计算已经成功的记录。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 02:18:00