Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
73bf0a0a2b |
@@ -255,8 +255,6 @@ def save_distillation_checkpoint(
|
||||
if generator_scheduler is not None:
|
||||
generator_states["scheduler"] = SchedulerWrapper(
|
||||
generator_scheduler)
|
||||
if generator_ema is not None:
|
||||
generator_states["ema"] = generator_ema.state_dict()
|
||||
|
||||
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
|
||||
"generator")
|
||||
@@ -288,8 +286,6 @@ def save_distillation_checkpoint(
|
||||
if generator_scheduler_2 is not None:
|
||||
generator_2_states["scheduler"] = SchedulerWrapper(
|
||||
generator_scheduler_2)
|
||||
if generator_ema_2 is not None:
|
||||
generator_2_states["ema"] = generator_ema_2.state_dict()
|
||||
|
||||
generator_2_dcp_dir = os.path.join(save_dir,
|
||||
"distributed_checkpoint",
|
||||
@@ -415,6 +411,62 @@ def save_distillation_checkpoint(
|
||||
rank,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Persist EMA separately to avoid shape mismatches across ranks.
|
||||
# Supports:
|
||||
# - mode="rank0_full": save consolidated EMA only on rank 0
|
||||
# - mode="local_shard": save per-rank EMA shard for each rank
|
||||
try:
|
||||
if generator_ema is not None and getattr(generator_ema, "mode",
|
||||
None) == "rank0_full":
|
||||
if rank == 0:
|
||||
ema_dir = os.path.join(save_dir, "ema")
|
||||
os.makedirs(ema_dir, exist_ok=True)
|
||||
ema_path = os.path.join(ema_dir, "generator_ema.pt")
|
||||
torch.save(generator_ema.state_dict(), ema_path)
|
||||
logger.info("rank: %s, saved generator EMA (rank0_full) to %s",
|
||||
rank,
|
||||
ema_path,
|
||||
local_main_process_only=False)
|
||||
elif generator_ema is not None and getattr(generator_ema, "mode",
|
||||
None) == "local_shard":
|
||||
ema_dir = os.path.join(save_dir, "ema_local_shard")
|
||||
os.makedirs(ema_dir, exist_ok=True)
|
||||
ema_path = os.path.join(ema_dir, f"generator_ema_rank{rank}.pt")
|
||||
torch.save(generator_ema.state_dict(), ema_path)
|
||||
logger.info(
|
||||
"rank: %s, saved generator EMA shard (local_shard) to %s",
|
||||
rank,
|
||||
ema_path,
|
||||
local_main_process_only=False)
|
||||
if generator_ema_2 is not None and getattr(generator_ema_2, "mode",
|
||||
None) == "rank0_full":
|
||||
if rank == 0:
|
||||
ema_dir = os.path.join(save_dir, "ema")
|
||||
os.makedirs(ema_dir, exist_ok=True)
|
||||
ema2_path = os.path.join(ema_dir, "generator_ema_2.pt")
|
||||
torch.save(generator_ema_2.state_dict(), ema2_path)
|
||||
logger.info(
|
||||
"rank: %s, saved generator_2 EMA (rank0_full) to %s",
|
||||
rank,
|
||||
ema2_path,
|
||||
local_main_process_only=False)
|
||||
elif generator_ema_2 is not None and getattr(generator_ema_2, "mode",
|
||||
None) == "local_shard":
|
||||
ema_dir = os.path.join(save_dir, "ema_local_shard")
|
||||
os.makedirs(ema_dir, exist_ok=True)
|
||||
ema2_path = os.path.join(ema_dir, f"generator_ema_2_rank{rank}.pt")
|
||||
torch.save(generator_ema_2.state_dict(), ema2_path)
|
||||
logger.info(
|
||||
"rank: %s, saved generator_2 EMA shard (local_shard) to %s",
|
||||
rank,
|
||||
ema2_path,
|
||||
local_main_process_only=False)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed saving EMA separately: %s",
|
||||
rank,
|
||||
str(e),
|
||||
local_main_process_only=False)
|
||||
|
||||
# Save generator model weights (consolidated) for inference
|
||||
cpu_state = gather_state_dict_on_cpu_rank0(generator_transformer,
|
||||
device=None)
|
||||
@@ -641,18 +693,37 @@ def load_distillation_checkpoint(
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA state if available and generator_ema is provided
|
||||
# Load EMA separately if saved in rank0_full mode
|
||||
if generator_ema is not None:
|
||||
try:
|
||||
ema_state = generator_states.get("ema")
|
||||
if ema_state is not None:
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info("rank: %s, generator EMA state loaded successfully",
|
||||
rank)
|
||||
else:
|
||||
logger.info("rank: %s, no EMA state found in checkpoint", rank)
|
||||
if getattr(generator_ema, "mode", None) == "rank0_full":
|
||||
ema_path = os.path.join(checkpoint_path, "ema",
|
||||
"generator_ema.pt")
|
||||
if rank == 0 and os.path.exists(ema_path):
|
||||
ema_state = torch.load(ema_path, map_location="cpu")
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info(
|
||||
"rank: %s, generator EMA (rank0_full) loaded from %s",
|
||||
rank, ema_path)
|
||||
elif rank == 0:
|
||||
logger.info(
|
||||
"rank: %s, generator EMA file not found at %s; skipping",
|
||||
rank, ema_path)
|
||||
elif getattr(generator_ema, "mode", None) == "local_shard":
|
||||
ema_path = os.path.join(checkpoint_path, "ema_local_shard",
|
||||
f"generator_ema_rank{rank}.pt")
|
||||
if os.path.exists(ema_path):
|
||||
ema_state = torch.load(ema_path, map_location="cpu")
|
||||
generator_ema.load_state_dict(ema_state)
|
||||
logger.info(
|
||||
"rank: %s, generator EMA shard (local_shard) loaded from %s",
|
||||
rank, ema_path)
|
||||
else:
|
||||
logger.info(
|
||||
"rank: %s, generator EMA shard file not found at %s; skipping",
|
||||
rank, ema_path)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA state: %s", rank,
|
||||
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
|
||||
str(e))
|
||||
|
||||
# Load generator_2 distributed checkpoint (MoE support)
|
||||
@@ -692,22 +763,43 @@ def load_distillation_checkpoint(
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA_2 state if available and generator_ema_2 is provided
|
||||
# Load EMA_2 separately if saved in rank0_full mode
|
||||
if generator_ema_2 is not None:
|
||||
try:
|
||||
ema_2_state = generator_2_states.get("ema")
|
||||
if ema_2_state is not None:
|
||||
generator_ema_2.load_state_dict(ema_2_state)
|
||||
logger.info(
|
||||
"rank: %s, generator_2 EMA state loaded successfully",
|
||||
rank)
|
||||
else:
|
||||
logger.info(
|
||||
"rank: %s, no EMA_2 state found in checkpoint",
|
||||
rank)
|
||||
if getattr(generator_ema_2, "mode", None) == "rank0_full":
|
||||
ema2_path = os.path.join(checkpoint_path, "ema",
|
||||
"generator_ema_2.pt")
|
||||
if rank == 0 and os.path.exists(ema2_path):
|
||||
ema_2_state = torch.load(ema2_path,
|
||||
map_location="cpu")
|
||||
generator_ema_2.load_state_dict(ema_2_state)
|
||||
logger.info(
|
||||
"rank: %s, generator_2 EMA (rank0_full) loaded from %s",
|
||||
rank, ema2_path)
|
||||
elif rank == 0:
|
||||
logger.info(
|
||||
"rank: %s, generator_2 EMA file not found at %s; skipping",
|
||||
rank, ema2_path)
|
||||
elif getattr(generator_ema_2, "mode",
|
||||
None) == "local_shard":
|
||||
ema2_path = os.path.join(
|
||||
checkpoint_path, "ema_local_shard",
|
||||
f"generator_ema_2_rank{rank}.pt")
|
||||
if os.path.exists(ema2_path):
|
||||
ema_2_state = torch.load(ema2_path,
|
||||
map_location="cpu")
|
||||
generator_ema_2.load_state_dict(ema_2_state)
|
||||
logger.info(
|
||||
"rank: %s, generator_2 EMA shard (local_shard) loaded from %s",
|
||||
rank, ema2_path)
|
||||
else:
|
||||
logger.info(
|
||||
"rank: %s, generator_2 EMA shard file not found at %s; skipping",
|
||||
rank, ema2_path)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA_2 state: %s",
|
||||
rank, str(e))
|
||||
logger.warning(
|
||||
"rank: %s, failed to load generator_2 EMA: %s", rank,
|
||||
str(e))
|
||||
else:
|
||||
logger.info("rank: %s, generator_2 checkpoint not found, skipping",
|
||||
rank)
|
||||
|
||||
Reference in New Issue
Block a user