Add the seed parameter to the TrainUnetSequence node

This commit is contained in:
sylym
2023-04-09 19:24:02 +08:00
parent 24266cc4ef
commit d8cbc8fddd
3 changed files with 7 additions and 4 deletions
+3 -1
View File
@@ -189,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.
@@ -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.
+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