diff --git a/README.md b/README.md index 5a950bf..86239e6 100644 --- a/README.md +++ b/README.md @@ -39,16 +39,18 @@ Model names: | `v1` | Stable Diffusion v1.x | | `xl` | SDXL | | `v3` | Stable Diffusion 3 | +| `fx` | Flux.1 | | `ca` | Stable Cascade (Stage A/B) | Available models: -| From | to `v1` | to `xl` | to `v3` | to `ca` | -|:----:|:-------:|:-------:|:-------:|:-------:| -| `v1` | - | v4.0 | v4.0 | No | -| `xl` | v4.0 | - | v4.0 | No | -| `v3` | v4.0 | v4.0 | - | No | -| `ca` | v4.0 | v4.0 | v4.0 | - | +| From | to `v1` | to `xl` | to `v3` | to `fx` | to `ca` | +|:----:|:-------:|:-------:|:-------:|:-------:|:-------:| +| `v1` | - | v4.0 | v4.0 | No | No | +| `xl` | v4.0 | - | v4.0 | No | No | +| `v3` | v4.0 | v4.0 | - | No | No | +| `fx` | v4.0 | v4.0 | v4.0 | - | No | +| `ca` | v4.0 | v4.0 | v4.0 | No | - | ## Training diff --git a/comfy_latent_interposer.py b/comfy_latent_interposer.py index a5a15cf..958123e 100644 --- a/comfy_latent_interposer.py +++ b/comfy_latent_interposer.py @@ -7,6 +7,7 @@ from huggingface_hub import hf_hub_download # v1 = Stable Diffusion 1.x # xl = Stable Diffusion Extra Large (SDXL) # v3 = Stable Diffusion Version Three (SD3) +# fx = Black Forest Labs Flux dot One # cc = Stable Cascade (Stage C) [not used] # ca = Stable Cascade (Stage A/B) config = { @@ -16,6 +17,9 @@ config = { "xl-to-v3": {"ch_in": 4, "ch_out":16, "ch_mid": 64, "scale": 1.0, "blocks": 12}, "v3-to-v1": {"ch_in":16, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12}, "v3-to-xl": {"ch_in":16, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12}, + "fx-to-v1": {"ch_in":16, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12}, + "fx-to-xl": {"ch_in":16, "ch_out": 4, "ch_mid": 64, "scale": 1.0, "blocks": 12}, + "fx-to-v3": {"ch_in":16, "ch_out":16, "ch_mid": 64, "scale": 1.0, "blocks": 12}, "ca-to-v1": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12}, "ca-to-xl": {"ch_in": 4, "ch_out": 4, "ch_mid": 64, "scale": 0.5, "blocks": 12}, "ca-to-v3": {"ch_in": 4, "ch_out":16, "ch_mid": 64, "scale": 0.5, "blocks": 12}, @@ -95,7 +99,7 @@ class ComfyLatentInterposer: return { "required": { "samples": ("LATENT", ), - "latent_src": (["v1", "xl", "v3", "ca"],), + "latent_src": (["v1", "xl", "v3", "fx", "ca"],), "latent_dst": (["v1", "xl", "v3"],), } } diff --git a/config/fx-to-v1.yaml b/config/fx-to-v1.yaml new file mode 100644 index 0000000..88e7d44 --- /dev/null +++ b/config/fx-to-v1.yaml @@ -0,0 +1,40 @@ +steps: 35000 +batch: 48 +fconst: 0 +cosine: False +resume: False +device: "cuda" +p_loss_weight: 1.0 +r_loss_weight: 0.1 +b_loss_weight: 1.0 +h_loss_weight: 0.0 +save_image: 100 +eval_model: 10 + +model: + src: fx # Black Forest Labs Flux dot One + dst: v1 # Stable Diffusion 1.x + rev: "v4.0-rc1" + args: + scale: 1.0 + ch_in: 16 + ch_out: 4 + ch_mid: 64 + blocks: 12 + +optim: + lr: 5.0e-4 + beta1: 0.5 + beta2: 0.95 + +dataset: + src: "./latents/fx_256px_combined.bin" + dst: "./latents/v1_256px_combined.bin" + preload: False + evals: + main: + src: "./latents/test_eru/test_fx_768px.npy" + dst: "./latents/test_eru/test_v1_768px.npy" + aux: + src: "./latents/test_bga/test_fx_768px.npy" + dst: "./latents/test_bga/test_v1_768px.npy" diff --git a/config/fx-to-v3.yaml b/config/fx-to-v3.yaml new file mode 100644 index 0000000..6dabbcd --- /dev/null +++ b/config/fx-to-v3.yaml @@ -0,0 +1,40 @@ +steps: 35000 +batch: 48 +fconst: 0 +cosine: False +resume: False +device: "cuda" +p_loss_weight: 1.0 +r_loss_weight: 0.1 +b_loss_weight: 1.0 +h_loss_weight: 0.0 +save_image: 100 +eval_model: 10 + +model: + src: fx # Black Forest Labs Flux dot One + dst: v3 # Stable Diffusion Version three point oh + rev: "v4.0-rc2" + args: + scale: 1.0 + ch_in: 16 + ch_out: 16 + ch_mid: 64 + blocks: 12 + +optim: + lr: 5.0e-4 + beta1: 0.5 + beta2: 0.95 + +dataset: + src: "./latents/fx_256px_combined.bin" + dst: "./latents/v3_256px_combined.bin" + preload: False + evals: + main: + src: "./latents/test_eru/test_fx_768px.npy" + dst: "./latents/test_eru/test_v3_768px.npy" + aux: + src: "./latents/test_bga/test_fx_768px.npy" + dst: "./latents/test_bga/test_v3_768px.npy" diff --git a/config/fx-to-xl.yaml b/config/fx-to-xl.yaml new file mode 100644 index 0000000..0d7c382 --- /dev/null +++ b/config/fx-to-xl.yaml @@ -0,0 +1,40 @@ +steps: 35000 +batch: 48 +fconst: 0 +cosine: False +resume: False +device: "cuda" +p_loss_weight: 1.0 +r_loss_weight: 0.1 +b_loss_weight: 1.0 +h_loss_weight: 0.0 +save_image: 100 +eval_model: 10 + +model: + src: fx # Black Forest Labs Flux dot One + dst: xl # Stable Diffusion Extra Large + rev: "v4.0-rc3" + args: + scale: 1.0 + ch_in: 16 + ch_out: 4 + ch_mid: 64 + blocks: 12 + +optim: + lr: 5.0e-4 + beta1: 0.5 + beta2: 0.95 + +dataset: + src: "./latents/fx_256px_combined.bin" + dst: "./latents/xl_256px_combined.bin" + preload: False + evals: + main: + src: "./latents/test_eru/test_fx_768px.npy" + dst: "./latents/test_eru/test_xl_768px.npy" + aux: + src: "./latents/test_bga/test_fx_768px.npy" + dst: "./latents/test_bga/test_xl_768px.npy" diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..0c675f7 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,13 @@ +[project] +name = "sd-latent-interposer" +description = "Custom node to convert the lantents between SDv1/SDXL/SD3/cascade/flux directly without the VAE decoding/encoding step." +version = "4.0.15" +license = { file = "LICENSE" } + +[project.urls] +Repository = "https://github.com/city96/SD-Latent-Interposer" + +[tool.comfy] +PublisherId = "city96" +DisplayName = "SD-Latent-Interposer" +Icon = "" diff --git a/train.py b/train.py index f2e6e5a..ee4855a 100644 --- a/train.py +++ b/train.py @@ -238,7 +238,7 @@ if __name__ == "__main__": # from torchvision.utils import save_image # save_file(model.state_dict(), f"{base_name}_{pbar.n:07}.safetensors") # for name, image in eval_images(model, vae, evals).items(): - # name = f"models/{name.replace('/', '_')}_{pbar.n:07}.png" + # name = f"{base_name}_{name.replace('/', '_')}_{pbar.n:07}.png" # save_image(image, name) # final save/cleanup diff --git a/vae.py b/vae.py index 713ffa9..295df59 100644 --- a/vae.py +++ b/vae.py @@ -62,6 +62,20 @@ class SDv3_VAE(SDv1_VAE): if dec_only: del self.model.encoder +class FLUX_VAE(SDv1_VAE): + scale = 1/8 + channels = 16 + def __init__(self, device=DEVICE, dtype=DTYPE, dec_only=False): + self.device = device + self.dtype = torch.bfloat16 if dec_only else dtype # decoder NaNs randomly + self.model = AutoencoderKL.from_pretrained( + "black-forest-labs/FLUX.1-dev", + subfolder="vae" + ) + self.model.eval().to(self.dtype).to(self.device) + if dec_only: + del self.model.encoder + class CascadeC_VAE(SDv1_VAE): scale = 1/32 channels = 16 @@ -135,6 +149,7 @@ vae_vers = { "v3": SDv3_VAE, "cc": CascadeC_VAE, "ca": CascadeA_VAE, + "fx": FLUX_VAE, } def load_vae(ver, *args, **kwargs):