尧图网络 高端网站定制 · 原创设计
免费咨询热线
400-888-6620
免费获取方案
DoWhy causal_prediction.algorithms:ERM 与 CACM 因果预测算法包的源码级详解
机器学习数据分析【免费下载链接】dowhyDoWhy is a Python library for causal inference that supports explicit modeling and testing of causal assumptions. DoWhy is based on a unified language for causal inference, combining causal graphical models and potential outcomes frameworks.项目地址https://gitcode.com/gh_mirrors/do/dowhy点击查看免费下载本文以 DoWhy 官方 API 文档中dowhy.causal_prediction.algorithms包docs/source/dowhy.causal_prediction.algorithms.rst为主体系统讲解该包下五个模块——base_algorithm、erm、cacm、regularization、utils——的类结构、构造参数、训练循环实现与 MMD 正则化原理。读完后你将理解 DoWhy 因果预测Causal Prediction任务中 ERM 与 CACM 两种算法的设计差异、各超参数的含义与默认值并能结合 示例 notebook 与 单元测试 在自己的数据上完成因果表征学习。1. 包概览五个模块的分工API 文档以 Sphinxautomodule指令声明了该包的完整构成即以下五个子模块加上包内容本身模块文件核心内容base_algorithmdowhy/causal_prediction/algorithms/base_algorithm.py基类PredictionAlgorithm继承自pl.LightningModule定义默认验证/测试流程与优化器工厂ermdowhy/causal_prediction/algorithms/erm.pyERM经验风险最小化算法标准跨域联合训练cacmdowhy/causal_prediction/algorithms/cacm.pyCACMCausally Adaptive Constraint Minimization基于属性类型施加 MMD 约束regularizationdowhy/causal_prediction/algorithms/regularization.pyRegularizer无条件 / 条件独立正则化的具体实现utilsdowhy/causal_prediction/algorithms/utils.pygaussian_kernel、mmd_compute两个底层核函数从源码结构看整个包的依赖方向是单向的ERM和CACM都继承PredictionAlgorithmCACM再组合RegularizerRegularizer调用utils中的 MMD 工具函数。这套结构服务于 DoWhy 的因果预测任务——与只使用全部特征的标准机器学习不同这里的预测算法试图利用用户提供的辅助属性auxiliary features的因果关系知识学习一个因果特征表征。官方背景文档见 dowhy.causal_prediction.rst 及用户指南 causal_tasks/causal_prediction/index.rst。2. 基类 PredictionAlgorithmLightning 训练骨架base_algorithm.py 中定义了包的基类class PredictionAlgorithm(pl.LightningModule): def __init__(self, model, optimizer, lr, weight_decay, betas, momentum):文档字符串声明了它实现pl.LightningModule的默认方法这些方法在调用fit()时被 PyTorch Lightning 回调。构造参数及取值如下以源码默认值和签名核对参数含义默认值约束model训练用的神经网络模块约定为torch.nn.Sequential(featurizer, classifier)必填特征器与分类器均为torch.nn.Moduleoptimizer优化算法名称子类默认Adam仅支持Adam、AdamW、SGDlr学习率子类默认1e-3—weight_decay权重衰减子类默认0.0—betasAdam 的 (beta1, beta2) 动量系数子类默认(0.9, 0.999)—momentumSGD 动量子类默认0.9—源码在第 29 行对优化器名称做白名单校验非法名称会抛出ValueError这一点被 test_causal_prediction.py 中的test_prediction_algorithm_rejects_unsupported_optimizer显式验证传入Adagrad期望报not implemented。基类还固化了三个训练协议方法training_step(train_batch, batch_idx)基类中只raise NotImplementedError作为子算法必须覆写的扩展点L33-L39validation_step/test_step将 batch 中所有环境的数据torch.cat拼接后前向计算交叉熵损失与准确率并按 epoch 记录val_acc/val_loss或test_acc/test_loss到进度条L41-L81configure_optimizers()依据构造时的字符串名称构造torch.optim.Adam/AdamW/SGD并把lr、weight_decay、betas、momentum透传进去L83-L101。测试分别断言了 Adam 与 SGD 两种配置返回对应类型优化器。一个值得注意的实现细节configure_optimizers使用self.parameters()而非self.model.parameters()来收集可训练参数。由于模型被存为 LightningModule 属性这一写法在 Lightning 参数管理下是等价的但阅读源码时应留意这一点。3. ERM跨域联合经验风险最小化erm.py 中的ERM是基类之上的最薄一层实现。其__init__签名含默认值为ERM(model, optimizerAdam, lr1e-3, weight_decay0.0, betas(0.9, 0.999), momentum0.9)training_step的核心逻辑L23-L39只有四步从train_batch一个 list每个元素是某个环境的(x, y, a)元组中把所有环境的样本x、y拼成一个大张量整体前向out self.model(x)即直接调用 Sequential 模型featurizer → classifier用F.cross_entropy(out, y)计算单一交叉熵损失记录train_acc、train_loss并返回 loss。可以看到 ERM 完全丢弃了属性张量a与环境信息——它把所有环境视作同一个数据源做池化训练。这正是后续 CACM 要修正的偏差来源池化目标隐含假设各环境数据同分布而 CACM 的目标恰恰是在训练时消除特征表征与特定属性之间的条件依赖。4. CACM带属性类型正则的因果预测CACMCausally Adaptive Constraint Minimization论文Modeling the Data-Generating Process is Necessary for Out-of-Distribution Generalization, Kaur et al., arXiv:2206.07837见 cacm.py 的 docstring在 ERM 池化损失之外按每个属性的“与标签 Y 的因果关系类型”施加不同的 MMD 惩罚。完整构造参数表如下参数默认值说明model必填同基类torch.nn.Sequential(featurizer, classifier)optimizerAdam仅支持Adam/AdamW/SGD基类白名单lr1e-3学习率weight_decay0.0权重衰减betas(0.9, 0.999)Adam 一阶/二阶矩指数衰减率momentum0.9SGD 动量kernel_typegaussianMMD 惩罚的核类型None时退化为均值 二阶统计量协方差距离ci_testmmd条件独立度量目前仅支持 MMDattr_typesNone属性类型列表必须与数据集中属性的列顺序一致支持causal因果、conf混杂、ind独立、sel选择偏倚。单偏移数据集用[causal]或[ind]多偏移数据集用[causal, ind]E_conditionedTrue是否应用 E-条件按环境条件正则E_eq_ANone指明哪些属性索引与环境 E 的定义重合gamma1e-6MMD 核带宽参数注意源码 docstring 特别说明由于实现方式实际带宽是 gamma 的倒数即gamma1e-6对应带宽 1e6lambda_causal1.0因果偏移的 MMD 惩罚系数lambda_conf1.0混杂偏移的 MMD 惩罚系数lambda_ind1.0独立偏移的 MMD 惩罚系数lambda_sel1.0选择偏移的 MMD 惩罚系数4.1 training_step目标函数逐项拆解CACM.training_step 实现了 CACM 的目标函数流程如下把model拆成featurizer self.model[0]、classifier self.model[1]L78-L79因此 CACM 必须拿到分层模型而不是 ERM 那种黑盒 Sequential 整体调用对每个环境的输入分别过 featurizer 与 classifier得到各环境的特征表征features与 logitsclassifs池化交叉熵对所有环境的(classifs[i], targets[i])求和再除以环境数nmb作为经验风险项L92-L98按属性类型分派正则L100-L142causal→CACMRegularizer.conditional_reg(...)即条件正则 φ(x) ⊥⊥ A_i | A_s条件子集里包含标签 Yconf与ind→unconditional_reg(...)即无条件正则 φ(x) ⊥⊥ A_isel→ 也走conditional_reg在 Y 条件下独立若一个 batch 里有多个环境nmb 1四项惩罚除以nmb*(nmb-1)/2做环境对归一化L132-L136最后以lambda_causal/conf/ind/sel加权累加进总损失若attr_types is None代码目前会走到elif self.graph is not None: pass # TODO分支并最终抛出ValueError(No attribute types or graph provided.)L157-L161——从源码结构看基于因果图graph的正则路径仍是未完成实现当前必须显式提供attr_types各项惩罚标量通过self.log(..., prog_barTrue)打到训练进度条键为penalty_causal/penalty_conf/penalty_ind/penalty_sel便于训练时监控约束是否在下降。官方 因果预测 demo notebook 给出了三种典型配置分别对应单一因果偏移、单一独立偏移、以及因果 独立混合偏移from dowhy.causal_prediction.algorithms.cacm import CACM # 单偏移因果属性 algorithm CACM(model, lr1e-3, gamma1e-2, attr_types[causal], lambda_causal100.) # 单偏移独立属性属性 0 与环境定义重合 algorithm CACM(model, lr1e-3, gamma1e-2, attr_types[ind], lambda_ind10., E_eq_A[0]) # 多偏移因果 独立 algorithm CACM(model, lr1e-3, gamma1e-2, attr_types[causal, ind], lambda_causal100., lambda_ind10., E_eq_A[1])注意 demo 中使用的gamma1e-2、lambda_causal100.与代码默认值gamma1e-6、lambda_causal1.0不同——这些惩罚系数与核带宽对 MMD 项的幅度影响很大需要按数据集调参不能直接套用默认值。5. Regularizer无条件与条件 MMD 正则regularization.py 的Regularizer是 CACM 的约束引擎构造函数接收E_conditioned、ci_test、kernel_type、gamma四个参数并直接保存为实例属性L13-L30测试 验证了这些属性被原样存储。5.1 unconditional_regφ(x) ⊥⊥ A_iunconditional_reg(classifs, attribute_labels, num_envs, E_eq_AFalse, use_optimizationFalse)L94-L148实现表征与属性 A_i 的无条件独立内部按三种情况分支E 与 A 重合且E_conditionedFalse直接对各环境的特征张量两两求 MMD环境即属性的取值划分E 与 A 不重合且E_conditionedTrue在每个环境内部用_split_by_attribute按属性标签把样本分组再对组间两两求 MMDE_conditionedFalse把所有环境的特征与属性标签各自拼接后整体按属性值分组再求组间 MMD。5.2 conditional_regφ(x) ⊥⊥ A_i | A_sconditional_reg(...)L150-L239在给定条件子集 A_s使 (X_c, A_i) 条件 D-分离的观测变量例如属性 标签 Y上施加独立。分组规则在 docstring 中有具体示例若conditioning_subset [A1, Y]A1 ∈ {0,1}、Y ∈ {0,1,2}则组合出 6 个组A10,Y0 → group 0A11,Y0 → group 1……。docstring 注明该分组索引代码改编自 WILDS 基准Koh et al., ICML 2021。核心计算落在_compute_conditional_penaltyL241-L262torch.unique(..., return_inverseTrue)枚举条件组在每个条件组内再按属性值切分对至少含两个属性值组的两两 MMD 求和。5.3 use_optimization 与 GPU 吞吐优化两个 reg 方法都带use_optimization开关True时走_optimized_mmd_penaltyL38-L85它把 k 组张量两两 MMD 的求和代数化简为penalty (k-1) * Σ K_ii − 2 * Σ K_ij一次计算所有组内/组间核均值False时保留原始嵌套双循环扩展性更好便于插入自定义 MMD 变体但更慢。docstring 明确建议标准 MMD 追求性能选True验证正确性或扩展新 MMD 类型选False。高斯核路径下张量会被提升到float64以保证精度返回时再转回原 dtypeL48、L85。6. utilsMMD 的两种计算路径utils.py 文件头注明函数借鉴自 Meta 的 DomainBed 项目Gulrajani Lopez-Paz, ICLR 2021仅两个函数def gaussian_kernel(x, y, gamma): return torch.exp(torch.cdist(x, y, p2.0).pow(2).clamp_min_(1e-30).mul(-gamma)) def mmd_compute(x, y, kernel_type, gamma): if kernel_type gaussian: Kxx gaussian_kernel(x, x, gamma).mean() Kyy gaussian_kernel(y, y, gamma).mean() Kxy gaussian_kernel(x, y, gamma).mean() return Kxx Kyy - 2 * Kxy else: # 非高斯核均值平方差 协方差矩阵平方差线性 MMD / 二阶统计量距离 ...两点源码细节值得注意高斯核先clamp_min_(1e-30)再做负指数避免平方距离为 0 时的下溢且由于gamma乘在距离平方上gamma 实际扮演“带宽倒数”的角色与 CACM docstring 的说明一致非高斯分支用样本均值差的平方均值加样本协方差差cova cent.T cent / (n-1)的平方均值近似分布距离当某一侧样本数为 1 时协方差退化为 0 矩阵Regularizer._optimized_mmd_penalty中有同样的n1判断。单元测试对该模块做了性质级验证test_causal_prediction.pymmd_compute(x, x, gaussian, gamma1.0)对相同张量应近似为 0而对偏移 10 个单位的分布gaussian与linear两种核都应返回正值。Regularizer.mmd同样被验证了对相同输入的近零性。7. 端到端用法从数据集到 trainer该算法包并不独立工作而是与datasets多域数据集、dataloaders数据加载配合。以官方 demo notebook 的 ERM 流程为例from dowhy.causal_prediction.datasets.mnist import MNISTCausalAttribute from dowhy.causal_prediction.dataloaders.get_data_loader import DataLoader # 具体以 notebook 为准 data_dir data dataset MNISTCausalAttribute(data_dir, downloadTrue) # 初始化必须提供 data_dir # 特征器 分类器组装为 Sequential 模型 model torch.nn.Sequential(featurizer, classifier) from dowhy.causal_prediction.algorithms.erm import ERM algorithm ERM(model, lr1e-3) trainer pl.Trainer(devices1, max_epochs5) trainer.fit(algorithm, loaders[train_loaders], loaders[val_loaders]) if test_loaders in loaders: trainer.test(dataloadersloaders[test_loaders], ckpt_pathbest)内置的 MNIST 变体数据集datasets/mnist.py恰好覆盖三种attr_types场景MNISTCausalAttribute属性 颜色与标签 Y 存在直接因果关系各环境ENVIRONMENTS [90%, 80%, -90%, -90%]颜色与标签相关性 0.1/0.2/0.9MNISTIndAttribute属性 旋转角度15°/60°/90°与标签独立ENVIRONMENTS [15, 60, 90, 90]MNISTCausalIndAttribute颜色 旋转双属性对应attr_types[causal, ind]的多偏移配置。这些数据集均继承自 base_dataset.py 的MultipleDomainDataset类级别默认N_STEPS5001、CHECKPOINT_FREQ100、N_WORKERS8子类覆写ENVIRONMENTS/INPUT_SHAPE每个环境是一个TensorDataset(x, y, a)其中a是形状(n, k)的属性张量——这正是算法层training_step中解包(x, y, a)三元组、以及attr_types必须“与数据集中属性列顺序一致”的原因。8. 实践要点与适用边界选择 ERM 还是 CACM若只是跨域分类且没有可标注的辅助属性ERM 足够当数据存在与 Y 因果相关causal/sel或独立ind/conf的辅助属性、且希望在分布偏移下保持泛化时使用 CACM 并显式传入attr_types。参数核对优化器白名单为 Adam/AdamW/SGDgamma是带宽倒数demo 用1e-2而非默认1e-6四个lambda_*惩罚系数默认均为 1.0在 demo 中被放大到 10~100 量级属于必须按数据集调参的项。限制CACM 要求模型可拆分为 featurizer/classifier 两层attr_typesNone时基于 graph 的正则分支在源码中仅为TODO占位ci_test目前仅支持 MMDE_eq_A需要用户准确指出与环境定义重合的属性索引如 notebook 中的E_eq_A[0]/E_eq_A[1]。验证手段除 notebook 外可运行 tests/causal_prediction/test_causal_prediction.py 中的相关用例依赖 torch/pytorch-lightning 的用例会通过pytest.importorskip在缺失依赖时自动跳过来核对优化器构造、MMD 性质与 Regularizer 行为。综上dowhy.causal_prediction.algorithms包用一个 Lightning 基类统一了训练协议用 ERM 给出池化训练基线用 CACM Regularizer MMD 工具链把“按属性类型做条件独立约束”这一因果假设直接编码进损失函数是 DoWhy 因果预测能力的算法核心。赞分享机器学习数据分析【免费下载链接】dowhyDoWhy is a Python library for causal inference that supports explicit modeling and testing of causal assumptions. DoWhy is based on a unified language for causal inference, combining causal graphical models and potential outcomes frameworks.项目地址https://gitcode.com/gh_mirrors/do/dowhy点击查看免费下载相关推荐DoWhy 因果反驳Refutation体系详解causal_refuters 包的全部反驳方法、敏感性与重叠分析DoWhy 因果反驳Refutation体系详解causal_refuters 包的全部反驳方法、敏感性与重叠分析 在 DoWhy 的识别—估计—反驳机器学习数据分析DoWhy dowhy.causal_prediction.models 深度解析因果预测任务中的预置神经网络构件DoWhy dowhy.causal_prediction.models 深度解析因果预测任务中的预置神经网络构件 DoWhy 的因果预测causal pr机器学习数据分析DoWhy causal prediction.datasets 数据集包解析面向分布外因果预测的 MNIST 多环境数据集DoWhy causal prediction.datasets 数据集包解析面向分布外因果预测的 MNIST 多环境数据集 本文基于 API 参考文档 do机器学习数据分析上一篇5分钟快速掌握UniHacker跨平台Unity破解工具终极指南下一篇BtrbkBtrfs快照与远程备份利器创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED

相关推荐

IronClaw Kernel 权威边界解析:九阶段安全外围的架构设计与源码实现

IronClaw Kernel 权威边界解析:九阶段安全外围的架构设计与源码实现

人工智能AI 应用交互助手AI Agent 【免费下载链接】ironclaw IronClaw is an Agent OS focused on privacy, security and extensibility 项目地址: https://gitcode.com/gh_mirrors/iro/ironclaw 点击查看 免费下载 本文以 docs/internal/reborn/target-architect…

📅 2026/9/25 5:16:15
STM32开发踩坑实录:从时钟树到调试救砖的实战指南

STM32开发踩坑实录:从时钟树到调试救砖的实战指南

开篇先唠叨两句。搞STM32这些年,从标准库一路折腾到HAL库,从Keil MDK换到VSCode,从F1玩到H7,踩过的坑比吃过的盐还多。尤其是刚入门那阵子,一个延时函数卡死能折腾一晚上,一个芯片包装不对能让你怀疑人生。…

📅 2026/9/25 5:16:15
掌握 AI 时代的「提示词」技巧:用 TaoToken 统一 Key 打通多工具配置

掌握 AI 时代的「提示词」技巧:用 TaoToken 统一 Key 打通多工具配置

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

📅 2026/9/25 5:11:14
MORE NEWS

更多资讯

📰

深入解析 AWS SDK for Java 2.x 的 DynamoDB 异步编程实战(附测试与分页原理)

示例工程教程后端 【免费下载链接】aws-doc-sdk-examples Welcome to the AWS Code Examples Repository. This repo contains code examples used in the AWS documentation, AWS SDK Developer Guides, and more. For more information, see the Readme.md file below. 项目地…

📰

Erlang/OTP 端口(Ports)与端口驱动(Port Drivers)完全指南:从消息协议到 C 语言实战

编程语言语言运行时标准库编译器并发编程 【免费下载链接】otp Erlang/OTP 项目地址: https://gitcode.com/gh_mirrors/ot/otp 点击查看 免费下载 Ports 是 Erlang 与外部世界通信的基础机制,它为 Erlang 进程提供了一条字节导向(byte-orien…

📰

RFID仓库管理系统从零搭建:设备选型、系统架构与避坑指南

简介:基于射频识别技术的仓库管理系统是一套面向仓储信息化与物联网初学者的完整项目资源,针对传统仓库在到货检验、入库、库位分配、库存变动及出库等环节数据录入慢、易出错的问题,利用RFID自动识别技术实现作业数据实时采集,帮…

📰

jc 的 lsblk 解析器:把 Linux 块设备树输出转换为结构化 JSON 的完整指南

开发工具 【免费下载链接】jc CLI tool and python library that converts the output of popular command-line tools, file-types, and common strings to JSON, YAML, or Dictionaries. This allows piping of output to tools like jq and simplifying automation scripts.…

📰

Atlas 300V 24G部署YOLO完整指南:模型转换、推理优化与排错

“atlas部署yolo”这个话题,最近在AI推理圈子里确实热得不行。很多人手里拿到一块Atlas 300V 24G,第一反应就是“这卡到底能不能跑YOLO?跑起来有多快?跟GPU比到底是啥水平?”我自己的答案是:能跑&#xff0…

📰

网御星云安全集中管理系统实战:日志接入、告警配置与运维避坑指南

简介:面向网御星云安全集中管理系统的运维人员与安全管理员,这份PDF手册围绕V3.0.7版本,系统梳理了从登录到主页模块的使用方法。内容涵盖产品特点、软件描述、License控制,并重点展开安全等级、24小时安全趋势、服务器状态和设备…

TODAY

今日更新

THIS WEEK

本周精选

THIS MONTH

本月热门

读完文章,想聊聊您的网站?

告诉我们您的行业与需求,资深顾问一对一梳理方案与报价,全程免费。

📞 💬