mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
fixes based on full skills run
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+11
-11
@@ -44,7 +44,7 @@ class Api:
|
||||
# server api
|
||||
self.add_api_route("/sdapi/v1/motd", server.get_motd, methods=["GET"], response_model=str)
|
||||
self.add_api_route("/sdapi/v1/log", server.get_log, methods=["GET"], response_model=list[str])
|
||||
self.add_api_route("/sdapi/v1/log", server.post_log, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/log", server.post_log, methods=["POST"], status_code=204)
|
||||
self.add_api_route("/sdapi/v1/start", self.get_session_start, methods=["GET"])
|
||||
self.add_api_route("/sdapi/v1/version", server.get_version, methods=["GET"])
|
||||
self.add_api_route("/sdapi/v1/torch", server.get_torch, methods=["GET"])
|
||||
@@ -52,9 +52,9 @@ class Api:
|
||||
self.add_api_route("/sdapi/v1/platform", server.get_platform, methods=["GET"])
|
||||
self.add_api_route("/sdapi/v1/progress", server.get_progress, methods=["GET"], response_model=models.ResProgress)
|
||||
self.add_api_route("/sdapi/v1/history", server.get_history, methods=["GET"], response_model=list[models.ResHistory])
|
||||
self.add_api_route("/sdapi/v1/interrupt", server.post_interrupt, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/skip", server.post_skip, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/shutdown", server.post_shutdown, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/interrupt", server.post_interrupt, methods=["POST"], status_code=204)
|
||||
self.add_api_route("/sdapi/v1/skip", server.post_skip, methods=["POST"], status_code=204)
|
||||
self.add_api_route("/sdapi/v1/shutdown", server.post_shutdown, methods=["POST"], status_code=204)
|
||||
self.add_api_route("/sdapi/v1/memory", server.get_memory, methods=["GET"], response_model=models.ResMemory)
|
||||
self.add_api_route("/sdapi/v1/cmd-flags", server.get_cmd_flags, methods=["GET"], response_model=models.FlagsModel)
|
||||
self.add_api_route("/sdapi/v1/gpu", gpu.get_gpu, methods=["GET"])
|
||||
@@ -101,12 +101,12 @@ class Api:
|
||||
self.add_api_route("/sdapi/v1/png-info", endpoints.post_pnginfo, methods=["POST"], response_model=models.ResImageInfo, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/checkpoint", endpoints.get_checkpoint, methods=["GET"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/checkpoint", endpoints.set_checkpoint, methods=["POST"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/refresh-checkpoints", endpoints.post_refresh_checkpoints, methods=["POST"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/unload-checkpoint", endpoints.post_unload_checkpoint, methods=["POST"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/reload-checkpoint", endpoints.post_reload_checkpoint, methods=["POST"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/lock-checkpoint", endpoints.post_lock_checkpoint, methods=["POST"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/refresh-vae", endpoints.post_refresh_vae, methods=["POST"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/refresh-unets", endpoints.post_refresh_unets, methods=["POST"], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/refresh-checkpoints", endpoints.post_refresh_checkpoints, methods=["POST"], status_code=204, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/unload-checkpoint", endpoints.post_unload_checkpoint, methods=["POST"], status_code=204, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/reload-checkpoint", endpoints.post_reload_checkpoint, methods=["POST"], status_code=204, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/lock-checkpoint", endpoints.post_lock_checkpoint, methods=["POST"], status_code=204, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/refresh-vae", endpoints.post_refresh_vae, methods=["POST"], status_code=204, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/refresh-unets", endpoints.post_refresh_unets, methods=["POST"], status_code=204, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/latents", endpoints.get_latent_history, methods=["GET"], response_model=list[str], tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/latents", endpoints.post_latent_history, methods=["POST"], response_model=int, tags=["Functional"])
|
||||
self.add_api_route("/sdapi/v1/modules", endpoints.get_modules, methods=["GET"], tags=["Functional"])
|
||||
@@ -131,7 +131,7 @@ class Api:
|
||||
|
||||
# gallery api
|
||||
from modules.api import gallery
|
||||
gallery.register_api(self.app)
|
||||
gallery.register_api(self)
|
||||
|
||||
# nudenet api
|
||||
from modules.api import nudenet
|
||||
|
||||
+2
-2
@@ -66,7 +66,7 @@ def create_docs(app: FastAPI):
|
||||
"dom_id": "#swagger-ui",
|
||||
}
|
||||
|
||||
@app.get("/docs", include_in_schema=True)
|
||||
@app.get("/docs", include_in_schema=True) # override for the default fastapi swagger route
|
||||
async def custom_swagger_html():
|
||||
res = get_swagger_ui_html(
|
||||
title=f'{app.title}: Swagger UI',
|
||||
@@ -79,7 +79,7 @@ def create_docs(app: FastAPI):
|
||||
|
||||
|
||||
def create_redocs(app: FastAPI):
|
||||
@app.get("/redocs", include_in_schema=True)
|
||||
@app.get("/redocs", include_in_schema=True) # override for the default fastapi redocs route
|
||||
async def custom_redoc_html():
|
||||
res = get_redoc_html(
|
||||
title=f'{app.title}: ReDoc',
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from fastapi.exceptions import HTTPException
|
||||
from fastapi.responses import JSONResponse, Response
|
||||
from modules import shared
|
||||
from modules.logger import log
|
||||
from modules.api import models, helpers
|
||||
@@ -216,7 +217,7 @@ def post_unload_checkpoint():
|
||||
from modules import sd_models
|
||||
sd_models.unload_model_weights(op='model')
|
||||
sd_models.unload_model_weights(op='refiner')
|
||||
return {}
|
||||
return Response(status_code=204)
|
||||
|
||||
def post_reload_checkpoint(force:bool=False):
|
||||
"""Reload the selected checkpoint. Set ``force=True`` to unload first and do a clean reload."""
|
||||
@@ -224,18 +225,19 @@ def post_reload_checkpoint(force:bool=False):
|
||||
if force:
|
||||
sd_models.unload_model_weights(op='model')
|
||||
sd_models.reload_model_weights()
|
||||
return {}
|
||||
return Response(status_code=204)
|
||||
|
||||
def post_lock_checkpoint(lock:bool=False):
|
||||
"""Lock or unlock the current model to prevent automatic model swaps."""
|
||||
from modules import modeldata
|
||||
modeldata.model_data.locked = lock
|
||||
return {}
|
||||
return Response(status_code=204)
|
||||
|
||||
def post_refresh_unets():
|
||||
"""Rescan UNet directories and update the available UNet list."""
|
||||
import modules.sd_unet
|
||||
return modules.sd_unet.refresh_unet_list()
|
||||
modules.sd_unet.refresh_unet_list()
|
||||
return Response(status_code=204)
|
||||
|
||||
def get_checkpoint():
|
||||
"""Return information about the currently loaded checkpoint including type, class, title, and hash."""
|
||||
@@ -273,12 +275,12 @@ def set_checkpoint(sd_model_checkpoint: str, dtype: str | None = None, force: bo
|
||||
def post_refresh_checkpoints():
|
||||
"""Rescan checkpoint directories and update the available models list."""
|
||||
shared.refresh_checkpoints()
|
||||
return {}
|
||||
return Response(status_code=204)
|
||||
|
||||
def post_refresh_vae():
|
||||
"""Rescan VAE directories and update the available VAE list."""
|
||||
shared.refresh_vaes()
|
||||
return {}
|
||||
return Response(status_code=204)
|
||||
|
||||
def get_modules():
|
||||
"""Analyze the loaded model and return its sub-module breakdown with device, dtype, and parameter info."""
|
||||
|
||||
@@ -71,7 +71,7 @@ class ConnectionManager:
|
||||
|
||||
### api definitions
|
||||
|
||||
def register_api(app: FastAPI): # register api
|
||||
def register_api(api): # register api
|
||||
manager = ConnectionManager()
|
||||
|
||||
def get_video_thumbnail(filepath):
|
||||
@@ -208,11 +208,11 @@ def register_api(app: FastAPI): # register api
|
||||
log.error(f'Gallery: {folder} {e}')
|
||||
return []
|
||||
|
||||
shared.api.add_api_route("/sdapi/v1/browser/folders", get_folders, methods=["GET"], response_model=list[str])
|
||||
shared.api.add_api_route("/sdapi/v1/browser/thumb", get_thumb, methods=["GET"], response_model=dict)
|
||||
shared.api.add_api_route("/sdapi/v1/browser/files", ht_files, methods=["GET"], response_model=list)
|
||||
api.add_api_route("/sdapi/v1/browser/folders", get_folders, methods=["GET"], response_model=list[str])
|
||||
api.add_api_route("/sdapi/v1/browser/thumb", get_thumb, methods=["GET"], response_model=dict)
|
||||
api.add_api_route("/sdapi/v1/browser/files", ht_files, methods=["GET"], response_model=list)
|
||||
|
||||
@app.websocket("/sdapi/v1/browser/files")
|
||||
@api.app.websocket("/sdapi/v1/browser/files")
|
||||
async def ws_files(ws: WebSocket):
|
||||
try:
|
||||
await manager.connect(ws)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
import time
|
||||
from fastapi import Request, Depends
|
||||
from fastapi import Request, Depends, BackgroundTasks, Response
|
||||
from fastapi.exceptions import HTTPException
|
||||
from fastapi.responses import FileResponse
|
||||
import installer
|
||||
@@ -79,12 +79,12 @@ def post_log(req: models.ReqPostLog):
|
||||
log.debug(f'UI: {req.debug}')
|
||||
elif req.error is not None:
|
||||
log.error(f'UI: {req.error}')
|
||||
return {}
|
||||
return Response(status_code=204)
|
||||
|
||||
def post_shutdown():
|
||||
def post_shutdown(background_tasks: BackgroundTasks):
|
||||
log.info("Shutdown request received")
|
||||
import sys
|
||||
sys.exit(0)
|
||||
background_tasks.add_task(os._exit, 0)
|
||||
return Response(status_code=204)
|
||||
|
||||
def get_cmd_flags():
|
||||
return vars(shared.cmd_opts)
|
||||
@@ -126,10 +126,11 @@ def get_status():
|
||||
|
||||
def post_interrupt():
|
||||
shared.state.interrupt()
|
||||
return {}
|
||||
return Response(status_code=204)
|
||||
|
||||
def post_skip():
|
||||
shared.state.skip()
|
||||
return Response(status_code=204)
|
||||
|
||||
def get_memory():
|
||||
try:
|
||||
|
||||
@@ -557,7 +557,7 @@ def control_run(state: str = '', # pylint: disable=keyword-arg-before-vararg
|
||||
grading_clahe_clip=grading_clahe_clip, grading_clahe_grid=grading_clahe_grid,
|
||||
grading_shadows_tint=grading_shadows_tint, grading_highlights_tint=grading_highlights_tint, grading_split_tone_balance=grading_split_tone_balance,
|
||||
grading_vignette=grading_vignette, grading_grain=grading_grain,
|
||||
grading_lut_file=grading_lut_file.name if hasattr(grading_lut_file, 'name') else (grading_lut_file or ''), grading_lut_strength=grading_lut_strength,
|
||||
grading_lut_file=getattr(grading_lut_file, 'name', grading_lut_file) if grading_lut_file else '', grading_lut_strength=grading_lut_strength,
|
||||
# path
|
||||
outpath_samples=resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_control_samples),
|
||||
outpath_grids=resolve_output_path(shared.opts.outdir_grids, shared.opts.outdir_control_grids),
|
||||
|
||||
@@ -154,9 +154,9 @@ def guess_by_name(fn, current_guess):
|
||||
elif 'longcat-image' in fn.lower():
|
||||
new_guess = 'LongCat'
|
||||
elif 'ovis-image' in fn.lower():
|
||||
new_guess = 'Ovis-Image'
|
||||
new_guess = 'OvisImage'
|
||||
elif 'glm-image' in fn.lower():
|
||||
new_guess = 'GLM-Image'
|
||||
new_guess = 'GLMImage'
|
||||
elif 'sdxs-1b' in fn.lower():
|
||||
new_guess = 'SDXS'
|
||||
elif 'step1x-edit' in fn.lower():
|
||||
|
||||
@@ -351,9 +351,13 @@ def load_diffuser_force(detected_model_type: str, checkpoint_info: CheckpointInf
|
||||
from pipelines.model_segmoe import load_segmoe
|
||||
sd_model = load_segmoe(checkpoint_info, diffusers_load_config)
|
||||
allow_post_quant = True
|
||||
elif model_type in ['PixArt Sigma']:
|
||||
from pipelines.model_pixart import load_pixart
|
||||
sd_model = load_pixart(checkpoint_info, diffusers_load_config)
|
||||
elif model_type in ['PixArtAlpha']:
|
||||
from pipelines.model_pixart import load_pixart_alpha
|
||||
sd_model = load_pixart_alpha(checkpoint_info, diffusers_load_config)
|
||||
allow_post_quant = False
|
||||
elif model_type in ['PixArtSigma']:
|
||||
from pipelines.model_pixart import load_pixart_sigma
|
||||
sd_model = load_pixart_sigma(checkpoint_info, diffusers_load_config)
|
||||
allow_post_quant = False
|
||||
elif model_type in ['Sana']:
|
||||
from pipelines.model_sana import load_sana
|
||||
@@ -543,11 +547,11 @@ def load_diffuser_force(detected_model_type: str, checkpoint_info: CheckpointInf
|
||||
from pipelines.model_longcat import load_longcat
|
||||
sd_model = load_longcat(checkpoint_info, diffusers_load_config)
|
||||
allow_post_quant = False
|
||||
elif model_type in ['Ovis-Image', 'Overfit']:
|
||||
elif model_type in ['OvisImage']:
|
||||
from pipelines.model_ovis import load_ovis
|
||||
sd_model = load_ovis(checkpoint_info, diffusers_load_config)
|
||||
allow_post_quant = False
|
||||
elif model_type in ['GLM-Image']:
|
||||
elif model_type in ['GLMImage']:
|
||||
from pipelines.model_glm import load_glm_image
|
||||
sd_model = load_glm_image(checkpoint_info, diffusers_load_config)
|
||||
allow_post_quant = False
|
||||
|
||||
@@ -36,7 +36,7 @@ pipelines = {
|
||||
'FLUX': getattr(diffusers, 'FluxPipeline', None),
|
||||
'FLUX2Klein': getattr(diffusers, 'Flux2KleinPipeline', None),
|
||||
'FLUX2': getattr(diffusers, 'Flux2Pipeline', None),
|
||||
'GLM-Image': getattr(diffusers, 'GlmImagePipeline', None),
|
||||
'GLMImage': getattr(diffusers, 'GlmImagePipeline', None),
|
||||
'HiDream': getattr(diffusers, 'HiDreamImagePipeline', None),
|
||||
'HunyuanDiT': getattr(diffusers, 'HunyuanDiTPipeline', None),
|
||||
'HunyuanImage': getattr(diffusers, 'HunyuanImagePipeline', None),
|
||||
@@ -50,6 +50,7 @@ pipelines = {
|
||||
'Lumina2': getattr(diffusers, 'Lumina2Pipeline', None),
|
||||
'LuminaNext': getattr(diffusers, 'LuminaText2ImgPipeline', None),
|
||||
'NucleusImage': getattr(diffusers, 'NucleusMoEImagePipeline', None),
|
||||
'OvisImage': getattr(diffusers, 'OvisImagePipeline', None),
|
||||
'OmniGen': getattr(diffusers, 'OmniGenPipeline', None),
|
||||
'PixArtAlpha': getattr(diffusers, 'PixArtAlphaPipeline', None),
|
||||
'PixArtSigma': getattr(diffusers, 'PixArtSigmaPipeline', None),
|
||||
|
||||
@@ -6,7 +6,43 @@ from modules.logger import log
|
||||
from pipelines import generic
|
||||
|
||||
|
||||
def load_pixart(checkpoint_info, diffusers_load_config=None):
|
||||
def load_pixart_alpha(checkpoint_info, diffusers_load_config=None):
|
||||
if diffusers_load_config is None:
|
||||
diffusers_load_config = {}
|
||||
repo_id = sd_models.path_to_repo(checkpoint_info)
|
||||
sd_models.hf_auth_check(checkpoint_info)
|
||||
|
||||
repo_id_tenc = repo_id
|
||||
repo_id_pipe = repo_id
|
||||
|
||||
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False)
|
||||
log.debug(f'Load model: type=PixArtAlpha repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}')
|
||||
|
||||
from pipelines.pixart import PIXART_SPEC
|
||||
transformer = generic.load_transformer(repo_id, cls_name=diffusers.PixArtTransformer2DModel, load_config=diffusers_load_config, native_spec=PIXART_SPEC)
|
||||
text_encoder = generic.load_text_encoder(repo_id_tenc, cls_name=transformers.T5EncoderModel, load_config=diffusers_load_config)
|
||||
|
||||
if repo_id is None or repo_id.lower() == 'none':
|
||||
return None
|
||||
pipe = diffusers.PixArtAlphaPipeline.from_pretrained(
|
||||
repo_id_pipe,
|
||||
transformer=transformer,
|
||||
text_encoder=text_encoder,
|
||||
cache_dir=shared.opts.diffusers_dir,
|
||||
**load_args,
|
||||
)
|
||||
|
||||
generic.load_vae_override(pipe, diffusers_load_config)
|
||||
|
||||
del text_encoder
|
||||
del transformer
|
||||
sd_hijack_te.init_hijack(pipe)
|
||||
|
||||
devices.torch_gc(force=True, reason='load')
|
||||
return pipe
|
||||
|
||||
|
||||
def load_pixart_sigma(checkpoint_info, diffusers_load_config=None):
|
||||
if diffusers_load_config is None:
|
||||
diffusers_load_config = {}
|
||||
repo_id = sd_models.path_to_repo(checkpoint_info)
|
||||
|
||||
@@ -10,6 +10,9 @@ class PromptMatrixScript(scripts_manager.Script):
|
||||
def title(self):
|
||||
return "Prompt matrix"
|
||||
|
||||
def show(self, is_img2img): # pylint: disable=unused-argument
|
||||
return True
|
||||
|
||||
def ui(self, _is_img2img):
|
||||
with gr.Row():
|
||||
gr.HTML('<span">  Prompt matrix</span><br>')
|
||||
|
||||
@@ -97,6 +97,9 @@ class PromptsFromFileScript(scripts_manager.Script):
|
||||
def title(self):
|
||||
return "Prompts from file"
|
||||
|
||||
def show(self, is_img2img): # pylint: disable=unused-argument
|
||||
return True
|
||||
|
||||
def ui(self, _is_img2img):
|
||||
with gr.Row():
|
||||
gr.HTML('<span">  Prompt from file</span><br>')
|
||||
|
||||
@@ -32,6 +32,9 @@ class XYZGridScript(scripts_manager.Script):
|
||||
def title(self):
|
||||
return "XYZ Grid Script"
|
||||
|
||||
def show(self, is_img2img): # pylint: disable=unused-argument
|
||||
return True
|
||||
|
||||
def ui(self, is_img2img):
|
||||
self.current_axis_options = [x for x in axis_options if type(x) == AxisOption or x.is_img2img == is_img2img]
|
||||
with gr.Row():
|
||||
|
||||
Reference in New Issue
Block a user