
人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载本篇技术指南聚焦 modded-nanogpt 的第三赛道records/track_3_optimization/中引入的 PyTorch Distributed Shampoo 优化器深入讲解其底层utils工具模块的设计与使用从类型安全的配置基类、优化器模块的状态管理到张量维度合并、量化压缩与分布式负载均衡再到该模块在 GPT 训练脚本中的真实接入方式。读完本文你将能够独立复用这些工具组件并为自己的分布式优化器实验搭建出可扩展、可断点续训的基础设施。目录工具模块在优化实验中的定位模块总览与文件地图快速上手六大组件的独立用法抽象基类与通用工具配置与迭代的基础设施OptimizerModule面向优化器的状态管理基类张量数学工具维度合并、分块与分布式分配模型工具CombinedLinear 组合线性层状态字典管理复杂嵌套状态的存取量化支持内存高效的低精度张量压缩负载均衡工具计算与内存成本模型综合示例Shampoo 参数处理流水线在 modded-nanogpt 训练脚本中的实际接入贡献与测试约定工具模块在优化实验中的定位Distributed Shampoo 是一种自适应梯度Adagrad 家族预条件优化器它利用神经网络参数的张量结构构造块对角预条件子从而在更少的迭代内达到同等模型质量代价是额外的 FLOPs 与内存开销。在records/track_3_optimization/results/20260513_shampoo_1_4_power/目录中该优化器被作为 modded-nanogpt 优化实验的候选方案引入从目录命名可以推断该实验聚焦于 Shampoo 的 1/4 次幂变体并配套了一份完整的训练脚本 train_gpt_shampoo.py 与多份运行日志如 503575c5-6dde-425a-b461-2df4d99db974.txt日志按 modded-nanogpt 惯例在开头完整记录训练脚本源码。整个distributed_shampoo/包采用分层结构顶层的distributed_shampoo.py暴露优化器入口preconditioner/存放各类预条件子实现distributor/负责 DDP/FSDP/HSDP 等并行策略下的张量分块而utils/则是被上述所有层共享的地基。官方 READMEdistributed_shampoo/README.md说明该实现曾赢得 MLCommons AlgoPerf 训练算法基准竞赛这也解释了它被选入 modded-nanogpt 优化赛道的原因。模块总览与文件地图utils模块位于 distributed_shampoo/utils/按职责划分为以下组件文件职责abstract_dataclass.py类型安全的抽象 dataclass 基类用于约束配置对象commons.py类内省、批处理迭代等通用函数dict_zip_iterator.py对字典内多个迭代器做同步遍历并校验长度optimizer_modules.py面向优化器组件的轻量nn.Module状态字典递归构建/加载shampoo_utils.py张量维度合并、多维多块切分、上下文管理器、分布式缓冲分配shampoo_model_utils.py面向 Shampoo 的CombinedLinear组合线性层shampoo_state_dict_utils.py复杂嵌套状态字典的提取与就地更新shampoo_quantization.pyBF16/FP16/FP32 量化的张量压缩与解压框架load_balancing_utils.py计算/内存成本模型支撑分布式负载均衡每个文件均配套独立测试位于 utils/tests/命名规范为{module_name}_test.py另有 GPU 专项测试 utils/gpu_tests/shampoo_utils_test.py。快速上手六大组件的独立用法该模块的每个工具都是独立设计、可单独引用的。以下是官方 README 提供的六类核心用法均可直接复制运行。抽象配置类from dataclasses import dataclass dataclass class MyOptimizerConfig(AbstractDataclass): learning_rate: float 0.001 momentum: float 0.9 def __init__(self, learning_rate: float 0.001, momentum: float 0.9) - None: self.learning_rate learning_rate self.momentum momentum # Usage config MyOptimizerConfig(learning_rate0.01, momentum0.95)张量数学工具# 优化张量形状以获得更好性能 tensor_shape (32, 3, 64, 64) # 典型卷积张量 merged_shape merge_small_dims( tensor_shapetensor_shape, threshold1024, target_tensor_dimensionality2 ) print(fOriginal: {tensor_shape}, Merged: {merged_shape}) # Output: Original: (32, 3, 64, 64), Merged: (96, 4096) # 将大张量切分为更小的块 large_tensor torch.randn(100, 200) tensor_blocks multi_dim_split(large_tensor, split_size50) print(fSplit into {len(tensor_blocks)} blocks)优化器模块基础设施# 创建带状态的自定义优化器组件 class AdaptivePreconditioner(OptimizerModule): def __init__(self, dim: int): self.squared_gradients torch.zeros(dim) self.step_count 0 def update(self, grad: torch.Tensor) - torch.Tensor: self.step_count 1 self.squared_gradients grad ** 2 return grad / (self.squared_gradients.sqrt() 1e-8) # 自动状态管理 preconditioner AdaptivePreconditioner(100) state preconditioner.state_dict() preconditioner.load_state_dict(state)内存高效的量化# 为节省内存压缩张量 tensor_list [torch.randn(500, 500, dtypetorch.float32) for _ in range(3)] quantized_list QuantizedTensorList( quantized_data[(t, None, None) for t in tensor_list], quantized_dtypetorch.bfloat16, computation_dtypetorch.float32 ) # 自动精度管理 with DequantizeQuantizedTensorListContext(quantized_list): # 以全精度进行计算 results [t.sum() for t in quantized_list.dequantized_value] # 退出上下文后自动压缩回 bfloat16模型工具# 内存高效的线性层权重与偏置合并为一个参数 combined_layer CombinedLinear(in_features256, out_features128, biasTrue) input_data torch.randn(32, 256) output combined_layer(input_data) # Shape: (32, 128) # 访问合并后的权重与偏置参数 combined_param combined_layer.combined_weight # Shape: (128, 257)字典同步迭代# 同步遍历多个数据源 training_data { samples: [sample1, sample2, sample3], labels: [label1, label2, label3], weights: [1.0, 0.8, 1.2] } for batch in DictZipIterator(training_data): # batch {samples: sample1, labels: label1, weights: 1.0} process_batch(batch)抽象基类与通用工具配置与迭代的基础设施AbstractDataclassabstract_dataclass.py该基类为配置对象提供健壮的抽象 dataclass 模式。其核心技巧是dataclass(initFalse)与默认的initTrue相反它禁止dataclass自动生成__init__从而迫使子类自行实现构造逻辑基类通过abstractmethod声明的__init__保证所有子类都必须实现初始化。这样既保留了 dataclass 的字段声明能力又保证了类型安全与继承层次的可控性。dataclass(initFalse) class MyBaseConfig(AbstractDataclass): Base configuration class. abstractmethod def __init__(self, *args, **kwargs) - None: pass dataclass class ConcreteConfig(MyBaseConfig): Concrete implementation. learning_rate: float 0.001 def __init__(self, learning_rate: float 0.001): self.learning_rate learning_rate关键特性强制正确的抽象 dataclass 模式阻止抽象类被实例化支持多层嵌套继承类型安全的配置对象。从源码注释可知若希望子类继续保持抽象需要在子类上也声明dataclass(initFalse)否则dataclass会为子类自动生成__init__使其成为具体类。这是整个distributed_shampoo中大量*Config对象的共同基类。通用工具commons.pyget_all_non_abstract_subclasses(cls)递归收集某抽象类的全部可实例化子类。源码实现非常精炼先用reduce(or_, map(...))对cls.__subclasses__()递归求并集再通过检查__abstractmethods__是否为空集来过滤抽象类。这一工具在注册/枚举不同预条件子实现如 SGD、AdaGrad、Shampoo时非常实用class BaseOptimizer(ABC): pass class SGDOptimizer(BaseOptimizer): pass class AdamOptimizer(BaseOptimizer): pass concrete_optimizers list(get_all_non_abstract_subclasses(BaseOptimizer)) # Returns: [SGDOptimizer, AdamOptimizer]batched(iterable, n)则是 Python 3.12itertools.batched的复刻实现源码注释建议 Python 3.12 之后直接改用标准库版本当n 1时抛出ValueErrordata range(10) for batch in batched(data, 3): print(batch) # Output: (0, 1, 2), (3, 4, 5), (6, 7, 8), (9,)DictZipIteratordict_zip_iterator.py该类实现字典压缩同步遍历输入是dict[str, Iterator]每次迭代产出一个同键的新字典值取自各迭代器的当前位置。源码在__next__中逐一尝试next()各迭代器并记录已耗尽/仍活跃的键若部分耗尽而部分仍有值则抛出带详细信息的ValueError指明哪些迭代器耗尽、哪些仍活跃从而在长度不匹配时给出清晰报错全部耗尽时正常触发StopIteration。它通过Generic[_DictValType]提供类型安全的泛型支持。data { gradients: [grad1, grad2, grad3], parameters: [param1, param2, param3], learning_rates: [0.1, 0.01, 0.001] } iterator DictZipIterator(data) for batch in iterator: # batch {gradients: grad1, parameters: param1, learning_rates: 0.1} update_parameter(batch[parameters], batch[gradients], batch[learning_rates])OptimizerModule面向优化器的状态管理基类optimizer_modules.py 提供与nn.Module相似但去脂肪的优化器组件基类。它只保留状态管理能力专为梯度累积器、预条件子等优化器内部件设计。class PreconditionerModule(OptimizerModule): def __init__(self, shape: tuple[int, ...]): self.shape shape self.accumulator torch.zeros(shape) self.step_count 0 def update(self, gradient: torch.Tensor): self.accumulator gradient ** 2 self.step_count 1 def precondition(self, gradient: torch.Tensor) - torch.Tensor: return gradient / (self.accumulator.sqrt() 1e-8) # Usage preconditioner PreconditionerModule((100, 50)) # State management state_dict preconditioner.state_dict() preconditioner.load_state_dict(state_dict)从源码看state_dict(destination, keep_vars, store_non_tensors)的实现要点包括递归构建save_to_state_dict遍历self.__dict__遇到torch.Tensor直接存储keep_varsFalse时detach()遇到嵌套OptimizerModule递归调用其state_dict遇到dict、list/tuple/set分别展开存储空容器会被remove_empty_entry剔除仅当store_non_tensorsTrue时才保存普通非张量对象。严格加载load_state_dict要求加载的状态与已有状态完全初始化。加载时对张量执行old_state.detach().copy_(new_state)直接原地拷贝避免状态复制对集合容器加载时会调用_convert_state_key_from_str_to_int把 PyTorch 展平flatten后变为字符串的索引键还原为整数保证与flatten_optimizer_state_dict生态兼容非张量对象则通过deepcopy恢复。Tensor 感知的序列化keep_vars选项控制是否保留 autograd 图。这套严格语义允许在断点续训时直接optimizer.load_state_dict(sd)完成状态恢复源码注释中的示例流程optimizer.step()→sd optimizer.state_dict()→load_checkpoint(sd)→optimizer.load_state_dict(sd)同时通过类型/键的校验提升可用性。张量数学工具维度合并、分块与分布式分配shampoo_utils.py 是模块中数学计算最密集的文件承担 Shampoo 预条件子所需的一切张量预处理。维度合并merge_small_dims将相邻小维度合并到阈值以下以提升算子效率。源码细节值得注意函数带cache装饰相同形状输入直接命中缓存反向合并按 PyTorch 张量布局从后往前合并对卷积核置于形状末端的场景尤其重要合并前先squeeze掉大小为 1 的维度若全为 1 则保留(1,)空张量返回(0,)0D 张量会被提升为 1Dtarget_tensor_dimensionality为 float 时仅允许math.inf否则触发断言表示不限制维度、不做合并合并会提前停止即使阈值允许继续合并一旦达到目标维度数便不再合并。# 优化张量形状小维度全部合并 original_shape (1, 2, 5, 1) # Small dimensions merged_shape merge_small_dims( tensor_shapeoriginal_shape, threshold10, target_tensor_dimensionality1 ) print(merged_shape) # (10,) - all dimensions merged # 卷积类张量 conv_shape (32, 3, 64, 64) merged_conv merge_small_dims( tensor_shapeconv_shape, threshold8192, target_tensor_dimensionality2 ) print(merged_conv) # (96, 4096) - optimal for Muon-style optimizers第二个示例正是 Muon 类谱系降算法spectral descent的典型用法把卷积参数重塑为 2D 后做半正交化。这一点与distributed_shampooREADME 中支持基于 reduced SVD / Newton-Schulz 迭代实现 Muon的特性相互印证。多维切分multi_dim_split沿张量所有维度依次执行torch.split源码用reduce对每个维度依次展平切分结果。当split_size为math.inf时返回(tensor,)原样不动当某维度尺寸不超过split_size时该维度不切分。# 沿所有维度切分张量 tensor torch.randn(5, 3) split_tensors multi_dim_split(tensor, split_size2) # Returns tuple of smaller tensors after splitting along each dimension # 尺寸不超过切分大小时不切分 large_split multi_dim_split(tensor, split_sizemath.inf) # Returns (tensor,) unchanged其他常用函数# 按布尔选择器压缩序列 data [a, b, c, d] selector [True, False, True, False] compressed compress_list(data, selector) print(compressed) # (a, c) # 获取数据类型的内存占用 float32_size get_dtype_size(torch.float32) # 4 bytes bool_size get_dtype_size(torch.bool) # 1 byte # 为分区生成累计区间索引 partitions [2, 3, 1] # Partition sizes indices list(generate_pairwise_indices(partitions)) print(indices) # [(0, 2), (2, 5), (5, 6)] # 在各 rank 间分配缓冲区大小 buffer_sizes (128, 64, 500, 256) distribution distribute_buffer_sizes(buffer_sizes, group_size2) # Balances memory allocation across 2 rankscompress_list断言两序列等长内部用itertools.compress实现统一返回 tuple 以保证下游兼容get_dtype_sizebool恒为 1 字节其余用(bits 7) // 8向上取整float 走torch.finfo、整型走torch.iinfogenerate_pairwise_indices一行实现pairwise(accumulate(chain([0], input_list)))用于根据每个参数的分块数量生成区间。上下文管理器与分布式缓冲分配ParameterizeEnterExitContext接受任意对象与一对进入/退出方法回调用partial绑定对象后包装成with上下文class StatefulObject: def __init__(self): self.active False def activate(self): self.active True def deactivate(self): self.active False obj StatefulObject() with ParameterizeEnterExitContext( input_with_enter_exit_contextobj, enter_method_callerlambda x: x.activate(), exit_method_callerlambda x: x.deactivate() ): assert obj.active # True inside context assert not obj.active # False after contextdistribute_buffer_sizes(blocked_params, group_size, load_balancing_config)是分布式预条件子分配的核心默认以DefaultCostModel.cost()计算对齐后的缓冲字节数同时按load_balancing_config.cost_model.cost()计算每个块的负载随后用heapq实现贪心分配——每次把负载最重的块交给当前累计负载最小的 rank最终返回(buffer_size, rank)元组列表保证各 rank 总负载尽可能均匀。源码注释也指出更优的策略应显式最小化最大/最小分配负载之间的差值。同文件还提供prepare_update_param_buffers分配参数更新的持久影子缓冲处理参数数少于 group_size时的零张量填充与redistribute_and_update_params通过多轮dist.all_to_all集合通信交换更新后的参数并用torch._foreach_copy_一次性写回各参数本地切片支撑 DDP Shampoo 的 ZeRO-1 式通信路径。模型工具CombinedLinear 组合线性层shampoo_model_utils.py 中的CombinedLinear是为利用张量结构的优化器如 Shampoo而特化的线性层它把权重与偏置合并进同一个参数combined_weight形状为(out_features, in_features_with_bias)biasTrue时in_features_with_bias in_features 1从而减少需单独预条件处理的参数个数。# 创建组合线性层 layer CombinedLinear( in_features512, out_features256, biasTrue ) # 前向传播 input_tensor torch.randn(32, 512) # (batch_size, in_features) output layer(input_tensor) # (32, 256) # 访问合并参数权重与偏置拼接 combined_param layer.combined_weight # Shape: (256, 513) when biasTrue weight_part layer.combined_weight[:, :-1] # (256, 512) bias_part layer.combined_weight[:, -1] # (256,)源码实现要点标准初始化reset_parameters对权重部分用kaiming_uniform_(asqrt(5))等价于uniform(-1/sqrt(in), 1/sqrt(in))偏置列单独用uniform(-bound, bound)其中bound 1/sqrt(fan_in)与nn.Linear的初始化语义一致前向解耦forward在biasTrue时通过F.linear(input, combined_weight[:, :-1], combined_weight[:, -1])同时使用权重与偏置biasFalse时退化为无偏置线性变换即插即用可作为nn.Linear的替代层in_features/out_features语义不变。状态字典管理复杂嵌套状态的存取shampoo_state_dict_utils.py 为含嵌套OptimizerModule的复杂状态字典提供健壮的检查点存取支持。# 从嵌套对象中提取状态字典 nested_objects { preconditioner: some_optimizer_module, parameters: {weight: tensor} } state_dict extract_state_dict_content(nested_objects) # 就地更新状态字典对象 current_state {param1: tensor1, param2: {nested: tensor2}} new_state {param1: new_tensor1, param2: {nested: new_tensor2}} update_param_state_dict_object(current_state, new_state) # current_state is updated in place源码细节extract_state_dict_content递归遍历输入字典遇到OptimizerModule调用其state_dict()遇到嵌套 dict 递归展开其余值原样保留update_param_state_dict_object(current, to_load, enable_missing_key_checkTrue)就地更新dict 递归处理带load_state_dict方法的对象调用其加载逻辑torch.Tensor执行detach().copy_()其余标量值deepcopy覆盖。enable_missing_key_check为True时缺失键直接抛KeyError为False时仅记录 warning 并跳过——这为严格/宽松两种检查点恢复策略提供了开关。量化支持内存高效的低精度张量压缩shampoo_quantization.py 提供张量的压缩与解压框架支持 FP16/BF16/FP32/FP64 等浮点格式并与分布式分块信息BlockInfo集成。QuantizedTensorQuantizedTensor继承自OptimizerModule内部持有quantized_values、min_value/max_value元数据与block_info。源码中_FLOAT_DTYPES (float16, bfloat16, float32, float64)当前仅支持浮点类型间的直接转换dest.copy_(src)其余量化格式会抛NotImplementedError。# 从全精度张量创建量化张量 full_precision_tensor torch.randn(1000, 1000, dtypetorch.float32) quantized QuantizedTensor.init_from_dequantized_tensor( dequantized_valuesfull_precision_tensor, quantized_dtypetorch.bfloat16, block_infoblock_info ) # 解压用于计算 result quantized.dequantize(torch.float32)QuantizedTensorList 与自动上下文QuantizedTensorList管理一批量化张量构造函数接受(tensor, min, max)元组列表或QuantizedTensor列表两种输入并断言存储精度一致、computation_dtype必须是_FLOAT_DTYPES之一。dequantize_()将量化值展开为计算精度并缓存quantize_()在退出时压缩回存储精度并释放缓存torch.cuda.empty_cache()dequantized_value属性只在解压缓存存在时可用否则触发断言防止误用。# 创建量化张量列表 tensor_list [torch.randn(100, 100) for _ in range(5)] quantized_list QuantizedTensorList( quantized_data[(t, None, None) for t in tensor_list], quantized_dtypetorch.bfloat16, computation_dtypetorch.float32 ) # 自动解压上下文 with DequantizeQuantizedTensorListContext(quantized_list): # 访问解压后的张量进行计算 dequantized_tensors quantized_list.dequantized_value # Perform computations... results [torch.matmul(t, t.T) for t in dequantized_tensors] # 退出上下文后自动重新量化 # 手动控制 quantized_list.dequantize_() # 缓存解压版本 # ... 执行计算 ... quantized_list.quantize_() # 转回量化格式DequantizeQuantizedTensorListContext复用了前文所述的ParameterizeEnterExitContext进入时调用dequantize_、退出时调用quantize_——这是以类组合替代继承复用上下文逻辑的典型示例。压缩与选择# 基于布尔选择器压缩 selector (True, False, True, False, True) compressed_list quantized_list.compress(selector) # 返回仅含被选中张量的新 QuantizedTensorListcompress要求当前未缓存解压值断言内部用compress_list分别压缩量化值、min/max 元数据后重建列表常用于按激活掩码裁剪部分状态。负载均衡工具计算与内存成本模型load_balancing_utils.py 为分布式训练中的张量负载均衡提供成本估算模型。所有模型继承自CostModel本身继承AbstractDataclass只需实现cost(tensor) - float。计算成本模型PolynomialComputationalCostModel以多项式函数估算计算成本系数个数即多项式阶数对张量的每个维度分别求值后求和并以min_cost兜底# 二次成本模型: cost a b*x c*x² cost_model PolynomialComputationalCostModel( coefficients(1.0, 0.1, 0.01), # a1.0, b0.1, c0.01 min_cost10.0 # 最小成本阈值 ) # 示例张量的计算成本 tensor torch.randn(100, 200) total_cost cost_model.cost(tensor) # Cost max(10.0, (1.0 0.1*100 0.01*100²)) max(10.0, (1.0 0.1*200 0.01*200²))源码用numpy.polynomial.polynomial.polyval求值min_cost默认值为 0。内存成本模型AlignedMemoryCostModel按对齐字节数计算内存成本buffer_size tensor.numel() * tensor.element_size()再向上取整到alignment_bytes的倍数默认 64 字节对齐返回对齐后的缓冲大小字节# 默认 64 字节对齐 memory_model AlignedMemoryCostModel(alignment_bytes64) # 内存成本计算 tensor torch.randn(100, 100, dtypetorch.float32) memory_cost memory_model.cost(tensor) # Calculates: aligned_size ceil(100*100*4 / 64) * 64 bytes # 自定义对齐 cache_aligned_model AlignedMemoryCostModel(alignment_bytes128) cost_128 cache_aligned_model.cost(tensor)模块底部定义DefaultCostModel AlignedMemoryCostModel()作为分布式分配的默认成本口径源码注释明确其为向后兼容的默认选择。在负载均衡中的使用# 基于计算成本分配张量 tensors [torch.randn(size) for size in [(100, 50), (200, 100), (50, 200)]] comp_model PolynomialComputationalCostModel(coefficients(0, 1, 0)) # 维度线性 # 计算成本用于负载均衡 costs [comp_model.cost(t) for t in tensors] # 用 costs 跨设备/进程分配 # 内存感知分配 mem_model AlignedMemoryCostModel(alignment_bytes64) memory_costs [mem_model.cost(t) for t in tensors] # 跨设备均衡内存占用典型应用场景包括分布式预条件子分配、内存感知的张量分块、计算负载均衡、资源分配优化。综合示例Shampoo 参数处理流水线官方 README 给出的两个综合示例把上述工具串联成完整的 Shampoo 参数处理与分布式调度流程这里完整保留并补充注释。自定义配置模式from dataclasses import dataclass dataclass(initFalse) class BaseOptimizerConfig(AbstractDataclass): Abstract base for optimizer configurations. abstractmethod def __init__(self) - None: pass dataclass class ClassicShampooPreconditionerConfig(BaseOptimizerConfig): Concrete Shampoo configuration. max_preconditioner_dim: int 8192 precondition_frequency: int 100 epsilon: float 1e-8 def __init__( self, max_preconditioner_dim: int 8192, precondition_frequency: int 100, epsilon: float 1e-8 ): self.max_preconditioner_dim max_preconditioner_dim self.precondition_frequency precondition_frequency self.epsilon epsilon带状态管理的优化器模块class AdaptivePreconditioner(OptimizerModule): def __init__(self, param_shape: tuple[int, ...], beta2: float 0.999): self.param_shape param_shape self.beta2 beta2 self.step 0 self.squared_avg torch.zeros(param_shape) def update(self, gradient: torch.Tensor) - None: self.step 1 self.squared_avg.mul_(self.beta2).addcmul_( gradient, gradient, value1 - self.beta2 ) def precondition(self, gradient: torch.Tensor) - torch.Tensor: bias_correction 1 - self.beta2 ** self.step corrected_avg self.squared_avg / bias_correction return gradient / (corrected_avg.sqrt() 1e-8) # 使用自动状态管理 preconditioner AdaptivePreconditioner((1000, 500)) state preconditioner.state_dict() # 包含 step, squared_avg 等 preconditioner.load_state_dict(state) # 恢复状态张量处理的数学工具组合# 为 Shampoo 处理卷积参数 def process_conv_params(param_tensor: torch.Tensor) - tuple[torch.Tensor, ...]: # 合并小维度以获得更好的数值性质 original_shape param_tensor.shape merged_shape merge_small_dims( tensor_shapeoriginal_shape, threshold1024, target_tensor_dimensionality2 ) # 重塑参数 reshaped_param param_tensor.view(merged_shape) # 必要时切分为可管理的块 if max(merged_shape) 8192: blocks multi_dim_split(reshaped_param, split_size4096) else: blocks (reshaped_param,) return blocks # 跨设备分配计算 def setup_distributed_computation( param_blocks: list[torch.Tensor], num_devices: int ) - dict[int, list[torch.Tensor]]: # 计算内存需求 buffer_sizes tuple(block.numel() * 4 for block in param_blocks) # 每个 float32 占 4 字节 # 跨设备分配 assignments distribute_buffer_sizes(buffer_sizes, num_devices) # 按设备分组 device_assignments {} for i, (size, device_id) in enumerate(assignments): if device_id not in device_assignments: device_assignments[device_id] [] device_assignments[device_id].append(param_blocks[i]) return device_assignments在 modded-nanogpt 训练脚本中的实际接入工具模块的价值最终体现在真实训练流程中。实验目录下的 train_gpt_shampoo.py 是 modded-nanogpt speedrun 训练脚本train_gpt_simple.py的 Shampoo 改造版其优化器部分展示了distributed_shampoo高层 API 与底层utils的配合from distributed_shampoo import ( AdamPreconditionerConfig, DDPDistributedConfig, DistributedShampoo, SingleDeviceDistributedConfig, WeightDecayType, ) def shampoo_distributed_config(): if dist.get_world_size() 1: return SingleDeviceDistributedConfig() return DDPDistributedConfig( communication_dtypetorch.float32, num_trainers_per_groupdist.get_world_size(), communicate_paramsTrue, )该脚本采用双优化器分工策略AdamW负责 embedding、输出投影与全部一维参数偏置、norm gains分别配置lr0.3、1/320、0.01betas(0.8, 0.95)DistributedShampoo负责所有ndim 2的 Transformer Block 参数配置为lr0.0015、betas(0.9, 0.95)、weight_decay0.2WeightDecayType.DECOUPLED即 AdamW 式解耦权重衰减、max_preconditioner_dim8192、precondition_frequency5、start_preconditioning_step-1从第一步就开始预条件并通过grafting_configAdamPreconditionerConfig(beta20.95, epsilon1e-10)从 Adam 的矩估计中嫁接学习率行为。学习率采用稳定期 冷却期调度cooldown_frac0.7后 30% 步数线性衰减到零配合 4150 步训练目标与8 × 64 × 1024的全局 batch size、64 的微批大小。脚本开头会断言两份优化器的参数集合恰好覆盖模型全部参数避免遗漏。实验中采用的max_preconditioner_dim8192、precondition_frequency5等取值正对应distributed_shampooREADME 中关于调参的指导以接近纯 Shampoo8192 / 1为起点再按内存与性能约束逐步放宽预条件频率。运行日志 503575c5-6dde-425a-b461-2df4d99db974.txt 的开头完整记录了训练脚本源码modded-nanogpt 的日志惯例说明该实验确实以这份脚本在 track_3 优化赛道中进行了实证验证完整实验结果与对比图可进一步查看 track_3_optimization/README.md 及同目录results/下的其他记录。贡献与测试约定官方 README 为后续向工具模块贡献代码给出了明确约定摘要如下保持向后兼容工具被全代码库广泛引用添加全面测试所有工具必须有充分测试覆盖记录边界情形清晰文档化边界条件与错误场景性能优先针对常见用例优化尤其是数学工具类型安全使用规范的类型标注并兼容pyre-strict内存效率特别是量化工具的访存模式跨平台兼容确保在不同硬件配置下可用。测试规范方面测试文件放入tests/目录遵循{module_name}_test.py命名覆盖边界条件、错误条件与性能测试优先使用参数化测试涉及分布式场景时必须编写分布式测试。文档规范则要求所有公开函数/类带 docstringdocstring 内提供用法示例记录参数约束与返回值格式并在新增工具时同步更新本 README。这些约定同样适用于本实验目录utils/tests/下已有abstract_dataclass_test.py、commons_test.py、dict_zip_iterator_test.py、optimizer_modules_test.py、shampoo_model_utils_test.py、shampoo_quantization_test.py、shampoo_state_dict_utils_test.py、shampoo_utils_test.py八个测试文件全部遵循上述命名与覆盖规范可作为自定义扩展时的参照模板。赞分享人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载相关推荐modded-nanogpt 实战用 PyTorch Distributed Shampoo 二阶优化器替换 AdamW 加速 GPT 训练modded nanogpt 实战用 PyTorch Distributed Shampoo 二阶优化器替换 AdamW 加速 GPT 训练 本文以 modd人工智能大模型预训练分布式训练模型优化深度学习modded-nanogpt 优化研究Distributed Shampoo Distributor 分布式引擎架构与多卡并行实战指南modded nanogpt 优化研究Distributed Shampoo Distributor 分布式引擎架构与多卡并行实战指南 本篇文章以 modde人工智能大模型预训练分布式训练模型优化深度学习modded-nanogpt 优化实验解析为 MLP 权重引入 Shampoo 基预条件的 Contra-MuonTrack 3 结果 14modded nanogpt 优化实验解析为 MLP 权重引入 Shampoo 基预条件的 Contra MuonTrack 3 结果 14 本文围绕 m人工智能大模型预训练分布式训练模型优化深度学习上一篇Argo CD 通知触发器调试指南argocd admin notifications trigger 命令详解与源码剖析下一篇polkadot/apps 终极指南探索 Polkadot 生态系统的未来路线图与发展机遇创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考