Add the seed parameter to the TrainUnetSequence node
This commit is contained in:
@@ -189,7 +189,7 @@ Same function as `LoraLoader` node, but acts on UNet3DConditionModel. Used after
|
|||||||
|
|
||||||
---
|
---
|
||||||
### TrainUnetSequence
|
### 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.
|
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.
|
- The fine-tuned model. This model is ready for inference.
|
||||||
|
|
||||||
**Parameters:**
|
**Parameters:**
|
||||||
|
- seed
|
||||||
|
- The seed used in model fine-tuning.
|
||||||
- steps
|
- steps
|
||||||
- The number of steps to fine-tune the model. If the steps is 0, the model will not be fine-tuned.
|
- The number of steps to fine-tune the model. If the steps is 0, the model will not be fine-tuned.
|
||||||
|
|
||||||
|
|||||||
+4
-3
@@ -315,7 +315,8 @@ class TrainUnetSequence:
|
|||||||
return {"required": {"samples": ("LATENT",),
|
return {"required": {"samples": ("LATENT",),
|
||||||
"model": ("ORIGINAL_MODEL",),
|
"model": ("ORIGINAL_MODEL",),
|
||||||
"context": ("CONDITIONING",),
|
"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",)
|
RETURN_TYPES = ("MODEL",)
|
||||||
@@ -323,12 +324,12 @@ class TrainUnetSequence:
|
|||||||
|
|
||||||
CATEGORY = "vid2vid"
|
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()
|
device = model_management.get_torch_device()
|
||||||
noise_scheduler = convert_scheduler_checkpoint(model)
|
noise_scheduler = convert_scheduler_checkpoint(model)
|
||||||
samples = rearrange(samples["samples"], "f c h w -> c f h w")
|
samples = rearrange(samples["samples"], "f c h w -> c f h w")
|
||||||
with torch.inference_mode(mode=False):
|
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():
|
if model_management.should_use_fp16():
|
||||||
model_train.model = model_train.model.half()
|
model_train.model = model_train.model.half()
|
||||||
return (model_train,)
|
return (model_train,)
|
||||||
|
|||||||
Binary file not shown.
|
Before Width: | Height: | Size: 24 KiB After Width: | Height: | Size: 54 KiB |
Reference in New Issue
Block a user