MCPcopy Create free account
hub / github.com/QWTforGithub/T2LDM / load_checkpoint

Function load_checkpoint

utils/common.py:1590–1678  ·  view source on GitHub ↗

恢复训练的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"
)

Source from the content-addressed store, hash-verified

1588 return total_params
1589
1590def 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

Callers

nothing calls this directly

Calls 5

weights_numFunction · 0.85
print_load_reportFunction · 0.85
load_state_dictMethod · 0.45
state_dictMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected