persist control units state on restarts

This commit is contained in:
Vladimir Mandic
2024-09-19 20:10:53 -04:00
parent 271b493921
commit 4acf3c00a2
10 changed files with 117 additions and 34 deletions
+7 -1
View File
@@ -181,7 +181,7 @@ class ControlNet():
cls = self.get_class()
self.model = cls.from_single_file(model_path, **self.load_config)
def load(self, model_id: str = None) -> str:
def load(self, model_id: str = None, force: bool = True) -> str:
try:
t0 = time.time()
model_id = model_id or self.model_id
@@ -197,6 +197,9 @@ class ControlNet():
if model_path is None:
log.error(f'Control {what} model load failed: id="{model_id}" error=unknown model id')
return
if model_id == self.model_id and not force:
log.debug(f'Control {what} model: id="{model_id}" path="{model_path}" already loaded')
return
log.debug(f'Control {what} model loading: id="{model_id}" path="{model_path}"')
if model_path.endswith('.safetensors'):
self.load_safetensors(model_path)
@@ -205,6 +208,9 @@ class ControlNet():
model_path = model_path.replace('/bin', '')
self.load_config['use_safetensors'] = False
cls = self.get_class()
if cls is None:
log.error(f'Control {what} model load failed: id="{model_id}" unknown base model')
return
self.model = cls.from_pretrained(model_path, **self.load_config)
if self.dtype is not None:
self.model.to(self.dtype)
+4 -1
View File
@@ -78,7 +78,7 @@ class ControlLLLite():
self.model = None
self.model_id = None
def load(self, model_id: str = None) -> str:
def load(self, model_id: str = None, force: bool = True) -> str:
try:
t0 = time.time()
model_id = model_id or self.model_id
@@ -94,6 +94,9 @@ class ControlLLLite():
if model_path is None:
log.error(f'Control {what} model load failed: id="{model_id}" error=unknown model id')
return
if model_id == self.model_id and not force:
log.debug(f'Control {what} model: id="{model_id}" path="{model_path}" already loaded')
return
log.debug(f'Control {what} model loading: id="{model_id}" path="{model_path}" {self.load_config}')
if model_path.endswith('.safetensors'):
self.model = ControlNetLLLite(model_path)
+4 -1
View File
@@ -86,7 +86,7 @@ class Adapter():
self.model = None
self.model_id = None
def load(self, model_id: str = None) -> str:
def load(self, model_id: str = None, force: bool = True) -> str:
try:
t0 = time.time()
model_id = model_id or self.model_id
@@ -100,6 +100,9 @@ class Adapter():
if model_path is None:
log.error(f'Control {what} model load failed: id="{model_id}" error=unknown model id')
return
if model_id == self.model_id and not force:
log.debug(f'Control {what} model: id="{model_id}" path="{model_path}" already loaded')
return
log.debug(f'Control {what} model loading: id="{model_id}" path="{model_path}"')
if model_path.endswith('.pth') or model_path.endswith('.pt') or model_path.endswith('.safetensors') or model_path.endswith('.bin'):
from huggingface_hub import hf_hub_download
+4 -1
View File
@@ -74,7 +74,7 @@ class ControlNetXS():
self.model = None
self.model_id = None
def load(self, model_id: str = None, time_embedding_mix: float = 0.0) -> str:
def load(self, model_id: str = None, time_embedding_mix: float = 0.0, force: bool = True) -> str:
try:
t0 = time.time()
model_id = model_id or self.model_id
@@ -90,6 +90,9 @@ class ControlNetXS():
if model_path is None:
log.error(f'Control {what} model load failed: id="{model_id}" error=unknown model id')
return
if model_id == self.model_id and not force:
log.debug(f'Control {what} model: id="{model_id}" path="{model_path}" already loaded')
return
self.load_config['time_embedding_mix'] = time_embedding_mix
log.debug(f'Control {what} model loading: id="{model_id}" path="{model_path}" {self.load_config}')
if model_path.endswith('.safetensors'):