ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

Representation Clustering:用 JAX/Flax 评估神经网络表示空间聚类质量的实战指南

2026/9/23 1:39:52 拓冰建站 浏览量
Representation Clustering:用 JAX/Flax 评估神经网络表示空间聚类质量的实战指南 Representation Clustering用 JAX/Flax 评估神经网络表示空间聚类质量的实战指南【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research本指南围绕representation_clustering/目录中的代码系统讲解如何训练 ResNet 系列模型ImageNet/BREEDS/CelebA并在训练过程中对神经网络内部表示空间执行层次聚类分析量化表示是否按语义结构聚成一团这一关键问题。读完本文你将掌握该仓库的完整训练管线、三类配置文件的使用方法、BREEDS 数据管道的构建原理以及如何复用其聚类纯度purity评估流程分析自己的模型 checkpoint。一、项目定位与核心思路representation_clustering是 Google Research 开源仓库中的一个独立子项目其目标非常聚焦评估神经网络表示空间的聚类特性。代码由 Thao Nguyen 在学生研究员期间编写见 README.md配套三个 Colab 分析脚本构成训练 → 提特征 → 聚类 → 量化评估的完整闭环。项目的核心假设是一个训练良好的神经网络其倒数第二层或某个中间层的表示应该天然形成与语义类别对应的簇结构。通过在不同训练阶段不同 checkpoint对这些表示做聚类并计算纯度指标可以观察表示空间结构随训练的演化过程也可以用于分析分布外OOD数据下的表示退化情况。从源码结构看仓库包含三条主要链路链路关键文件作用训练main.py、train.py、resnet_v1.pyJAX/Flax 实现的 ResNet-18/50 训练与评估数据input_pipeline_breeds.py、input_pipeline_celebA.pyImageNet/BREEDS/CelebA 数据读取、过滤、预处理分析cluster_checkpoints.py 及三个.ipynb从 checkpoint 提取表示、层次聚类、计算纯度/AMI二、环境依赖与数据准备2.1 外部依赖项目本身是研究原型代码依赖项集中在 requirements.txtJAX 生态jax0.3.4、jaxlib0.3.2cuda11.cudnn82注释明确指出需让 CUDA 版本匹配基础镜像、flax实验基础设施absl-pyflag 与日志、clucheckpoint、metrics、preprocess_spec、ml_collections配置数据与数值tensorflow-cpu、tensorflow-datasets、numpy聚类分析scipyscipy.stats.mode用于纯度计算、scikit-learnsklearn.cluster.AgglomerativeClustering注意两个特殊依赖robustness 库input_pipeline_breeds.py与cluster_checkpoints.py中直接from robustness.tools import breeds_helpers因此必须安装 MadryLab 的 robustness 库其中的breeds_helpers.py负责生成 BREEDS 的子类-超类划分。BREEDS 层次文件README 明确要求从 BREEDS-Benchmarks 获取的imagenet_class_hierarchy/modified下的层次文件必须放在源码目录下的breeds子目录中。这一点在代码里也有硬编码约束input_pipeline_breeds.py中BREEDS_INFO_DIR被定义为os.path.join(os.path.dirname(os.path.abspath(__file__)), breeds)cluster_checkpoints.py同样如此。如果文件缺失make_breeds_dataset将无法读取超类-子类映射。2.2 目录约定训练输出写入--workdir指定的工作目录其中包含checkpoints/子目录train.py 中checkpoint_dir os.path.join(workdir, checkpoints)。聚类分析脚本通过--exp_dir指定存放 checkpoint 的实验目录cluster_checkpoints.py。三、训练管线深入解析3.1 入口main.pymain.py 刻意保持简短只做三件事解析参数通过config_flags.DEFINE_config_file接收--configml_collections 配置文件通过--workdir指定工作目录两者均为必填flags.mark_flags_as_required。隐藏 TF 的 GPU 可见性tf.config.experimental.set_visible_devices([], GPU)避免 TensorFlow 抢占显存导致 JAX 无设备可用——这是 JAX/Flax 训练中常见的坑。调用训练循环train.train_and_evaluate(FLAGS.config, FLAGS.workdir)。同时它还通过jax.config.config_with_absl()暴露了--jax_backend_target与--jax_xla_backend两个 JAX 自带 flag便于在 XCloud/Borg 等分布式环境指定后端。启动训练的命令形如python main.py \ --configconfigs/default_breeds.py \ --workdir/path/to/experiment_dir3.2 训练核心train.pytrain.py 是整个训练循环的实现要点如下模型与优化器支持resnet18与resnet50两种模型create_train_state中按config.model_name分派到resnet_v1.ResNet18/resnet_v1.ResNet50。优化器为optax.sgdconfig.sgd_momentum默认 0.9。支持EMA指数滑动平均若config.ema_decay非空则用optax.ema(decay, debiasTrue)维护参数滑动平均评估时优先使用ema_paramseval_step中state.params if state.ema_params is None else state.ema_params并在每次合并 batch stats 时做 bias correctionmerge_batch_stats。损失函数支持两种由config.loss_fn切换cross_entropy标准 softmax 交叉熵cross_entropy_loss。squared带类权重因子的平方损失squared_loss中k9, M60损失额外除以 10。学习率调度BREEDS/ImageNet 使用余弦衰减 5 epoch warmupget_learning_rate→cosine_decay且按 batch size 线性缩放base_learning_rate config.learning_rate * global_batch_size / 256.0256 是 ImageNet 的参考 batch size。CelebA 路径则使用固定学习率learning_rate_fn lambda x: base_learning_rate缩放分母为 128。正则化L2 weight decay 只作用在维度大于 1 的参数上weight_l2 sum(jnp.sum(x**2) for x in params if x.ndim 1)然后loss weight_decay * 0.5 * weight_l2。分布式与评估训练与评估均通过jax.pmap(..., axis_namebatch)做数据并行梯度用jax.lax.pmean跨设备平均。eval_step以trainFalse前向计算评估指标用 CLU 的EvalMetricsaccuracy loss。checkpoint 用clu.checkpoint.MultihostCheckpointmax_to_keep1000每steps_per_epoch * 5步保存一次若 workdir 中已有 checkpoint 会自动恢复restore_or_initialize。3.3 网络实现resnet_v1.pyresnet_v1.py 是 Flax 版 ResNet V1He et al., 2015。实现要点Conv1x1/Conv3x3通过functools.partial(nn.Conv, ...)预定义均use_biasFalseBN 前置模式。ResNetBlock用于 ResNet-18/34BottleneckResNetBlock用于 ResNet-50 及以上其最后一个 BN 的 scale 初始化为 0注释中提及 Fixup 初始化思想保证残差分支初始为零映射。模型接受batch_norm_decay参数对应配置中的config.batch_norm_decay默认 0.99。四、配置文件详解三份 default_*.py配置文件基于ml_collections.ConfigDict三份配置分别对应三类实验场景。以下是各参数的含义与默认值对照参数default_breeds.pydefault_imagenet.pydefault_celebA.py说明dataset_namebreedsimagenet2012celeb_a数据源dataset_typeentity13--BREEDS 层次划分类型shuffle_subclassesFalse--是否打乱子类分配num_classes-1--超类数量-1 表示由数据自动推断num_subclasses4--每个超类包含的子类数model_nameresnet50resnet50resnet50网络结构loss_fncross_entropycross_entropycross_entropy损失类型learning_rate0.10.10.0001参考 batch size 下的学习率learning_rate_schedulecosinecosine-调度方式warmup_epochs55-线性 warmup 轮数sgd_momentum0.90.90.9SGD 动量ema_decay0.990.990.99EMA 衰减系数batch_norm_decay0.990.990.99BN 滑动平均衰减weight_decay0.00010.00010.0001L2 系数num_epochs4009050训练轮数num_train_steps-1-1-1-1 表示由 epochs 推算num_eval_steps-1-1-1-1 表示评估完整 epochper_device_batch_size6464128每设备 batch sizeeval_pad_last_batchTrueTrueTrue评估时补齐最后 batchlog_loss_every_steps5050050日志频率eval_every_steps50500050评估频率checkpoint_every_steps-5000100checkpoint 频率shuffle_buffer_size100001000010000shuffle bufferseed4222随机种子trial000重复实验编号几点实操建议来自配置注释与源码default_breeds.py注释说明其面向 4x4 TPU slicenum_epochs400是研究级训练量在 Colab 中快速验证时应把num_eval_steps调小并可通过get_hyper中注释掉的 sweep 行如config.learning_rate、config.weight_decay做超参网格搜索。default_imagenet.py的训练目标是达到约 76% top-1 准确率配置注释中明确写出。default_celebA.py的学习率显著更低0.0001且不使用余弦调度与 warmup因为 CelebA 二分类任务更简单。五、数据管道BREEDS 与 CelebA5.1 BREEDS 管道input_pipeline_breeds.pyBREEDSBias Reduction via Evaluation of Encoder-Decoder Systems的核心思想是把 ImageNet 的细粒度子类按 WordNet 层次合并成超类superclass从而构造具有可控语义粒度与偏置的 benchmark。该管道的关键逻辑生成子类-超类映射调用breeds_helpers.make_breeds_dataset(dataset_type, ...)传入config.dataset_type如entity13、num_classes、num_subclasses、shuffle_subclasses得到训练子类划分train_subclasses。构造哈希表LabelMappingOp用tf.lookup.StaticHashTable将每个原始子类标签映射到超类索引默认值设为 -1用于过滤不属于当前划分的样本。过滤样本predicate通过判断样本标签是否在all_subclasses中过滤掉不属于该超类集合的 ImageNet 图片。训练/评估预处理训练DecodeAndRandomResizedCrop(resize_size224)随机裁剪面积占比 0.05~1.0→RandomFlipLeftRight()→LabelMappingOp评估RescaleValues()0~255 归一化到 0~1→ResizeSmall(256)短边缩放至 256保持长宽比→CentralCrop(224)→LabelMappingOp。多主机分片与 padding数据集先按jax.process_count()shard再按[jax.local_device_count(), config.per_device_batch_size]两层 batch评估集通过pad_dataset补齐到整数个 batch并为填充样本打上maskFalse保证完整评估全部验证样本。5.2 CelebA 管道input_pipeline_celebA.pyCelebA 管道复用同一套预处理算子差异在于标签来源LabelMapping直接取features[attributes][Blond_Hair]作为二分类标签金发/非金发num_classes 2。同时该文件也支持imagenet20121000 类此时不做标签映射。六、聚类分析与 checkpoint 评估6.1 聚类评估脚本cluster_checkpoints.pycluster_checkpoints.py 是分析环节的核心实现通过--exp_dir指向实验目录其主流程cluster_each_checkpoint用与训练相同的get_config()重建模型与状态checkpoints.restore_checkpoint(checkpoint_path, state)载入指定 checkpoint。对每个超类superclass内部的表示进行聚类——注意聚类是每个超类内部进行的因此衡量的是子类语义是否在表示空间中彼此分离。使用sklearn.cluster.AgglomerativeClustering层次聚类而非 KMeans源码中 KMeans/PCA 已被注释掉并引入overcluster_factor5参数即把簇数设为子类数的 5 倍进行过度聚类观察子类表示能否被更细的簇进一步切分。6.2 纯度指标compute_puritycompute_purity(clusters, classes)的实现对每个簇用scipy.stats.mode求出簇内占比最高的类别并把该类的样本数累加为n_cluster_points最终purity n_cluster_points / len(clusters)。纯度越接近 1说明每个簇内部越纯只包含单一子类表示空间与语义结构的对齐程度越高。这是论文中最常用的聚类质量指标之一。6.3 Colab 分析脚本仓库附带三个 Jupyter Notebook构成完整的分析工作流Analyze_clusters.ipynb对 checkpoint 提取的表示进行聚类可视化与结构分析Compute_purity_ami.ipynb计算纯度purity与 AMIAdjusted Mutual Information等聚类指标对比不同训练阶段/不同模型Evaluate_OOD_data.ipynb在分布外数据上评估表示质量观察 OOD 样本的表示聚类行为。这组脚本与cluster_checkpoints.py配合使用先用脚本批量导出各 checkpoint 的表示与聚类结果再在 Colab 中完成指标对比与可视化。七、端到端使用流程总结综合上述源码完整的使用流程如下1. 准备环境 pip install -r requirements.txt # 含 jax/flax/clu/ml_collections 等 pip install robustness # 提供 breeds_helpers # 将 BREEDS 层次文件放入 representation_clustering/breeds/ 目录 2. 训练模型 python representation_clustering/main.py \ --configrepresentation_clustering/configs/default_breeds.py \ --workdirexperiments/entity13_run1 3. 聚类分析可选需在可访问 checkpoint 的目录执行 python representation_clustering/cluster_checkpoints.py \ --exp_direxperiments/entity13_run1 4. 深入分析 用 Analyze_clusters / Compute_purity_ami / Evaluate_OOD_data 三个 Notebook 做可视化与指标计算八、注意事项与已知约束版本锁定requirements.txt明确锁定jax0.3.4与jaxlib0.3.2cuda11.cudnn82这是 2022 年左右的 JAX 版本若使用新版 JAX/Flaxflax.optim、cluAPI 可能存在迁移问题如cluster_checkpoints.py仍使用旧版flax.optim.Optimizer而train.py已迁移到optaxtrain_state需要按需适配。硬件前提default_breeds.py/default_imagenet.py注释标明面向 4x4 TPU slice 设计在单卡 GPU 上运行需自行调小 batch size 并注意学习率按global_batch_size / 256的线性缩放关系同步调整。TensorFlow 与 JAX 共存main.py强制隐藏 TF 的 GPU训练设备完全交给 JAX 管理切勿自行修改该行为否则可能因显存竞争导致初始化失败。数据读取ImageNet/CelebA 均通过tensorflow_datasets的try_gcsTrue读取需要可访问 GCS 或已本地缓存对应的 TFDS 数据集。总而言之representation_clustering提供了一套轻量、可复用的训练 表示聚类评估研究框架。其核心价值不在于追求 SOTA 精度而在于把表示空间的簇结构与语义对齐程度这一抽象问题落成可量化的 purity/AMI 指标与完整的工程链路方便研究者复现实验、快速验证自己的模型在聚类视角下的表现。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考