MLX 随机采样完全指南:从隐式全局 PRNG 到可分裂 Threefry 密钥管理 MLX 随机采样完全指南从隐式全局 PRNG 到可分裂 Threefry 密钥管理【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx本篇指南围绕 MLX 的随机采样模块mlx.random展开系统讲解其默认使用隐式全局 PRNG 状态、同时为每个采样函数提供可选key关键字参数的双轨设计并结合仓库源码剖析 Threefry 可分裂 PRNG 的底层实现。读完本文你将掌握key/split/seed的正确用法、全部 12 个采样 APIuniform、normal、bernoulli、categorical、randint、truncated_normal、gumbel、laplace、multivariate_normal、permutation等的参数与语义以及如何在可复现实验与并行/分布式场景之间自由切换。MLX 随机数生成的双轨设计隐式全局状态与显式 KeyMLX 的随机采样函数random.rst默认使用一个隐式的全局 PRNG 状态每次调用采样函数而不指定key时都会从内部维护的全局状态中取出一份新的随机数序列因此连续调用会得到互不相同的结果。与此同时所有采样函数都接受一个可选的key关键字参数用于需要更细粒度控制或显式状态管理的场景。这一设计的典型对比在文档中给出了最直观的例子。隐式模式for _ in range(3): print(mx.random.uniform())每次迭代都会打印一个不同的伪随机数因为全局状态在不断推进。而显式模式key mx.random.key(0) for _ in range(3): print(mx.random.uniform(keykey))由于每次都传入同一个key三次打印的将是完全相同的伪随机数——这正是同一个 key 产生同一个输出的可复现语义。从源码看隐式全局状态的载体是mlx/random.h中的KeySequence类。它用**线程本地存储thread-local**维护每个线程各自的 key以避免多线程竞争// mlx/random.h static KeySequence default_() { static auto time_seed []() { auto now std::chrono::system_clock::now(); return std::chrono::duration_caststd::chrono::milliseconds( now.time_since_epoch()) .count(); }(); static thread_local KeySequence ks(time_seed); return ks; }这意味着隐式模式下初始种子取自系统时钟毫秒级时间戳且每个线程拥有独立的 PRNG 状态见 mlx/random.h。KeySequence::next()的实现则是分裂语义的体现每取一次随机数就把当前 key 分裂为两个一个推进自身状态一个作为本次输出// mlx/random.cpp array KeySequence::next() { auto out split(key_); key_ out.first; // 推进内部状态 return out.second; // 交给本次采样 }显式 Key 管理key、seed 与 split当需要可复现性或精细的状态控制时显式 key 管理是首选。核心函数有三个定义见 mlx/random.h函数签名核心参数作用keykey(seed: int) - array由 64 位整数种子生成一个形状为(2,)、类型为uint32的 PRNG keyseedseed(seed: int)为默认隐式PRNG 重新播种影响之后所有不带key的采样splitsplit(key, num: int) - array或split(key) - (array, array)将一个 key 分裂为num个或一对独立 key用于并行分支key的实现mlx/random.cpp把 64 位种子拆成高 32 位与低 32 位两个uint32值array key(uint64_t seed) { uint32_t k1 static_castuint32_t(seed 32); uint32_t k2 static_castuint32_t(seed); return array({k1, k2}); }split的实现mlx/random.cpp则直接复用bits原语split(key, num)等价于生成num × 2的随机位块split(key)则把结果切成两个(2,)形状的 key 对返回。seed的实现mlx/random.cpp只是对线程本地KeySequence重新播种因此它只影响调用线程的隐式状态。实战用 key split 组织可复现实验MLX 官方推荐的习惯用法是只从seed或key出发用split为每个独立的数据流分支派生新 key。例如在分布式/并行测试中见 nccl_test_distributed.pyimport mlx.core as mx mx.random.seed(0xF0F0F0F0) kx, ky mx.random.split(mx.random.key(0)) x mx.random.normal((4, 128), dtypemx.float32, keykx) y mx.random.normal((4, 128), keyky)每个进程用mx.random.key(rank)获得与自身 rank 绑定的独立 key见 nccl_test_distributed.py再用split继续派生既保证确定性又避免不同分支之间的相关性。底层核心遵循 JAX 设计的可分裂 Threefry PRNGMLX 的 PRNG 设计遵循 JAX 的 PRNG 设计JEP 263采用Threefry 的可分裂splittable计数版counter-based PRNG。这一点在文档中有明确说明其核心优势是可复现同一 key 同一计数器必然产生同一输出可并行/可分裂key 可以无限分裂成相互独立的子 key天然适配多线程、多设备、分布式场景无全局可变状态污染显式 key 模式下随机数的生成完全不依赖调用顺序或线程调度。Threefry 2x32 哈希函数的实现位于 mlx/backend/cpu/threefry.cpp核心是 5 轮基于旋转常数的混合运算// mlx/backend/cpu/threefry.cpp constexpr static uint32_t rotations[2][4] { {13, 15, 26, 6}, {17, 29, 16, 24}}; uint32_t ks[3] {key.first, key.second, key.first ^ key.second ^ 0x1BD11BDA}; // ... for (int i 0; i 5; i) { for (auto r : rotations[i % 2]) { count.first count.second; count.second (count.second r) | (count.second (32 - r)); count.second ^ count.first; } count.first ks[(i 1) % 3]; count.second ks[(i 2) % 3] i 1; }该实现基于 JAX 的 JEP 263 说明与 Random123 论文中的 Threefry 参考实现头文件注释见 mlx/backend/cpu/threefry.h。GPUMetal侧同样有对应 kernel整个随机原语以RandomBitsprimitive 的形式进入 MLX 的计算图从而支持惰性求值与流调度。采样原语bits 与 uniform 的构建方式所有连续分布采样都建立在两个基础原语之上。bitsmlx/random.h生成类型为uint32默认width4、充满随机位的数组是所有分布的熵来源。其实现mlx/random.cpp会校验key 类型必须是uint32、形状必须是(2,)且width只允许{1, 2, 4}分别对应uint8/uint16/uint32维度不允许为负。uniformmlx/random.cpp通过随机位除以UINT32_MAX再截断到nextafter(1.0, 0.0)的技巧把整数随机位映射到[0, 1)区间再经out * (high - low) low缩放到[low, high)。源码中专门处理了半精度类型的上界问题——fp16/bf16用below_one()找到1.0的前一个可表示值mlx/random.cpp确保半精度下采样结果严格小于 1。12 个采样 API 全解以下 API 均在 mlx/random.h 中声明Python 侧通过mlx.core.random模块暴露文档 autosummary 见 random.rst。所有函数均支持可选的key与StreamOrDevice参数。连续分布uniform(low, high, shape, dtypefloat32, keyNone)均匀分布采样区间[low, high)low/high支持标量或数组广播。仅支持实数浮点类型。也提供uniform(shape, dtypefloat32, keyNone)简写等价于uniform(0, 1, shape, dtype)。normal(shape, loc0.0, scale1.0, dtypefloat32, keyNone)标准正态采样通过逆误差函数erfinv从均匀分布变换而来mlx/random.cpp。支持complex64类型复数正态通过把实部/虚部视为两个独立正态实现见complex_normal。loc/scale可为标量或数组。truncated_normal(lower, upper, shape, dtypefloat32, keyNone)截断正态算法与 JAX 的truncated_normal一致源码注释见 mlx/random.cpp先对区间端点做erf变换在变换后的均匀区间采样再经erfinv逆变换并裁剪回[lower, upper]。laplace(shape, loc0.0, scale1.0, dtypefloat32, keyNone)拉普拉斯分布双指数分布通过逆 CDF 生成mlx/random.cpp常用于差分隐私等需要重尾噪声的场景。gumbel(shape, dtypefloat32, keyNone)Gumbel 分布实现为-log(-log(uniform(shape)))mlx/random.cpp是 Gumbel-max 技巧见下文categorical的基础。multivariate_normal(mean, cov, shape, dtypefloat32, keyNone)多元正态。要求dtype为float32mean至少一维、cov至少二维且最后两维相等mlx/random.cpp。实现通过 SVD 求协方差矩阵的平方根Σ^{1/2}再对标准正态样本做线性变换mean z Σ^{1/2}shape与mean/cov的前导维度做广播。离散分布与整数采样randint(low, high, shape, dtypeint32, keyNone)均匀整数采样区间[low, high)。仅接受整数 dtype 与bool。实现先用floor代替astype的截断避免-1.7 → -1的错误再把结果钳制回[low, high-1]防止low/high超出 float32 表示范围导致越界mlx/random.cpp。bernoulli(p0.5, shapeNone, keyNone)伯努利二项采样输出布尔数组p为取True的概率支持标量或数组。实现把p放大到[0, nextafter(UINT32_MAX, inf)]标度后与随机位比较保证p1全真、p0全假mlx/random.cpp。categorical(logits, axis-1, shapeNone, num_samplesNone, keyNone)按 logits 权重进行类别采样输出uint32索引。底层采用两条路径mlx/random.cpp一维且指定shape时用O(NM) 的逆 CDF searchsorted方法避免 Gumbel-max 的 O(N×M) 内存开销一般情况用Gumbel-max 技巧argmax(gumbel_shape logits)。num_samples参数允许一次采样多个样本此时会在末尾追加一个采样维度。permutation(x, axis0, keyNone)随机排列。对整数x返回arange(x)的一个随机排列对数组x沿axis轴打乱元素。实现基于argsort(bits({x}))与takemlx/random.cpp。在真实项目中的典型用法初始化模型参数是 MLX 生态中最常见的用法。例如用uniform做 Xavier 风格初始化或像 mlx_distributed_tests.py 那样用固定 key 生成可复现的输入数据import mlx.core as mx key mx.random.key(0) x (mx.random.uniform(shape(4, 1024), keykey) * 10).astype(mx.float32)随机初始化权重参考 test_array.py 的 bf16 用例w mx.random.uniform(low-0.1, high0.1, shape(128, 128), dtypemx.bfloat16) b mx.random.normal((128,), dtypemx.float32)打乱数据集顺序indices mx.random.permutation(mx.arange(n_samples)) shuffled mx.take(dataset, indices)Gumbel-Softmax 风格采样离散动作采样logits mx.random.normal((batch, n_actions)) action mx.random.categorical(logits, axis-1, num_samples1)分布式确定性初始化每个 rank 用mx.random.key(rank)派生独立序列见 nccl_test_distributed.py保证全集群数据可复现且互不相关。使用建议与易错点可复现实验用显式 key 管理把mx.random.key(seed)作为根用split为每个分支派生 key绝不手动复用同一 key 采样想要的不同结果——同一 key 必然产生同一序列。隐式模式与线程默认 PRNG 是线程本地的跨线程共享隐式状态不成立需要并行确定性时请用显式 key 并split。dtype 约束uniform/truncated_normal/laplace仅接受实数浮点类型normal额外支持complex64randint仅整数与boolmultivariate_normal仅float32bernoulli的p必须是浮点类型。不满足时会抛出带[函数名]前缀的std::invalid_argument。半精度上界fp16/bf16的uniform结果严格小于 1这与 float32 行为一致mlx/random.cpp无需额外防御性处理。categorical的内存权衡单维大类别时走逆 CDF 路径O(NM)多维走 Gumbel-maxO(N×M)理解这一点有助于为大规模 softmax 采样选择合适写法。如需查看完整 API 签名与更多实现细节可继续阅读 mlx/random.h、mlx/random.cpp、底层哈希 mlx/backend/cpu/threefry.cpp以及 Python 侧测试 test_random.py 与 mlx_tests.py。【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考