COSTAR 反事实自监督 Transformer 完整实战指南:环境搭建、数据集配置与实验复现

发布时间:2026/9/19 17:52:41
COSTAR 反事实自监督 Transformer 完整实战指南:环境搭建、数据集配置与实验复现 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载本文以 Google Research 开源仓库中的 COSTAR/README.md 为主线系统讲解 COSTARCOunterfactual Self-supervised TrAnsformeR对应论文 arXiv 2311.00886在时序反事实推断场景中的完整使用流程。读者将掌握从 conda 环境搭建、预生成数据导入、三种内置数据集肿瘤生长 / 半合成 MIMIC-III / M5的接入方式到基于 costar.sh 与 baselines.sh 复现 zero-shot 迁移、数据高效迁移与标准监督学习三种实验设置以及如何通过 Hydra 配置系统与源码级实现理解其内部机理。一、COSTAR 项目概览COSTAR 是一套面向基于时序观测数据评估治疗策略效果counterfactual treatment outcome estimation的自监督表示学习框架。它借助反事实counterfactual自监督信号训练编码器再通过轻量估计头estimation head完成 1 步与多步2~6 步反事实结果预测目标是在少量标注/治疗数据下也能做出可靠的时序反事实推断。从源码结构看仓库主体由四部分构成相对路径均以仓库根目录为起点目录/文件作用COSTAR/runnables各方法的训练/评估入口脚本COSTAR/src/data数据集实现cancer_sim、mimic_iii、m5与统一的 DatasetCollection 抽象COSTAR/src/models模型实现COSTAR 的 COTrans、CT、CRN、RMSN、G-Net、MSM 等COSTAR/configHydra 配置文件dataset、backbone、exp 参数COSTAR/scripts_release一键复现全部实验的 shell 脚本下文所有命令均需在仓库根目录执行示例中以PYTHONPATH.将仓库根目录加入模块搜索路径。二、环境搭建Environment Setup按照 README环境搭建分为三步# 1. 基于 environment.yml 创建 conda 环境 conda env create -f environment.yml conda activate costar # 2. 为兼容已保存的数据将 pandas 固定到 1.5.3 conda install -c conda-forge pandas1.5.3 # 3. 配置 wandb 用于实验日志 wandb login关键依赖解读来自 COSTAR/environment.yml该环境文件锁定了完整可复现的依赖矩阵其中与训练直接相关的核心组件包括Python 3.9.15PyTorch 1.12.1py3.9_cuda11.3_cudnn8.3.2配套 torchvision 0.13.1、torchaudio 0.12.1CUDA toolkit 11.3.1——这是经测试的 CUDA 11.3 组合注意不要随意升级pytorch-lightning 1.4.5所有 runnable 脚本中的Trainer、EarlyStopping、ModelCheckpoint均来自该版本Hydra 1.3.2 OmegaConf 2.3.0支撑实验的配置组合系统dataset...、backbone...语法wandb 0.15.3实验日志记录mlflow 2.3.2作为备选/兼容组件numpy 1.24.3、scikit-learn 1.2.2、scipy 1.10.1、einops 0.6.1、tables 3.8.0、pyarrow 11.0.0数据处理与序列存储ray 2.4.0、tqdm、torch-ema 0.3分别用于并行调参与权重 EMA。README 特别提示创建环境后必须执行conda install -c conda-forge pandas1.5.3因为 environment.yml 中 pip 段默认安装了 pandas 2.0.1而降级到 1.5.3 是为了兼容仓库随附的预生成数据pickle 序列化格式。三、数据集准备Dataset Setup3.1 预生成数据下载tar xzvf all_needed_data.tar.gz将预生成数据包解压到项目根目录即可压缩包对应预生成的 MIMIC-III 半合成数据等。若要引入自定义数据集请遵循 COSTAR/dataset_instructions.md 中的规范详见本文第七节。3.2 三种内置数据集数据集类型是否需要额外配置对应实现Tumor growth肿瘤生长合成无需数据即时生成src/data/cancer_sim/dataset.pySemi-synthetic MIMIC-III半合成无需加载预生成数据可重跑仿真较慢src/data/mimic_iii/semi_synthetic_dataset.pyM5真实世界无需数据即时生成src/data/m5/real_dataset.py肿瘤生长与 M5两种数据均在首次运行时按配置即时生成无需下载。肿瘤生长为合成数据具备测试阶段反事实治疗结局的真实标签ground truth支持完整的 one-step 与多步反事实评估M5 为真实世界零售销量数据foods_household_5k类别子集只能以事实治疗结局评估。半合成 MIMIC-III默认直接加载预生成数据若需重新运行仿真过程速度较慢使用 README 提供的两条命令PYTHONPATH. python runnables/train_enc_dec.py -m dataset/mimic3_syn_age0-3_all backbonecrn_noncausal_troff backbone/crn_noncausal_troff_hparams/mimic3_syntheticall dataset.data_gen_n_jobs8 exp.gen_data_onlyTrue exp.seed17,43,44,91,95 exp.tags230814_mimicsyn_all_data PYTHONPATH. python runnables/train_enc_dec.py -m dataset/mimic3_syn_age0-3_all_srctest backbonecrn_noncausal_troff backbone/crn_noncausal_troff_hparams/mimic3_syntheticall dataset.data_gen_n_jobs8 exp.gen_data_onlyTrue exp.seed17,43,44,91,95 exp.tags230910_mimicsyn_all_data参数要点-mHydra 多任务multirun模式配合逗号分隔的种子列表逐一运行dataset/mimic3_syn_age0-3_all选择 0~3 岁年龄段的全量数据配置0-3_all_srctest为源域训练 源域测试评估配置backbonecrn_noncausal_troff使用非因果 CRN 骨干此处仅用于数据生成管线exp.gen_data_onlyTrue只生成数据并保存对应 train_enc_dec.py 中exp.gen_data_only分支最终调用dataset_collection.save_to_pkl不训练dataset.data_gen_n_jobs8并行数据生成进程数exp.seed17,43,44,91,95官方实验统一使用的 5 个随机种子MIMIC-III 相关实验。四、运行 COSTAR 实验COSTAR 主实验的全部命令收录在 COSTAR/scripts_release/costar.sh 中涵盖三类设置zero-shot 迁移、数据高效迁移few-shot与标准监督学习。4.1 肿瘤生长数据集# zero-shot transfer /># zero-shot transfer /># zero-shot transfer /># 标准监督学习先运行为迁移实验提供预训练 checkpoint PYTHONPATH. python3 runnables/train_msm.py -m datasetcancer_sim target_datasetcancer_sim dataset.tr_te_typesrc_src backbonemsm backbone/msm_hparams/cancer_sim_dseed_10\10_tuned\ exp.seed10,101,1010,10101,101010 exp.tags231005_msm_tcs_srctest PYTHONPATH. python3 runnables/train_enc_dec.py -m datasetcancer_sim backbonecrn_noncausal_troff backbone/crn_noncausal_troff_hparams/cancer_sim_dseed_1010_tuned exp.seed10,101,1010,10101,101010 exp.tagsconclude_tcs_cancer_sim_seed_var # zero-shot transfer加载预训练不更新权重 PYTHONPATH. python3 runnables/train_rmsn.py -m backbonermsn backbone/rmsn_hparams/cancer_sim_dseed_1010_tuned datasetcancer_sim dataset.coeff0 dataset.src_coeff10 dataset.few_shot_sample_num100 exp.finetune_tagconclude_tcs_cancer_sim_seed_var exp.max_epochs0 exp.finetune_ckpt_typelast exp.seed10,101,1010,10101,101010 exp.tags231005_bs_tcs_10_0shot #>dataset: _target_: src.data.SyntheticCancerDatasetCollection name: tumor_generator coeff: ??? # 混杂系数 gamma必须由命令行提供 chemo_coeff: ${dataset.coeff} radio_coeff: ${dataset.coeff} seed: ${exp.seed} num_patients: {train: 10000, val: 1000, test: 1000} window_size: 15 # 有偏治疗分配窗口 lag: 0 # 治疗分配窗口滞后 max_seq_length: 60 # 序列最大长度 projection_horizon: 5 # τ 步预测范围τ projection_horizon 1 cf_seq_mode: sliding_treatment # sliding_treatment / random_trajectories / no_cf_treatment val_batch_size: 512 treatment_mode: multiclass # multiclass / multilabelRMSN 需要 model: dim_treatments: 4 dim_vitals: 0 dim_static_features: 1 dim_outcomes: 1config.yaml中通用 dataset 参数还包括val_batch_size、treatment_mode、few_shot_sample_num0 时重采样训练样本构造 few-shot、data_gen_n_jobs: 8、use_few_shot: False、max_number: -1限制样本量的调试开关。6.3 骨干模型配置backboneconfig/backbone/cotrans_tcn_comp_contrast.yaml 是 COSTAR 主模型的配置骨架model.name: CT_TCN模型标识model.rep_encoder._target_: src.models.rep_est.cotrans.COTransMoCoEncoderCOTrans MoCo 风格对比学习编码器use_comp_contrast: True启用成分component对比——即对时间维/特征维拆分的子表示也做对比对齐对应源码中_encode(..., return_comp_repsTrue)返回enc comp_reps列表cotrans.pytemporal_positional_encoding: 绝对/相对位置编码max_relative_position: 15feature_positional_encoding特征位置编码num_layer: 1、num_heads: 2Transformer 层数与注意力头数momentum/temperatureMoCo 动量与 InfoNCE 温度由具体 hparams 文件覆盖如cotrans_tcn_hparams/cancer_sim/10.yamlmodel.est_head._target_: src.models.rep_est.MoCo.TCNEstHeadTCN 估计头step_mse_loss_weights_type: avg命令行可覆盖为inverse即按步数倒数加权多步损失tune_hparams: False是否启用 Ray 超参搜索。各数据集对应具体超参文件位于config/backbone/cotrans_tcn_hparams/{cancer_sim,m5_real,mimic3_synthetic}/例如 MIMIC-III 用all_es、M5 用sales_es。七、源码级理解COSTAR 是如何工作的7.1 训练主流程COSTAR/runnables/train_rep_est.py 是 COSTAR 的训练/评估入口核心流程为seed_everything固定种子通过load_saved_data尝试加载已保存数据否则instantiate(args.dataset)即时生成process_data_rep_est()完成数据预处理并从训练集推导dim_outcomes、dim_treatments、dim_vitals、dim_static_features等维度train_rep_est.py实例化rep_encoderCOTransMoCoEncoder并训练ModelCheckpoint监控rep_encoder-val_metric实例化est_headTCNEstHead在源域训练src头随后切换到target_dataset或替换为 test 子集分别训练/评估dst-zero-shotmax_epochs0零更新与dst头依次在test_cf_one_step1 步反事实与test_cf_treatment_seq多步反事实序列上计算归一化 RMSE 并记录。值得注意的是COSTAR 的对比预训练通过 MoCo 动量编码器对同一样本做两次数据增强augment由 scale → shift → jitter 组合而成见 cotrans.py并同时最大化全局表示与成分表示的互信息use_comp_contrast这为下游 few-shot 估计头提供了鲁棒的时间表示。7.2 数据集抽象层COSTAR/src/data/dataset_collection.py 定义了两个基类SyntheticDatasetCollection适用于具备反事实标签的合成/半合成数据持有train_f、val_f、test_cf_one_step、test_cf_treatment_seq四个子集并提供process_data_encoder、process_data_decoder、process_data_rep_est等预处理编排方法RealDatasetCollection适用于真实世界数据如 M5仅含test_f事实评估子集。dataset_instructions.md要求自定义数据集在初始化时提供train_f/val_f事实结局、test_cf_one_step与test_cf_treatment_seq合成数据、test_f真实数据、seed、projection_horizon、autoregressive是否把上一步结局并入协变量、has_vitals是否存在结局之外的协变量、train_scaling_params归一化参数、max_seq_length。7.3 数据格式化与预处理约定子集类需继承torch.utils.data.Dataset并实现get_scaling_params与process_data。process_data须把原始数据整理为字典至少包含键含义与形状prev_treatments时间[0, T-1]的治疗变量[N, T, treatment_feature_num]current_treatments时间[1, T]的治疗变量[N, T, treatment_feature_num]prev_outputs时间[0, T-1]的输出[N, T, output_feature_num]outputs时间[1, T]的输出[N, T, output_feature_num]vitals除历史结局外的协变量[N, T, vital_feature_num]static_features不随时间变化的主体属性[N, static_feature_num]sequence_lengthspadding 前各序列真实长度[N]active_entries有效时间步二进制掩码[N, T, 1]此外不同方法还需实现特定预处理MSM/RMSN/CRN 需要process_sequential滚动起点将序列爆炸为多段子序列并附加init_state、original_index、active_encoder_r、unscaled_outputs等键、process_sequential_test、process_autoregressive_test以预测占位prev_outputs实现自回归推演与explode_trajectoriesCT 需要process_sequential_multi利用future_past_split整数标记历史段终点而 COSTCOSTAR 论文中的另一实现变体将数据重组逻辑直接内联在训练/评估代码中无需额外预处理步骤。八、自定义数据集接入指南若要把自有数据接入 COSTAR 流程按 COSTAR/dataset_instructions.md 的步骤执行明确数据范式合成数据测试时存在反事实治疗结局真值继承SyntheticDatasetCollection真实数据继承RealDatasetCollection可参考 cancer_sim/dataset.py、semi_synthetic_dataset.py 与 m5/real_dataset.py 三个范例实现集合类在初始化中构造上节列出的各成员变量与process_data字典实现子集类复用示例中的__getitem__/__len__补齐get_scaling_params与process_data含归一化与时间对齐t 时刻观测协变量须与其后紧邻的治疗/结局对齐按方法实现预处理根据目标方法MSM/RMSN/CRN、CT、COST实现对应的process_sequential系列方法在config/dataset下登记配置将构造参数写入 YAML即可通过datasetyour_dataset接入训练脚本。结语与注意事项复现时请严格按脚本注释的执行顺序运行标准监督学习 → zero-shot 迁移 → 数据高效迁移因为后两者依赖前者的预训练 checkpoint通过exp.finetune_tag自动衔接若使用 GPU 训练确认 CUDA 11.3 与 PyTorch 1.12.1 组合与本地驱动匹配exp.gpus支持多卡指标以 wandb 中带-test_rmse后缀的键名为准三个数据集的前缀与步数规则不同见 4.4 节表半合成 MIMIC-III 默认使用预生成数据重新仿真仅在你需要调整仿真过程时执行且耗时较长。如需更深入地理解 COSTAR 的模型细节与完整实验配置可直接查阅仓库中的 runnables 入口脚本、src/models/rep_est 实现与 config 配置目录。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐NeMo Speech 自监督预训练数据集构建与配置实战指南NeMo Speech 自监督预训练数据集构建与配置实战指南 自监督学习Self Supervised Learning, SSL无需人工标注即可从语音数据人工智能语音音频大模型深度学习Mill构建工具为JVM生态提供高性能构建解决方案Mill构建工具为JVM生态提供高性能构建解决方案 Mill是一个专为Java、Scala和Kotlin设计的现代化构建工具旨在通过创新的架构设计解决传统Jgh_mirrors/di/dino论文复现指南Emerging Properties in Self-Supervised Vision Transformers实验重现gh_mirrors/di/dino论文复现指南Emerging Properties in Self Supervised Vision Transform人工智能深度学习计算机视觉预训练创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考