diff --git a/modules/api/api.py b/modules/api/api.py index 57be1c8e4..bb7d78bd6 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -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 diff --git a/modules/api/docs.py b/modules/api/docs.py index 8f4db8cd1..78186f164 100644 --- a/modules/api/docs.py +++ b/modules/api/docs.py @@ -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', diff --git a/modules/api/endpoints.py b/modules/api/endpoints.py index c189e7104..5e767b6cc 100644 --- a/modules/api/endpoints.py +++ b/modules/api/endpoints.py @@ -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.""" diff --git a/modules/api/gallery.py b/modules/api/gallery.py index eaf22c319..b585ad896 100644 --- a/modules/api/gallery.py +++ b/modules/api/gallery.py @@ -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) diff --git a/modules/api/server.py b/modules/api/server.py index db1e75f8d..1dee86409 100644 --- a/modules/api/server.py +++ b/modules/api/server.py @@ -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: diff --git a/modules/control/run.py b/modules/control/run.py index 665afedba..22029b541 100644 --- a/modules/control/run.py +++ b/modules/control/run.py @@ -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), diff --git a/modules/sd_detect.py b/modules/sd_detect.py index 0717eed70..84d793876 100644 --- a/modules/sd_detect.py +++ b/modules/sd_detect.py @@ -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(): diff --git a/modules/sd_models.py b/modules/sd_models.py index 2bf3d0868..935e07361 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -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 diff --git a/modules/shared_items.py b/modules/shared_items.py index 8ca401be9..4fd642458 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -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), diff --git a/pipelines/model_pixart.py b/pipelines/model_pixart.py index b74680c6d..8c280ec79 100644 --- a/pipelines/model_pixart.py +++ b/pipelines/model_pixart.py @@ -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) diff --git a/scripts/prompt_matrix.py b/scripts/prompt_matrix.py index 7df52598e..b9876e5b2 100644 --- a/scripts/prompt_matrix.py +++ b/scripts/prompt_matrix.py @@ -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('  Prompt matrix
') diff --git a/scripts/prompts_from_file.py b/scripts/prompts_from_file.py index fb192253b..64474be6a 100644 --- a/scripts/prompts_from_file.py +++ b/scripts/prompts_from_file.py @@ -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('  Prompt from file
') diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index d86e22af6..8e0a3bda6 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -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():