Anomalib 预处理模块参考:PreProcessor 如何统一 PyTorch 推理与 Lightning 训练中的数据变换
【免费下载链接】anomalibAn anomaly detection library comprising state-of-the-art algorithms and features such as experiment management, hyper-parameter optimization, and edge inference.项目地址: https://gitcode.com/GitHub_Trending/an/anomalib
本文围绕 Anomalib 的预处理 API 参考页(docs/source/markdown/guides/reference/pre_processing/index.md所渲染的anomalib.pre_processing模块)展开,结合仓库源码讲解PreProcessor类的双角色设计(torch.nn.Module+ LightningCallback)、其在训练/验证/测试/预测各阶段注入变换的具体位置,以及模型导出(ONNX/OpenVINO)场景下 exportable transform 的兼容处理机制。读完本文,你可以掌握如何在自定义模型中正确接入数据变换、理解导出的推理图里变换是如何被“改写”的,并知道这些行为在单元测试中的验证方式。
模块定位与导出内容
参考页对应的模块是 anomalib.pre_processing。模块级 docstring 明确了它的职责边界:
- 在流水线不同阶段对数据应用 transforms;
- 管理阶段相关的 transforms(train/val/test);
- 同时对接 PyTorch 与 Lightning 两种工作流。
模块只导出一个公共对象:
from .pre_processor import PreProcessor __all__ = ["PreProcessor"]也就是说,PreProcessor是该模块对外的唯一 API 入口,其余实现细节(如导出变换的构造工具)位于内部子模块 utils/transform.py 中。
PreProcessor 类:一个 nn.Module,也是一个 Lightning Callback
核心实现在 pre_processor.py:
class PreProcessor(nn.Module, Callback): """Anomalib pre-processor. This class serves as both a PyTorch module and a Lightning callback, handling the application of transforms to data batches as a pre-processing step. Args: transform (Transform | None): Transform to apply to the data before passing it to the model. """ def __init__(self, transform: Transform | None = None) -> None: super().__init__() self.transform = transform self.export_transform = get_exportable_transform(self.transform)这里有三个设计要点:
- 构造参数只有一个
transform,类型是torchvision.transforms.v2.Transform或None,默认为None(即不施加任何变换,输入原样透传); - 双继承:
nn.Module身份用于模型导出后的推理前向,Callback身份用于 Lightning 训练循环中的批次级注入; - 构造时即派生
export_transform:通过get_exportable_transform(self.transform)生成一份导出兼容的变换副本(详见下文“可导出的变换”一节),训练用的self.transform与导出用的self.export_transform从此分道扬镳。
模块 docstring 给出的三类典型用法(摘自类文档,保持原样):
>>> from anomalib.pre_processing import PreProcessor >>> from torchvision.transforms.v2 import Resize >>> pre_processor = PreProcessor(transform=Resize(size=(256, 256))) >>> transformed_batch = pre_processor(batch)自定义变换组合:
>>> from torchvision.transforms.v2 import Compose, Resize, ToTensor >>> from anomalib.pre_processing import PreProcessor >>> # Define a custom set of transforms >>> transform = Compose([Resize((224, 224)), Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])]) >>> # Pass the custom set of transforms to a model >>> pre_processor = PreProcessor(transform=transform) >>> model = MyModel(pre_processor=pre_processor)在 Lightning 模块内以 hook 方式覆盖默认预处理:
>>> class MyModel(LightningModule): ... def __init__(self): ... super().__init__() ... ... def configure_pre_processor(self): ... transform = Compose([ ... Resize((224, 224)), ... Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ... ]) ... return PreProcessor(transform)训练、验证、测试、预测四个阶段的变换注入点
PreProcessor实现了 Lightning 的四个“批次开始”回调,分别在 pre_processor.py 中定义:
| 回调方法 | 触发阶段 | 行为 |
|---|---|---|
on_train_batch_start | 训练 | 若配置了 transform,原地改写batch.image与batch.gt_mask |
on_validation_batch_start | 验证 | 同上 |
on_test_batch_start | 测试(评估) | 同上 |
on_predict_batch_start | 预测(推理) | 同上 |
四者实现完全一致,核心逻辑均为:
def on_train_batch_start(self, trainer, pl_module, batch: Batch, batch_idx) -> None: del trainer, pl_module, batch_idx # Unused if self.transform: batch.image, batch.gt_mask = self.transform(batch.image, batch.gt_mask)两个值得注意的细节:
- 输入是
anomalib.data的Batch对象(如ImageBatch),而非裸 tensor。变换同时作用于image和gt_mask两个字段,保证异常掩码与图像在几何变换下严格对齐; - 变换在
forward之前完成:源码注释明确指出,Lightning 训练/验证/测试循环中,变换是在on_*_batch_start系列方法里施加的,模型的forward拿到的是已处理过的数据。
forward 接口与导出模型的推理路径
forward方法(pre_processor.py)承担了与训练循环完全不同的职责:
def forward(self, batch: torch.Tensor) -> torch.Tensor: """Apply transforms to the batch of tensors for inference. This forward-pass is only used after the model is exported. Within the Lightning training/validation/testing loops, the transforms are applied in the ``on_*_batch_start`` methods. """ return self.export_transform(batch) if self.export_transform else batch- 训练循环:走 Callback 路径(上文的四个
on_*_batch_start),输入是Batch; - 导出模型(ONNX/OpenVINO 推理图):走
nn.Module.forward路径,输入输出都是torch.Tensor,执行的是构造期生成的self.export_transform; - 未配置变换时,
forward直接透传输入 tensor。
可导出的变换:让 torchvision transform 兼容 ONNX/OpenVINO
__init__中的self.export_transform = get_exportable_transform(self.transform)指向 utils/transform.py 中的同名函数,它解决两类导出兼容性问题:
def get_exportable_transform(transform: Transform | None) -> Transform | None: if transform is None: return None transform = copy.deepcopy(transform) transform = disable_antialiasing(transform) return convert_center_crop_transform(transform)- 关闭
Resize的抗锯齿:disable_antialiasing递归遍历Compose子链,把所有Resize的antialias属性置为False,因为抗锯齿路径不被 ONNX 导出支持; - 把
CenterCrop替换为ExportableCenterCrop:convert_center_crop_transform递归扫描,将每个CenterCrop换成 Anomalib 自己实现的 ExportableCenterCrop(位于anomalib.data.transforms),因为 torchvision 原版CenterCrop无法直接导出。
另外,该过程先对传入 transform 做deepcopy,因此导出兼容化不会污染用户在训练侧使用的原始transform对象——这正是self.transform与self.export_transform分开存储的意义。
模型侧集成:configure_pre_processor 钩子与流水线位置
PreProcessor并不是孤立的:所有 Anomalib 图像模型都通过基类把它接入组件体系。在 AnomalibModule 中:
def __init__( self, pre_processor: nn.Module | bool = True, post_processor: nn.Module | bool = True, evaluator: Evaluator | bool = True, visualizer: Visualizer | bool = True, ) -> None: ... self.pre_processor = self._resolve_component(pre_processor, nn.Module, self.configure_pre_processor)其运作方式可以从源码结构看归纳为:
- 构造参数支持传入现成的
nn.Module、布尔开关,或留空由configure_pre_processor类方法生成默认实例(该基类方法带image_size: tuple[int, int] | None = None参数,供子类按模型输入尺寸定制); configure_pre_processor在大量模型中被重写,例如 Patchcore、Draem、EfficientAD、Glass、AnomalyDINO 等,各自给出与该模型输入尺寸、归一化要求匹配的Resize/Normalize组合;- 由于
PreProcessor同时是Callback,configure_callbacks 会自动把它(连同 post_processor、evaluator、visualizer 中属于 Callback 的成员)注册进模型回调列表,无需用户手动挂接。
推理侧的调用顺序同样在基类中固定:forward的文档说明输入批次依次经过“1. Pre-processor(若配置)→ 2. Model → 3. Post-processor(若配置)”,见 anomalib_module.py。
测试用例中的可验证行为
单元测试 tests/unit/pre_processing/test_pre_processing.py 对上述行为做了两条关键验证:
test_forward:对(3, 256, 256)的图像施加Compose([Resize((224, 224)), ToImage(), ToDtype(torch.float32, scale=True)])后,pre_processor(image)输出形状应为(1, 3, 224, 224);test_no_transform:PreProcessor()不传 transform 时,输入(3, 256, 256)的图像原样返回(1, 3, 256, 256),验证了透传分支。
这两条断言恰好覆盖了forward的两种分支(有 export_transform / 无 export_transform),与上文源码分析一致。
小结
anomalib.pre_processing参考页背后是一个小而职责清晰的模块:PreProcessor以transform为唯一配置项,用 Callback 钩子在 Lightning 训练循环的批次级注入几何与归一化变换(同时保持image与gt_mask对齐),又用nn.Module.forward承接导出模型的推理路径;get_exportable_transform在构造期深拷贝并改写Resize/CenterCrop,解决 ONNX/OpenVINO 导出的兼容性缺口。对于自定义模型,只需重写configure_pre_processor类方法返回一个PreProcessor,即可让 Anomalib 基类自动完成组件解析与回调注册。
【免费下载链接】anomalibAn anomaly detection library comprising state-of-the-art algorithms and features such as experiment management, hyper-parameter optimization, and edge inference.项目地址: https://gitcode.com/GitHub_Trending/an/anomalib
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考