|
16 | 16 |
|
17 | 17 | from abc import ABC, abstractmethod |
18 | 18 | import json |
19 | | -from typing import Optional, Tuple |
| 19 | +from typing import Optional, Tuple, Generic, TypeVar, Type |
20 | 20 | import jax |
21 | 21 | from flax import nnx |
22 | 22 | from maxdiffusion.checkpointing.checkpointing_utils import ( |
23 | 23 | add_sharding_to_struct, |
24 | 24 | create_orbax_checkpoint_manager, |
25 | 25 | get_cpu_mesh_and_sharding, |
26 | 26 | ) |
27 | | -from ..pipelines.wan.wan_pipeline_2_1 import WanPipeline2_1 |
28 | | -from ..pipelines.wan.wan_pipeline_2_2 import WanPipeline2_2 |
29 | | -from ..pipelines.wan.wan_pipeline_i2v_2p1 import WanPipelineI2V_2_1 |
30 | | -from ..pipelines.wan.wan_pipeline_i2v_2p2 import WanPipelineI2V_2_2 |
| 27 | +from ..pipelines.wan.wan_pipeline import WanPipeline |
31 | 28 | from .. import max_logging, max_utils |
32 | 29 | import orbax.checkpoint as ocp |
33 | 30 |
|
34 | 31 |
|
35 | 32 | WAN_CHECKPOINT = "WAN_CHECKPOINT" |
36 | 33 |
|
37 | 34 |
|
38 | | -class WanCheckpointer(ABC): |
| 35 | +T = TypeVar("T", bound=WanPipeline) |
| 36 | + |
| 37 | + |
| 38 | +class WanCheckpointer(Generic[T], ABC): |
| 39 | + pipeline_class: Optional[Type[T]] = None |
39 | 40 |
|
40 | 41 | def __init__(self, config, checkpoint_type: str = WAN_CHECKPOINT): |
41 | 42 | self.config = config |
@@ -176,16 +177,61 @@ def _pretrained_save_items(pipeline, pretrained_state_sources, pretrained_config |
176 | 177 | def load_wan_configs_from_orbax(self, step: Optional[int]) -> Tuple[Optional[dict], Optional[int]]: |
177 | 178 | raise NotImplementedError |
178 | 179 |
|
179 | | - @abstractmethod |
180 | | - def load_diffusers_checkpoint(self): |
181 | | - raise NotImplementedError |
| 180 | + def load_diffusers_checkpoint( |
| 181 | + self, |
| 182 | + vae_only=False, |
| 183 | + load_vae=None, |
| 184 | + load_text_encoder=None, |
| 185 | + load_transformer=None, |
| 186 | + load_scheduler=None, |
| 187 | + ) -> T: |
| 188 | + pipeline = self.pipeline_class.from_pretrained( |
| 189 | + self.config, |
| 190 | + vae_only=vae_only, |
| 191 | + load_vae=load_vae, |
| 192 | + load_text_encoder=load_text_encoder, |
| 193 | + load_transformer=load_transformer, |
| 194 | + load_scheduler=load_scheduler, |
| 195 | + ) |
| 196 | + return pipeline |
182 | 197 |
|
183 | | - @abstractmethod |
184 | 198 | def load_checkpoint( |
185 | | - self, step=None |
186 | | - ) -> Tuple[ |
187 | | - Optional[WanPipeline2_1 | WanPipeline2_2 | WanPipelineI2V_2_1 | WanPipelineI2V_2_2], Optional[dict], Optional[int] |
188 | | - ]: |
| 199 | + self, |
| 200 | + step=None, |
| 201 | + vae_only=False, |
| 202 | + load_vae=None, |
| 203 | + load_text_encoder=None, |
| 204 | + load_transformer=None, |
| 205 | + load_scheduler=None, |
| 206 | + ) -> Tuple[T, Optional[dict], Optional[int]]: |
| 207 | + restored_checkpoint, step = self.load_wan_configs_from_orbax(step) |
| 208 | + opt_state = None |
| 209 | + if restored_checkpoint: |
| 210 | + max_logging.log("Loading WAN pipeline from checkpoint") |
| 211 | + pipeline = self.pipeline_class.from_checkpoint( |
| 212 | + self.config, |
| 213 | + restored_checkpoint, |
| 214 | + vae_only=vae_only, |
| 215 | + load_vae=load_vae, |
| 216 | + load_text_encoder=load_text_encoder, |
| 217 | + load_transformer=load_transformer, |
| 218 | + load_scheduler=load_scheduler, |
| 219 | + ) |
| 220 | + opt_state = self._extract_opt_state(restored_checkpoint) |
| 221 | + else: |
| 222 | + max_logging.log("No checkpoint found, loading default pipeline.") |
| 223 | + pipeline = self.load_diffusers_checkpoint( |
| 224 | + vae_only=vae_only, |
| 225 | + load_vae=load_vae, |
| 226 | + load_text_encoder=load_text_encoder, |
| 227 | + load_transformer=load_transformer, |
| 228 | + load_scheduler=load_scheduler, |
| 229 | + ) |
| 230 | + |
| 231 | + return pipeline, opt_state, step |
| 232 | + |
| 233 | + @abstractmethod |
| 234 | + def _extract_opt_state(self, restored_checkpoint): |
189 | 235 | raise NotImplementedError |
190 | 236 |
|
191 | 237 | @abstractmethod |
|
0 commit comments