Compare commits

...
Author SHA1 Message Date
SolitaryThinker 73bf0a0a2b fix 2025-10-28 10:00:55 +00:00
+118 -26
View File
@@ -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)