训练库实战指南)
ALX基于 JAX 的大规模 TPU 矩阵分解ALS训练库实战指南【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research导读ALXAlternating Least Squares on XLA/TPU是 Google Research 开源的一个基于 JAX 的分布式矩阵分解库它使用交替最小二乘Alternating Least SquaresALS算法针对 TPU 架构做了深度优化能够通过横向扩展 TPU 核心数高效处理规模达到 O(B) 行/列的矩阵分解问题B 为十亿级别。本文将围绕 alx/README.md 给出的核心用法结合仓库内 als.py、dataset_utils.py、batching_utils.py、checkpoints.py、multihost_utils.py、topk.py 等源码完整讲解 ALX 的配置体系、数据流水线、ALS 求解原理、训练评估流程与多机检查点机制帮助你从能跑通示例进阶到理解其为何能在 TPU 上高效扩展。一、ALX 的设计目标面向 TPU 的分布式 ALS矩阵分解Matrix Factorization是推荐系统、协同过滤等场景的核心技术其目标是把稀疏的用户-物品交互矩阵分解为低维的用户嵌入表与物品嵌入表。ALS 通过交替固定一个因子、用最小二乘求解另一个因子的方式迭代收敛。ALX 的关键设计目标可以概括为两点见 alx/README.md高效利用 TPU 架构TPU 擅长大规模矩阵乘法和规约reduction运算ALS 中的 Gramian 计算、求解器、打分score步骤天然契合这一特点规模可扩展矩阵分解问题达到 O(B) 行/列规模时可以通过增加可用 TPU 核心数来线性扩展计算能力而不是把数据塞进单机内存。仓库的 requirements.txt 只声明了四个核心依赖jax、flax、numpy、tensorflow。其中 JAX 承担数值计算与自动并行pmapFlax 提供结构化状态与检查点接口TensorFlow 用于数据读取与 TFRecord 解析。二、核心用法训练主循环逐行拆解README 给出了一个rudimentary structure基础骨架示例完整展示了从构建数据集到训练、保存检查点、评估的整个流程。下面先复现原文代码再逐段解读其与源码的对应关系ds, tds, test_ds dataset_utils.build_datasets(cfgcfg) # Initialize model_dir and setup summary writer. if jax.process_index() 0: tf.io.gfile.makedirs(FLAGS.model_dir) summary_writer tensorboard.SummaryWriter( os.path.join(FLAGS.model_dir, eval)) summarize_gin_config(model_dirFLAGS.model_dir, summary_writersummary_writer) # Check if there are any intermediate checkpoints. state checkpoints.restore_checkpoint(FLAGS.model_dir) als_state None if state: als_state als.ALSState(**state) model als.ALS(cfgcfg, als_stateals_state) for epoch in range(model.als_state.step, cfg.num_epochs): model.train(ds, tds) # Save a checkpoint after every epoch. checkpoints.save_checkpoint(model.als_state, FLAGS.model_dir) metrics model.eval(test_ds) if jax.process_index() 0: for key, val in zip([ fRecall20/{jax.process_index()}, fRecall50/{jax.process_index()}, fNum valid examples/{jax.process_index()} ], list(metrics)): summary_writer.scalar(key, val, epoch) logging.info(str(metrics))这个示例中的每一步都能在源码中找到对应实现1. 数据集构建dataset_utils.build_datasets(cfg)在 dataset_utils.py 中实现。它会根据cfg.is_pre_batched决定走两条路径预分批pre-batched数据直接加载否则先做分批 序列化再加载。返回三个数据集训练集ds用户侧、转置训练集tds物品侧、测试集test_ds。2. 目录与日志只有jax.process_index() 0的主进程负责创建model_dir并初始化 TensorBoard SummaryWriter避免多进程重复建目录。3. 检查点恢复checkpoints.restore_checkpoint(FLAGS.model_dir)在 checkpoints.py 中实现它会先恢复每个 host 子目录下的状态再用jax.sharding.NamedSharding把数据按设备维度切分回各 TPU 设备。恢复成功后构造als.ALSState见 als.py其字段为step、col_embedding、row_embedding。4. 模型初始化与训练als.ALS(cfg, als_state)是顶层类als.py。训练循环从model.als_state.step支持断点续训开始到cfg.num_epochs结束。每个 epoch 调用model.train(ds, tds)als.py——先优化用户嵌入、更新用户 Gramian再优化物品嵌入、更新物品 Gramian。5. 每轮评估model.eval(test_ds)als.py返回三个标量Recall20、Recall50和num_valid_examples有效用户数并在主进程写入 TensorBoard。三、ALSConfig全部配置参数详解ALX 的全部行为由一个ALSConfigdataclass 控制als.pyREADME 的示例代码正是以cfg贯穿始终。这个配置类可以按职责分为四组3.1 问题规模与嵌入参数类型/取值说明num_colsint矩阵列数如物品数num_rowsint矩阵行数如用户数tie_rows_and_colsbool是否让行/列共享同一嵌入表。若为 Truenum_rows必须等于num_cols源码会显式校验见 als.pyembedding_dimint嵌入向量维度regfloat正则化系数参与求解方程(G_i reg*I) * U_i b_i中的对角加项unobserved_weightfloat未观测交互的权重用于把全局 Gramian 加权加入每个用户的局部 Gramianstddevfloat嵌入表随机初始化的标准差is_bfloat16bool嵌入表是否以 bfloat16 存储训练中会按需转换精度eval_topkint评估时每个用户返回的 TopK 数量3.2 训练与批处理参数说明seq_len用户历史的序列长度历史与 ground truth 都按此长度切块、不足补-1batch_size用户侧训练批次大小transpose_batch_size物品侧转置训练批次大小eval_batch_size评估批次大小num_rows_per_batch每个批次固定的用户行数保持张量 shape 静态不变transpose_num_rows_per_batch物品侧每批次行数eval_num_rows_per_batch评估时每批次行数ground_truth_batch_size测试集 ground truth 批次大小num_epochs训练轮数train_files/train_transpose_files/test_files训练、转置训练、测试的 TFRecord 文件 patternis_pre_batched数据是否已预先分批。为 True 时跳过分批并序列化步骤直接加载3.3 评估与求解器参数默认值说明num_eval_iterations必填评估步数上限设为-1表示跑完整评估集。源码注释als.py指出大模型通常不跑全量评估因为 TPU 上 TopK 目前是完整排序实现会成为瓶颈approx_topk_for_eval必填评估时是否用近似 TopK。精确 TopK 在 TPU 上很慢大词表建议开启近似 TopK 有损评估指标是真实性能的高概率下界linear_solverNone实际使用cg线性求解器可取lu、cholesky、qr、cg。源码注释说明 TPU 上cg共轭梯度最快local_device_countNone控制使用多少 TPU 核取值[0, 8]。仅限单进程多进程设置会直接抛ValueErrorals.pyloop_gatherTrue是否用循环方式 gather 嵌入。需要 gather 的嵌入数量随 TPU 数量线性增长循环 gather 可在不 OOM 的情况下完成 gather从 als.py 可以看到求解器分发的实现逻辑linear_solver lu用jnp.linalg.solveqr用 QR 分解加三角求解cholesky用cho_factor/cho_solve其余情况含None和cg都落到jax.scipy.sparse.linalg.cg共轭梯度。最终求解步骤user_embeddings solve_fn(post_lambda, mu_batch_summed)对应 ALS 的闭式更新U_i G_i^{-1} b_i。四、数据流水线从原始 TFRecord 到 TPU 批数据4.1 两条构建路径build_datasets依据is_pre_batched分派dataset_utils.py未预分批_batch_and_build_datasets先读取原始 TFRecord用tf_examples_to_examples解析出(row_id, history, ground_truth)三元组再经过batch_and_create_tf_examples分批并序列化为新的 TF Example最后通过 generator 加载已预分批_build_pre_batched_datasets直接对train_files、train_transpose_files、test_files三个 pattern 调用load_dataset_from_files加载。4.2 密集打包Batch 数据结构矩阵分解的数据天然稀疏每个用户交互的物品数远小于总物品数直接按原始形状送入 TPU 会浪费大量算力。batching_utils.Batchbatching_utils.py定义了密集打包densely packed batch的中间结构每个用户的长历史按seq_len切成多行不足的部分用-1填充保证整个数据集所有张量 shape 恒定batch_ids/ground_truth_batch_ids记录哪些行属于同一个原始样本供后续segment_sum归约与 recall 计算使用item_lengths记录每行真实长度用于计算正则化项reg * (item_lengths unobserved_weight * num_items)见 als.py。batch_with_batch_sizebatching_utils.py是这个打包过程的核心函数默认seq_len16。它有两个值得注意的约束若单个用户历史长度超过batch_size * seq_len会直接抛ValueError提示需要增大 batch size 或 seq_lenbatching_utils.py多进程jax.process_count() 1环境下禁止 shuffle否则dataset.shard切分会失效batching_utils.py。4.3 TFRecord 编码与加载分批后的数据通过create_tf_example_from_batchbatching_utils.py序列化为 TF Example特征包括batched_history、item_lengths、row_ids、batch_ids测试集额外携带batched_ground_truths、ground_truth_batch_ids。加载侧的关键实现在 batching_utils.py_decode_record会把tf.int64强转tf.int32因为 TPU 只支持 int32process_dataset先按进程 shard再map解码、按num_devices合并 batch、prefetch(AUTOTUNE)最后用tfds.as_numpy转成 numpy 迭代对象——这是 JAX 消费数据时推荐的形态每批次张量形状与num_devices对齐每台设备拿到一个 batch 样本并支持用padding_examples生成空样本补齐batching_utils.py。五、ALS 核心实现求解、Gramian 与嵌入表5.1 solve一次最小二乘更新solveals.py完成一次对一批用户的闭式求解是 ALS 每步更新的数学内核Gather 嵌入gather_embeddings从本设备分片的物品嵌入表中按用户历史索引取出对应嵌入als.py局部统计量lambda_batch einsum(bij,bik-bjk)计算局部 Gramian 累加项mu_batch einsum(bij-bj)计算局部向量累加项分段求和通过jax.ops.segment_sum按id_list把同一用户的多个切块历史归约到一起这正是密集打包的核心用途修正方程lambda_batch_summed unobserved_weight * item_gramian未观测项加权再加reg * I正则化对角项得到post_lambda求解用第 3.3 节所述的求解器解post_lambda U mu_batch_summed。值得注意的是代码刻意使用jax.lax.convert_element_type而非tensor.astype做精度转换注释als.py说明 XLA 若把转换熔合得不好会导致超过 50% 的 TPU 时间花在 convert 算子上。5.2 Gramian 与嵌入表的分片策略Gramiancompute_gramain先本地计算embedding_table.T embedding_table再通过jax.lax.psum做跨设备求和als.py得到全局 Gramian 供所有设备使用嵌入表分片device_embedding_table_size返回num_items // device_count 1多出的 1 用于处理num_items不能被设备数整除的情况als.py嵌入表创建create_embedding_tableals.py按 100 万行分块创建避免一次性申请大表导致 OOM。源码注释给出量级参考TPU v3 单核直接创建 128 维嵌入表约只能支撑 1000 万条嵌入而分块拼接策略可扩展到 6000 万条Gather 的两种模式direct_gather一次性对所有设备 gather内存峰值高loop_gather用jax.lax.scan循环逐设备 gather、只保留本设备需要的结果als.py。源码注释提到实践中串行 scan 几乎没有性能回退因此默认开启loop_gather。5.3 训练循环交替投影ALS.trainals.py一个 epoch 内完成两次投影for batch in user_batches: solve(batch, is_userTrue) # 固定物品解用户 update_user_gramian() for batch in item_batches: solve(batch, is_userFalse) # 固定用户解物品 update_item_gramian() step 1每次solve结束后会用user_embedding_table.at[users_from_batch, :].add(user_embeddings_residual)以残差方式就地更新嵌入表als.py其中 mask 确保只有本设备负责的用户行被更新。这正体现了 ALS交替固定一个因子、更新另一个的经典结构。六、评估机制Recall20 / Recall50 与 TopK 优化6.1 打分与召回评估入口ALS.evalals.py遍历测试批次逐批调用pmapped_eval_step_fn累计三个标量后求平均scoreals.py求解用户嵌入后通过user_embeddings item_embedding_table.T计算全量打分矩阵并用NINF -1e19掩掉用户历史中已交互过的物品防止推荐自己看过的内容。打分矩阵可能过大因此用jax.lax.map而非vmap逐设备串行计算als.pyrecallals.pysum(isin(top_ids[:r_at], ground_truth)) / min(r_at, num_valid_ground_truth)分母做了r_at与真实 ground truth 数的下界截断避免数据不足时指标失真多设备归约最终Recall20、Recall50、num_valid_examples都通过jax.lax.psum跨设备求和再在 host 侧汇总为均值als.py。6.2 精确 TopK 与近似 TopK评估的另一大开销来自 TopK。源码注释明确指出 TPU 上的jax.lax.top_k目前实现为完整排序词表很大时极慢因此提供了可选的近似版本精确模式直接调用jax.lax.top_k(scores, cfg.eval_topk)als.py近似模式top_k_approxtopk.py把每行打分按固定窗口长度分成k * num_windows_multiplier个窗口先用jnp.max取每个窗口的局部最大值相当于有损的 top-1再把窗口最大值 全局索引偏移汇总后做一次精确jax.lax.top_k。窗口越小近似越准、但越慢num_windows_multiplier默认 5越大召回越好。由于近似 TopK 是有损的README 与源码都强调开启approx_topk_for_eval后评估指标是真实性能的高概率下界。七、检查点与多机同步7.1 分 host 的检查点格式save_checkpointcheckpoints.py的实现要点每个 host 把状态存到自己的子目录work_dir/host_{process_index()}见get_host_dirmultihost_utils.py先jax.device_get把设备上的分片数组取回 host再用 Flax 的checkpoints.save_checkpoint(..., keep3)保存保留最近 3 个版本保存完毕后调用sync_devices()做一次跨设备屏障确保所有 host 都写完再进入下一轮。restore_checkpointcheckpoints.py则按相反方向从 host 子目录恢复后通过jax.sharding.NamedSharding把数组按设备维度重新切分、jax.device_put送回各设备从而无缝支持断点续训。7.2 多机同步原语multihost_utils.sync_devicesmultihost_utils.py实现了一个全集群屏障每个设备贡献一个全 1 向量经pmappsum(hosts)归约后校验总和等于jax.device_count()否则报错。这个简单却可靠的同步原语保证了检查点写入、日志输出等步骤的全局有序性。八、环境与运行前提依赖jax、flax、numpy、tensorflowrequirements.txt其中数据加载还间接依赖tensorflow_datasetstfds.as_numpy硬件代码默认面向 TPU 拓扑设计8 核/进程、pmap 轴i但local_device_count允许在单进程中缩减核数多进程多 host场景下禁用local_device_count与数据 shuffle数据格式训练/转置/测试数据均为 TFRecord原始数据特征为row_tag、col_tag测试集还有gt_tag见 batching_utils.py或使用 ALX 自定义的密集打包 TF Example 格式is_pre_batchedTrue配置方式README 示例以cfg对象贯穿全程实践中常用 Gin Config 承载ALSConfig字段并配合FLAGS.model_dir指定输出目录。九、License 与免责声明ALX 以Apache 2.0 License开源见 alx/README.md。同时需注意这不是 Google 官方支持的产品This is not an officially supported Google product适合作为研究参考与二次开发基础而非受 SLA 保障的生产组件。结语通过本文你已经掌握了 ALX 从配置、数据流水线、ALS 求解内核到评估与检查点的完整脉络ALSConfig控制一切行为batching_utils负责把稀疏历史密集打包成 TPU 友好的定形张量solve通过 Gramian 线性求解器完成交替最小二乘更新topk.py用近似 TopK 突破 TPU 上的排序瓶颈checkpoints.pymultihost_utils.py保障多机场景下的可靠持久化。如果要在自己的超大规模矩阵分解任务上复现或改进建议从修改ALSConfig的求解器与 TopK 策略入手并逐步深入 als.py 中的 pmap 分片逻辑。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考