- 人工智能
- 深度学习
- 机器学习
- 预训练
- 分布式训练
- 微调
【免费下载链接】pytorch-lightning
Pretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.
本篇文章以 PyTorch Lightning 官方仓库中的 Fabric Methods API 文档 为骨架,逐一剖析Fabric对象暴露的 18 组核心方法——涵盖进程启动(launch)、模型与数据装配(setup/setup_dataloaders)、梯度反向(backward/clip_gradients)、精度控制(autocast/init_module)、分布式通信(barrier/broadcast/all_gather/all_reduce)、断点续训(save/load/load_raw)以及回调与日志(call/log/log_dict)等。读完本文,你将掌握在不改动模型代码的前提下,用最小代价把一段普通 PyTorch 训练脚本改造成可在单机多卡、多机多卡、混合精度乃至超大模型分片场景下稳定运行的 Fabric 训练程序,并理解每个方法底层在 fabric.py 中是如何实现的。
方法总览:一份可直接查阅的速查表
Fabric的全部加速能力都收敛在一个类上,其实例通过Fabric(accelerator=..., strategy=..., devices=..., precision=...)构造。下表汇总了本文将要展开的全部方法及其一句话用途,方便速查:
| 方法 | 一句话用途 | 对应源码位置 |
|---|---|---|
launch(function=None) | 启动分布式进程,可传入函数在子进程中执行 | fabric.py |
setup(model, *optimizers, scheduler=None, move_to_device=True) | 装配模型与优化器,自动搬到设备并适配精度 | fabric.py |
setup_module(model)/setup_optimizers(*opts) | 分步装配模型或优化器(FSDP 等策略需要) | fabric.py |
setup_dataloaders(*dls, use_distributed_sampler=True, move_to_device=True) | 装配数据加载器,自动替换分布式采样器 | fabric.py |
backward(tensor, model=None) | 替代loss.backward(),兼容精度与策略 | fabric.py |
clip_gradients(model, optimizer, clip_val/max_norm, norm_type, error_if_nonfinite) | 梯度裁剪(按值或按范数) | fabric.py |
to_device(obj) | 把模型/张量/嵌套集合搬到当前设备 | fabric.py |
seed_everything(seed, workers=True) | 统一设置 torch/numpy/random 种子 | fabric.py |
init_module(empty_init=None) | 上下文管理器:按目标设备与精度实例化模型 | fabric.py |
autocast() | 上下文管理器:对指定代码块自动精度转换 | fabric.py |
print(*args) | 只在主进程打印,避免多进程刷屏 | fabric.py |
save(path, state, filter=None) | 替代torch.save()保存状态字典 | fabric.py |
load(path, state=None, strict=True, weights_only=None) | 替代torch.load()恢复状态 | fabric.py |
load_raw(path, obj, strict=True) | 从非 Fabric 保存的裸 PyTorch 权重恢复 | fabric.py |
barrier(name=None) | 让所有进程在此等待同步 | fabric.py |
broadcast(obj, src=0)/all_gather(data)/all_reduce(data, reduce_op="mean") | 进程间集合通信 | fabric.py |
no_backward_sync(model, enabled=True) | 梯度累积期跳过冗余梯度同步 | fabric.py |
call(hook_name, *args, **kwargs) | 触发已注册回调的同名钩子 | fabric.py |
log(name, value, step=None)/log_dict(metrics, step=None) | 向 Logger 写入标量指标 | fabric.py |
launch:以零样板代码启动多进程分布式训练
launch是进入分布式世界的入口。它的核心价值在于:你不需要手动调用torch.distributed.init_process_group,也不需要借助torch.multiprocessing.spawn或任何外部启动工具,Fabric 会在内部按所选策略(如 DDP)完成进程组初始化、环境变量设置与子进程派发。
文档给出了两种用法:
# 用法一:只初始化进程组,后续代码继续在“当前脚本”执行 fabric = Fabric(devices=2) fabric.launch() # 用法二:把训练逻辑放进函数,由 launch 在多个进程中执行 def run(fabric): # Your distributed code here ... fabric = Fabric(devices=2) fabric.launch(run)从 fabric.py 的实现可以看到几个关键行为:
- CLI 冲突保护:如果脚本是通过
fabric run ...这种 CLI 方式启动的(_is_using_cli()返回真),进程组已经就绪,此时再调用launch()会直接抛出RuntimeError,提示“进程已创建,不允许再次调用.launch()”。这意味着launch只适用于在代码里以编程方式指定 accelerator、devices 等参数的场景。 - 函数签名校验:传入的函数必须至少接受一个参数(
inspect.signature(function).parameters为空会抛TypeError),因为 Fabric 会把自身实例作为第一个参数传进去。常见笔误fabric.launch(your_fn())也会被拦下(HINT 提示应传your_fn而非其调用结果)。 - spawn 型策略强制要求传函数:当策略的 launcher 是
_MultiProcessingLauncher或_XLALauncher(例如 TPU/XLA 场景)时,必须以函数形式调用launch,否则无法知道子进程要执行什么代码。 - 返回值约定:如果传了函数,
launch返回在 rank 0 工作进程中执行该函数的结果;未传函数时仅做进程初始化。
文档特别提到,用法二可以用于在 Jupyter notebook 中进行多 GPU 训练(参见 notebooks 指南),而 CLI 启动、多节点集群、云环境等场景则参考 launch 指南。
一个完整的可运行骨架如下:
import torch from lightning.fabric import Fabric def train(fabric: Fabric): model = torch.nn.Linear(32, 64) optimizer = torch.optim.SGD(model.parameters(), lr=0.001) model, optimizer = fabric.setup(model, optimizer) # ... 训练循环 ... fabric = Fabric(accelerator="gpu", devices=2, strategy="ddp") fabric.launch(train) # 返回 rank 0 上 train 的返回值setup:模型与优化器的“一键加速装配”
setup是使用频率最高的方法。文档强调:模型和对应优化器必须经过setup之后才能用于训练;如果有多个模型,需要对每个模型分别调用setup()。其职责是把模型和优化器自动移动到正确的设备上,并按所选精度对模型进行适配。
model = nn.Linear(32, 64) optimizer = torch.optim.SGD(model.parameters(), lr=0.001) scheduler = torch.optim.lr_scheduler.LinearLR(optimizer, start_factor=1.0, end_factor=0.3, total_iters=10) # 标准用法:返回装配好的模型与优化器 model, optimizer = fabric.setup(model, optimizer) # 不想要 Fabric 自动搬移设备 model, optimizer = fabric.setup(model, optimizer, move_to_device=False) # 额外注册学习率调度器(DeepSpeed 等兼容策略可用) model, optimizer, scheduler = fabric.setup(model, optimizer, scheduler)对照 fabric.py 的源码,setup的内部流水线依次是:
- 参数校验(
_validate_setup):模型/优化器不能是已装配过的_FabricModule/_FabricOptimizer;若使用 FSDP 且优化器引用了 meta 设备参数,会提示改用setup_module+setup_optimizers分步装配。 - 精度转换:调用
self._precision.convert_module(module),这是“setup 会按所选精度准备模型、使forward()中的运算自动转换”的底层来源。 - 设备搬移(
move_to_device=True时):_move_model_to_device会把模型搬到目标设备;源码中还会检测模型参数是否横跨多个设备,并给出PossibleUserWarning提示。对于 XLA/TPU 策略,还会特殊处理参数引用映射。 - 策略包装:交给
self._strategy.setup_module_and_optimizers(...)或setup_module(...),由 DDP/FSDP/DeepSpeed 等策略完成各自包装。 - 封装返回:模型被包成
_FabricModule,优化器被包成_FabricOptimizer,并触发on_after_setup回调。
几点值得注意的边界行为:
- 模型可单独 setup:
fabric.setup(model)只传模型也是合法的,返回值是包装后的模型。 - 返回顺序与传入顺序一致:
model, opt1, opt2, scheduler = fabric.setup(model, opt1, opt2, scheduler=scheduler)。 move_to_device=False的含义:用于手动控制设备搬移,此时设备属性取自模型首个参数所在设备。- FSDP 等分片策略的例外:源码文档串明确提示,某些策略(如 FSDP)需要先用
setup_module建立模型、再创建优化器、最后用setup_optimizers装配优化器,因为优化器创建时需要引用已分片/已搬到 meta 设备的参数。而 DeepSpeed、XLA 策略则必须通过setup(model, optimizer, ...)联合装配,单独调用setup_optimizers会抛RuntimeError。 - 文档同时提醒:
setup还会让模型适配所选精度,使forward()内运算自动转换;高级读者应阅读 wrappers 相关说明 了解模型被包装后的细节。
setup_dataloaders:数据加载器的分布式装配
setup_dataloaders接收一个或多个torch.utils.data.DataLoader,为加速训练做准备。文档明确了两个核心行为:在分布式策略(如 DDP)下自动替换采样器,以及让返回的批次数据自动搬移到正确设备。
train_data = torch.utils.DataLoader(train_dataset, ...) test_data = torch.utils.DataLoader(test_dataset, ...) # 装配多个加载器 train_data, test_data = fabric.setup_dataloaders(train_data, test_data) # 关闭自动设备搬移 train_data, test_data = fabric.setup_dataloaders(train_data, test_data, move_to_device=False) # 关闭分布式采样器替换(用于自定义采样器场景) train_data, test_data = fabric.setup_dataloaders(train_data, test_data, use_distributed_sampler=False)源码层面(fabric.py 的_setup_dataloader)揭示了自动化的判断逻辑:
- 何时替换采样器(
_requires_distributed_sampler):只有当策略带有distributed_sampler_kwargs、当前采样器不是DistributedSampler、且数据集不是iterable 类型时才会替换。 - 替换策略:如果原采样器是
RandomSampler/SequentialSampler,直接替换为DistributedSampler(shuffle会继承原随机采样器的设定,seed默认取自环境变量PL_GLOBAL_SEED);如果是自定义采样器,则包装成DistributedSamplerWrapper。这一点与 test_fabric.py 中的test_setup_dataloaders_replace_standard_sampler等用例相互印证。 - 自动种子初始化:
_auto_add_worker_init_fn(dataloader, self.global_rank)会为 DataLoader 注入worker_init_fn,保证每个 worker 进程的随机状态可复现(这也是seed_everything提及“Fabric 会妥善初始化 DataLoader worker 进程种子”的落地实现)。 - 设备搬移:封装成
_FabricDataLoader,其迭代时自动把返回数据move_data_to_device到self.device。
setup_dataloaders的入参校验同样严格:不接受空参数(至少一个 DataLoader)、不接受非 DataLoader 对象、不能重复装配同一个已包装的加载器。返回规则是“单入单出、多入多出”:传一个返回一个,传多个返回同顺序的列表。
backward 与 clip_gradients:让反向传播与梯度裁剪与策略无关
backward:替代 loss.backward()
output = model(input) loss = loss_fn(output, target) # loss.backward() fabric.backward(loss)文档说得很清楚:backward取代代码中所有loss.backward()调用,使代码对加速器和精度(precision)无关。源码实现(fabric.py)的关键点:
- 最终调用
self._strategy.backward(tensor, module, *args, **kwargs),由策略决定如何在 DDP 梯度同步、DeepSpeed 引擎、AMP 缩放等场景下正确执行反向。 - DeepSpeed + 多模型:使用
strategy="deepspeed"且设置了多个模型时,必须显式传入fabric.backward(loss, model=model)指定由哪个模型执行优化;若一个模型都没 setup 会抛RuntimeError,setup 了多个却不传model会抛ValueError。 backward还接收*args/**kwargs并透传给底层反向函数。
clip_gradients:按值或按范数裁剪
# 按值裁剪:梯度限制在 ±0.5 内 fabric.clip_gradients(model, optimizer, clip_val=0.5) # 按范数裁剪:总范数不超过 2.0 fabric.clip_gradients(model, optimizer, max_norm=2.0) # 默认使用 2-范数 fabric.clip_gradients(model, optimizer, max_norm=2.0, norm_type=2) # 也可以使用无穷范数:裁剪所有元素中最大的那个 fabric.clip_gradients(model, optimizer, max_norm=2.0, norm_type="inf")对照 fabric.py 的实现,有几个容易忽略的细节:
clip_val与max_norm互斥:同时传两者会抛ValueError;两者都不传也会抛ValueError(“You have to specify eitherclip_valormax_normto do gradient clipping!”)。- 返回值语义:传
max_norm时,clip_gradients返回裁剪前梯度的总范数(标量张量);传clip_val时返回None。 norm_type默认 2.0,支持"inf"(无穷范数,即裁剪所有元素中的最大值)。error_if_nonfinite=True(默认):当总范数为 NaN 或无穷时抛错,仅在max_norm模式下生效——这有助于尽早暴露数值发散问题。- 策略无关:方法最终委托
self.strategy.clip_gradients_value或clip_gradients_norm,传入前会用_unwrap_objects解包_FabricModule/_FabricOptimizer,因此在单卡、DDP、FSDP 等不同策略下行为一致。
to_device 与 seed_everything:设备搬移与可复现性
to_device:手动搬移任意对象
data = torch.load("dataset.pt") data = fabric.to_device(data)to_device用于把模型、张量或张量集合(dict/list/tuple 等嵌套结构)搬到当前设备。由于setup和setup_dataloaders默认已自动搬移模型与数据,文档明确说明该方法仅在需要手动操作时才必须调用。实现上(fabric.py):对nn.Module走strategy.module_to_device,对其他对象走move_data_to_device(递归处理嵌套集合)。
seed_everything:一次调用覆盖三大随机源
# 替代 torch.manual_seed(...),直接调用: fabric.seed_everything(1234)文档指出该调用覆盖PyTorch、NumPy、Python 内置 random三个随机数生成器,并且 Fabric 会负责正确初始化 DataLoader worker 进程的种子(可通过workers=False关闭)。底层实现是 seed.py 的seed_everything函数,几个值得说明的源码事实:
- 调用后会设置环境变量
PL_GLOBAL_SEED(传给 spawn 出的子进程,例如ddp_spawn)和PL_SEED_WORKERS。 seed省略时从PL_GLOBAL_SEED读取,两者皆无则默认为 0;种子必须落在[0, 2^32-1](uint32 范围)内,否则抛ValueError。- Fabric 实例方法默认
workers=True(注释说明这是新版本有意为之,以确保更好的可复现性);而 worker 种子的实际注入依赖setup_dataloaders中_auto_add_worker_init_fn添加的worker_init_fn。
init_module 与 autocast:省内存的模型初始化与按需精度转换
init_module:直接在目标设备上按目标精度建模型
PyTorch 默认在 CPU 上以 float32 实例化nn.Module的所有参数,之后搬到 GPU 会产生一次等待时间。init_module上下文管理器允许你不修改模型代码,强制让模型直接创建在目标设备、目标精度上:
fabric = Fabric(accelerator="cuda", precision="16-true") with fabric.init_module(): # 这里创建的模型直接在 GPU 上、以 float16 精度实例化 model = MyModel()文档强调这消除了“从 CPU 搬移到设备”的等待时间;更重要的是,对处理超大分片模型的策略(FSDP、DeepSpeed),init_module会先把模型参数分配到 meta 设备上再进行分片,从而让单个设备显存放不下的超大模型也能工作。更多细节见 model_init 指南。
源码中(fabric.py),init_module(empty_init=None)委托strategy.module_init_context(empty_init=empty_init):empty_init传None时由策略自行决定;文档与源码都建议在“向大模型加载 checkpoint”时设empty_init=True(用未初始化内存分配参数,加载后立即覆盖)。另外,旧 APIsharded_model()已弃用,官方推荐统一使用init_module()。
autocast:给 forward 之外的运算也加上自动精度转换
model, optimizer = fabric.setup(model, optimizer) # Fabric 对模型的 forward 已自动处理精度 output = model(inputs) with fabric.autocast(): # 可选:为 forward 之外的运算开启自动转换 loss = loss_function(output, target) fabric.backward(loss)文档强调:autocast是可选的,因为模型一旦setup过,Fabric 已经为其forward方法启用了自动精度转换;只有当你想让 forward之外的更多运算(如损失函数计算)也走自动转换时才需要它。源码一行即可印证:return self._precision.forward_context()(fabric.py)。相关精度机制的完整说明见 precision 指南。
print、save、load、load_raw:进程安全地输出与断点续训
print:只在主进程打印
# 只在主进程打印,避免多设备/多节点时重复刷屏 fabric.print(f"{epoch}/{num_epochs}| Train Epoch Loss: {loss}")实现(fabric.py)是if self.local_rank == 0: print(*args, **kwargs)——注意是每台机器的 local rank 0 都会打印,参数原样透传给内置print。
save:替代 torch.save 的多进程安全保存
# 定义你的程序/训练循环状态 state = { "model1": model1, "model2": model2, "optimizer": optimizer, "iteration": iteration, } # 替代 torch.save(...) fabric.save("path/to/checkpoint.ckpt", state)文档强调应把模型和优化器对象直接放进字典,Fabric 会自动解包它们并提取各自的 state-dict。源码事实(fabric.py):
- “谁保存、怎么保存”由策略决定:例如 DDP 策略只在进程 0 保存,而 FSDP 策略可能每个 rank 都写文件(分片保存)。因此
save必须在所有进程上调用。 - 保存后内部会调用
self.barrier(),确保多进程写入完成后再继续。 - 支持
filter参数:filter={"model": lambda name, param: "bias" not in name}这样的可调用对象可对特定 state key 的 state-dict 做选择性过滤;filter的 key 必须与state的 key 一致,且值必须可调用,否则分别抛TypeError/ValueError。对应测试见 test_fabric.py 的test_save_filter。
load:替代 torch.load 的多进程安全恢复
state = { "model1": model1, "model2": model2, "optimizer": optimizer, "iteration": iteration, } # 恢复 state 中对象的状态(原地修改) fabric.load("path/to/checkpoint.ckpt", state) # 或者完整读回,手动恢复 checkpoint = fabric.load("./checkpoints/version_2/checkpoint.ckpt") model.load_state_dict(checkpoint["model"]) ...对照 fabric.py 的实现:
- 传
state时,其中的_FabricModule/_FabricOptimizer/_FabricDataLoader会被解包后按策略恢复(原地);未消耗完的键会作为返回值 remainder 返回(例如把epoch这类元数据读出来),源码示例即epoch = remainder.get("epoch", 0)。 - 不传
state时返回完整 checkpoint 字典。 - 与
save一样,必须在所有进程上调用,且加载后会执行barrier()同步。 strict=True(默认)强制state的键与 checkpoint 键一一对应;weights_only默认None,传True时仅允许加载纯张量 state-dict 与基本类型(适合不可信来源的 checkpoint),加载可信来源、含nn.Module的 checkpoint 时用weights_only=False。源码注释还引导读者参阅 PyTorch 官方序列化说明了解安全加载的更多细节。
load_raw:加载非 Fabric 保存的裸权重
model = MyModel() # 你的朋友不用 Fabric 保存的模型权重文件 fabric.load_raw("path/to/model.pt", model) # 等价于: # model.load_state_dict(torch.load("path/to/model.pt"))load_raw用于从非 Fabric 保存的裸 PyTorch checkpoint 中恢复模型或优化器的 state-dict(fabric.py)。它概念上等价于obj.load_state_dict(torch.load(path)),但对策略无关——例如在 FSDP 等分片策略下也能正确处理。参数strict控制 state-dict 键的严格匹配(仅对模型生效,不适用于优化器),同样支持weights_only。断点续训与 checkpoint 的完整讨论见 checkpoint 指南。
barrier、broadcast、all_gather、all_reduce:进程间同步与集合通信
barrier:进程栅栏
if fabric.global_rank == 0: print("Downloading dataset. This can take a while ...") download_dataset("http://...") # 所有其他进程在此等待 rank 0 下载完成 fabric.barrier() # 所有人越过栅栏后,即可访问已下载的文件 load_dataset()barrier让所有进程等待、直到全部进入该调用后才继续执行——典型场景就是文档示例中的“rank 0 下载数据,其余进程等待数据落盘”。实现(fabric.py)委托strategy.barrier(name=name),源码注释强调:必须在所有进程上调用,否则程序会永远卡死,且只有在必要时才应使用(同步本身有开销)。下载数据的场景也可考虑更高级的rank_zero_first()上下文管理器(同一文件 fabric.py 提供),它天然处理“rank 0 先执行、其他人等待”的逻辑。
集合通信三件套
# 把 rank 0 的张量值广播给所有进程 result = fabric.broadcast(tensor, src=0) # 每个进程都拿到所有人的张量堆叠结果 all_tensors = fabric.all_gather(tensor) # 跨进程对张量做规约(求和等),每个人都拿到结果 reduced_tensor = fabric.all_reduce(tensor, reduce_op="sum") # 同样支持张量集合(dict、list、tuple): collection = {"loss": torch.tensor(...), "data": ...} gathered_collection = fabric.all_gather(collection, ...) reduced_collection = fabric.all_reduce(collection, ...)三个方法的语义与源码细节如下:
broadcast(obj, src=0)(fabric.py):从src指定的全局 rank 向其他所有进程发送数据。任何可序列化对象都支持,但张量效率最高。返回的是每个 rank 上相同的值。all_gather(data, group=None, sync_grads=False)(fabric.py):收集每个进程的张量并堆叠。输入为形状(batch, ...)的张量时,返回形状(world_size, batch, ...);world_size == 1时不加额外维度。内部先convert_to_tensors再apply_to_collection,因此 dict/list/tuple 嵌套集合也支持。sync_grads=True时梯度会随操作同步。all_reduce(data, group=None, reduce_op="mean")(fabric.py):对多进程张量做逐点规约,默认"mean",也支持"sum"(字符串或ReduceOp枚举)。注意规约是原地执行的(结果写回输入张量),某些策略可能限制可用的规约操作。
重要警告(文档原文强调):每个进程都必须进入集合通信调用,且各进程的张量形状必须一致,否则程序会挂起(hang)!
更完整的分布式通信模式(如按进程组通信)参见 distributed_communication 指南。对应测试覆盖见 test_fabric.py 的test_broadcast、test_all_reduce等用例。
no_backward_sync:梯度累积场景下的通信优化
在分布式策略(如 DDP)下做梯度累积时,每个微批次的反向传播都会触发一次跨进程梯度同步,而这在累积阶段是冗余的——只有累积满后真正optimizer.step()那次才需要同步。no_backward_sync上下文管理器正是用来在累积阶段跳过这些多余通信、加速训练:
# 每 8 个 batch 累积一次梯度 is_accumulating = batch_idx % 8 != 0 with fabric.no_backward_sync(model, enabled=is_accumulating): output = model(input) loss = ... fabric.backward(loss) ... # 每 8 个 batch 更新一次优化器 if not is_accumulating: optimizer.step() optimizer.zero_grad()源码实现(fabric.py)揭示了几个严格约束:
- 模型必须先 setup:传入的
module必须是_FabricModule(即经过setup/setup_module装配),否则抛TypeError并提示“先调用model = fabric.setup(model, ...)”。 - forward 与 backward 都必须在上下文内:文档和源码 docstring 都明确要求
model.forward()和fabric.backward()同时处于该上下文内,否则跳过同步的目的无法达成。 - 单设备策略是 no-op:
SingleDeviceStrategy(单卡)和XLAStrategy直接返回nullcontext()。 - 不支持的策略会降级并告警:
deepspeed、dp、xla三种策略不支持控制梯度同步,此时返回nullcontext()并发出PossibleUserWarning,建议“从代码中移除.no_backward_sync()或换用其他策略”。 enabled参数控制开关:True表示跳过同步,False表示不跳过——这也解释了为什么可以用enabled=is_accumulating来动态控制。
call、log 与 log_dict:回调钩子与指标日志
call:按名字触发回调钩子
call面向“自己构建 Trainer”的高级用户,让你在训练循环的固定位置执行任意注册的回调代码:
class MyCallback: def on_train_start(self): ... def on_train_epoch_end(self, model, results): ... fabric = Fabric(callbacks=[MyCallback()]) # 按名字触发任意钩子 fabric.call("on_train_start") # 传入钩子所需的其他参数 fabric.call("on_train_epoch_end", model=..., results={...}) # 只有定义了该方法的回调会被执行 fabric.call("undefined")实现(fabric.py)会遍历self._callbacks,通过getattr(callback, hook_name, None)查找同名方法:未定义该方法(None)或方法不可调用时跳过(后者发警告);存在时调用method(*args, **filtered_kwargs)。一个贴心设计是kwargs 自动过滤(_filter_kwargs_for_callback):会检查回调方法的签名,只传入其参数列表中出现的 kwargs,从而允许不同回调对同一钩子持有不同签名(例如有的回调on_train_epoch_end(self, model)、有的on_train_epoch_end(self, model, results))。若方法声明了**kwargs则原样透传全部 kwargs。对应测试见 test_fabric.py 的test_call与test_callback_kwargs_filtering。更完整的回调设计讨论见 callbacks 指南,仓库还提供了一个完整的“用 Fabric 自己写 Trainer”的参考实现:examples/fabric/build_your_own_trainer/trainer.py。
log 与 log_dict:向 Logger 写指标
# 在 Fabric 中设置 Logger fabric = Fabric(loggers=TensorBoardLogger(...)) # 在训练循环或模型中的任意位置: fabric.log("loss", loss) # 或一次发送多个指标: fabric.log_dict({"loss": loss, "accuracy": acc})源码(fabric.py)的关键事实:
- 未配置 Logger 时是安全的 no-op:
Fabric默认不绑定 logger,此时log/log_dict什么都不做,不会报错。 - 底层循环(文档以伪代码给出):
log最终调用log_dict,后者先把张量指标convert_tensors_to_scalars(自动从计算图中 detach),然后遍历所有 logger 执行logger.log_metrics(metrics=metrics, step=step)。 step参数可选:多数 Logger 实现每次 log 会自动把 step 加一,需要自定义时再显式传入。- Tensor 值会自动 detach,不会残留计算图引用。
更完整的日志接入与指标组织方式参见 logging 指南。
实战串联:一个覆盖全部核心方法的训练骨架
把本文的方法串起来,可以得到一个最小但完整的 Fabric 分布式训练程序(对应仓库示例见 examples/fabric/build_your_own_trainer/run.py):
import torch from lightning.fabric import Fabric from torch.utils.data import DataLoader def train(fabric: Fabric): fabric.seed_everything(1234) # 可复现性:覆盖 torch/numpy/random with fabric.init_module(): # 直接按目标设备/精度建模型,省去搬移等待 model = torch.nn.Linear(32, 64) optimizer = torch.optim.SGD(model.parameters(), lr=0.01) model, optimizer = fabric.setup(model, optimizer) # 装配:搬设备 + 精度 + 策略包装 train_loader = fabric.setup_dataloaders( # 装配数据:分布式采样器 + 自动搬数据 DataLoader(dataset, batch_size=32, shuffle=True) ) for batch_idx, (inputs, targets) in enumerate(train_loader): is_accumulating = batch_idx % 8 != 0 with fabric.no_backward_sync(model, enabled=is_accumulating): # 累积期跳过梯度同步 output = model(inputs) loss = torch.nn.functional.mse_loss(output, targets) fabric.backward(loss) if not is_accumulating: fabric.clip_gradients(model, optimizer, max_norm=2.0) # 按 2-范数裁剪 optimizer.step() optimizer.zero_grad() fabric.log("loss", loss) # 写指标(未配置 logger 时安全 no-op) if batch_idx % 100 == 0: fabric.print(f"step {batch_idx}: loss = {loss.item()}") # 仅主进程打印 fabric.save("checkpoint.ckpt", {"model": model, "optimizer": optimizer, "step": batch_idx}) fabric.load("checkpoint.ckpt", {"model": model, "optimizer": optimizer}) # 原地恢复 fabric = Fabric(accelerator="gpu", devices=2, strategy="ddp", precision="bf16-mixed") fabric.launch(train)结语
Fabric的方法设计遵循一条主线:把分布式、精度、设备搬移等基础设施从业务代码中剥离——launch接管进程管理,setup/setup_dataloaders接管设备与数据装配,backward/autocast/init_module接管精度与反向传播,save/load接管跨策略的断点续训,barrier/broadcast/all_gather/all_reduce接管进程通信。理解这些方法的底层实现(全部集中在 src/lightning/fabric/fabric.py)能帮助你判断在何种策略、何种硬件与何种规模下选用哪个方法、哪些参数,以及哪些操作必须“所有进程同时参与”。本文对应的 API 文档原文位于 docs/source-fabric/api/fabric_methods.rst,可结合 wrappers 说明、model_init 指南、precision 指南、checkpoint 指南、distributed_communication 指南、callbacks 指南 与 logging 指南 深入研读。
- 人工智能
- 深度学习
- 机器学习
- 预训练
- 分布式训练
- 微调
【免费下载链接】pytorch-lightning
Pretrain, finetune ANY AI model of ANY size on 1 or 10,000+ GPUs with zero code changes.
相关推荐
PyTorch Lightning Fabric 策略体系全解析:从单设备到大规模分布式训练
PyTorch Lightning Fabric 策略体系全解析:从单设备到大规模分布式训练 导读 lightning.fabric.strategies 是
人工智能深度学习机器学习预训练分布式训练微调Lightning Fabric 核心 API 全解析:Fabric 类构造参数与关键方法实战指南
Lightning Fabric 核心 API 全解析:Fabric 类构造参数与关键方法实战指南 lightning.fabric.Fabric 是 PyTo
人工智能深度学习机器学习预训练分布式训练微调Lightning Fabric 使用指南:5 行改造 PyTorch 代码,迈向多卡分布式训练
Lightning Fabric 使用指南:5 行改造 PyTorch 代码,迈向多卡分布式训练 Fabric 是 Lightning 生态中面向 PyTorc
人工智能深度学习机器学习预训练分布式训练微调
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考