(self, state_dict)
| 438 | pass |
| 439 | |
| 440 | def from_diffusers(self, state_dict): |
| 441 | rename_dict = { |
| 442 | "blocks.0.attn1.norm_k.weight": "blocks.0.self_attn.norm_k.weight", |
| 443 | "blocks.0.attn1.norm_q.weight": "blocks.0.self_attn.norm_q.weight", |
| 444 | "blocks.0.attn1.to_k.bias": "blocks.0.self_attn.k.bias", |
| 445 | "blocks.0.attn1.to_k.weight": "blocks.0.self_attn.k.weight", |
| 446 | "blocks.0.attn1.to_out.0.bias": "blocks.0.self_attn.o.bias", |
| 447 | "blocks.0.attn1.to_out.0.weight": "blocks.0.self_attn.o.weight", |
| 448 | "blocks.0.attn1.to_q.bias": "blocks.0.self_attn.q.bias", |
| 449 | "blocks.0.attn1.to_q.weight": "blocks.0.self_attn.q.weight", |
| 450 | "blocks.0.attn1.to_v.bias": "blocks.0.self_attn.v.bias", |
| 451 | "blocks.0.attn1.to_v.weight": "blocks.0.self_attn.v.weight", |
| 452 | "blocks.0.attn2.norm_k.weight": "blocks.0.cross_attn.norm_k.weight", |
| 453 | "blocks.0.attn2.norm_q.weight": "blocks.0.cross_attn.norm_q.weight", |
| 454 | "blocks.0.attn2.to_k.bias": "blocks.0.cross_attn.k.bias", |
| 455 | "blocks.0.attn2.to_k.weight": "blocks.0.cross_attn.k.weight", |
| 456 | "blocks.0.attn2.to_out.0.bias": "blocks.0.cross_attn.o.bias", |
| 457 | "blocks.0.attn2.to_out.0.weight": "blocks.0.cross_attn.o.weight", |
| 458 | "blocks.0.attn2.to_q.bias": "blocks.0.cross_attn.q.bias", |
| 459 | "blocks.0.attn2.to_q.weight": "blocks.0.cross_attn.q.weight", |
| 460 | "blocks.0.attn2.to_v.bias": "blocks.0.cross_attn.v.bias", |
| 461 | "blocks.0.attn2.to_v.weight": "blocks.0.cross_attn.v.weight", |
| 462 | "blocks.0.ffn.net.0.proj.bias": "blocks.0.ffn.0.bias", |
| 463 | "blocks.0.ffn.net.0.proj.weight": "blocks.0.ffn.0.weight", |
| 464 | "blocks.0.ffn.net.2.bias": "blocks.0.ffn.2.bias", |
| 465 | "blocks.0.ffn.net.2.weight": "blocks.0.ffn.2.weight", |
| 466 | "blocks.0.norm2.bias": "blocks.0.norm3.bias", |
| 467 | "blocks.0.norm2.weight": "blocks.0.norm3.weight", |
| 468 | "blocks.0.scale_shift_table": "blocks.0.modulation", |
| 469 | "condition_embedder.text_embedder.linear_1.bias": "text_embedding.0.bias", |
| 470 | "condition_embedder.text_embedder.linear_1.weight": "text_embedding.0.weight", |
| 471 | "condition_embedder.text_embedder.linear_2.bias": "text_embedding.2.bias", |
| 472 | "condition_embedder.text_embedder.linear_2.weight": "text_embedding.2.weight", |
| 473 | "condition_embedder.time_embedder.linear_1.bias": "time_embedding.0.bias", |
| 474 | "condition_embedder.time_embedder.linear_1.weight": "time_embedding.0.weight", |
| 475 | "condition_embedder.time_embedder.linear_2.bias": "time_embedding.2.bias", |
| 476 | "condition_embedder.time_embedder.linear_2.weight": "time_embedding.2.weight", |
| 477 | "condition_embedder.time_proj.bias": "time_projection.1.bias", |
| 478 | "condition_embedder.time_proj.weight": "time_projection.1.weight", |
| 479 | "patch_embedding.bias": "patch_embedding.bias", |
| 480 | "patch_embedding.weight": "patch_embedding.weight", |
| 481 | "scale_shift_table": "head.modulation", |
| 482 | "proj_out.bias": "head.head.bias", |
| 483 | "proj_out.weight": "head.head.weight", |
| 484 | } |
| 485 | state_dict_ = {} |
| 486 | for name, param in state_dict.items(): |
| 487 | if name in rename_dict: |
| 488 | state_dict_[rename_dict[name]] = param |
| 489 | else: |
| 490 | name_ = ".".join(name.split(".")[:1] + ["0"] + name.split(".")[2:]) |
| 491 | if name_ in rename_dict: |
| 492 | name_ = rename_dict[name_] |
| 493 | name_ = ".".join(name_.split(".")[:1] + [name.split(".")[1]] + name_.split(".")[2:]) |
| 494 | state_dict_[name_] = param |
| 495 | if hash_state_dict_keys(state_dict) == "cb104773c6c2cb6df4f9529ad5c60d0b": |
| 496 | config = { |
| 497 | "model_type": "t2v", |
nothing calls this directly
no test coverage detected