| 34 | } |
| 35 | |
| 36 | |
| 37 | def get_key_mapping_rules(direction, model_type): |
| 38 | if model_type == "wan_dit": |
| 39 | unified_rules = [ |
| 40 | { |
| 41 | "forward": (r"^head\.head$", "proj_out"), |
| 42 | "backward": (r"^proj_out$", "head.head"), |
| 43 | }, |
| 44 | { |
| 45 | "forward": (r"^head\.modulation$", "scale_shift_table"), |
| 46 | "backward": (r"^scale_shift_table$", "head.modulation"), |
| 47 | }, |
| 48 | { |
| 49 | "forward": ( |
| 50 | r"^text_embedding\.0\.", |
| 51 | "condition_embedder.text_embedder.linear_1.", |
| 52 | ), |
| 53 | "backward": ( |
| 54 | r"^condition_embedder.text_embedder.linear_1\.", |
| 55 | "text_embedding.0.", |
| 56 | ), |
| 57 | }, |
| 58 | { |
| 59 | "forward": ( |
| 60 | r"^text_embedding\.2\.", |
| 61 | "condition_embedder.text_embedder.linear_2.", |
| 62 | ), |
| 63 | "backward": ( |
| 64 | r"^condition_embedder.text_embedder.linear_2\.", |
| 65 | "text_embedding.2.", |
| 66 | ), |
| 67 | }, |
| 68 | { |
| 69 | "forward": ( |
| 70 | r"^time_embedding\.0\.", |
| 71 | "condition_embedder.time_embedder.linear_1.", |
| 72 | ), |
| 73 | "backward": ( |
| 74 | r"^condition_embedder.time_embedder.linear_1\.", |
| 75 | "time_embedding.0.", |
| 76 | ), |
| 77 | }, |
| 78 | { |
| 79 | "forward": ( |
| 80 | r"^time_embedding\.2\.", |
| 81 | "condition_embedder.time_embedder.linear_2.", |
| 82 | ), |
| 83 | "backward": ( |
| 84 | r"^condition_embedder.time_embedder.linear_2\.", |
| 85 | "time_embedding.2.", |
| 86 | ), |
| 87 | }, |
| 88 | { |
| 89 | "forward": (r"^time_projection\.1\.", "condition_embedder.time_proj."), |
| 90 | "backward": (r"^condition_embedder.time_proj\.", "time_projection.1."), |
| 91 | }, |
| 92 | { |
| 93 | "forward": (r"blocks\.(\d+)\.self_attn\.q\.", r"blocks.\1.attn1.to_q."), |