diff --git a/README.md b/README.md index 721f047..d86e573 100644 --- a/README.md +++ b/README.md @@ -189,7 +189,7 @@ Same function as `LoraLoader` node, but acts on UNet3DConditionModel. Used after --- ### TrainUnetSequence -Local Image +Local Image Fine-tune the incoming model using latent vector and context, and convert the model to inference mode. @@ -208,6 +208,8 @@ Fine-tune the incoming model using latent vector and context, and convert the mo - 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. diff --git a/__init__.py b/__init__.py index 5514f0f..20c003c 100644 --- a/__init__.py +++ b/__init__.py @@ -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,) diff --git a/images/nodes/TrainUnetSequence.png b/images/nodes/TrainUnetSequence.png index 09f0d15..77e9e1a 100644 Binary files a/images/nodes/TrainUnetSequence.png and b/images/nodes/TrainUnetSequence.png differ