4 Commits
Author SHA1 Message Date
sylym 65626c4c4a Merge pull request #10 from rttt1093/main
fix #9
2023-08-12 12:35:23 +08:00
rttt1093 22a8bb0d7c fix #9 2023-08-12 13:23:00 +09:00
sylym d8cbc8fddd Add the seed parameter to the TrainUnetSequence node 2023-04-09 19:24:02 +08:00
sylym 24266cc4ef Add examples 2023-04-05 23:27:32 +08:00
5 changed files with 14 additions and 11 deletions
+7 -2
View File
@@ -25,6 +25,9 @@ For ComfyUI portable standalone build:
## Usage
All nodes are classified under the vid2vid category.
For some workflow examples you can check out:
### [vid2vid workflow examples](https://github.com/sylym/comfy_vid2vid/releases/tag/v1.0.0)
## Nodes
@@ -186,7 +189,7 @@ Same function as `LoraLoader` node, but acts on UNet3DConditionModel. Used after
---
### TrainUnetSequence
<img alt="Local Image" src="images/nodes/TrainUnetSequence.png" width="518" height="264"/>
<img alt="Local Image" src="images/nodes/TrainUnetSequence.png" width="518" height="400"/>
Fine-tune the incoming model using latent vector and context, and convert the model to inference mode.
@@ -198,13 +201,15 @@ Fine-tune the incoming model using latent vector and context, and convert the mo
- The model that will be fine-tuned.
- context: CONDITIONING
- The context that will be used to fine-tune the incoming model.
- The context used for fine-tuning the input model, typically consists of words or sentences describing the subject of the action in the latent vector and its behavior.
**Outputs:**
- MODEL
- The fine-tuned model. This model is ready for inference.
**Parameters:**
- seed
- The seed used in model fine-tuning.
- steps
- The number of steps to fine-tune the model. If the steps is 0, the model will not be fine-tuned.
+4 -3
View File
@@ -315,7 +315,8 @@ class TrainUnetSequence:
return {"required": {"samples": ("LATENT",),
"model": ("ORIGINAL_MODEL",),
"context": ("CONDITIONING",),
"steps": ("INT", {"default": 20, "min": 0, "max": 10000}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffff}),
"steps": ("INT", {"default": 100, "min": 0, "max": 10000}),
}}
RETURN_TYPES = ("MODEL",)
@@ -323,12 +324,12 @@ class TrainUnetSequence:
CATEGORY = "vid2vid"
def train_unet(self, samples, model, context, steps):
def train_unet(self, samples, model, context, seed, steps):
device = model_management.get_torch_device()
noise_scheduler = convert_scheduler_checkpoint(model)
samples = rearrange(samples["samples"], "f c h w -> c f h w")
with torch.inference_mode(mode=False):
model_train = train(copy.deepcopy(model), noise_scheduler, samples, context[0][0].squeeze(0), device, max_train_steps=steps)
model_train = train(copy.deepcopy(model), noise_scheduler, samples, context[0][0].squeeze(0), device, max_train_steps=steps, seed=seed)
if model_management.should_use_fp16():
model_train.model = model_train.model.half()
return (model_train,)
Binary file not shown.

Before

Width:  |  Height:  |  Size: 24 KiB

After

Width:  |  Height:  |  Size: 54 KiB

+3 -3
View File
@@ -1,6 +1,6 @@
import torch
from comfy import model_management
from comfy.sd import load_model_weights, ModelPatcher, VAE, CLIP, model_lora_keys
from comfy.sd import load_model_weights, ModelPatcher, VAE, CLIP, model_lora_keys_unet, model_lora_keys_clip
from comfy import utils
from comfy import clip_vision
from comfy.ldm.util import instantiate_from_config
@@ -261,8 +261,8 @@ def use_lora(pretrained_LoRA_path, model, alpha):
def load_lora_for_models(model, clip, lora_path, strength_model, strength_clip):
key_map = model_lora_keys(model.model)
key_map = model_lora_keys(clip.cond_stage_model, key_map)
key_map = model_lora_keys_unet(model.model)
key_map = model_lora_keys_clip(clip.cond_stage_model, key_map)
loaded = load_lora(lora_path, key_map)
new_modelpatcher = model.clone()
new_modelpatcher = use_lora(lora_path, new_modelpatcher, strength_model)
-3
View File
@@ -8,7 +8,6 @@ import torch.utils.checkpoint
from torch.utils.data import Dataset
from accelerate import Accelerator
from accelerate.logging import get_logger
from accelerate.utils import set_seed
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
@@ -20,8 +19,6 @@ from einops import rearrange
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.10.0.dev0")
logger = get_logger(__name__, log_level="INFO")
class TuneAVideoDataset(Dataset):
def __init__(self):
self.prompt_ids = None