refactor lora load/unload

This commit is contained in:
Vladimir Mandic
2023-10-08 12:00:51 -04:00
parent c0e4605dfb
commit f2fc41cfc2
10 changed files with 344 additions and 297 deletions
+32
View File
@@ -92,6 +92,8 @@ class ImageGridLoopParams:
ScriptCallback = namedtuple("ScriptCallback", ["script", "callback"])
callback_map = dict(
callbacks_app_started=[],
callbacks_before_process=[],
callbacks_after_process=[],
callbacks_model_loaded=[],
callbacks_ui_tabs=[],
callbacks_ui_train_tabs=[],
@@ -134,6 +136,26 @@ def app_started_callback(demo: Optional[Blocks], app: FastAPI):
report_exception(e, c, 'app_started_callback')
def before_process_callback(p):
for c in callback_map['callbacks_before_process']:
try:
t0 = time.time()
c.callback(p)
timer(t0, c.script, 'before_process')
except Exception as e:
report_exception(e, c, 'before_process_callback')
def after_process_callback(p):
for c in callback_map['callbacks_after_process']:
try:
t0 = time.time()
c.callback(p)
timer(t0, c.script, 'after_process')
except Exception as e:
report_exception(e, c, 'after_process_callback')
def app_reload_callback():
for c in callback_map['callbacks_on_reload']:
try:
@@ -334,6 +356,16 @@ def on_app_started(callback):
add_callback(callback_map['callbacks_app_started'], callback)
def on_before_process(callback):
"""register a function to be called just before processing starts"""
add_callback(callback_map['callbacks_before_process'], callback)
def on_after_process(callback):
"""register a function to be called just after processing ends"""
add_callback(callback_map['callbacks_after_process'], callback)
def on_before_reload(callback):
"""register a function to be called just before the server reloads."""
add_callback(callback_map['callbacks_on_reload'], callback)