恢复训练的checkpoint加载函数 Args: checkpoint_path (str): checkpoint文件路径 model (torch.nn.Module): 主模型 ema_model (torch.nn.Module): EMA模型 optimizer (torch.optim.Optimizer, optional): 优化器 lr_scheduler (torch.optim.lr_scheduler._LRScheduler, optional): 学
(
checkpoint_path,
ema_model,
optimizer=None,
lr_scheduler=None,
strict=True,
print_info=True,
map_location="cuda"
)
| 1588 | return total_params |
| 1589 | |
| 1590 | def load_checkpoint( |
| 1591 | checkpoint_path, |
| 1592 | ema_model, |
| 1593 | optimizer=None, |
| 1594 | lr_scheduler=None, |
| 1595 | strict=True, |
| 1596 | print_info=True, |
| 1597 | map_location="cuda" |
| 1598 | ): |
| 1599 | """ |
| 1600 | 恢复训练的checkpoint加载函数 |
| 1601 | |
| 1602 | Args: |
| 1603 | checkpoint_path (str): checkpoint文件路径 |
| 1604 | model (torch.nn.Module): 主模型 |
| 1605 | ema_model (torch.nn.Module): EMA模型 |
| 1606 | optimizer (torch.optim.Optimizer, optional): 优化器 |
| 1607 | lr_scheduler (torch.optim.lr_scheduler._LRScheduler, optional): 学习率调度器 |
| 1608 | map_location (str, optional): 加载设备 |
| 1609 | |
| 1610 | Returns: |
| 1611 | cfg (dict): 保存的配置字典 |
| 1612 | global_step (int): 当前训练步数 |
| 1613 | """ |
| 1614 | print(f"Loading checkpoint from {checkpoint_path} ...") |
| 1615 | checkpoint = torch.load(checkpoint_path, map_location=map_location) |
| 1616 | |
| 1617 | # 恢复配置 |
| 1618 | cfg = checkpoint["cfg"] |
| 1619 | |
| 1620 | # 加载模型参数 |
| 1621 | weights = checkpoint["weights"] |
| 1622 | ema_weights = checkpoint["ema_weights"] |
| 1623 | load_info_online = ema_model.online_model.load_state_dict(weights, strict=strict) |
| 1624 | load_info_ema = ema_model.ema_model.load_state_dict(ema_weights, strict=strict) |
| 1625 | print(f"✅ EMA weights loaded : online_weights {weights_num(weights)}, ema_weights : {weights_num(ema_weights)}") |
| 1626 | |
| 1627 | print_load_report(load_info_online, "online_model", len(weights), print_info=print_info) |
| 1628 | print_load_report(load_info_ema, "ema_model", len(ema_weights), print_info=print_info) |
| 1629 | |
| 1630 | # 加载优化器和学习率调度器 |
| 1631 | if optimizer is not None and "optimizer" in checkpoint: |
| 1632 | |
| 1633 | for i, param_group in enumerate(optimizer.param_groups): |
| 1634 | print(f"optimizer参数load之前:Param group {i}:") |
| 1635 | for key, value in param_group.items(): |
| 1636 | if key != "params": # params 太长 |
| 1637 | print(f" {key}: {value}") |
| 1638 | |
| 1639 | optimizer.load_state_dict(checkpoint["optimizer"]) |
| 1640 | |
| 1641 | for i, param_group in enumerate(optimizer.param_groups): |
| 1642 | print(f"optimizer参数load之后:Param group {i}:") |
| 1643 | for key, value in param_group.items(): |
| 1644 | if key != "params": # params 太长 |
| 1645 | print(f" {key}: {value}") |
| 1646 | |
| 1647 | removed_cnt = 0 |
nothing calls this directly
no test coverage detected