如何简洁向列表值字典追加数据以适配DataFrame及JSON存储需求?
解决方案
两种优化方案均完全兼容你提出的两个核心需求:无需额外适配即可直接转DataFrame、直接导出JSON。
方案1:轻量化辅助函数(零侵入适配现有代码)
仅需新增一个全局通用辅助函数,不需要修改原有数据结构,就能把多行append简化为单行调用,是性价比最高的改造方案:
from typing import Dict def append_row(result_dict: Dict, row_data: Dict): """ 批量给对应键的列表追加值,自动匹配已有键名 """ for key, value in row_data.items(): result_dict[key].append(value)
调用时只需要传入单行的键值对即可,不管有多少个字段都只需要写一次调用逻辑,还能避免漏写某个键导致后续列表长度不一致的报错。
方案2:dataclass类型约束(适合字段多的复杂场景)
如果你的实验字段固定且数量较多,用dataclass可以获得IDE自动补全、类型校验能力,避免手敲键名出错:
from dataclasses import dataclass, asdict from typing import List import pandas as pd # 提前定义实验记录的固定字段 @dataclass class ExpRecord: model_name: str seed: int identifier: str val_mse: float # 存储所有记录的列表 records: List[ExpRecord] = [] # 循环内添加记录 records.append(ExpRecord( model_name=model_name, seed=seed, identifier=f"fold {fold}", val_mse=val_score )) # 最终转成你需要的字典结构,直接对接后续逻辑 result_dict = pd.DataFrame([asdict(r) for r in records]).to_dict(orient="list")
优化后业务代码示例(采用方案1)
from typing import Dict import numpy as np import pandas as pd from sklearn import metrics # 全局通用辅助函数,定义一次可在所有同类场景复用 def append_row(result_dict: Dict, row_data: Dict): for key, value in row_data.items(): result_dict[key].append(value) # 原有初始化逻辑完全不变 result_dict: Dict = {"model_name": [], "seed": [], "identifier": [], "val_mse": []} model_name = model.__class__.__name__ for fold in range(1, num_folds + 1): train_df = df_folds[df_folds["fold"] != fold].reset_index(drop=True) val_df = df_folds[df_folds["fold"] == fold].reset_index(drop=True) X_train, y_train = train_df[predictor_col].values, train_df[target_col].values X_val, y_val = val_df[predictor_col].values, val_df[target_col].values model.fit(X_train, y_train) y_val_pred = model.predict(X_val) val_score = metrics.mean_squared_error(y_true=y_val, y_pred=y_val_pred) # 原4行append简化为1次调用 append_row(result_dict, { "model_name": model_name, "seed": seed, "identifier": f"fold {fold}", "val_mse": val_score }) avg_val_score = np.mean(result_dict["val_mse"], axis=None) standard_error = np.std(result_dict["val_mse"], axis=None) / np.sqrt(num_folds) # 追加平均结果同样可以用该方法,避免漏写字段 append_row(result_dict, { "model_name": model_name, "seed": seed, "identifier": "average_score", "val_mse": avg_val_score })
内容的提问来源于stack exchange,提问作者ilovewt
相关产品推荐
相关产品推荐

