如何实现返回Dask DataFrame函数的分块确定性缓存至Parquet文件
实现方案
核心逻辑调整
要实现行级失败重试的缓存机制,我们需要对原有装饰器做以下改造:
- 新增
primary_key_cols参数指定唯一标识每行的主键列,failure_col参数指定用于判断请求是否失败的列 - 缓存逻辑从「首次写入后只读」调整为「读缓存→筛选失败记录→仅重算失败记录→合并结果更新缓存」
- 兼容原有装饰器的全量缓存能力,未传入主键和失败列时默认退化为原有全量缓存逻辑
修改后完整代码
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
相关产品推荐
相关产品推荐

