[ PROMPT_NODE_22496 ]
Infrastructure Lambda Labs 高级用法
[ SKILL_DOCUMENTATION ]
# Lambda Labs 高级用法指南
## 多节点分布式训练
### 跨节点的 PyTorch DDP
python
# train_multi_node.py
import os
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
def setup_distributed():
# 由启动器设置的环境变量
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
local_rank = int(os.environ["LOCAL_RANK"])
dist.init_process_group(
backend="nccl",
rank=rank,
world_size=world_size
)
torch.cuda.set_device(local_rank)
return rank, world_size, local_rank
def main():
rank, world_size, local_rank = setup_distributed()
model = MyModel().cuda(local_rank)
model = DDP(model, device_ids=[local_rank])
# 带有同步梯度的训练循环
for epoch in range(num_epochs):
train_one_epoch(model, dataloader)
# 仅在 rank 0 上保存检查点
if rank == 0:
torch.save(model.module.state_dict(), f"checkpoint_{epoch}.pt")
dist.destroy_process_group()
if __name__ == "__main__":
main()
### 在多个实例上启动
bash
# 在节点 0 (主节点) 上
export MASTER_ADDR=
export MASTER_PORT=29500
torchrun
--nnodes=2
--nproc_per_node=8
--node_rank=0
--master_addr=$MASTER_ADDR
--master_port=$MASTER_PORT
train_multi_node.py
# 在节点 1 上
export MASTER_ADDR=
export MASTER_PORT=29500
torchrun
--nnodes=2
--nproc_per_node=8
--node_rank=1
--master_addr=$MASTER_ADDR
--master_port=$MASTER_PORT
train_multi_node.py
### 用于大模型的 FSDP
python
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
from transformers.models.llama.modeling_llama import LlamaDecoderLayer
# Transformer 模型的自动封装策略
auto_wrap_policy = functools.partial(
transformer_auto_wrap_policy,
transformer_layer_cls={LlamaDecoderLayer}
)
model = FSDP(
model,
auto_wrap_policy=auto_wrap_policy,
mixed_precision=MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16,
),
device_id=local_rank,
)
### DeepSpeed ZeRO
python
# ds_config.json
{
"train_batch_size": 64,
"gradient_accumulation_steps": 4,
"fp16": {"enabled": true},
"zero_optimization": "..."
}