Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion src/maxdiffusion/checkpointing/checkpointing_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,9 +74,16 @@ def create_orbax_checkpoint_manager(
"text_encoder_state": ocp.StandardCheckpointHandler(),
}
elif checkpoint_type == WAN_CHECKPOINT:
item_names = ("low_noise_transformer_state", "high_noise_transformer_state", "wan_state", "wan_config")
item_names = (
"low_noise_transformer_state",
"high_noise_transformer_state",
"wan_state",
"wan_config",
"wan_config_high",
)
item_handlers = {
"wan_config": ocp.JsonCheckpointHandler(),
"wan_config_high": ocp.JsonCheckpointHandler(),
"wan_state": ocp.StandardCheckpointHandler(),
"low_noise_transformer_state": ocp.StandardCheckpointHandler(),
"high_noise_transformer_state": ocp.StandardCheckpointHandler(),
Expand Down
96 changes: 71 additions & 25 deletions src/maxdiffusion/checkpointing/wan_checkpointer_2_2.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,15 +18,40 @@
import jax
from typing import Optional, Tuple
from ..pipelines.wan.wan_pipeline_2_2 import WanPipeline2_2
from .. import max_logging
from .. import max_logging, max_utils
import orbax.checkpoint as ocp
from maxdiffusion.checkpointing.checkpointing_utils import add_sharding_to_struct, get_cpu_mesh_and_sharding
from maxdiffusion.checkpointing.wan_checkpointer import WanCheckpointer


def _get_item(obj, key, default=None):
"""Safely gets an item from obj supporting both dictionary mapping and attribute access."""
if obj is None:
return default
if isinstance(obj, dict):
return obj.get(key, default)
if hasattr(obj, key):
return getattr(obj, key, default)
if hasattr(obj, "get") and callable(obj.get):
try:
return obj.get(key, default)
except Exception:
pass
return default


class WanCheckpointer2_2(WanCheckpointer[WanPipeline2_2]):
pipeline_class = WanPipeline2_2

def _create_optimizer(self, model, config, learning_rate, scale_factor: float = 1.0):
total_steps = max(1, int(config.max_train_steps * scale_factor))
schedule_steps = max(1, int(config.learning_rate_schedule_steps * scale_factor))
learning_rate_scheduler = max_utils.create_learning_rate_schedule(
learning_rate, schedule_steps, config.warmup_steps_fraction, total_steps
)
tx = max_utils.create_optimizer(config, learning_rate_scheduler)
return tx, learning_rate_scheduler

def load_wan_configs_from_orbax(self, step: Optional[int]) -> Tuple[Optional[dict], Optional[int]]:
if step is None:
step = self.checkpoint_manager.latest_step()
Expand All @@ -40,52 +65,72 @@ def load_wan_configs_from_orbax(self, step: Optional[int]) -> Tuple[Optional[dic
metadatas = self.checkpoint_manager.item_metadata(step)

# Handle low_noise_transformer
low_noise_transformer_metadata = metadatas.low_noise_transformer_state
low_noise_transformer_metadata = _get_item(metadatas, "low_noise_transformer_state")
target_shardings = jax.tree_util.tree_map(lambda x: replicated_sharding, low_noise_transformer_metadata)
with mesh:
abstract_tree_structure_low_params = jax.tree_util.tree_map(
add_sharding_to_struct, low_noise_transformer_metadata, target_shardings
)

# Handle high_noise_transformer
high_noise_transformer_metadata = metadatas.high_noise_transformer_state
high_noise_transformer_metadata = _get_item(metadatas, "high_noise_transformer_state")
target_shardings = jax.tree_util.tree_map(lambda x: replicated_sharding, high_noise_transformer_metadata)
with mesh:
abstract_tree_structure_high_params = jax.tree_util.tree_map(
add_sharding_to_struct, high_noise_transformer_metadata, target_shardings
)

max_logging.log("Restoring WAN 2.2 checkpoint")
restore_items = {
"low_noise_transformer_state": ocp.args.StandardRestore(abstract_tree_structure_low_params),
"high_noise_transformer_state": ocp.args.StandardRestore(abstract_tree_structure_high_params),
"wan_config": ocp.args.JsonRestore(),
}
has_high_config = _get_item(metadatas, "wan_config_high") is not None

if has_high_config:
restore_items["wan_config_high"] = ocp.args.JsonRestore()

restored_checkpoint = self.checkpoint_manager.restore(
step=step,
args=ocp.args.Composite(
low_noise_transformer_state=ocp.args.StandardRestore(abstract_tree_structure_low_params),
high_noise_transformer_state=ocp.args.StandardRestore(abstract_tree_structure_high_params),
wan_config=ocp.args.JsonRestore(),
),
)
max_logging.log(f"restored checkpoint {restored_checkpoint.keys()}")
max_logging.log(
f"restored checkpoint low_noise_transformer_state {restored_checkpoint.low_noise_transformer_state.keys()}"
)
max_logging.log(
f"restored checkpoint high_noise_transformer_state {restored_checkpoint.high_noise_transformer_state.keys()}"
args=ocp.args.Composite(**restore_items),
)
max_logging.log(
f"optimizer found in low_noise checkpoint {'opt_state' in restored_checkpoint.low_noise_transformer_state.keys()}"
)
max_logging.log(
f"optimizer found in high_noise checkpoint {'opt_state' in restored_checkpoint.high_noise_transformer_state.keys()}"
keys_fn = _get_item(restored_checkpoint, "keys")
keys_list = (
list(keys_fn())
if callable(keys_fn)
else list(restored_checkpoint.keys())
if hasattr(restored_checkpoint, "keys")
else []
)
max_logging.log(f"restored checkpoint {keys_list}")

low_state = _get_item(restored_checkpoint, "low_noise_transformer_state", {})
high_state = _get_item(restored_checkpoint, "high_noise_transformer_state", {})
low_keys = list(low_state.keys()) if hasattr(low_state, "keys") else []
high_keys = list(high_state.keys()) if hasattr(high_state, "keys") else []
max_logging.log(f"restored checkpoint low_noise_transformer_state {low_keys}")
max_logging.log(f"restored checkpoint high_noise_transformer_state {high_keys}")
max_logging.log(f"optimizer found in low_noise checkpoint {'opt_state' in low_keys}")
max_logging.log(f"optimizer found in high_noise checkpoint {'opt_state' in high_keys}")
max_logging.log(f"optimizer state saved in attribute self.opt_state {self.opt_state}")
return restored_checkpoint, step

def _extract_opt_state(self, restored_checkpoint):
if "opt_state" in restored_checkpoint.low_noise_transformer_state.keys():
return restored_checkpoint.low_noise_transformer_state["opt_state"]
elif "opt_state" in restored_checkpoint.high_noise_transformer_state.keys():
return restored_checkpoint.high_noise_transformer_state["opt_state"]
return None
low_state = _get_item(restored_checkpoint, "low_noise_transformer_state", {})
high_state = _get_item(restored_checkpoint, "high_noise_transformer_state", {})
low_opt = _get_item(low_state, "opt_state")
high_opt = _get_item(high_state, "opt_state")
low_step = _get_item(low_state, "step")
high_step = _get_item(high_state, "step")
if low_opt is None and high_opt is None:
return None
return {
"low_noise_transformer": low_opt,
"high_noise_transformer": high_opt,
"low_noise_step": low_step,
"high_noise_step": high_step,
}

def save_checkpoint(self, train_step, pipeline: WanPipeline2_2, train_states: dict):
"""Saves the training state and model configurations."""
Expand All @@ -96,6 +141,7 @@ def config_to_json(model_or_config):
max_logging.log(f"Saving checkpoint for step {train_step}")
items = {
"wan_config": ocp.args.JsonSave(config_to_json(pipeline.low_noise_transformer)),
"wan_config_high": ocp.args.JsonSave(config_to_json(pipeline.high_noise_transformer)),
Comment thread
Toshi-31 marked this conversation as resolved.
}

items["low_noise_transformer_state"] = ocp.args.StandardSave(train_states["low_noise_transformer"])
Expand Down
6 changes: 6 additions & 0 deletions src/maxdiffusion/configs/base_wan_27b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -410,6 +410,12 @@ num_inference_steps: 40
fps: 16
save_final_checkpoint: False

# Staging directory for downloading pretrained weights from Google Cloud Storage (gs://)
# prior to loading into device memory. Note: On Cloud TPU VMs with limited root disk/tmpfs,
# ensure this directory resides on a persistent disk or large volume to avoid ENOSPC.
# Differs from pretrained_model_name_or_path which specifies the remote/local model URI.
checkpoint_save_location: "/tmp"
Comment thread
Toshi-31 marked this conversation as resolved.

# SDXL Lightning parameters
lightning_from_pt: True
# Empty or "ByteDance/SDXL-Lightning" to enable lightning.
Expand Down
15 changes: 14 additions & 1 deletion src/maxdiffusion/pipelines/wan/wan_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -317,7 +317,20 @@ def create_model(rngs: nnx.Rngs, wan_config: dict):

# 1. Load config.
if restored_checkpoint:
wan_config = restored_checkpoint["wan_config"]
wan_config_high = (
restored_checkpoint.get("wan_config_high")
if isinstance(restored_checkpoint, dict)
else getattr(restored_checkpoint, "wan_config_high", None)
)
if subfolder == "transformer" and wan_config_high is not None:
wan_config = dict(wan_config_high)
else:
raw_config = (
restored_checkpoint["wan_config"]
if isinstance(restored_checkpoint, dict)
else getattr(restored_checkpoint, "wan_config")
)
wan_config = dict(raw_config)
else:
with _HF_METADATA_LOCK:
wan_config = WanModel.load_config(config.pretrained_model_name_or_path, subfolder=subfolder)
Expand Down
13 changes: 9 additions & 4 deletions src/maxdiffusion/pyconfig.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@
)

_ALLOWED_MODEL_NAMES = {WAN2_1, WAN2_2, LTX2_VIDEO, LTX2_3, Z_IMAGE}
_ALLOWED_TRAINING_MODEL_NAMES = {WAN2_1}
_ALLOWED_TRAINING_MODEL_NAMES = {WAN2_1, WAN2_2}


def _validate_model_name(model_name: str | None):
Expand Down Expand Up @@ -287,12 +287,17 @@ def user_init(raw_keys):

# Orbax doesn't save the tokenizer params, instead it loads them from the pretrained_model_name_or_path
raw_keys["tokenizer_model_name_or_path"] = raw_keys["pretrained_model_name_or_path"]
ckpt_save_loc = raw_keys.get("checkpoint_save_location", "/tmp")
if "gs://" in raw_keys["pretrained_model_name_or_path"]:
raw_keys["pretrained_model_name_or_path"] = max_utils.download_blobs(raw_keys["pretrained_model_name_or_path"], "/tmp")
raw_keys["pretrained_model_name_or_path"] = max_utils.download_blobs(
raw_keys["pretrained_model_name_or_path"], ckpt_save_loc
)
if "gs://" in raw_keys["unet_checkpoint"]:
raw_keys["unet_checkpoint"] = max_utils.download_blobs(raw_keys["unet_checkpoint"], "/tmp")
raw_keys["unet_checkpoint"] = max_utils.download_blobs(raw_keys["unet_checkpoint"], ckpt_save_loc)
if "gs://" in raw_keys["tokenizer_model_name_or_path"]:
raw_keys["tokenizer_model_name_or_path"] = max_utils.download_blobs(raw_keys["tokenizer_model_name_or_path"], "/tmp")
raw_keys["tokenizer_model_name_or_path"] = max_utils.download_blobs(
raw_keys["tokenizer_model_name_or_path"], ckpt_save_loc
)
if "gs://" in raw_keys["dataset_name"]:
raw_keys["dataset_name"] = max_utils.download_blobs(raw_keys["dataset_name"], raw_keys["dataset_save_location"])
raw_keys["dataset_save_location"] = raw_keys["dataset_name"]
Expand Down
94 changes: 90 additions & 4 deletions src/maxdiffusion/tests/wan/wan_checkpointer_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -387,7 +387,8 @@ def test_load_checkpoint_with_optimizer_in_low_noise(self, mock_from_checkpoint,
)
self.assertEqual(pipeline, mock_pipeline_instance)
self.assertIsNotNone(opt_state)
self.assertEqual(opt_state["learning_rate"], 0.001)
self.assertEqual(opt_state["low_noise_transformer"]["learning_rate"], 0.001)
self.assertIsNone(opt_state["high_noise_transformer"])
self.assertEqual(step, 1)

@patch("maxdiffusion.checkpointing.wan_checkpointer.create_orbax_checkpoint_manager")
Expand Down Expand Up @@ -429,9 +430,39 @@ def test_load_checkpoint_with_optimizer_in_high_noise(self, mock_from_checkpoint
)
self.assertEqual(pipeline, mock_pipeline_instance)
self.assertIsNotNone(opt_state)
self.assertEqual(opt_state["learning_rate"], 0.002)
self.assertIsNone(opt_state["low_noise_transformer"])
self.assertEqual(opt_state["high_noise_transformer"]["learning_rate"], 0.002)
self.assertEqual(step, 1)

@patch("maxdiffusion.pipelines.wan.wan_pipeline.nnx.eval_shape")
def test_create_sharded_logical_transformer_reads_wan_config_high(self, mock_eval_shape):
"""Test that create_sharded_logical_transformer uses wan_config_high for high expert."""
from maxdiffusion.pipelines.wan.wan_pipeline import create_sharded_logical_transformer

captured_configs = []

def fake_eval_shape(factory, *args, **kwargs):
captured_configs.append(factory.keywords["wan_config"])
raise StopIteration("Verified")

mock_eval_shape.side_effect = fake_eval_shape

restored = {
"wan_config": {"dim": 128},
"wan_config_high": {"dim": 256},
}
with self.assertRaises(StopIteration):
create_sharded_logical_transformer(
MagicMock(), MagicMock(), MagicMock(), self.config, restored_checkpoint=restored, subfolder="transformer"
)
self.assertEqual(captured_configs[-1]["dim"], 256)

with self.assertRaises(StopIteration):
create_sharded_logical_transformer(
MagicMock(), MagicMock(), MagicMock(), self.config, restored_checkpoint=restored, subfolder="transformer_2"
)
self.assertEqual(captured_configs[-1]["dim"], 128)


class WanCheckpointerI2V_2_1Test(unittest.TestCase):
"""Tests for WAN 2.1 I2V checkpointer."""
Expand Down Expand Up @@ -758,9 +789,64 @@ def test_load_checkpoint_both_optimizers_present(self, mock_from_checkpoint, moc
checkpointer = WanCheckpointer2_2(config=self.config)
pipeline, opt_state, step = checkpointer.load_checkpoint(step=1)

# Should prioritize low_noise_transformer's optimizer state
# Should preserve both low_noise_transformer and high_noise_transformer optimizer states
self.assertIsNotNone(opt_state)
self.assertEqual(opt_state["learning_rate"], 0.001)
self.assertEqual(opt_state["low_noise_transformer"]["learning_rate"], 0.001)
self.assertEqual(opt_state["high_noise_transformer"]["learning_rate"], 0.002)

@patch("maxdiffusion.checkpointing.wan_checkpointer.create_orbax_checkpoint_manager")
@patch.object(WanPipeline2_2, "from_checkpoint", autospec=True)
def test_load_checkpoint_with_dict_mapping_and_wan_config_high(self, mock_from_checkpoint, mock_create_manager):
"""Test loading checkpoint when Orbax returns standard dictionaries for metadata and checkpoint."""
mock_manager = MagicMock()
mock_manager.latest_step.return_value = 5
mock_manager.item_metadata.return_value = {
"low_noise_transformer_state": {},
"high_noise_transformer_state": {},
"wan_config_high": {"dim": 256},
}

mock_manager.restore.return_value = {
"low_noise_transformer_state": {"params": {}, "opt_state": {"lr": 0.0001}, "step": 10},
"high_noise_transformer_state": {"params": {}, "opt_state": {"lr": 0.0002}, "step": 20},
"wan_config": {"dim": 128},
"wan_config_high": {"dim": 256},
}

mock_create_manager.return_value = mock_manager
mock_pipeline_instance = MagicMock()
mock_from_checkpoint.return_value = mock_pipeline_instance

checkpointer = WanCheckpointer2_2(config=self.config)
pipeline, opt_state, step = checkpointer.load_checkpoint(step=5)

self.assertEqual(step, 5)
self.assertIsNotNone(opt_state)
self.assertEqual(opt_state["low_noise_transformer"]["lr"], 0.0001)
self.assertEqual(opt_state["high_noise_transformer"]["lr"], 0.0002)
self.assertEqual(opt_state["low_noise_step"], 10)
self.assertEqual(opt_state["high_noise_step"], 20)

def test_get_item_helper(self):
"""Test _get_item handles dicts, objects, None, and default fallbacks."""
from maxdiffusion.checkpointing.wan_checkpointer_2_2 import _get_item

self.assertIsNone(_get_item(None, "key"))
self.assertEqual(_get_item(None, "key", default=42), 42)

d = {"a": 1, "b": None}
self.assertEqual(_get_item(d, "a"), 1)
self.assertIsNone(_get_item(d, "b"))
self.assertEqual(_get_item(d, "missing", default="def"), "def")

class Obj:
x = 100
y = None

obj = Obj()
self.assertEqual(_get_item(obj, "x"), 100)
self.assertIsNone(_get_item(obj, "y"))
self.assertEqual(_get_item(obj, "z", default="fallback"), "fallback")


if __name__ == "__main__":
Expand Down
Loading
Loading