64 lines
1.7 KiB
Python
64 lines
1.7 KiB
Python
from dataclasses import dataclass, field
|
|
from typing import Union
|
|
|
|
import torch
|
|
|
|
|
|
class KeyframePart:
|
|
def __init__(self, batch_index: int, image: torch.Tensor, denoise: float) -> None:
|
|
self.batch_index = batch_index
|
|
self.denoise = denoise
|
|
self.image = image
|
|
|
|
|
|
class KeyframePartGroup:
|
|
def __init__(self) -> None:
|
|
self.keyframes: list[KeyframePart] = []
|
|
|
|
def add(self, keyframe: KeyframePart) -> None:
|
|
added = False
|
|
for i in range(len(self.keyframes)):
|
|
if self.keyframes[i].batch_index == keyframe.batch_index:
|
|
self.keyframes[i] = keyframe
|
|
added = True
|
|
break
|
|
if not added:
|
|
self.keyframes.append(keyframe)
|
|
self.keyframes.sort(key=lambda k: k.batch_index)
|
|
|
|
def get_index(self, index: int) -> Union[KeyframePart, None]:
|
|
try:
|
|
return self.keyframes[index]
|
|
except IndexError:
|
|
return None
|
|
|
|
def __getitem__(self, index) -> KeyframePart:
|
|
return self.keyframes[index]
|
|
|
|
def is_empty(self) -> bool:
|
|
return len(self.keyframes) == 0
|
|
|
|
def clone(self) -> 'KeyframePartGroup':
|
|
cloned = KeyframePartGroup()
|
|
for k in self.keyframes:
|
|
cloned.add(k)
|
|
return cloned
|
|
|
|
|
|
@dataclass
|
|
class ModelInjectParam:
|
|
keyframe_part_group: KeyframePartGroup
|
|
latent: dict = field(default_factory=dict, repr=False)
|
|
seed: int = 0
|
|
steps: int = 0
|
|
scheduler: str = 'normal'
|
|
denoise: float = 0
|
|
noise: torch.Tensor = field(default=None, repr=False)
|
|
|
|
def reset(self):
|
|
self.seed: int = 0
|
|
self.steps: int = 0
|
|
self.scheduler: str = 'normal'
|
|
self.denoise: float = 0
|
|
self.noise: torch.Tensor = None
|