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
-
+
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