mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
Merge branch 'dev' into add-noobaixl-controlnet
This commit is contained in:
+6
-3
@@ -5,7 +5,7 @@ from fastapi import FastAPI, APIRouter, Depends, Request
|
||||
from fastapi.security import HTTPBasic, HTTPBasicCredentials
|
||||
from fastapi.exceptions import HTTPException
|
||||
from modules import errors, shared, postprocessing
|
||||
from modules.api import models, endpoints, script, helpers, server, nvml, generate, process, control, gallery
|
||||
from modules.api import models, endpoints, script, helpers, server, nvml, generate, process, control, gallery, docs
|
||||
|
||||
|
||||
errors.install()
|
||||
@@ -23,8 +23,10 @@ class Api:
|
||||
for line in file.readlines():
|
||||
user, password = line.split(":")
|
||||
self.credentials[user.replace('"', '').strip()] = password.replace('"', '').strip()
|
||||
|
||||
self.router = APIRouter()
|
||||
if shared.cmd_opts.docs:
|
||||
docs.create_docs(app)
|
||||
docs.create_redocs(app)
|
||||
self.app = app
|
||||
self.queue_lock = queue_lock
|
||||
self.generate = generate.APIGenerate(queue_lock)
|
||||
@@ -36,6 +38,7 @@ class Api:
|
||||
self.add_api_route("/sdapi/v1/log", server.get_log_buffer, methods=["GET"], response_model=List[str])
|
||||
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/status", server.get_status, methods=["GET"], response_model=models.ResStatus)
|
||||
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/interrupt", server.post_interrupt, methods=["POST"])
|
||||
@@ -55,7 +58,7 @@ class Api:
|
||||
self.add_api_route("/sdapi/v1/extra-batch-images", self.extras_batch_images_api, methods=["POST"], response_model=models.ResProcessBatch)
|
||||
self.add_api_route("/sdapi/v1/preprocess", self.process.post_preprocess, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/mask", self.process.post_mask, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/faces", self.process.post_face, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/detect", self.process.post_detect, methods=["POST"])
|
||||
|
||||
# api dealing with optional scripts
|
||||
self.add_api_route("/sdapi/v1/scripts", script.get_scripts_list, methods=["GET"], response_model=models.ResScripts)
|
||||
|
||||
@@ -31,6 +31,7 @@ ReqControl = models.create_model_from_signature(
|
||||
{"key": "ip_adapter", "type": Optional[List[models.ItemIPAdapter]], "default": None, "exclude": True},
|
||||
{"key": "face", "type": Optional[models.ItemFace], "default": None, "exclude": True},
|
||||
{"key": "control", "type": Optional[List[ItemControl]], "default": [], "exclude": True},
|
||||
{"key": "extra", "type": Optional[dict], "default": {}, "exclude": True},
|
||||
]
|
||||
)
|
||||
|
||||
@@ -103,7 +104,7 @@ class APIControl():
|
||||
args['ip_adapter_scales'].append(ipadapter.scale)
|
||||
args['ip_adapter_starts'].append(ipadapter.start)
|
||||
args['ip_adapter_ends'].append(ipadapter.end)
|
||||
args['ip_adapter_crops'].append(ipadapter.end)
|
||||
args['ip_adapter_crops'].append(ipadapter.crop)
|
||||
args['ip_adapter_images'].append([helpers.decode_base64_to_image(x) for x in ipadapter.images])
|
||||
if ipadapter.masks:
|
||||
args['ip_adapter_masks'].append([helpers.decode_base64_to_image(x) for x in ipadapter.masks])
|
||||
@@ -159,6 +160,7 @@ class APIControl():
|
||||
output_processed = []
|
||||
output_info = ''
|
||||
run.control_set({ 'do_not_save_grid': not req.save_images, 'do_not_save_samples': not req.save_images, **self.prepare_ip_adapter(req) })
|
||||
run.control_set(getattr(req, "extra", {}))
|
||||
res = run.control_run(**args)
|
||||
for item in res:
|
||||
if len(item) > 0 and (isinstance(item[0], list) or item[0] is None): # output_images
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
import json
|
||||
from starlette.responses import HTMLResponse
|
||||
from fastapi import FastAPI
|
||||
from fastapi.openapi.docs import get_redoc_html, swagger_ui_default_parameters
|
||||
from fastapi.encoders import jsonable_encoder
|
||||
|
||||
|
||||
def get_swagger_ui_html(*,
|
||||
openapi_url: str,
|
||||
title: str,
|
||||
swagger_js_url: str = "https://cdn.jsdelivr.net/npm/swagger-ui-dist@5/swagger-ui-bundle.js",
|
||||
swagger_css_url: str = "https://cdn.jsdelivr.net/npm/swagger-ui-dist@5/swagger-ui.css",
|
||||
swagger_extra_css_url: str = None,
|
||||
swagger_favicon_url: str = "https://fastapi.tiangolo.com/img/favicon.png",
|
||||
oauth2_redirect_url: str = None,
|
||||
init_oauth: dict = None,
|
||||
swagger_ui_parameters: dict = None,
|
||||
) -> HTMLResponse:
|
||||
current_swagger_ui_parameters = swagger_ui_default_parameters.copy()
|
||||
if swagger_ui_parameters:
|
||||
current_swagger_ui_parameters.update(swagger_ui_parameters)
|
||||
html = f"""
|
||||
<!DOCTYPE html>
|
||||
<html>
|
||||
<head>
|
||||
<link type="text/css" rel="stylesheet" href="{swagger_css_url}">
|
||||
<link rel="shortcut icon" href="{swagger_favicon_url}">
|
||||
<title>{title}</title>
|
||||
</head>
|
||||
<body>
|
||||
<div id="swagger-ui"></div>
|
||||
<script src="{swagger_js_url}"></script>
|
||||
<script>
|
||||
const ui = SwaggerUIBundle({{
|
||||
url: '{openapi_url}',
|
||||
"""
|
||||
if swagger_extra_css_url is not None:
|
||||
html = html.replace('</head>', f'<link type="text/css" rel="stylesheet" href="{swagger_extra_css_url}"></head>')
|
||||
for key, value in current_swagger_ui_parameters.items():
|
||||
html += f"{json.dumps(key)}: {json.dumps(jsonable_encoder(value))},\n"
|
||||
if oauth2_redirect_url:
|
||||
html += f"oauth2RedirectUrl: window.location.origin + '{oauth2_redirect_url}',"
|
||||
html += """
|
||||
presets: [
|
||||
SwaggerUIBundle.presets.apis,
|
||||
SwaggerUIBundle.SwaggerUIStandalonePreset
|
||||
],
|
||||
})"""
|
||||
if init_oauth:
|
||||
html += f"""
|
||||
ui.initOAuth({json.dumps(jsonable_encoder(init_oauth))})
|
||||
"""
|
||||
html += """
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
"""
|
||||
return HTMLResponse(html)
|
||||
|
||||
|
||||
def create_docs(app: FastAPI):
|
||||
swagger_ui_parameters = {
|
||||
"displayOperationId": True,
|
||||
"layout": "BaseLayout",
|
||||
"showExtensions": True,
|
||||
"showCommonExtensions": True,
|
||||
"deepLinking": False,
|
||||
"dom_id": "#swagger-ui",
|
||||
}
|
||||
|
||||
@app.get("/docs", include_in_schema=True)
|
||||
async def custom_swagger_html():
|
||||
res = get_swagger_ui_html(
|
||||
title=f'{app.title}: Swagger UI',
|
||||
openapi_url=app.openapi_url,
|
||||
swagger_favicon_url='/file=html/favicon.svg',
|
||||
swagger_ui_parameters=swagger_ui_parameters,
|
||||
swagger_extra_css_url='file=html/swagger.css',
|
||||
)
|
||||
# res = inject_css(html.content, 'html/swagger.css')
|
||||
return res
|
||||
|
||||
|
||||
def create_redocs(app: FastAPI):
|
||||
@app.get("/redocs", include_in_schema=True)
|
||||
async def custom_redoc_html():
|
||||
res = get_redoc_html(
|
||||
title=f'{app.title}: ReDoc',
|
||||
openapi_url=app.openapi_url,
|
||||
redoc_favicon_url='/file=html/favicon.svg',
|
||||
)
|
||||
return res
|
||||
+13
-3
@@ -40,7 +40,7 @@ class APIGenerate():
|
||||
sanitize_str(request.script_args)
|
||||
|
||||
def prepare_face_module(self, request):
|
||||
if hasattr(request, "face") and request.face and not request.script_name and (not request.alwayson_scripts or "face" not in request.alwayson_scripts.keys()):
|
||||
if getattr(request, "face", None) is not None and (not request.alwayson_scripts or "face" not in request.alwayson_scripts.keys()):
|
||||
request.script_name = "face"
|
||||
request.script_args = [
|
||||
request.face.mode,
|
||||
@@ -106,6 +106,8 @@ class APIGenerate():
|
||||
p.scripts = script_runner
|
||||
p.outpath_grids = shared.opts.outdir_grids or shared.opts.outdir_txt2img_grids
|
||||
p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_txt2img_samples
|
||||
for key, value in getattr(txt2imgreq, "extra", {}).items():
|
||||
setattr(p, key, value)
|
||||
shared.state.begin('API TXT', api=True)
|
||||
script_args = script.init_script_args(p, txt2imgreq, self.default_script_arg_txt2img, selectable_scripts, selectable_script_idx, script_runner)
|
||||
if selectable_scripts is not None:
|
||||
@@ -114,7 +116,10 @@ class APIGenerate():
|
||||
p.script_args = tuple(script_args) # Need to pass args as tuple here
|
||||
processed = process_images(p)
|
||||
shared.state.end(api=False)
|
||||
b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else []
|
||||
if processed.images is None or len(processed.images) == 0:
|
||||
b64images = []
|
||||
else:
|
||||
b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else []
|
||||
self.sanitize_b64(txt2imgreq)
|
||||
return models.ResTxt2Img(images=b64images, parameters=vars(txt2imgreq), info=processed.js())
|
||||
|
||||
@@ -150,6 +155,8 @@ class APIGenerate():
|
||||
p.scripts = script_runner
|
||||
p.outpath_grids = shared.opts.outdir_img2img_grids
|
||||
p.outpath_samples = shared.opts.outdir_img2img_samples
|
||||
for key, value in getattr(img2imgreq, "extra", {}).items():
|
||||
setattr(p, key, value)
|
||||
shared.state.begin('API-IMG', api=True)
|
||||
script_args = script.init_script_args(p, img2imgreq, self.default_script_arg_img2img, selectable_scripts, selectable_script_idx, script_runner)
|
||||
if selectable_scripts is not None:
|
||||
@@ -158,7 +165,10 @@ class APIGenerate():
|
||||
p.script_args = tuple(script_args) # Need to pass args as tuple here
|
||||
processed = process_images(p)
|
||||
shared.state.end(api=False)
|
||||
b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else []
|
||||
if processed.images is None or len(processed.images) == 0:
|
||||
b64images = []
|
||||
else:
|
||||
b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else []
|
||||
if not img2imgreq.include_init_images:
|
||||
img2imgreq.init_images = None
|
||||
img2imgreq.mask = None
|
||||
|
||||
@@ -14,7 +14,7 @@ def validate_sampler_name(name):
|
||||
return name
|
||||
|
||||
|
||||
def decode_base64_to_image(encoding):
|
||||
def decode_base64_to_image(encoding, quiet=False):
|
||||
if encoding.startswith("data:image/"):
|
||||
encoding = encoding.split(";")[1].split(",")[1]
|
||||
try:
|
||||
@@ -22,7 +22,9 @@ def decode_base64_to_image(encoding):
|
||||
return image
|
||||
except Exception as e:
|
||||
shared.log.warning(f'API cannot decode image: {e}')
|
||||
raise HTTPException(status_code=500, detail="Invalid encoded image") from e
|
||||
if not quiet:
|
||||
raise HTTPException(status_code=500, detail="Invalid encoded image") from e
|
||||
return None
|
||||
|
||||
|
||||
def encode_pil_to_base64(image):
|
||||
|
||||
@@ -90,4 +90,4 @@ def setup_middleware(app: FastAPI, cmd_opts):
|
||||
return handle_exception(req, e)
|
||||
|
||||
app.build_middleware_stack() # rebuild middleware stack on-the-fly
|
||||
log.debug(f'FastAPI middleware: {[m.__class__.__name__ for m in app.user_middleware]}')
|
||||
log.debug(f'API middleware: {[m.cls for m in app.user_middleware]}')
|
||||
|
||||
+28
-11
@@ -11,14 +11,6 @@ API_NOT_ALLOWED = [
|
||||
"sd_model",
|
||||
"outpath_samples",
|
||||
"outpath_grids",
|
||||
"sampler_index",
|
||||
"extra_generation_params",
|
||||
"overlay_images",
|
||||
"do_not_reload_embeddings",
|
||||
"seed_enable_extras",
|
||||
"prompt_for_display",
|
||||
"sampler_noise_scheduler_override",
|
||||
"ddim_discretize"
|
||||
]
|
||||
|
||||
class ModelDef(BaseModel):
|
||||
@@ -202,14 +194,17 @@ ReqTxt2Img = PydanticModelGenerator(
|
||||
"StableDiffusionProcessingTxt2Img",
|
||||
StableDiffusionProcessingTxt2Img,
|
||||
[
|
||||
{"key": "sampler_index", "type": str, "default": "UniPC"},
|
||||
{"key": "script_name", "type": str, "default": None},
|
||||
{"key": "sampler_index", "type": int, "default": 0},
|
||||
{"key": "sampler_name", "type": str, "default": "UniPC"},
|
||||
{"key": "hr_sampler_name", "type": str, "default": "Same as primary"},
|
||||
{"key": "script_name", "type": str, "default": "none"},
|
||||
{"key": "script_args", "type": list, "default": []},
|
||||
{"key": "send_images", "type": bool, "default": True},
|
||||
{"key": "save_images", "type": bool, "default": False},
|
||||
{"key": "alwayson_scripts", "type": dict, "default": {}},
|
||||
{"key": "ip_adapter", "type": Optional[List[ItemIPAdapter]], "default": None, "exclude": True},
|
||||
{"key": "face", "type": Optional[ItemFace], "default": None, "exclude": True},
|
||||
{"key": "extra", "type": Optional[dict], "default": {}, "exclude": True},
|
||||
]
|
||||
).generate_model()
|
||||
StableDiffusionTxt2ImgProcessingAPI = ReqTxt2Img
|
||||
@@ -223,7 +218,11 @@ ReqImg2Img = PydanticModelGenerator(
|
||||
"StableDiffusionProcessingImg2Img",
|
||||
StableDiffusionProcessingImg2Img,
|
||||
[
|
||||
{"key": "sampler_index", "type": str, "default": "UniPC"},
|
||||
{"key": "sampler_index", "type": int, "default": 0},
|
||||
{"key": "sampler_name", "type": str, "default": "UniPC"},
|
||||
{"key": "hr_sampler_name", "type": str, "default": "Same as primary"},
|
||||
{"key": "script_name", "type": str, "default": "none"},
|
||||
{"key": "script_args", "type": list, "default": []},
|
||||
{"key": "init_images", "type": list, "default": None},
|
||||
{"key": "denoising_strength", "type": float, "default": 0.5},
|
||||
{"key": "mask", "type": str, "default": None},
|
||||
@@ -235,6 +234,7 @@ ReqImg2Img = PydanticModelGenerator(
|
||||
{"key": "alwayson_scripts", "type": dict, "default": {}},
|
||||
{"key": "ip_adapter", "type": Optional[List[ItemIPAdapter]], "default": None, "exclude": True},
|
||||
{"key": "face_id", "type": Optional[ItemFace], "default": None, "exclude": True},
|
||||
{"key": "extra", "type": Optional[dict], "default": {}, "exclude": True},
|
||||
]
|
||||
).generate_model()
|
||||
StableDiffusionImg2ImgProcessingAPI = ReqImg2Img
|
||||
@@ -300,6 +300,23 @@ class ResProgress(BaseModel):
|
||||
current_image: str = Field(default=None, title="Current image", description="The current image in base64 format. opts.show_progress_every_n_steps is required for this to work.")
|
||||
textinfo: str = Field(default=None, title="Info text", description="Info text used by WebUI.")
|
||||
|
||||
class ResStatus(BaseModel):
|
||||
status: str = Field(title="Status", description="Current status")
|
||||
task: str = Field(title="Task", description="Current task")
|
||||
timestamp: Optional[str] = Field(title="Timestamp", description="Timestamp of the current job")
|
||||
id: str = Field(title="ID", description="ID of the current task")
|
||||
job: int = Field(title="Job", description="Current job")
|
||||
jobs: int = Field(title="Jobs", description="Total jobs")
|
||||
total: int = Field(title="Total Jobs", description="Total jobs")
|
||||
step: int = Field(title="Step", description="Current step")
|
||||
steps: int = Field(title="Steps", description="Total steps")
|
||||
queued: int = Field(title="Queued", description="Number of queued tasks")
|
||||
uptime: int = Field(title="Uptime", description="Uptime of the server")
|
||||
elapsed: Optional[float] = Field(title="Elapsed time")
|
||||
eta: Optional[float] = Field(title="ETA in secs")
|
||||
progress: Optional[float] = Field(title="Progress", description="The progress with a range of 0 to 1")
|
||||
|
||||
|
||||
class ReqInterrogate(BaseModel):
|
||||
image: str = Field(default="", title="Image", description="Image to work on, must be a Base64 string containing the image's data.")
|
||||
clip_model: str = Field(default="", title="CLiP Model", description="The interrogate model used.")
|
||||
|
||||
+17
-7
@@ -28,8 +28,12 @@ class ReqMask(BaseModel):
|
||||
|
||||
class ReqFace(BaseModel):
|
||||
image: str = Field(title="Image", description="The base64 encoded image")
|
||||
model: Optional[str] = Field(title="Model", description="The model to use for detection")
|
||||
|
||||
class ResFace(BaseModel):
|
||||
classes: List[int] = Field(title="Class", description="The class of detected item")
|
||||
labels: List[str] = Field(title="Label", description="The label of detected item")
|
||||
boxes: List[List[int]] = Field(title="Box", description="The bounding box of detected item")
|
||||
images: List[str] = Field(title="Image", description="The base64 encoded images of detected faces")
|
||||
scores: List[float] = Field(title="Scores", description="The scores of the detected faces")
|
||||
|
||||
@@ -106,16 +110,22 @@ class APIProcess():
|
||||
image = encode_pil_to_base64(processed)
|
||||
return ResMask(mask=image)
|
||||
|
||||
def post_face(self, req: ReqFace):
|
||||
from shared import yolo # pylint: disable=no-name-in-module
|
||||
def post_detect(self, req: ReqFace):
|
||||
from modules.shared import yolo # pylint: disable=no-name-in-module
|
||||
image = decode_base64_to_image(req.image)
|
||||
shared.state.begin('API-FACE', api=True)
|
||||
images = []
|
||||
scores = []
|
||||
classes = []
|
||||
boxes = []
|
||||
labels = []
|
||||
with self.queue_lock:
|
||||
faces = yolo.predict('face-yolo8n', image)
|
||||
for face in faces:
|
||||
images.append(encode_pil_to_base64(face.item))
|
||||
scores.append(face.score)
|
||||
items = yolo.predict(req.model, image)
|
||||
for item in items:
|
||||
images.append(encode_pil_to_base64(item.item))
|
||||
scores.append(item.score)
|
||||
classes.append(item.cls)
|
||||
labels.append(item.label)
|
||||
boxes.append(item.box)
|
||||
shared.state.end(api=False)
|
||||
return ResFace(images=images, scores=scores)
|
||||
return ResFace(classes=classes, labels=labels, scores=scores, boxes=boxes, images=images)
|
||||
|
||||
+23
-6
@@ -3,27 +3,39 @@ from fastapi.exceptions import HTTPException
|
||||
import gradio as gr
|
||||
from modules.api import models
|
||||
from modules import scripts
|
||||
from modules.errors import log
|
||||
|
||||
|
||||
def script_name_to_index(name, scripts_list):
|
||||
try:
|
||||
return [script.title().lower() for script in scripts_list].index(name.lower())
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=422, detail=f"Script '{name}' not found") from e
|
||||
if name is None or len(name) == 0 or name == 'none':
|
||||
return None
|
||||
available = [script.title().lower() for script in scripts_list]
|
||||
if name.lower() in available:
|
||||
return available.index(name.lower())
|
||||
short = [available.split(':')[0] for available in available]
|
||||
if name.lower() in short:
|
||||
return short.index(name.lower())
|
||||
log.error(f'API: script={name} available={available} not found')
|
||||
return None
|
||||
|
||||
|
||||
def get_selectable_script(script_name, script_runner):
|
||||
if script_name is None or script_name == "":
|
||||
if script_name is None or script_name == "" or script_name == 'none':
|
||||
return None, None
|
||||
script_idx = script_name_to_index(script_name, script_runner.selectable_scripts)
|
||||
if script_idx is None:
|
||||
return None, None
|
||||
script = script_runner.selectable_scripts[script_idx]
|
||||
return script, script_idx
|
||||
|
||||
|
||||
def get_scripts_list():
|
||||
t2ilist = [script.name for script in scripts.scripts_txt2img.scripts if script.name is not None]
|
||||
i2ilist = [script.name for script in scripts.scripts_img2img.scripts if script.name is not None]
|
||||
control = [script.name for script in scripts.scripts_control.scripts if script.name is not None]
|
||||
return models.ResScripts(txt2img = t2ilist, img2img = i2ilist, control = control)
|
||||
|
||||
|
||||
def get_script_info(script_name: Optional[str] = None):
|
||||
res = []
|
||||
for script_list in [scripts.scripts_txt2img.scripts, scripts.scripts_img2img.scripts, scripts.scripts_control.scripts]:
|
||||
@@ -32,12 +44,16 @@ def get_script_info(script_name: Optional[str] = None):
|
||||
res.append(script.api_info)
|
||||
return res
|
||||
|
||||
|
||||
def get_script(script_name, script_runner):
|
||||
if script_name is None or script_name == "":
|
||||
if script_name is None or script_name == "" or script_name == 'none':
|
||||
return None, None
|
||||
script_idx = script_name_to_index(script_name, script_runner.scripts)
|
||||
if script_idx is None:
|
||||
return None
|
||||
return script_runner.scripts[script_idx]
|
||||
|
||||
|
||||
def init_default_script_args(script_runner):
|
||||
# find max idx from the scripts in runner and generate a none array to init script_args
|
||||
last_arg_index = 1
|
||||
@@ -60,6 +76,7 @@ def init_default_script_args(script_runner):
|
||||
script_args[script.args_from:script.args_to] = ui_default_values
|
||||
return script_args
|
||||
|
||||
|
||||
def init_script_args(p, request, default_script_args, selectable_scripts, selectable_script_idx, script_runner):
|
||||
script_args = default_script_args.copy()
|
||||
# position 0 in script_arg is the idx+1 of the selectable script that is going to be run when using scripts.scripts_*2img.run()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import time
|
||||
from typing import Any, Dict
|
||||
from fastapi import Depends
|
||||
from modules import shared
|
||||
@@ -66,7 +67,6 @@ def get_cmd_flags():
|
||||
return vars(shared.cmd_opts)
|
||||
|
||||
def get_progress(req: models.ReqProgress = Depends()):
|
||||
import time
|
||||
if shared.state.job_count == 0:
|
||||
return models.ResProgress(progress=0, eta_relative=0, state=shared.state.dict(), textinfo=shared.state.textinfo)
|
||||
shared.state.do_set_current_image()
|
||||
@@ -85,6 +85,9 @@ def get_progress(req: models.ReqProgress = Depends()):
|
||||
res = models.ResProgress(progress=progress, eta_relative=eta_relative, state=shared.state.dict(), current_image=current_image, textinfo=shared.state.textinfo)
|
||||
return res
|
||||
|
||||
def get_status():
|
||||
return shared.state.status()
|
||||
|
||||
def post_interrupt():
|
||||
shared.state.interrupt()
|
||||
return {}
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
original code from <https://github.com/NVlabs/consistory>
|
||||
"""
|
||||
from .consistory_pipeline import ConsistoryExtendAttnSDXLPipeline
|
||||
from .consistory_unet_sdxl import ConsistorySDXLUNet2DConditionModel
|
||||
from .consistory_run import run_anchor_generation, run_extra_generation
|
||||
@@ -0,0 +1,287 @@
|
||||
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# Not a contribution
|
||||
# Changes made by NVIDIA CORPORATION & AFFILIATES enabling ConsiStory or otherwise documented as NVIDIA-proprietary
|
||||
# are not a contribution and subject to the license under the LICENSE file located at the root directory.
|
||||
|
||||
|
||||
from typing import Callable, Optional
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.utils import USE_PEFT_BACKEND
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from .consistory_utils import AnchorCache, FeatureInjector, QueryStore
|
||||
|
||||
|
||||
class ConsistoryAttnStoreProcessor:
|
||||
def __init__(self, attnstore, place_in_unet):
|
||||
super().__init__()
|
||||
self.attnstore = attnstore
|
||||
self.place_in_unet = place_in_unet
|
||||
|
||||
def __call__(self, attn: Attention, hidden_states, encoder_hidden_states=None, attention_mask=None, record_attention=True, **kwargs):
|
||||
batch_size, sequence_length, _ = hidden_states.shape
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
is_cross = encoder_hidden_states is not None
|
||||
encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
|
||||
# only need to store attention maps during the Attend and Excite process
|
||||
# if attention_probs.requires_grad:
|
||||
if record_attention:
|
||||
self.attnstore(attention_probs, is_cross, self.place_in_unet, attn.heads)
|
||||
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class ConsistoryExtendedAttnXFormersAttnProcessor:
|
||||
r"""
|
||||
Processor for implementing memory efficient attention using xFormers.
|
||||
|
||||
Args:
|
||||
attention_op (`Callable`, *optional*, defaults to `None`):
|
||||
The base
|
||||
[operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to
|
||||
use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best
|
||||
operator.
|
||||
"""
|
||||
|
||||
def __init__(self, place_in_unet, attnstore, extended_attn_kwargs, attention_op: Optional[Callable] = None):
|
||||
self.attention_op = attention_op
|
||||
self.t_range = extended_attn_kwargs.get('t_range', [])
|
||||
self.extend_kv_unet_parts = extended_attn_kwargs.get('extend_kv_unet_parts', ['down', 'mid', 'up'])
|
||||
|
||||
self.place_in_unet = place_in_unet
|
||||
self.curr_unet_part = self.place_in_unet.split('_')[0]
|
||||
self.attnstore = attnstore
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: Optional[torch.FloatTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
temb: Optional[torch.FloatTensor] = None,
|
||||
scale: float = 1.0,
|
||||
perform_extend_attn: bool = False,
|
||||
query_store: Optional[QueryStore] = None,
|
||||
feature_injector: Optional[FeatureInjector] = None,
|
||||
anchors_cache: Optional[AnchorCache] = None,
|
||||
**kwargs
|
||||
) -> torch.FloatTensor:
|
||||
residual = hidden_states
|
||||
|
||||
args = () if USE_PEFT_BACKEND else (scale,)
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
else:
|
||||
batch_size, wh, channel = hidden_states.shape
|
||||
height = width = int(wh ** 0.5)
|
||||
|
||||
is_cross = encoder_hidden_states is not None
|
||||
perform_extend_attn = perform_extend_attn and (not is_cross) and \
|
||||
any([self.attnstore.curr_iter >= x[0] and self.attnstore.curr_iter <= x[1] for x in self.t_range]) and \
|
||||
self.curr_unet_part in self.extend_kv_unet_parts
|
||||
|
||||
batch_size, key_tokens, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, key_tokens, batch_size)
|
||||
if attention_mask is not None:
|
||||
# expand our mask's singleton query_tokens dimension:
|
||||
# [batch*heads, 1, key_tokens] ->
|
||||
# [batch*heads, query_tokens, key_tokens]
|
||||
# so that it can be added as a bias onto the attention scores that xformers computes:
|
||||
# [batch*heads, query_tokens, key_tokens]
|
||||
# we do this explicitly because xformers doesn't broadcast the singleton dimension for us.
|
||||
_, query_tokens, _ = hidden_states.shape
|
||||
attention_mask = attention_mask.expand(-1, query_tokens, -1)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states, *args)
|
||||
|
||||
if (self.curr_unet_part in self.extend_kv_unet_parts) and query_store and query_store.mode == 'cache':
|
||||
query_store.cache_query(query, self.place_in_unet)
|
||||
elif perform_extend_attn and query_store and query_store.mode == 'inject':
|
||||
query = query_store.inject_query(query, self.place_in_unet, self.attnstore.curr_iter)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states, *args)
|
||||
value = attn.to_v(encoder_hidden_states, *args)
|
||||
|
||||
query = attn.head_to_batch_dim(query).contiguous()
|
||||
|
||||
if perform_extend_attn:
|
||||
# Anchor Caching
|
||||
if anchors_cache and anchors_cache.is_cache_mode():
|
||||
if self.place_in_unet not in anchors_cache.input_h_cache:
|
||||
anchors_cache.input_h_cache[self.place_in_unet] = {}
|
||||
|
||||
# Hidden states inside the mask, for uncond (index 0) and cond (index 1) prompts
|
||||
subjects_hidden_states = torch.stack([x[self.attnstore.last_mask_dropout[width]] for x in hidden_states.chunk(2)])
|
||||
anchors_cache.input_h_cache[self.place_in_unet][self.attnstore.curr_iter] = subjects_hidden_states
|
||||
|
||||
if anchors_cache and anchors_cache.is_inject_mode():
|
||||
# We make extended key and value by concatenating the original key and value with the query.
|
||||
anchors_hidden_states = anchors_cache.input_h_cache[self.place_in_unet][self.attnstore.curr_iter]
|
||||
|
||||
anchors_keys = attn.to_k(anchors_hidden_states, *args)
|
||||
anchors_values = attn.to_v(anchors_hidden_states, *args)
|
||||
|
||||
extended_key = torch.cat([torch.cat([key.chunk(2, dim=0)[x], anchors_keys[x].unsqueeze(0)], dim=1) for x in range(2)])
|
||||
extended_value = torch.cat([torch.cat([value.chunk(2, dim=0)[x], anchors_values[x].unsqueeze(0)], dim=1) for x in range(2)])
|
||||
|
||||
extended_key = attn.head_to_batch_dim(extended_key).contiguous()
|
||||
extended_value = attn.head_to_batch_dim(extended_value).contiguous()
|
||||
|
||||
# attn_masks needs to be of shape [batch_size, query_tokens, key_tokens]
|
||||
# hidden_states = xformers.ops.memory_efficient_attention(query, extended_key, extended_value, op=self.attention_op, scale=attn.scale)
|
||||
hidden_states = F.scaled_dot_product_attention(query, extended_key, extended_value, scale=attn.scale)
|
||||
else:
|
||||
# # We make extended key and value by concatenating the original key and value with the query.
|
||||
# attention_mask_bias = self.attnstore.get_attn_mask_bias(tgt_size = width, bsz = batch_size)
|
||||
|
||||
# if attention_mask_bias is not None:
|
||||
# attention_mask_bias = torch.cat([x.unsqueeze(0).expand(attn.heads, -1, -1) for x in attention_mask_bias])
|
||||
|
||||
# Pre-allocate the output tensor
|
||||
ex_out = torch.empty_like(query)
|
||||
|
||||
for i in range(batch_size):
|
||||
start_idx = i * attn.heads
|
||||
end_idx = start_idx + attn.heads
|
||||
|
||||
attention_mask = self.attnstore.get_extended_attn_mask_instance(width, i%(batch_size//2))
|
||||
|
||||
curr_q = query[start_idx:end_idx]
|
||||
|
||||
if i < batch_size//2:
|
||||
curr_k = key[:batch_size//2]
|
||||
curr_v = value[:batch_size//2]
|
||||
else:
|
||||
curr_k = key[batch_size//2:]
|
||||
curr_v = value[batch_size//2:]
|
||||
|
||||
curr_k = curr_k.flatten(0,1)[attention_mask].unsqueeze(0)
|
||||
curr_v = curr_v.flatten(0,1)[attention_mask].unsqueeze(0)
|
||||
|
||||
curr_k = attn.head_to_batch_dim(curr_k).contiguous()
|
||||
curr_v = attn.head_to_batch_dim(curr_v).contiguous()
|
||||
|
||||
# hidden_states = xformers.ops.memory_efficient_attention(curr_q, curr_k, curr_v, op=self.attention_op, scale=attn.scale)
|
||||
hidden_states = F.scaled_dot_product_attention(curr_q, curr_k, curr_v, scale=attn.scale)
|
||||
|
||||
ex_out[start_idx:end_idx] = hidden_states
|
||||
|
||||
hidden_states = ex_out
|
||||
else:
|
||||
key = attn.head_to_batch_dim(key).contiguous()
|
||||
value = attn.head_to_batch_dim(value).contiguous()
|
||||
|
||||
# attn_masks needs to be of shape [batch_size, query_tokens, key_tokens]
|
||||
# hidden_states = xformers.ops.memory_efficient_attention(query, key, value, op=self.attention_op, scale=attn.scale)
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, scale=attn.scale)
|
||||
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states, *args)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if feature_injector is not None:
|
||||
output_res = int(hidden_states.shape[1] ** 0.5)
|
||||
|
||||
if anchors_cache and anchors_cache.is_inject_mode():
|
||||
hidden_states[batch_size//2:] = feature_injector.inject_anchors(hidden_states[batch_size//2:], self.attnstore.curr_iter, output_res, self.attnstore.extended_mapping, self.place_in_unet, anchors_cache)
|
||||
else:
|
||||
hidden_states[batch_size//2:] = feature_injector.inject_outputs(hidden_states[batch_size//2:], self.attnstore.curr_iter, output_res, self.attnstore.extended_mapping, self.place_in_unet, anchors_cache)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
def register_extended_self_attn(unet, attnstore, extended_attn_kwargs):
|
||||
DICT_PLACE_TO_RES = {'down_0': 64, 'down_1': 64, 'down_2': 64, 'down_3': 64, 'down_4': 64, 'down_5': 64, 'down_6': 64, 'down_7': 64,
|
||||
'down_8': 32, 'down_9': 32, 'down_10': 32, 'down_11': 32, 'down_12': 32, 'down_13': 32, 'down_14': 32, 'down_15': 32,
|
||||
'down_16': 32, 'down_17': 32, 'down_18': 32, 'down_19': 32, 'down_20': 32, 'down_21': 32, 'down_22': 32, 'down_23': 32,
|
||||
'down_24': 32, 'down_25': 32, 'down_26': 32, 'down_27': 32, 'down_28': 32, 'down_29': 32, 'down_30': 32, 'down_31': 32,
|
||||
'down_32': 32, 'down_33': 32, 'down_34': 32, 'down_35': 32, 'down_36': 32, 'down_37': 32, 'down_38': 32, 'down_39': 32,
|
||||
'down_40': 32, 'down_41': 32, 'down_42': 32, 'down_43': 32, 'down_44': 32, 'down_45': 32, 'down_46': 32, 'down_47': 32,
|
||||
'mid_120': 32, 'mid_121': 32, 'mid_122': 32, 'mid_123': 32, 'mid_124': 32, 'mid_125': 32, 'mid_126': 32, 'mid_127': 32,
|
||||
'mid_128': 32, 'mid_129': 32, 'mid_130': 32, 'mid_131': 32, 'mid_132': 32, 'mid_133': 32, 'mid_134': 32, 'mid_135': 32,
|
||||
'mid_136': 32, 'mid_137': 32, 'mid_138': 32, 'mid_139': 32, 'up_49': 32, 'up_51': 32, 'up_53': 32, 'up_55': 32, 'up_57': 32,
|
||||
'up_59': 32, 'up_61': 32, 'up_63': 32, 'up_65': 32, 'up_67': 32, 'up_69': 32, 'up_71': 32, 'up_73': 32, 'up_75': 32,
|
||||
'up_77': 32, 'up_79': 32, 'up_81': 32, 'up_83': 32, 'up_85': 32, 'up_87': 32, 'up_89': 32, 'up_91': 32, 'up_93': 32,
|
||||
'up_95': 32, 'up_97': 32, 'up_99': 32, 'up_101': 32, 'up_103': 32, 'up_105': 32, 'up_107': 32, 'up_109': 64, 'up_111': 64,
|
||||
'up_113': 64, 'up_115': 64, 'up_117': 64, 'up_119': 64}
|
||||
attn_procs = {}
|
||||
for i, name in enumerate(unet.attn_processors.keys()):
|
||||
is_self_attn = i % 2 == 0
|
||||
if name.startswith("mid_block"):
|
||||
place_in_unet = f"mid_{i}"
|
||||
elif name.startswith("up_blocks"):
|
||||
place_in_unet = f"up_{i}"
|
||||
elif name.startswith("down_blocks"):
|
||||
place_in_unet = f"down_{i}"
|
||||
else:
|
||||
continue
|
||||
|
||||
if is_self_attn:
|
||||
attn_procs[name] = ConsistoryExtendedAttnXFormersAttnProcessor(place_in_unet, attnstore, extended_attn_kwargs)
|
||||
else:
|
||||
attn_procs[name] = ConsistoryAttnStoreProcessor(attnstore, place_in_unet)
|
||||
|
||||
unet.set_attn_processor(attn_procs)
|
||||
@@ -0,0 +1,519 @@
|
||||
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# Not a contribution
|
||||
# Changes made by NVIDIA CORPORATION & AFFILIATES enabling ConsiStory or otherwise documented as NVIDIA-proprietary
|
||||
# are not a contribution and subject to the license under the LICENSE file located at the root directory.
|
||||
|
||||
import torch
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import StableDiffusionXLPipeline, \
|
||||
rescale_noise_cfg, EXAMPLE_DOC_STRING
|
||||
from diffusers.utils import (
|
||||
deprecate,
|
||||
is_torch_xla_available,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
)
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from .attention_processor import register_extended_self_attn
|
||||
from .consistory_utils import FeatureInjector, AnchorCache, QueryStore
|
||||
from .utils.ptp_utils import AttentionStore
|
||||
|
||||
if is_torch_xla_available():
|
||||
# import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
T = torch.Tensor
|
||||
|
||||
class ConsistoryExtendAttnSDXLPipeline(
|
||||
StableDiffusionXLPipeline
|
||||
):
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_inference_steps: int = 50,
|
||||
denoising_end: Optional[float] = None,
|
||||
guidance_scale: float = 5.0,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
negative_prompt_2: Optional[Union[str, List[str]]] = None,
|
||||
num_images_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
guidance_rescale: float = 0.0,
|
||||
original_size: Optional[Tuple[int, int]] = None,
|
||||
crops_coords_top_left: Tuple[int, int] = (0, 0),
|
||||
target_size: Optional[Tuple[int, int]] = None,
|
||||
negative_original_size: Optional[Tuple[int, int]] = None,
|
||||
negative_crops_coords_top_left: Tuple[int, int] = (0, 0),
|
||||
negative_target_size: Optional[Tuple[int, int]] = None,
|
||||
clip_skip: Optional[int] = None,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
|
||||
attention_store_kwargs: Optional[Dict] = None,
|
||||
extended_attn_kwargs: Optional[Dict] = None,
|
||||
share_queries: bool = False,
|
||||
query_store_kwargs: Optional[Dict] = {},
|
||||
feature_injector: Optional[FeatureInjector] = None,
|
||||
anchors_cache: Optional[AnchorCache] = None,
|
||||
|
||||
instance_latents: Optional[torch.FloatTensor] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
used in both text-encoders
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The height in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
Anything below 512 pixels won't work well for
|
||||
[stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0)
|
||||
and checkpoints that are not specifically fine-tuned on low resolutions.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The width in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
Anything below 512 pixels won't work well for
|
||||
[stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0)
|
||||
and checkpoints that are not specifically fine-tuned on low resolutions.
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
denoising_end (`float`, *optional*):
|
||||
When specified, determines the fraction (between 0.0 and 1.0) of the total denoising process to be
|
||||
completed before it is intentionally prematurely terminated. As a result, the returned sample will
|
||||
still retain a substantial amount of noise as determined by the discrete timesteps selected by the
|
||||
scheduler. The denoising_end parameter should ideally be utilized when this pipeline forms a part of a
|
||||
"Mixture of Denoisers" multi-pipeline setup, as elaborated in [**Refining the Image
|
||||
Output**](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_xl#refining-the-image-output)
|
||||
guidance_scale (`float`, *optional*, defaults to 5.0):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||
less than `1`).
|
||||
negative_prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and
|
||||
`text_encoder_2`. If not defined, `negative_prompt` is used in both text-encoders
|
||||
num_images_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
eta (`float`, *optional*, defaults to 0.0):
|
||||
Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to
|
||||
[`schedulers.DDIMScheduler`], will be ignored for others.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting.
|
||||
If not provided, pooled text embeddings will be generated from `prompt` input argument.
|
||||
negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt`
|
||||
input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] instead
|
||||
of a plain tuple.
|
||||
cross_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
guidance_rescale (`float`, *optional*, defaults to 0.0):
|
||||
Guidance rescale factor proposed by [Common Diffusion Noise Schedules and Sample Steps are
|
||||
Flawed](https://arxiv.org/pdf/2305.08891.pdf) `guidance_scale` is defined as `φ` in equation 16. of
|
||||
[Common Diffusion Noise Schedules and Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf).
|
||||
Guidance rescale factor should fix overexposure when using zero terminal SNR.
|
||||
original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||
If `original_size` is not the same as `target_size` the image will appear to be down- or upsampled.
|
||||
`original_size` defaults to `(height, width)` if not specified. Part of SDXL's micro-conditioning as
|
||||
explained in section 2.2 of
|
||||
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||
crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)):
|
||||
`crops_coords_top_left` can be used to generate an image that appears to be "cropped" from the position
|
||||
`crops_coords_top_left` downwards. Favorable, well-centered images are usually achieved by setting
|
||||
`crops_coords_top_left` to (0, 0). Part of SDXL's micro-conditioning as explained in section 2.2 of
|
||||
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||
target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||
For most cases, `target_size` should be set to the desired height and width of the generated image. If
|
||||
not specified it will default to `(height, width)`. Part of SDXL's micro-conditioning as explained in
|
||||
section 2.2 of [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952).
|
||||
negative_original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||
To negatively condition the generation process based on a specific image resolution. Part of SDXL's
|
||||
micro-conditioning as explained in section 2.2 of
|
||||
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more
|
||||
information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208.
|
||||
negative_crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)):
|
||||
To negatively condition the generation process based on a specific crop coordinates. Part of SDXL's
|
||||
micro-conditioning as explained in section 2.2 of
|
||||
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more
|
||||
information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208.
|
||||
negative_target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)):
|
||||
To negatively condition the generation process based on a target image resolution. It should be as same
|
||||
as the `target_size` for most cases. Part of SDXL's micro-conditioning as explained in section 2.2 of
|
||||
[https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more
|
||||
information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208.
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeine class.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] or `tuple`:
|
||||
[`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] if `return_dict` is True, otherwise a
|
||||
`tuple`. When returning a tuple, the first element is a list with the generated images.
|
||||
"""
|
||||
callback = kwargs.pop("callback", None)
|
||||
callback_steps = kwargs.pop("callback_steps", None)
|
||||
|
||||
if callback is not None:
|
||||
deprecate(
|
||||
"callback",
|
||||
"1.0.0",
|
||||
"Passing `callback` as an input argument to `__call__` is deprecated, consider use `callback_on_step_end`",
|
||||
)
|
||||
if callback_steps is not None:
|
||||
deprecate(
|
||||
"callback_steps",
|
||||
"1.0.0",
|
||||
"Passing `callback_steps` as an input argument to `__call__` is deprecated, consider use `callback_on_step_end`",
|
||||
)
|
||||
|
||||
# 0. Default height and width to unet
|
||||
height = height or self.default_sample_size * self.vae_scale_factor
|
||||
width = width or self.default_sample_size * self.vae_scale_factor
|
||||
|
||||
original_size = original_size or (height, width)
|
||||
target_size = target_size or (height, width)
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
callback_steps,
|
||||
negative_prompt,
|
||||
negative_prompt_2,
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._guidance_rescale = guidance_rescale
|
||||
self._clip_skip = clip_skip
|
||||
self._cross_attention_kwargs = cross_attention_kwargs
|
||||
self._denoising_end = denoising_end
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 3. Encode input prompt
|
||||
lora_scale = (
|
||||
self.cross_attention_kwargs.get("scale", None) if self.cross_attention_kwargs is not None else None
|
||||
)
|
||||
|
||||
(
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
) = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt_2,
|
||||
device=device,
|
||||
num_images_per_prompt=num_images_per_prompt,
|
||||
do_classifier_free_guidance=self.do_classifier_free_guidance,
|
||||
negative_prompt=negative_prompt,
|
||||
negative_prompt_2=negative_prompt_2,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds=negative_pooled_prompt_embeds,
|
||||
lora_scale=lora_scale,
|
||||
clip_skip=self.clip_skip,
|
||||
)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_images_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||
|
||||
if share_queries:
|
||||
query_store = QueryStore(**query_store_kwargs)
|
||||
else:
|
||||
query_store = None
|
||||
|
||||
self.attention_store = AttentionStore(attention_store_kwargs)
|
||||
register_extended_self_attn(self.unet, self.attention_store, extended_attn_kwargs)
|
||||
|
||||
# 7. Prepare added time ids & embeddings
|
||||
add_text_embeds = pooled_prompt_embeds
|
||||
if self.text_encoder_2 is None:
|
||||
text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1])
|
||||
else:
|
||||
text_encoder_projection_dim = self.text_encoder_2.config.projection_dim
|
||||
|
||||
add_time_ids = self._get_add_time_ids(
|
||||
original_size,
|
||||
crops_coords_top_left,
|
||||
target_size,
|
||||
dtype=prompt_embeds.dtype,
|
||||
text_encoder_projection_dim=text_encoder_projection_dim,
|
||||
)
|
||||
if negative_original_size is not None and negative_target_size is not None:
|
||||
negative_add_time_ids = self._get_add_time_ids(
|
||||
negative_original_size,
|
||||
negative_crops_coords_top_left,
|
||||
negative_target_size,
|
||||
dtype=prompt_embeds.dtype,
|
||||
text_encoder_projection_dim=text_encoder_projection_dim,
|
||||
)
|
||||
else:
|
||||
negative_add_time_ids = add_time_ids
|
||||
|
||||
if self.do_classifier_free_guidance:
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
add_text_embeds = torch.cat([negative_pooled_prompt_embeds, add_text_embeds], dim=0)
|
||||
add_time_ids = torch.cat([negative_add_time_ids, add_time_ids], dim=0)
|
||||
|
||||
prompt_embeds = prompt_embeds.to(device)
|
||||
add_text_embeds = add_text_embeds.to(device)
|
||||
add_time_ids = add_time_ids.to(device).repeat(batch_size * num_images_per_prompt, 1)
|
||||
|
||||
# 8. Denoising loop
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
|
||||
# 8.1 Apply denoising_end
|
||||
if (
|
||||
self.denoising_end is not None
|
||||
and isinstance(self.denoising_end, float)
|
||||
and self.denoising_end > 0
|
||||
and self.denoising_end < 1
|
||||
):
|
||||
discrete_timestep_cutoff = int(
|
||||
round(
|
||||
self.scheduler.config.num_train_timesteps
|
||||
- (self.denoising_end * self.scheduler.config.num_train_timesteps)
|
||||
)
|
||||
)
|
||||
num_inference_steps = len(list(filter(lambda ts: ts >= discrete_timestep_cutoff, timesteps)))
|
||||
timesteps = timesteps[:num_inference_steps]
|
||||
|
||||
# 9. Optionally get Guidance Scale Embedding
|
||||
timestep_cond = None
|
||||
if self.unet.config.time_cond_proj_dim is not None:
|
||||
guidance_scale_tensor = torch.tensor(self.guidance_scale - 1).repeat(batch_size * num_images_per_prompt)
|
||||
timestep_cond = self.get_guidance_scale_embedding(
|
||||
guidance_scale_tensor, embedding_dim=self.unet.config.time_cond_proj_dim
|
||||
).to(device=device, dtype=latents.dtype)
|
||||
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
if instance_latents is not None:
|
||||
n_instances = instance_latents.shape[0]
|
||||
instance_noise = latents[:n_instances].clone()
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
self.attention_store.curr_iter = i
|
||||
|
||||
if instance_latents is not None:
|
||||
noised_instances = self.scheduler.add_noise(instance_latents, instance_noise, t.repeat(n_instances).long())
|
||||
latents[:n_instances] = noised_instances
|
||||
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# predict the noise residual
|
||||
added_cond_kwargs = {"text_embeds": add_text_embeds, "time_ids": add_time_ids}
|
||||
|
||||
if share_queries and (i >= query_store.t_range[0] and i <= query_store.t_range[1]):
|
||||
query_store.set_mode('cache')
|
||||
noise_pred_vanilla = self.unet(
|
||||
latent_model_input,
|
||||
t,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep_cond=timestep_cond,
|
||||
cross_attention_kwargs={'query_store': query_store,
|
||||
'perform_extend_attn': False,
|
||||
'record_attention': False},
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
query_store.set_mode('inject')
|
||||
|
||||
noise_pred = self.unet(
|
||||
latent_model_input,
|
||||
t,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep_cond=timestep_cond,
|
||||
cross_attention_kwargs={'query_store': query_store,
|
||||
'perform_extend_attn': True,
|
||||
'record_attention': True,
|
||||
'feature_injector': feature_injector,
|
||||
'anchors_cache': anchors_cache},
|
||||
added_cond_kwargs=added_cond_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
if self.do_classifier_free_guidance and self.guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=self.guidance_rescale)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
add_text_embeds = callback_outputs.pop("add_text_embeds", add_text_embeds)
|
||||
negative_pooled_prompt_embeds = callback_outputs.pop(
|
||||
"negative_pooled_prompt_embeds", negative_pooled_prompt_embeds
|
||||
)
|
||||
add_time_ids = callback_outputs.pop("add_time_ids", add_time_ids)
|
||||
negative_add_time_ids = callback_outputs.pop("negative_add_time_ids", negative_add_time_ids)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
||||
callback(step_idx, t, latents)
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
# xm.mark_step()
|
||||
pass
|
||||
|
||||
# Update attention store mask
|
||||
self.attention_store.aggregate_last_steps_attention()
|
||||
|
||||
if not output_type == "latent":
|
||||
# make sure the VAE is in float32 mode, as it overflows in float16
|
||||
needs_upcasting = self.vae.dtype == torch.float16 and self.vae.config.force_upcast
|
||||
|
||||
if needs_upcasting:
|
||||
self.upcast_vae()
|
||||
latents = latents.to(next(iter(self.vae.post_quant_conv.parameters())).dtype)
|
||||
|
||||
image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0]
|
||||
|
||||
# cast back to fp16 if needed
|
||||
if needs_upcasting:
|
||||
self.vae.to(dtype=torch.float16)
|
||||
else:
|
||||
image = latents
|
||||
|
||||
if not output_type == "latent":
|
||||
# apply watermark if available
|
||||
if self.watermark is not None:
|
||||
image = self.watermark.apply_watermark(image)
|
||||
|
||||
image = self.image_processor.postprocess(image, output_type=output_type)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (image,)
|
||||
|
||||
return StableDiffusionXLPipelineOutput(images=image)
|
||||
@@ -0,0 +1,260 @@
|
||||
# Copyright (C) 2024 NVIDIA Corporation. All rights reserved.
|
||||
#
|
||||
# This work is licensed under the LICENSE file
|
||||
# located at the root directory.
|
||||
|
||||
import torch
|
||||
from diffusers import DDIMScheduler
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from .consistory_unet_sdxl import ConsistorySDXLUNet2DConditionModel
|
||||
from .consistory_pipeline import ConsistoryExtendAttnSDXLPipeline
|
||||
from .consistory_utils import FeatureInjector, AnchorCache
|
||||
# from .utils.general_utils import *
|
||||
from .utils.general_utils import gaussian_smooth, cyclic_nn_map, anchor_nn_map
|
||||
|
||||
|
||||
LATENT_RESOLUTIONS = [32, 64]
|
||||
|
||||
|
||||
def load_pipeline(gpu_id=0):
|
||||
float_type = torch.float16
|
||||
sd_id = "stabilityai/stable-diffusion-xl-base-1.0"
|
||||
device = torch.device(f'cuda:{gpu_id}') if torch.cuda.is_available() else torch.device('cpu')
|
||||
unet = ConsistorySDXLUNet2DConditionModel.from_pretrained(sd_id, subfolder="unet", torch_dtype=float_type)
|
||||
scheduler = DDIMScheduler.from_pretrained(sd_id, subfolder="scheduler")
|
||||
story_pipeline = ConsistoryExtendAttnSDXLPipeline.from_pretrained(sd_id, unet=unet, torch_dtype=float_type, variant="fp16", use_safetensors=True, scheduler=scheduler).to(device)
|
||||
story_pipeline.enable_freeu(s1=0.6, s2=0.4, b1=1.1, b2=1.2)
|
||||
return story_pipeline
|
||||
|
||||
|
||||
def create_anchor_mapping(bsz, anchor_indices=[0]):
|
||||
anchor_mapping = torch.eye(bsz, dtype=torch.bool)
|
||||
for anchor_idx in anchor_indices:
|
||||
anchor_mapping[:, anchor_idx] = True
|
||||
return anchor_mapping
|
||||
|
||||
|
||||
def create_token_indices(prompts, batch_size, concept_token, tokenizer):
|
||||
if isinstance(concept_token, str):
|
||||
concept_token = [concept_token]
|
||||
concept_token_id = [tokenizer.encode(x, add_special_tokens=False)[0] for x in concept_token]
|
||||
tokens = tokenizer.batch_encode_plus(prompts, padding=True, return_tensors='pt')['input_ids']
|
||||
token_indices = torch.full((len(concept_token), batch_size), -1, dtype=torch.int64)
|
||||
for i, token_id in enumerate(concept_token_id):
|
||||
batch_loc, token_loc = torch.where(tokens == token_id)
|
||||
token_indices[i, batch_loc] = token_loc
|
||||
return token_indices
|
||||
|
||||
|
||||
def create_latents(story_pipeline, seed, batch_size, same_latent, device, float_type):
|
||||
# if seed is int
|
||||
if isinstance(seed, int):
|
||||
g = torch.Generator('cuda').manual_seed(seed)
|
||||
shape = (batch_size, story_pipeline.unet.config.in_channels, 128, 128)
|
||||
latents = randn_tensor(shape, generator=g, device=device, dtype=float_type)
|
||||
elif isinstance(seed, list):
|
||||
shape = (batch_size, story_pipeline.unet.config.in_channels, 128, 128)
|
||||
latents = torch.empty(shape, device=device, dtype=float_type)
|
||||
for i, seed_i in enumerate(seed):
|
||||
g = torch.Generator('cuda').manual_seed(seed_i)
|
||||
curr_latent = randn_tensor(shape, generator=g, device=device, dtype=float_type)
|
||||
latents[i] = curr_latent[i]
|
||||
if same_latent:
|
||||
latents = latents[:1].repeat(batch_size, 1, 1, 1)
|
||||
return latents, g
|
||||
|
||||
|
||||
# Batch inference
|
||||
def run_batch_generation(story_pipeline, prompts, concept_token,
|
||||
seed=40, n_steps=50, mask_dropout=0.5,
|
||||
same_latent=False, share_queries=True,
|
||||
perform_sdsa=True, perform_injection=True,
|
||||
inject_range_alpha=(10,20,0.8),
|
||||
n_achors=2):
|
||||
device = story_pipeline.device
|
||||
tokenizer = story_pipeline.tokenizer
|
||||
float_type = story_pipeline.dtype
|
||||
unet = story_pipeline.unet
|
||||
batch_size = len(prompts)
|
||||
token_indices = create_token_indices(prompts, batch_size, concept_token, tokenizer)
|
||||
anchor_mappings = create_anchor_mapping(batch_size, anchor_indices=list(range(n_achors)))
|
||||
default_attention_store_kwargs = {
|
||||
'token_indices': token_indices,
|
||||
'mask_dropout': mask_dropout,
|
||||
'extended_mapping': anchor_mappings
|
||||
}
|
||||
default_extended_attn_kwargs = {'extend_kv_unet_parts': ['up']}
|
||||
query_store_kwargs= {'t_range': [0,n_steps//10], 'strength_start': 0.9, 'strength_end': 0.81836735}
|
||||
latents, g = create_latents(story_pipeline, seed, batch_size, same_latent, device, float_type)
|
||||
|
||||
# ------------------ #
|
||||
# Extended attention First Run #
|
||||
if perform_sdsa:
|
||||
extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': [(1, n_steps)]}
|
||||
else:
|
||||
extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': []}
|
||||
out = story_pipeline(prompt=prompts, generator=g, latents=latents,
|
||||
attention_store_kwargs=default_attention_store_kwargs,
|
||||
extended_attn_kwargs=extended_attn_kwargs,
|
||||
share_queries=share_queries,
|
||||
query_store_kwargs=query_store_kwargs,
|
||||
num_inference_steps=n_steps)
|
||||
last_masks = story_pipeline.attention_store.last_mask
|
||||
dift_features = unet.latent_store.dift_features['261_0'][batch_size:]
|
||||
dift_features = torch.stack([gaussian_smooth(x, kernel_size=3, sigma=1) for x in dift_features], dim=0)
|
||||
nn_map, nn_distances = cyclic_nn_map(dift_features, last_masks, LATENT_RESOLUTIONS, device)
|
||||
|
||||
# ------------------ #
|
||||
# Extended attention with nn_map #
|
||||
if perform_injection:
|
||||
feature_injector = FeatureInjector(
|
||||
nn_map,
|
||||
nn_distances,
|
||||
last_masks,
|
||||
inject_range_alpha=[inject_range_alpha],
|
||||
swap_strategy='min', inject_unet_parts=['up', 'down'], dist_thr='dynamic')
|
||||
out = story_pipeline(prompt=prompts, generator=g, latents=latents,
|
||||
attention_store_kwargs=default_attention_store_kwargs,
|
||||
extended_attn_kwargs=extended_attn_kwargs,
|
||||
share_queries=share_queries,
|
||||
query_store_kwargs=query_store_kwargs,
|
||||
feature_injector=feature_injector,
|
||||
num_inference_steps=n_steps)
|
||||
# display_attn_maps(story_pipeline.attention_store.last_mask, out.images)
|
||||
return out.images
|
||||
|
||||
|
||||
# Anchors
|
||||
def run_anchor_generation(story_pipeline, prompts, concept_token,
|
||||
seed=40, n_steps=50, mask_dropout=0.5,
|
||||
inject_range_alpha=(10,20,0.8),
|
||||
same_latent=False, share_queries=True,
|
||||
perform_sdsa=True, perform_injection=True):
|
||||
device = story_pipeline.device
|
||||
tokenizer = story_pipeline.tokenizer
|
||||
float_type = story_pipeline.dtype
|
||||
unet = story_pipeline.unet
|
||||
batch_size = len(prompts)
|
||||
token_indices = create_token_indices(prompts, batch_size, concept_token, tokenizer)
|
||||
default_attention_store_kwargs = {
|
||||
'token_indices': token_indices,
|
||||
'mask_dropout': mask_dropout
|
||||
}
|
||||
default_extended_attn_kwargs = {'extend_kv_unet_parts': ['up']}
|
||||
query_store_kwargs={'t_range': [0,n_steps//10], 'strength_start': 0.9, 'strength_end': 0.81836735}
|
||||
latents, g = create_latents(story_pipeline, seed, batch_size, same_latent, device, float_type)
|
||||
anchor_cache_first_stage = AnchorCache()
|
||||
anchor_cache_second_stage = AnchorCache()
|
||||
|
||||
# ------------------ #
|
||||
# Extended attention First Run #
|
||||
if perform_sdsa:
|
||||
extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': [(1, n_steps)]}
|
||||
else:
|
||||
extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': []}
|
||||
out = story_pipeline(prompt=prompts, generator=g, latents=latents,
|
||||
attention_store_kwargs=default_attention_store_kwargs,
|
||||
extended_attn_kwargs=extended_attn_kwargs,
|
||||
share_queries=share_queries,
|
||||
query_store_kwargs=query_store_kwargs,
|
||||
anchors_cache=anchor_cache_first_stage,
|
||||
num_inference_steps=n_steps)
|
||||
last_masks = story_pipeline.attention_store.last_mask
|
||||
dift_features = unet.latent_store.dift_features['261_0'][batch_size:]
|
||||
dift_features = torch.stack([gaussian_smooth(x, kernel_size=3, sigma=1) for x in dift_features], dim=0)
|
||||
anchor_cache_first_stage.dift_cache = dift_features
|
||||
anchor_cache_first_stage.anchors_last_mask = last_masks
|
||||
nn_map, nn_distances = cyclic_nn_map(dift_features, last_masks, LATENT_RESOLUTIONS, device)
|
||||
|
||||
# ------------------ #
|
||||
# Extended attention with nn_map #
|
||||
if perform_injection:
|
||||
feature_injector = FeatureInjector(
|
||||
nn_map,
|
||||
nn_distances,
|
||||
last_masks,
|
||||
inject_range_alpha=[inject_range_alpha],
|
||||
swap_strategy='min',
|
||||
inject_unet_parts=['up', 'down'],
|
||||
dist_thr='dynamic')
|
||||
out = story_pipeline(prompt=prompts, generator=g, latents=latents,
|
||||
attention_store_kwargs=default_attention_store_kwargs,
|
||||
extended_attn_kwargs=extended_attn_kwargs,
|
||||
share_queries=share_queries,
|
||||
query_store_kwargs=query_store_kwargs,
|
||||
feature_injector=feature_injector,
|
||||
anchors_cache=anchor_cache_second_stage,
|
||||
num_inference_steps=n_steps)
|
||||
# display_attn_maps(story_pipeline.attention_store.last_mask, out.images)
|
||||
anchor_cache_second_stage.dift_cache = dift_features
|
||||
anchor_cache_second_stage.anchors_last_mask = last_masks
|
||||
return out.images, anchor_cache_first_stage, anchor_cache_second_stage
|
||||
|
||||
|
||||
def run_extra_generation(story_pipeline, prompts, concept_token,
|
||||
anchor_cache_first_stage, anchor_cache_second_stage,
|
||||
seed=40, n_steps=50, mask_dropout=0.5,
|
||||
inject_range_alpha=(10,20,0.8),
|
||||
same_latent=False, share_queries=True,
|
||||
perform_sdsa=True, perform_injection=True):
|
||||
device = story_pipeline.device
|
||||
tokenizer = story_pipeline.tokenizer
|
||||
float_type = story_pipeline.dtype
|
||||
unet = story_pipeline.unet
|
||||
batch_size = len(prompts)
|
||||
token_indices = create_token_indices(prompts, batch_size, concept_token, tokenizer)
|
||||
default_attention_store_kwargs = {
|
||||
'token_indices': token_indices,
|
||||
'mask_dropout': mask_dropout
|
||||
}
|
||||
default_extended_attn_kwargs = {'extend_kv_unet_parts': ['up']}
|
||||
query_store_kwargs={'t_range': [0,n_steps//10], 'strength_start': 0.9, 'strength_end': 0.81836735}
|
||||
extra_batch_size = batch_size + 2
|
||||
if isinstance(seed, list):
|
||||
seed = [seed[0], seed[0], *seed]
|
||||
latents, g = create_latents(story_pipeline, seed, extra_batch_size, same_latent, device, float_type)
|
||||
latents = latents[2:]
|
||||
anchor_cache_first_stage.set_mode_inject()
|
||||
anchor_cache_second_stage.set_mode_inject()
|
||||
|
||||
# ------------------ #
|
||||
# Extended attention First Run #
|
||||
if perform_sdsa:
|
||||
extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': [(1, n_steps)]}
|
||||
else:
|
||||
extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': []}
|
||||
out = story_pipeline(prompt=prompts, generator=g, latents=latents,
|
||||
attention_store_kwargs=default_attention_store_kwargs,
|
||||
extended_attn_kwargs=extended_attn_kwargs,
|
||||
share_queries=share_queries,
|
||||
query_store_kwargs=query_store_kwargs,
|
||||
anchors_cache=anchor_cache_first_stage,
|
||||
num_inference_steps=n_steps)
|
||||
last_masks = story_pipeline.attention_store.last_mask
|
||||
dift_features = unet.latent_store.dift_features['261_0'][batch_size:]
|
||||
dift_features = torch.stack([gaussian_smooth(x, kernel_size=3, sigma=1) for x in dift_features], dim=0)
|
||||
anchor_dift_features = anchor_cache_first_stage.dift_cache
|
||||
anchor_last_masks = anchor_cache_first_stage.anchors_last_mask
|
||||
nn_map, nn_distances = anchor_nn_map(dift_features, anchor_dift_features, last_masks, anchor_last_masks, LATENT_RESOLUTIONS, device)
|
||||
|
||||
# ------------------ #
|
||||
# Extended attention with nn_map #
|
||||
if perform_injection:
|
||||
feature_injector = FeatureInjector(
|
||||
nn_map,
|
||||
nn_distances,
|
||||
last_masks,
|
||||
inject_range_alpha=[inject_range_alpha],
|
||||
swap_strategy='min',
|
||||
inject_unet_parts=['up', 'down'],
|
||||
dist_thr='dynamic')
|
||||
out = story_pipeline(prompt=prompts, generator=g, latents=latents,
|
||||
attention_store_kwargs=default_attention_store_kwargs,
|
||||
extended_attn_kwargs=extended_attn_kwargs,
|
||||
share_queries=share_queries,
|
||||
query_store_kwargs=query_store_kwargs,
|
||||
feature_injector=feature_injector,
|
||||
anchors_cache=anchor_cache_second_stage,
|
||||
num_inference_steps=n_steps)
|
||||
# display_attn_maps(story_pipeline.attention_store.last_mask, out.images)
|
||||
return out.images
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,192 @@
|
||||
# Copyright (C) 2024 NVIDIA Corporation. All rights reserved.
|
||||
#
|
||||
# This work is licensed under the LICENSE file
|
||||
# located at the root directory.
|
||||
|
||||
from typing import List
|
||||
from collections import defaultdict
|
||||
import numpy as np
|
||||
import torch
|
||||
from .utils.general_utils import get_dynamic_threshold
|
||||
|
||||
|
||||
class FeatureInjector:
|
||||
def __init__(self, nn_map, nn_distances, attn_masks, inject_range_alpha=[(10,20,0.8)], swap_strategy='min', dist_thr='dynamic', inject_unet_parts=['up']):
|
||||
self.nn_map = nn_map
|
||||
self.nn_distances = nn_distances
|
||||
self.attn_masks = attn_masks
|
||||
self.inject_range_alpha = inject_range_alpha if isinstance(inject_range_alpha, list) else [inject_range_alpha]
|
||||
self.swap_strategy = swap_strategy # 'min / 'mean' / 'first'
|
||||
self.dist_thr = dist_thr
|
||||
self.inject_unet_parts = inject_unet_parts
|
||||
self.inject_res = [64]
|
||||
|
||||
def inject_outputs(self, output, curr_iter, output_res, extended_mapping, place_in_unet, anchors_cache=None):
|
||||
curr_unet_part = place_in_unet.split('_')[0]
|
||||
|
||||
# Inject only in the specified unet parts (up, mid, down)
|
||||
if (curr_unet_part not in self.inject_unet_parts) or output_res not in self.inject_res:
|
||||
return output
|
||||
|
||||
bsz = output.shape[0]
|
||||
nn_map = self.nn_map[output_res]
|
||||
nn_distances = self.nn_distances[output_res]
|
||||
attn_masks = self.attn_masks[output_res]
|
||||
vector_dim = output_res**2
|
||||
|
||||
alpha = next((alpha for min_range, max_range, alpha in self.inject_range_alpha if min_range <= curr_iter <= max_range), None)
|
||||
if alpha:
|
||||
old_output = output#.clone()
|
||||
for i in range(bsz):
|
||||
other_outputs = []
|
||||
|
||||
if self.swap_strategy == 'min':
|
||||
curr_mapping = extended_mapping[i]
|
||||
|
||||
# If the current image is not mapped to any other image, skip
|
||||
if not torch.any(torch.cat([curr_mapping[:i], curr_mapping[i+1:]])):
|
||||
continue
|
||||
|
||||
min_dists = nn_distances[i][curr_mapping].argmin(dim=0)
|
||||
curr_nn_map = nn_map[i][curr_mapping][min_dists, torch.arange(vector_dim)]
|
||||
|
||||
curr_nn_distances = nn_distances[i][curr_mapping][min_dists, torch.arange(vector_dim)]
|
||||
dist_thr = get_dynamic_threshold(curr_nn_distances) if self.dist_thr == 'dynamic' else self.dist_thr
|
||||
dist_mask = curr_nn_distances < dist_thr
|
||||
final_mask_tgt = attn_masks[i] & dist_mask
|
||||
|
||||
other_outputs = old_output[curr_mapping][min_dists, curr_nn_map][final_mask_tgt]
|
||||
|
||||
output[i][final_mask_tgt] = alpha * other_outputs + (1 - alpha)*old_output[i][final_mask_tgt]
|
||||
|
||||
if anchors_cache and anchors_cache.is_cache_mode():
|
||||
if place_in_unet not in anchors_cache.h_out_cache:
|
||||
anchors_cache.h_out_cache[place_in_unet] = {}
|
||||
|
||||
anchors_cache.h_out_cache[place_in_unet][curr_iter] = output
|
||||
|
||||
return output
|
||||
|
||||
def inject_anchors(self, output, curr_iter, output_res, extended_mapping, place_in_unet, anchors_cache):
|
||||
curr_unet_part = place_in_unet.split('_')[0]
|
||||
|
||||
# Inject only in the specified unet parts (up, mid, down)
|
||||
if (curr_unet_part not in self.inject_unet_parts) or output_res not in self.inject_res:
|
||||
return output
|
||||
|
||||
bsz = output.shape[0]
|
||||
nn_map = self.nn_map[output_res]
|
||||
nn_distances = self.nn_distances[output_res]
|
||||
attn_masks = self.attn_masks[output_res]
|
||||
vector_dim = output_res**2
|
||||
|
||||
alpha = next((alpha for min_range, max_range, alpha in self.inject_range_alpha if min_range <= curr_iter <= max_range), None)
|
||||
if alpha:
|
||||
|
||||
anchor_outputs = anchors_cache.h_out_cache[place_in_unet][curr_iter]
|
||||
|
||||
old_output = output#.clone()
|
||||
for i in range(bsz):
|
||||
other_outputs = []
|
||||
|
||||
if self.swap_strategy == 'min':
|
||||
min_dists = nn_distances[i].argmin(dim=0)
|
||||
curr_nn_map = nn_map[i][min_dists, torch.arange(vector_dim)]
|
||||
|
||||
curr_nn_distances = nn_distances[i][min_dists, torch.arange(vector_dim)]
|
||||
dist_thr = get_dynamic_threshold(curr_nn_distances) if self.dist_thr == 'dynamic' else self.dist_thr
|
||||
dist_mask = curr_nn_distances < dist_thr
|
||||
final_mask_tgt = attn_masks[i] & dist_mask
|
||||
|
||||
other_outputs = anchor_outputs[min_dists, curr_nn_map][final_mask_tgt]
|
||||
|
||||
output[i][final_mask_tgt] = alpha * other_outputs + (1 - alpha)*old_output[i][final_mask_tgt]
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class AnchorCache:
|
||||
def __init__(self):
|
||||
self.input_h_cache = {} # place_in_unet, iter, h_in
|
||||
self.h_out_cache = {} # place_in_unet, iter, h_out
|
||||
self.anchors_last_mask = None
|
||||
self.dift_cache = None
|
||||
|
||||
self.mode = 'cache' # mode can be 'cache' or 'inject'
|
||||
|
||||
def set_mode(self, mode):
|
||||
self.mode = mode
|
||||
|
||||
def set_mode_inject(self):
|
||||
self.mode = 'inject'
|
||||
|
||||
def set_mode_cache(self):
|
||||
self.mode = 'cache'
|
||||
|
||||
def is_inject_mode(self):
|
||||
return self.mode == 'inject'
|
||||
|
||||
def is_cache_mode(self):
|
||||
return self.mode == 'cache'
|
||||
|
||||
|
||||
def to_device(self, device):
|
||||
for key, value in self.input_h_cache.items():
|
||||
self.input_h_cache[key] = {k: v.to(device) for k, v in value.items()}
|
||||
|
||||
for key, value in self.h_out_cache.items():
|
||||
self.h_out_cache[key] = {k: v.to(device) for k, v in value.items()}
|
||||
|
||||
if self.anchors_last_mask:
|
||||
self.anchors_last_mask = {k: v.to(device) for k, v in self.anchors_last_mask.items()}
|
||||
|
||||
if self.dift_cache is not None:
|
||||
self.dift_cache = self.dift_cache.to(device)
|
||||
|
||||
|
||||
class QueryStore:
|
||||
def __init__(self, mode='store', t_range=[0, 1000], strength_start=1, strength_end=1):
|
||||
"""
|
||||
Initialize an empty ActivationsStore
|
||||
"""
|
||||
self.query_store = defaultdict(list)
|
||||
self.mode = mode
|
||||
self.t_range = t_range
|
||||
self.strengthes = np.linspace(strength_start, strength_end, (t_range[1] - t_range[0])+1)
|
||||
|
||||
def set_mode(self, mode): # mode can be 'cache' or 'inject'
|
||||
self.mode = mode
|
||||
|
||||
def cache_query(self, query, place_in_unet: str):
|
||||
self.query_store[place_in_unet] = query
|
||||
|
||||
def inject_query(self, query, place_in_unet, t):
|
||||
if t >= self.t_range[0] and t <= self.t_range[1]:
|
||||
relative_t = t - self.t_range[0]
|
||||
strength = self.strengthes[relative_t]
|
||||
new_query = strength * self.query_store[place_in_unet] + (1 - strength) * query
|
||||
else:
|
||||
new_query = query
|
||||
|
||||
return new_query
|
||||
|
||||
class DIFTLatentStore:
|
||||
def __init__(self, steps: List[int], up_ft_indices: List[int]):
|
||||
self.steps = steps
|
||||
self.up_ft_indices = up_ft_indices
|
||||
self.dift_features = {}
|
||||
|
||||
def __call__(self, features: torch.Tensor, t: int, layer_index: int):
|
||||
if t in self.steps and layer_index in self.up_ft_indices:
|
||||
self.dift_features[f'{int(t)}_{layer_index}'] = features
|
||||
|
||||
def copy(self):
|
||||
copy_dift = DIFTLatentStore(self.steps, self.up_ft_indices)
|
||||
|
||||
for key, value in self.dift_features.items():
|
||||
copy_dift.dift_features[key] = value.clone()
|
||||
|
||||
return copy_dift
|
||||
|
||||
def reset(self):
|
||||
self.dift_features = {}
|
||||
@@ -0,0 +1,117 @@
|
||||
# Copyright (C) 2024 NVIDIA Corporation. All rights reserved.
|
||||
#
|
||||
# This work is licensed under the LICENSE file
|
||||
# located at the root directory.
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
from skimage import filters
|
||||
|
||||
|
||||
## Attention Utils
|
||||
def get_dynamic_threshold(tensor):
|
||||
return filters.threshold_otsu(tensor.float().cpu().numpy())
|
||||
|
||||
|
||||
def attn_map_to_binary(attention_map, scaler=1.):
|
||||
attention_map_np = attention_map.float().cpu().numpy()
|
||||
threshold_value = filters.threshold_otsu(attention_map_np) * scaler
|
||||
binary_mask = (attention_map_np > threshold_value).astype(np.uint8)
|
||||
|
||||
return binary_mask
|
||||
|
||||
|
||||
## Features
|
||||
|
||||
def gaussian_smooth(input_tensor, kernel_size=3, sigma=1):
|
||||
"""
|
||||
Function to apply Gaussian smoothing on each 2D slice of a 3D tensor.
|
||||
"""
|
||||
kernel = np.fromfunction(
|
||||
lambda x, y: (1/ (2 * np.pi * sigma ** 2)) *
|
||||
np.exp(-((x - (kernel_size - 1) / 2) ** 2 + (y - (kernel_size - 1) / 2) ** 2) / (2 * sigma ** 2)),
|
||||
(kernel_size, kernel_size)
|
||||
)
|
||||
kernel = torch.Tensor(kernel / kernel.sum()).to(input_tensor.dtype).to(input_tensor.device)
|
||||
# Add batch and channel dimensions to the kernel
|
||||
kernel = kernel.unsqueeze(0).unsqueeze(0)
|
||||
# Iterate over each 2D slice and apply convolution
|
||||
smoothed_slices = []
|
||||
for i in range(input_tensor.size(0)):
|
||||
slice_tensor = input_tensor[i, :, :]
|
||||
slice_tensor = F.conv2d(slice_tensor.unsqueeze(0).unsqueeze(0), kernel, padding=kernel_size // 2)[0, 0]
|
||||
smoothed_slices.append(slice_tensor)
|
||||
# Stack the smoothed slices to get the final tensor
|
||||
smoothed_tensor = torch.stack(smoothed_slices, dim=0)
|
||||
return smoothed_tensor
|
||||
|
||||
|
||||
## Dense correspondence utils
|
||||
|
||||
def cos_dist(a, b):
|
||||
a_norm = F.normalize(a, dim=-1)
|
||||
b_norm = F.normalize(b, dim=-1)
|
||||
res = a_norm @ b_norm.T
|
||||
return 1 - res
|
||||
|
||||
|
||||
def gen_nn_map(src_features, src_mask, tgt_features, tgt_mask, device, batch_size=100, tgt_size=768):
|
||||
resized_src_features = F.interpolate(src_features.unsqueeze(0), size=tgt_size, mode='bilinear', align_corners=False).squeeze(0)
|
||||
resized_src_features = resized_src_features.permute(1,2,0).view(tgt_size**2, -1)
|
||||
resized_tgt_features = F.interpolate(tgt_features.unsqueeze(0), size=tgt_size, mode='bilinear', align_corners=False).squeeze(0)
|
||||
resized_tgt_features = resized_tgt_features.permute(1,2,0).view(tgt_size**2, -1)
|
||||
nearest_neighbor_indices = torch.zeros(tgt_size**2, dtype=torch.long, device=device)
|
||||
nearest_neighbor_distances = torch.zeros(tgt_size**2, dtype=src_features.dtype, device=device)
|
||||
if not batch_size:
|
||||
batch_size = tgt_size**2
|
||||
for i in range(0, tgt_size**2, batch_size):
|
||||
distances = cos_dist(resized_src_features, resized_tgt_features[i:i+batch_size])
|
||||
distances[~src_mask] = 2.
|
||||
min_distances, min_indices = torch.min(distances, dim=0)
|
||||
nearest_neighbor_indices[i:i+batch_size] = min_indices
|
||||
nearest_neighbor_distances[i:i+batch_size] = min_distances
|
||||
return nearest_neighbor_indices, nearest_neighbor_distances
|
||||
|
||||
|
||||
def cyclic_nn_map(features, masks, latent_resolutions, device):
|
||||
bsz = features.shape[0]
|
||||
nn_map_dict = {}
|
||||
nn_distances_dict = {}
|
||||
|
||||
for tgt_size in latent_resolutions:
|
||||
nn_map = torch.empty(bsz, bsz, tgt_size**2, dtype=torch.long, device=device)
|
||||
nn_distances = torch.full((bsz, bsz, tgt_size**2), float('inf'), dtype=features.dtype, device=device)
|
||||
|
||||
for i in range(bsz):
|
||||
for j in range(bsz):
|
||||
if i != j:
|
||||
nearest_neighbor_indices, nearest_neighbor_distances = gen_nn_map(features[j], masks[tgt_size][j], features[i], masks[tgt_size][i], device, batch_size=None, tgt_size=tgt_size)
|
||||
nn_map[i,j] = nearest_neighbor_indices
|
||||
nn_distances[i,j] = nearest_neighbor_distances
|
||||
|
||||
nn_map_dict[tgt_size] = nn_map
|
||||
nn_distances_dict[tgt_size] = nn_distances
|
||||
|
||||
return nn_map_dict, nn_distances_dict
|
||||
|
||||
|
||||
def anchor_nn_map(features, anchor_features, masks, anchor_masks, latent_resolutions, device):
|
||||
bsz = features.shape[0]
|
||||
anchor_bsz = anchor_features.shape[0]
|
||||
nn_map_dict = {}
|
||||
nn_distances_dict = {}
|
||||
|
||||
for tgt_size in latent_resolutions:
|
||||
nn_map = torch.empty(bsz, anchor_bsz, tgt_size**2, dtype=torch.long, device=device)
|
||||
nn_distances = torch.full((bsz, anchor_bsz, tgt_size**2), float('inf'), dtype=features.dtype, device=device)
|
||||
|
||||
for i in range(bsz):
|
||||
for j in range(anchor_bsz):
|
||||
nearest_neighbor_indices, nearest_neighbor_distances = gen_nn_map(anchor_features[j], anchor_masks[tgt_size][j], features[i], masks[tgt_size][i], device, batch_size=None, tgt_size=tgt_size)
|
||||
nn_map[i,j] = nearest_neighbor_indices
|
||||
nn_distances[i,j] = nearest_neighbor_distances
|
||||
nn_map_dict[tgt_size] = nn_map
|
||||
nn_distances_dict[tgt_size] = nn_distances
|
||||
|
||||
return nn_map_dict, nn_distances_dict
|
||||
@@ -0,0 +1,194 @@
|
||||
# Copyright 2022 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# MIT License
|
||||
#
|
||||
# Copyright (c) 2023 AttendAndExcite
|
||||
#
|
||||
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
# of this software and associated documentation files (the "Software"), to deal
|
||||
# in the Software without restriction, including without limitation the rights
|
||||
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
# copies of the Software, and to permit persons to whom the Software is
|
||||
# furnished to do so, subject to the following conditions:
|
||||
#
|
||||
# The above copyright notice and this permission notice shall be included in all
|
||||
# copies or substantial portions of the Software.
|
||||
#
|
||||
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
# SOFTWARE.
|
||||
|
||||
# Copyright 2022 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# Not a contribution
|
||||
# Changes made by NVIDIA CORPORATION & AFFILIATES enabling ConsiStory or otherwise documented as NVIDIA-proprietary
|
||||
# are not a contribution and subject to the license under the LICENSE file located at the root directory.
|
||||
|
||||
import torch
|
||||
from collections import defaultdict
|
||||
import numpy as np
|
||||
from typing import Union, List
|
||||
from PIL import Image
|
||||
|
||||
from modules.consistory.utils.general_utils import attn_map_to_binary
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class AttentionStore:
|
||||
def __init__(self, attention_store_kwargs):
|
||||
"""
|
||||
Initialize an empty AttentionStore :param step_index: used to visualize only a specific step in the diffusion
|
||||
process
|
||||
"""
|
||||
self.attn_res = attention_store_kwargs.get('attn_res', (32,32))
|
||||
self.token_indices = attention_store_kwargs['token_indices']
|
||||
bsz = self.token_indices.size(1)
|
||||
self.mask_background_query = attention_store_kwargs.get('mask_background_query', False)
|
||||
self.original_attn_masks = attention_store_kwargs.get('original_attn_masks', None)
|
||||
self.extended_mapping = attention_store_kwargs.get('extended_mapping', torch.ones(bsz, bsz).bool())
|
||||
self.mask_dropout = attention_store_kwargs.get('mask_dropout', 0.0)
|
||||
torch.manual_seed(0) # For dropout mask reproducibility
|
||||
|
||||
self.curr_iter = 0
|
||||
self.ALL_RES = [32, 64]
|
||||
self.step_store = defaultdict(list)
|
||||
self.attn_masks = {res: None for res in self.ALL_RES}
|
||||
self.last_mask = {res: None for res in self.ALL_RES}
|
||||
self.last_mask_dropout = {res: None for res in self.ALL_RES}
|
||||
|
||||
def __call__(self, attn, is_cross: bool, place_in_unet: str, attn_heads: int):
|
||||
if is_cross and attn.shape[1] == np.prod(self.attn_res):
|
||||
guidance_attention = attn[attn.size(0)//2:]
|
||||
batched_guidance_attention = guidance_attention.reshape([guidance_attention.shape[0]//attn_heads, attn_heads, *guidance_attention.shape[1:]])
|
||||
batched_guidance_attention = batched_guidance_attention.mean(dim=1)
|
||||
self.step_store[place_in_unet].append(batched_guidance_attention)
|
||||
|
||||
def reset(self):
|
||||
self.step_store = defaultdict(list)
|
||||
self.attn_masks = {res: None for res in self.ALL_RES}
|
||||
self.last_mask = {res: None for res in self.ALL_RES}
|
||||
self.last_mask_dropout = {res: None for res in self.ALL_RES}
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def aggregate_last_steps_attention(self) -> torch.Tensor:
|
||||
"""Aggregates the attention across the different layers and heads at the specified resolution."""
|
||||
attention_maps = torch.cat([torch.stack(x[-20:]) for x in self.step_store.values()]).mean(dim=0)
|
||||
bsz, wh, _ = attention_maps.shape
|
||||
|
||||
# Create attention maps for each concept token, for each batch item
|
||||
agg_attn_maps = []
|
||||
for i in range(bsz):
|
||||
curr_prompt_indices = []
|
||||
|
||||
for concept_token_indices in self.token_indices:
|
||||
if concept_token_indices[i] != -1:
|
||||
curr_prompt_indices.append(attention_maps[i, :, concept_token_indices[i]].view(*self.attn_res))
|
||||
|
||||
agg_attn_maps.append(torch.stack(curr_prompt_indices))
|
||||
|
||||
# Upsample the attention maps to the target resolution
|
||||
# and create the attention masks, unifying masks across the different concepts
|
||||
for tgt_size in self.ALL_RES:
|
||||
pixels = tgt_size ** 2
|
||||
tgt_agg_attn_maps = [F.interpolate(x.unsqueeze(1), size=tgt_size, mode='bilinear').squeeze(1) for x in agg_attn_maps]
|
||||
|
||||
attn_masks = []
|
||||
for batch_item_map in tgt_agg_attn_maps:
|
||||
concept_attn_masks = []
|
||||
|
||||
for concept_maps in batch_item_map:
|
||||
concept_attn_masks.append(torch.from_numpy(attn_map_to_binary(concept_maps, 1.)).to(attention_maps.device).bool().view(-1))
|
||||
|
||||
concept_attn_masks = torch.stack(concept_attn_masks, dim=0).max(dim=0).values
|
||||
attn_masks.append(concept_attn_masks)
|
||||
|
||||
attn_masks = torch.stack(attn_masks)
|
||||
self.last_mask[tgt_size] = attn_masks.clone()
|
||||
|
||||
# Add mask dropout
|
||||
if self.curr_iter < 1000:
|
||||
rand_mask = (torch.rand_like(attn_masks.float()) < self.mask_dropout)
|
||||
attn_masks[rand_mask] = False
|
||||
|
||||
self.last_mask_dropout[tgt_size] = attn_masks.clone()
|
||||
|
||||
# # Create subject driven extended self attention masks
|
||||
# output_attn_mask = torch.zeros((bsz, tgt_size**2, attn_masks.view(-1).size(0)), device=attn_masks.device).bool()
|
||||
|
||||
# for i in range(bsz):
|
||||
# for j in range(bsz):
|
||||
# if i==j:
|
||||
# output_attn_mask[i, :, j*pixels:(j+1)*pixels] = 1
|
||||
# else:
|
||||
# if self.extended_mapping[i,j]:
|
||||
# if not self.mask_background_query:
|
||||
# output_attn_mask[i, :, j*pixels:(j+1)*pixels] = attn_masks[j].unsqueeze(0).expand(pixels, -1)
|
||||
# else:
|
||||
# output_attn_mask[i, attn_masks[i], j*pixels:(j+1)*pixels] = attn_masks[j].unsqueeze(0).expand(attn_masks[i].sum(), -1)
|
||||
|
||||
# self.attn_masks[tgt_size] = output_attn_mask
|
||||
|
||||
def get_attn_mask_bias(self, tgt_size, bsz=None):
|
||||
attn_mask = self.attn_masks[tgt_size] if self.original_attn_masks is None else self.original_attn_masks[tgt_size]
|
||||
|
||||
if attn_mask is None:
|
||||
return None
|
||||
|
||||
attn_bias = torch.zeros_like(attn_mask, dtype=torch.float16)
|
||||
attn_bias[~attn_mask] = float('-inf')
|
||||
|
||||
if bsz and bsz != attn_bias.shape[0]:
|
||||
attn_bias = attn_bias.repeat(bsz // attn_bias.shape[0], 1, 1)
|
||||
|
||||
return attn_bias
|
||||
|
||||
def get_extended_attn_mask_instance(self, width, i):
|
||||
attn_mask = self.last_mask_dropout[width]
|
||||
if attn_mask is None:
|
||||
return None
|
||||
|
||||
n_patches = width**2
|
||||
|
||||
|
||||
output_attn_mask = torch.zeros((attn_mask.shape[0] * attn_mask.shape[1],), device=attn_mask.device, dtype=torch.bool)
|
||||
for j in range(attn_mask.shape[0]):
|
||||
if i==j:
|
||||
output_attn_mask[j*n_patches:(j+1)*n_patches] = 1
|
||||
else:
|
||||
if self.extended_mapping[i,j]:
|
||||
if not self.mask_background_query:
|
||||
output_attn_mask[j*n_patches:(j+1)*n_patches] = attn_mask[j].unsqueeze(0) #.expand(n_patches, -1)
|
||||
else:
|
||||
raise NotImplementedError('mask_background_query is not supported anymore')
|
||||
output_attn_mask[0, attn_mask[i], k*n_patches:(k+1)*n_patches] = attn_mask[j].unsqueeze(0).expand(attn_mask[i].sum(), -1)
|
||||
|
||||
return output_attn_mask
|
||||
@@ -66,11 +66,11 @@ def control_run(state: str = '',
|
||||
resize_mode_before: int = 0, resize_name_before: str = 'None', resize_context_before: str = 'None', width_before: int = 512, height_before: int = 512, scale_by_before: float = 1.0, selected_scale_tab_before: int = 0,
|
||||
resize_mode_after: int = 0, resize_name_after: str = 'None', resize_context_after: str = 'None', width_after: int = 0, height_after: int = 0, scale_by_after: float = 1.0, selected_scale_tab_after: int = 0,
|
||||
resize_mode_mask: int = 0, resize_name_mask: str = 'None', resize_context_mask: str = 'None', width_mask: int = 0, height_mask: int = 0, scale_by_mask: float = 1.0, selected_scale_tab_mask: int = 0,
|
||||
denoising_strength: float = 0, batch_count: int = 1, batch_size: int = 1,
|
||||
enable_hr: bool = False, hr_sampler_index: int = None, hr_denoising_strength: float = 0.3, hr_resize_mode: int = 0, hr_resize_context: str = 'None', hr_upscaler: str = None, hr_force: bool = False, hr_second_pass_steps: int = 20,
|
||||
denoising_strength: float = 0.3, batch_count: int = 1, batch_size: int = 1,
|
||||
enable_hr: bool = False, hr_sampler_index: int = None, hr_denoising_strength: float = 0.0, hr_resize_mode: int = 0, hr_resize_context: str = 'None', hr_upscaler: str = None, hr_force: bool = False, hr_second_pass_steps: int = 20,
|
||||
hr_scale: float = 1.0, hr_resize_x: int = 0, hr_resize_y: int = 0, refiner_steps: int = 5, refiner_start: float = 0.0, refiner_prompt: str = '', refiner_negative: str = '',
|
||||
video_skip_frames: int = 0, video_type: str = 'None', video_duration: float = 2.0, video_loop: bool = False, video_pad: int = 0, video_interpolate: int = 0,
|
||||
*input_script_args
|
||||
*input_script_args,
|
||||
):
|
||||
# handle optional initialization via ui
|
||||
for u in units:
|
||||
|
||||
+2
-2
@@ -59,7 +59,7 @@ def exception(suppress=[]):
|
||||
console.print_exception(show_locals=False, max_frames=16, extra_lines=2, suppress=suppress, theme="ansi_dark", word_wrap=False, width=min([console.width, 200]))
|
||||
|
||||
|
||||
def profile(profiler, msg: str):
|
||||
def profile(profiler, msg: str, n: int = 5):
|
||||
profiler.disable()
|
||||
import io
|
||||
import pstats
|
||||
@@ -83,7 +83,7 @@ def profile(profiler, msg: str):
|
||||
and 'rich' not in x
|
||||
and x.strip() != ''
|
||||
]
|
||||
txt = '\n'.join(lines[:min(5, len(lines))])
|
||||
txt = '\n'.join(lines[:min(n, len(lines))])
|
||||
log.debug(f'Profile {msg}: {txt}')
|
||||
|
||||
|
||||
|
||||
@@ -102,8 +102,9 @@ def activate(p, extra_network_data, step=0):
|
||||
except Exception as e:
|
||||
errors.display(e, f"Activating network: type={extra_network_name}")
|
||||
|
||||
p.extra_network_data = extra_network_data
|
||||
if stepwise:
|
||||
p.extra_network_data = extra_network_data
|
||||
p.stepwise_lora = True
|
||||
shared.opts.data['lora_functional'] = functional
|
||||
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ debug = shared.log.trace if os.environ.get('SD_FACE_DEBUG', None) is not None el
|
||||
|
||||
class Script(scripts.Script):
|
||||
def title(self):
|
||||
return 'Face'
|
||||
return 'Face: Multiple ID Transfers'
|
||||
|
||||
def show(self, is_img2img):
|
||||
return True if shared.native else False
|
||||
@@ -28,10 +28,10 @@ class Script(scripts.Script):
|
||||
elif hasattr(file, 'name'):
|
||||
image = Image.open(file.name) # _TemporaryFileWrapper from gr.Files
|
||||
else:
|
||||
raise ValueError(f'PhotoMaker unknown input: {file}')
|
||||
raise ValueError(f'Face: unknown input: {file}')
|
||||
init_images.append(image)
|
||||
except Exception as e:
|
||||
shared.log.warning(f'PhotoMaker failed to load image: {e}')
|
||||
shared.log.warning(f'Face: failed to load image: {e}')
|
||||
return init_images
|
||||
|
||||
def mode_change(self, mode):
|
||||
@@ -45,7 +45,7 @@ class Script(scripts.Script):
|
||||
# return signature is array of gradio components
|
||||
def ui(self, _is_img2img):
|
||||
with gr.Row():
|
||||
gr.HTML("<span>  Face module</span><br>")
|
||||
gr.HTML("<span>  Face: Multiple ID Transfers</span><br>")
|
||||
with gr.Row():
|
||||
mode = gr.Dropdown(label='Mode', choices=['None', 'FaceID', 'FaceSwap', 'InstantID', 'PhotoMaker'], value='None')
|
||||
with gr.Group(visible=False) as cfg_faceid:
|
||||
|
||||
@@ -68,7 +68,7 @@ def instant_id(p: processing.StableDiffusionProcessing, app, source_images, stre
|
||||
processing.process_init(p)
|
||||
p.init(p.all_prompts, p.all_seeds, p.all_subseeds)
|
||||
orig_prompt_attention = shared.opts.prompt_attention
|
||||
shared.opts.data['prompt_attention'] = 'Fixed attention' # otherwise need to deal with class_tokens_mask
|
||||
shared.opts.data['prompt_attention'] = 'fixed' # otherwise need to deal with class_tokens_mask
|
||||
p.task_args['image_embeds'] = face_embeds[0].shape # placeholder
|
||||
p.task_args['image'] = face_images[0]
|
||||
p.task_args['controlnet_conditioning_scale'] = float(conditioning)
|
||||
|
||||
@@ -49,7 +49,7 @@ def photo_maker(p: processing.StableDiffusionProcessing, input_images, trigger,
|
||||
shared.sd_model.to(dtype=devices.dtype)
|
||||
|
||||
orig_prompt_attention = shared.opts.prompt_attention
|
||||
shared.opts.data['prompt_attention'] = 'Fixed attention' # otherwise need to deal with class_tokens_mask
|
||||
shared.opts.data['prompt_attention'] = 'fixed' # otherwise need to deal with class_tokens_mask
|
||||
p.task_args['input_id_images'] = input_images
|
||||
p.task_args['start_merge_step'] = int(start * p.steps)
|
||||
p.task_args['prompt'] = p.all_prompts[0] if p.all_prompts is not None else p.prompt
|
||||
|
||||
+46
-42
@@ -40,8 +40,6 @@ def atomically_save_image():
|
||||
except Exception:
|
||||
shared.log.warning(f'Save: unknown image format: {extension}')
|
||||
image_format = 'JPEG'
|
||||
if shared.opts.image_watermark_enabled or (shared.opts.image_watermark_position != 'none' and shared.opts.image_watermark_image != ''):
|
||||
image = set_watermark(image, shared.opts.image_watermark)
|
||||
exifinfo = (exifinfo or "") if shared.opts.image_metadata else ""
|
||||
# additional metadata saved in files
|
||||
if shared.opts.save_txt and len(exifinfo) > 0:
|
||||
@@ -153,6 +151,11 @@ def save_image(image,
|
||||
info = image.info.get(pnginfo_section_name, '')
|
||||
if info is not None:
|
||||
pnginfo[pnginfo_section_name] = info
|
||||
|
||||
wm_text = getattr(p, 'watermark_text', shared.opts.image_watermark)
|
||||
wm_image = getattr(p, 'watermark_image', shared.opts.image_watermark_image)
|
||||
image = set_watermark(image, wm_text, wm_image)
|
||||
|
||||
params = script_callbacks.ImageSaveParams(image, p, filename, pnginfo)
|
||||
params.filename = namegen.sanitize(filename)
|
||||
dirname = os.path.dirname(params.filename)
|
||||
@@ -361,53 +364,54 @@ def flatten(img, bgcolor):
|
||||
return img.convert('RGB')
|
||||
|
||||
|
||||
def draw_overlay(im, text):
|
||||
def draw_overlay(im, text: str = '', y_offset: int = 0):
|
||||
d = ImageDraw.Draw(im)
|
||||
fontsize = (im.width + im.height) // 50
|
||||
font = get_font(fontsize)
|
||||
d.text((fontsize//2, fontsize//2), text, font=font, fill=shared.opts.font_color)
|
||||
d.text((fontsize//2, fontsize//2 + y_offset), text, font=font, fill=shared.opts.font_color)
|
||||
return im
|
||||
|
||||
|
||||
def set_watermark(image, watermark):
|
||||
if shared.opts.image_watermark_position != 'none': # visible watermark
|
||||
wm_image = None
|
||||
try:
|
||||
wm_image = Image.open(shared.opts.image_watermark_image)
|
||||
if wm_image.mode != 'RGBA':
|
||||
wm_image = wm_image.convert('RGBA')
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Set image watermark: fn="{shared.opts.image_watermark_image}" {e}')
|
||||
if wm_image is not None:
|
||||
if shared.opts.image_watermark_position == 'top/left':
|
||||
position = (0, 0)
|
||||
elif shared.opts.image_watermark_position == 'top/right':
|
||||
position = (image.width - wm_image.width, 0)
|
||||
elif shared.opts.image_watermark_position == 'bottom/left':
|
||||
position = (0, image.height - wm_image.height)
|
||||
elif shared.opts.image_watermark_position == 'bottom/right':
|
||||
position = (image.width - wm_image.width, image.height - wm_image.height)
|
||||
elif shared.opts.image_watermark_position == 'center':
|
||||
position = ((image.width - wm_image.width) // 2, (image.height - wm_image.height) // 2)
|
||||
else:
|
||||
position = (random.randint(0, image.width - wm_image.width), random.randint(0, image.height - wm_image.height))
|
||||
def set_watermark(image, wm_text: str = None, wm_image: Image.Image = None):
|
||||
if shared.opts.image_watermark_position != 'none' and wm_image is not None: # visible watermark
|
||||
if isinstance(wm_image, str):
|
||||
try:
|
||||
for x in range(wm_image.width):
|
||||
for y in range(wm_image.height):
|
||||
rgba = wm_image.getpixel((x, y))
|
||||
orig = image.getpixel((x+position[0], y+position[1]))
|
||||
# alpha blend
|
||||
a = rgba[3] / 255
|
||||
r = int(rgba[0] * a + orig[0] * (1 - a))
|
||||
g = int(rgba[1] * a + orig[1] * (1 - a))
|
||||
b = int(rgba[2] * a + orig[2] * (1 - a))
|
||||
if not a == 0:
|
||||
image.putpixel((x+position[0], y+position[1]), (r, g, b))
|
||||
shared.log.debug(f'Set image watermark: fn="{shared.opts.image_watermark_image}" image={wm_image} position={position}')
|
||||
wm_image = Image.open(wm_image)
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Set image watermark: image={wm_image} {e}')
|
||||
return image
|
||||
if isinstance(wm_image, Image.Image):
|
||||
if wm_image.mode != 'RGBA':
|
||||
wm_image = wm_image.convert('RGBA')
|
||||
if shared.opts.image_watermark_position == 'top/left':
|
||||
position = (0, 0)
|
||||
elif shared.opts.image_watermark_position == 'top/right':
|
||||
position = (image.width - wm_image.width, 0)
|
||||
elif shared.opts.image_watermark_position == 'bottom/left':
|
||||
position = (0, image.height - wm_image.height)
|
||||
elif shared.opts.image_watermark_position == 'bottom/right':
|
||||
position = (image.width - wm_image.width, image.height - wm_image.height)
|
||||
elif shared.opts.image_watermark_position == 'center':
|
||||
position = ((image.width - wm_image.width) // 2, (image.height - wm_image.height) // 2)
|
||||
else:
|
||||
position = (random.randint(0, image.width - wm_image.width), random.randint(0, image.height - wm_image.height))
|
||||
try:
|
||||
for x in range(wm_image.width):
|
||||
for y in range(wm_image.height):
|
||||
rgba = wm_image.getpixel((x, y))
|
||||
orig = image.getpixel((x+position[0], y+position[1]))
|
||||
# alpha blend
|
||||
a = rgba[3] / 255
|
||||
r = int(rgba[0] * a + orig[0] * (1 - a))
|
||||
g = int(rgba[1] * a + orig[1] * (1 - a))
|
||||
b = int(rgba[2] * a + orig[2] * (1 - a))
|
||||
if not a == 0:
|
||||
image.putpixel((x+position[0], y+position[1]), (r, g, b))
|
||||
shared.log.debug(f'Set image watermark: image={wm_image} position={position}')
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Set image watermark: image={wm_image} {e}')
|
||||
|
||||
if shared.opts.image_watermark_enabled: # invisible watermark
|
||||
if shared.opts.image_watermark_enabled and wm_text is not None: # invisible watermark
|
||||
from imwatermark import WatermarkEncoder
|
||||
wm_type = 'bytes'
|
||||
wm_method = 'dwtDctSvd'
|
||||
@@ -416,16 +420,16 @@ def set_watermark(image, watermark):
|
||||
info = image.info
|
||||
data = np.asarray(image)
|
||||
encoder = WatermarkEncoder()
|
||||
text = f"{watermark:<{length}}"[:length]
|
||||
text = f"{wm_text:<{length}}"[:length]
|
||||
bytearr = text.encode(encoding='ascii', errors='ignore')
|
||||
try:
|
||||
encoder.set_watermark(wm_type, bytearr)
|
||||
encoded = encoder.encode(data, wm_method)
|
||||
image = Image.fromarray(encoded)
|
||||
image.info = info
|
||||
shared.log.debug(f'Set invisible watermark: {watermark} method={wm_method} bits={wm_length}')
|
||||
shared.log.debug(f'Set invisible watermark: {wm_text} method={wm_method} bits={wm_length}')
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Set invisible watermark error: {watermark} method={wm_method} bits={wm_length} {e}')
|
||||
shared.log.warning(f'Set invisible watermark error: {wm_text} method={wm_method} bits={wm_length} {e}')
|
||||
|
||||
return image
|
||||
|
||||
|
||||
@@ -34,10 +34,10 @@ class FilenameGenerator:
|
||||
'timestamp': lambda self: getattr(self.p, "job_timestamp", shared.state.job_timestamp),
|
||||
'job_timestamp': lambda self: getattr(self.p, "job_timestamp", shared.state.job_timestamp),
|
||||
|
||||
'model': lambda self: shared.sd_model.sd_checkpoint_info.title if shared.sd_loaded else '',
|
||||
'model_shortname': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded else '',
|
||||
'model_name': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded else '',
|
||||
'model_hash': lambda self: shared.sd_model.sd_checkpoint_info.shorthash if shared.sd_loaded else '',
|
||||
'model': lambda self: shared.sd_model.sd_checkpoint_info.title if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '',
|
||||
'model_shortname': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '',
|
||||
'model_name': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '',
|
||||
'model_hash': lambda self: shared.sd_model.sd_checkpoint_info.shorthash if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '',
|
||||
|
||||
'prompt': lambda self: self.prompt_full(),
|
||||
'prompt_no_styles': lambda self: self.prompt_no_style(),
|
||||
|
||||
+19
-1
@@ -137,6 +137,7 @@ def img2img(id_task: str, state: str, mode: int,
|
||||
inpaint_full_res, inpaint_full_res_padding, inpainting_mask_invert,
|
||||
img2img_batch_files, img2img_batch_input_dir, img2img_batch_output_dir, img2img_batch_inpaint_mask_dir,
|
||||
hdr_mode, hdr_brightness, hdr_color, hdr_sharpen, hdr_clamp, hdr_boundary, hdr_threshold, hdr_maximize, hdr_max_center, hdr_max_boundry, hdr_color_picker, hdr_tint_ratio,
|
||||
enable_hr, hr_sampler_index, hr_denoising_strength, hr_resize_mode, hr_resize_context, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, hr_refiner_start, refiner_prompt, refiner_negative,
|
||||
override_settings_texts,
|
||||
*args): # pylint: disable=unused-argument
|
||||
|
||||
@@ -214,7 +215,6 @@ def img2img(id_task: str, state: str, mode: int,
|
||||
subseed_strength=subseed_strength,
|
||||
seed_resize_from_h=seed_resize_from_h,
|
||||
seed_resize_from_w=seed_resize_from_w,
|
||||
seed_enable_extras=True,
|
||||
sampler_name = processing.get_sampler_name(sampler_index, img=True),
|
||||
batch_size=batch_size,
|
||||
n_iter=n_iter,
|
||||
@@ -247,6 +247,23 @@ def img2img(id_task: str, state: str, mode: int,
|
||||
inpainting_mask_invert=inpainting_mask_invert,
|
||||
hdr_mode=hdr_mode, hdr_brightness=hdr_brightness, hdr_color=hdr_color, hdr_sharpen=hdr_sharpen, hdr_clamp=hdr_clamp,
|
||||
hdr_boundary=hdr_boundary, hdr_threshold=hdr_threshold, hdr_maximize=hdr_maximize, hdr_max_center=hdr_max_center, hdr_max_boundry=hdr_max_boundry, hdr_color_picker=hdr_color_picker, hdr_tint_ratio=hdr_tint_ratio,
|
||||
# refiner
|
||||
enable_hr=enable_hr,
|
||||
hr_denoising_strength=hr_denoising_strength,
|
||||
hr_scale=hr_scale,
|
||||
hr_resize_mode=hr_resize_mode,
|
||||
hr_resize_context=hr_resize_context,
|
||||
hr_upscaler=hr_upscaler,
|
||||
hr_force=hr_force,
|
||||
hr_second_pass_steps=hr_second_pass_steps,
|
||||
hr_resize_x=hr_resize_x,
|
||||
hr_resize_y=hr_resize_y,
|
||||
hr_sampler_name = processing.get_sampler_name(hr_sampler_index),
|
||||
refiner_steps=refiner_steps,
|
||||
hr_refiner_start=hr_refiner_start,
|
||||
refiner_prompt=refiner_prompt,
|
||||
refiner_negative=refiner_negative,
|
||||
# override
|
||||
override_settings=override_settings,
|
||||
)
|
||||
p.scripts = modules.scripts.scripts_img2img
|
||||
@@ -267,6 +284,7 @@ def img2img(id_task: str, state: str, mode: int,
|
||||
processed = modules.scripts.scripts_img2img.run(p, *args)
|
||||
if processed is None:
|
||||
processed = processing.process_images(p)
|
||||
processed = modules.scripts.scripts_img2img.after(p, processed, *args)
|
||||
p.close()
|
||||
generation_info_js = processed.js() if processed is not None else ''
|
||||
if processed is None:
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from .sdxl_instantir import InstantIRPipeline
|
||||
from .lcm_single_step_scheduler import LCMSingleStepScheduler
|
||||
from .ip_adapter.utils import init_adapter_in_unet, load_adapter_to_pipe
|
||||
@@ -0,0 +1,983 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.models.attention_processor import (
|
||||
ADDED_KV_ATTENTION_PROCESSORS,
|
||||
CROSS_ATTENTION_PROCESSORS,
|
||||
AttentionProcessor,
|
||||
AttnAddedKVProcessor,
|
||||
AttnProcessor,
|
||||
)
|
||||
from diffusers.models.embeddings import TextImageProjection, TextImageTimeEmbedding, TextTimeEmbedding, TimestepEmbedding, Timesteps
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.unets.unet_2d_blocks import (
|
||||
CrossAttnDownBlock2D,
|
||||
DownBlock2D,
|
||||
UNetMidBlock2D,
|
||||
UNetMidBlock2DCrossAttn,
|
||||
get_down_block,
|
||||
)
|
||||
from diffusers.models.unets.unet_2d_condition import UNet2DConditionModel
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class ZeroConv(nn.Module):
|
||||
def __init__(self, label_nc, norm_nc, mask=False):
|
||||
super().__init__()
|
||||
self.zero_conv = zero_module(nn.Conv2d(label_nc+norm_nc, norm_nc, 1, 1, 0))
|
||||
self.mask = mask
|
||||
|
||||
def forward(self, hidden_states, h_ori=None):
|
||||
# with torch.cuda.amp.autocast(enabled=False, dtype=torch.float32):
|
||||
c, h = hidden_states
|
||||
if not self.mask:
|
||||
h = self.zero_conv(torch.cat([c, h], dim=1))
|
||||
else:
|
||||
h = self.zero_conv(torch.cat([c, h], dim=1)) * torch.zeros_like(h)
|
||||
if h_ori is not None:
|
||||
h = torch.cat([h_ori, h], dim=1)
|
||||
return h
|
||||
|
||||
|
||||
class SFT(nn.Module):
|
||||
def __init__(self, label_nc, norm_nc, mask=False):
|
||||
super().__init__()
|
||||
|
||||
# param_free_norm_type = str(parsed.group(1))
|
||||
ks = 3
|
||||
pw = ks // 2
|
||||
|
||||
self.mask = mask
|
||||
|
||||
nhidden = 128
|
||||
|
||||
self.mlp_shared = nn.Sequential(
|
||||
nn.Conv2d(label_nc, nhidden, kernel_size=ks, padding=pw),
|
||||
nn.SiLU()
|
||||
)
|
||||
self.mul = nn.Conv2d(nhidden, norm_nc, kernel_size=ks, padding=pw)
|
||||
self.add = nn.Conv2d(nhidden, norm_nc, kernel_size=ks, padding=pw)
|
||||
|
||||
def forward(self, hidden_states, mask=False):
|
||||
|
||||
c, h = hidden_states
|
||||
mask = mask or self.mask
|
||||
assert mask is False
|
||||
|
||||
actv = self.mlp_shared(c)
|
||||
gamma = self.mul(actv)
|
||||
beta = self.add(actv)
|
||||
|
||||
if self.mask:
|
||||
gamma = gamma * torch.zeros_like(gamma)
|
||||
beta = beta * torch.zeros_like(beta)
|
||||
# gamma_ori, gamma_res = torch.split(gamma, [h_ori_c, h_c], dim=1)
|
||||
# beta_ori, beta_res = torch.split(beta, [h_ori_c, h_c], dim=1)
|
||||
# print(gamma_ori.mean(), gamma_res.mean(), beta_ori.mean(), beta_res.mean())
|
||||
h = h * (gamma + 1) + beta
|
||||
# sample_ori, sample_res = torch.split(h, [h_ori_c, h_c], dim=1)
|
||||
# print(sample_ori.mean(), sample_res.mean())
|
||||
|
||||
return h
|
||||
|
||||
|
||||
@dataclass
|
||||
class AggregatorOutput(BaseOutput):
|
||||
"""
|
||||
The output of [`Aggregator`].
|
||||
|
||||
Args:
|
||||
down_block_res_samples (`tuple[torch.Tensor]`):
|
||||
A tuple of downsample activations at different resolutions for each downsampling block. Each tensor should
|
||||
be of shape `(batch_size, channel * resolution, height //resolution, width // resolution)`. Output can be
|
||||
used to condition the original UNet's downsampling activations.
|
||||
mid_down_block_re_sample (`torch.Tensor`):
|
||||
The activation of the midde block (the lowest sample resolution). Each tensor should be of shape
|
||||
`(batch_size, channel * lowest_resolution, height // lowest_resolution, width // lowest_resolution)`.
|
||||
Output can be used to condition the original UNet's middle block activation.
|
||||
"""
|
||||
|
||||
down_block_res_samples: Tuple[torch.Tensor]
|
||||
mid_block_res_sample: torch.Tensor
|
||||
|
||||
|
||||
class ConditioningEmbedding(nn.Module):
|
||||
"""
|
||||
Quoting from https://arxiv.org/abs/2302.05543: "Stable Diffusion uses a pre-processing method similar to VQ-GAN
|
||||
[11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized
|
||||
training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the
|
||||
convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides
|
||||
(activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full
|
||||
model) to encode image-space conditions ... into feature maps ..."
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
conditioning_embedding_channels: int,
|
||||
conditioning_channels: int = 3,
|
||||
block_out_channels: Tuple[int, ...] = (16, 32, 96, 256),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
|
||||
|
||||
self.blocks = nn.ModuleList([])
|
||||
|
||||
for i in range(len(block_out_channels) - 1):
|
||||
channel_in = block_out_channels[i]
|
||||
channel_out = block_out_channels[i + 1]
|
||||
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
|
||||
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
|
||||
|
||||
self.conv_out = zero_module(
|
||||
nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
|
||||
)
|
||||
|
||||
def forward(self, conditioning):
|
||||
embedding = self.conv_in(conditioning)
|
||||
embedding = F.silu(embedding)
|
||||
|
||||
for block in self.blocks:
|
||||
embedding = block(embedding)
|
||||
embedding = F.silu(embedding)
|
||||
|
||||
embedding = self.conv_out(embedding)
|
||||
|
||||
return embedding
|
||||
|
||||
|
||||
class Aggregator(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
||||
"""
|
||||
Aggregator model.
|
||||
|
||||
Args:
|
||||
in_channels (`int`, defaults to 4):
|
||||
The number of channels in the input sample.
|
||||
flip_sin_to_cos (`bool`, defaults to `True`):
|
||||
Whether to flip the sin to cos in the time embedding.
|
||||
freq_shift (`int`, defaults to 0):
|
||||
The frequency shift to apply to the time embedding.
|
||||
down_block_types (`tuple[str]`, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`):
|
||||
The tuple of downsample blocks to use.
|
||||
only_cross_attention (`Union[bool, Tuple[bool]]`, defaults to `False`):
|
||||
block_out_channels (`tuple[int]`, defaults to `(320, 640, 1280, 1280)`):
|
||||
The tuple of output channels for each block.
|
||||
layers_per_block (`int`, defaults to 2):
|
||||
The number of layers per block.
|
||||
downsample_padding (`int`, defaults to 1):
|
||||
The padding to use for the downsampling convolution.
|
||||
mid_block_scale_factor (`float`, defaults to 1):
|
||||
The scale factor to use for the mid block.
|
||||
act_fn (`str`, defaults to "silu"):
|
||||
The activation function to use.
|
||||
norm_num_groups (`int`, *optional*, defaults to 32):
|
||||
The number of groups to use for the normalization. If None, normalization and activation layers is skipped
|
||||
in post-processing.
|
||||
norm_eps (`float`, defaults to 1e-5):
|
||||
The epsilon to use for the normalization.
|
||||
cross_attention_dim (`int`, defaults to 1280):
|
||||
The dimension of the cross attention features.
|
||||
transformer_layers_per_block (`int` or `Tuple[int]`, *optional*, defaults to 1):
|
||||
The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for
|
||||
[`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.CrossAttnUpBlock2D`],
|
||||
[`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`].
|
||||
encoder_hid_dim (`int`, *optional*, defaults to None):
|
||||
If `encoder_hid_dim_type` is defined, `encoder_hidden_states` will be projected from `encoder_hid_dim`
|
||||
dimension to `cross_attention_dim`.
|
||||
encoder_hid_dim_type (`str`, *optional*, defaults to `None`):
|
||||
If given, the `encoder_hidden_states` and potentially other embeddings are down-projected to text
|
||||
embeddings of dimension `cross_attention` according to `encoder_hid_dim_type`.
|
||||
attention_head_dim (`Union[int, Tuple[int]]`, defaults to 8):
|
||||
The dimension of the attention heads.
|
||||
use_linear_projection (`bool`, defaults to `False`):
|
||||
class_embed_type (`str`, *optional*, defaults to `None`):
|
||||
The type of class embedding to use which is ultimately summed with the time embeddings. Choose from None,
|
||||
`"timestep"`, `"identity"`, `"projection"`, or `"simple_projection"`.
|
||||
addition_embed_type (`str`, *optional*, defaults to `None`):
|
||||
Configures an optional embedding which will be summed with the time embeddings. Choose from `None` or
|
||||
"text". "text" will use the `TextTimeEmbedding` layer.
|
||||
num_class_embeds (`int`, *optional*, defaults to 0):
|
||||
Input dimension of the learnable embedding matrix to be projected to `time_embed_dim`, when performing
|
||||
class conditioning with `class_embed_type` equal to `None`.
|
||||
upcast_attention (`bool`, defaults to `False`):
|
||||
resnet_time_scale_shift (`str`, defaults to `"default"`):
|
||||
Time scale shift config for ResNet blocks (see `ResnetBlock2D`). Choose from `default` or `scale_shift`.
|
||||
projection_class_embeddings_input_dim (`int`, *optional*, defaults to `None`):
|
||||
The dimension of the `class_labels` input when `class_embed_type="projection"`. Required when
|
||||
`class_embed_type="projection"`.
|
||||
controlnet_conditioning_channel_order (`str`, defaults to `"rgb"`):
|
||||
The channel order of conditional image. Will convert to `rgb` if it's `bgr`.
|
||||
conditioning_embedding_out_channels (`tuple[int]`, *optional*, defaults to `(16, 32, 96, 256)`):
|
||||
The tuple of output channel for each block in the `conditioning_embedding` layer.
|
||||
global_pool_conditions (`bool`, defaults to `False`):
|
||||
TODO(Patrick) - unused parameter.
|
||||
addition_embed_type_num_heads (`int`, defaults to 64):
|
||||
The number of heads to use for the `TextTimeEmbedding` layer.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 4,
|
||||
conditioning_channels: int = 3,
|
||||
flip_sin_to_cos: bool = True,
|
||||
freq_shift: int = 0,
|
||||
down_block_types: Tuple[str, ...] = (
|
||||
"CrossAttnDownBlock2D",
|
||||
"CrossAttnDownBlock2D",
|
||||
"CrossAttnDownBlock2D",
|
||||
"DownBlock2D",
|
||||
),
|
||||
mid_block_type: Optional[str] = "UNetMidBlock2DCrossAttn",
|
||||
only_cross_attention: Union[bool, Tuple[bool]] = False,
|
||||
block_out_channels: Tuple[int, ...] = (320, 640, 1280, 1280),
|
||||
layers_per_block: int = 2,
|
||||
downsample_padding: int = 1,
|
||||
mid_block_scale_factor: float = 1,
|
||||
act_fn: str = "silu",
|
||||
norm_num_groups: Optional[int] = 32,
|
||||
norm_eps: float = 1e-5,
|
||||
cross_attention_dim: int = 1280,
|
||||
transformer_layers_per_block: Union[int, Tuple[int, ...]] = 1,
|
||||
encoder_hid_dim: Optional[int] = None,
|
||||
encoder_hid_dim_type: Optional[str] = None,
|
||||
attention_head_dim: Union[int, Tuple[int, ...]] = 8,
|
||||
num_attention_heads: Optional[Union[int, Tuple[int, ...]]] = None,
|
||||
use_linear_projection: bool = False,
|
||||
class_embed_type: Optional[str] = None,
|
||||
addition_embed_type: Optional[str] = None,
|
||||
addition_time_embed_dim: Optional[int] = None,
|
||||
num_class_embeds: Optional[int] = None,
|
||||
upcast_attention: bool = False,
|
||||
resnet_time_scale_shift: str = "default",
|
||||
projection_class_embeddings_input_dim: Optional[int] = None,
|
||||
controlnet_conditioning_channel_order: str = "rgb",
|
||||
conditioning_embedding_out_channels: Optional[Tuple[int, ...]] = (16, 32, 96, 256),
|
||||
global_pool_conditions: bool = False,
|
||||
addition_embed_type_num_heads: int = 64,
|
||||
pad_concat: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# If `num_attention_heads` is not defined (which is the case for most models)
|
||||
# it will default to `attention_head_dim`. This looks weird upon first reading it and it is.
|
||||
# The reason for this behavior is to correct for incorrectly named variables that were introduced
|
||||
# when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131
|
||||
# Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking
|
||||
# which is why we correct for the naming here.
|
||||
num_attention_heads = num_attention_heads or attention_head_dim
|
||||
self.pad_concat = pad_concat
|
||||
|
||||
# Check inputs
|
||||
if len(block_out_channels) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if isinstance(transformer_layers_per_block, int):
|
||||
transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types)
|
||||
|
||||
# input
|
||||
conv_in_kernel = 3
|
||||
conv_in_padding = (conv_in_kernel - 1) // 2
|
||||
self.conv_in = nn.Conv2d(
|
||||
in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding
|
||||
)
|
||||
|
||||
# time
|
||||
time_embed_dim = block_out_channels[0] * 4
|
||||
self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift)
|
||||
timestep_input_dim = block_out_channels[0]
|
||||
self.time_embedding = TimestepEmbedding(
|
||||
timestep_input_dim,
|
||||
time_embed_dim,
|
||||
act_fn=act_fn,
|
||||
)
|
||||
|
||||
if encoder_hid_dim_type is None and encoder_hid_dim is not None:
|
||||
encoder_hid_dim_type = "text_proj"
|
||||
self.register_to_config(encoder_hid_dim_type=encoder_hid_dim_type)
|
||||
logger.info("encoder_hid_dim_type defaults to 'text_proj' as `encoder_hid_dim` is defined.")
|
||||
|
||||
if encoder_hid_dim is None and encoder_hid_dim_type is not None:
|
||||
raise ValueError(
|
||||
f"`encoder_hid_dim` has to be defined when `encoder_hid_dim_type` is set to {encoder_hid_dim_type}."
|
||||
)
|
||||
|
||||
if encoder_hid_dim_type == "text_proj":
|
||||
self.encoder_hid_proj = nn.Linear(encoder_hid_dim, cross_attention_dim)
|
||||
elif encoder_hid_dim_type == "text_image_proj":
|
||||
# image_embed_dim DOESN'T have to be `cross_attention_dim`. To not clutter the __init__ too much
|
||||
# they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use
|
||||
# case when `addition_embed_type == "text_image_proj"` (Kandinsky 2.1)`
|
||||
self.encoder_hid_proj = TextImageProjection(
|
||||
text_embed_dim=encoder_hid_dim,
|
||||
image_embed_dim=cross_attention_dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
)
|
||||
|
||||
elif encoder_hid_dim_type is not None:
|
||||
raise ValueError(
|
||||
f"encoder_hid_dim_type: {encoder_hid_dim_type} must be None, 'text_proj' or 'text_image_proj'."
|
||||
)
|
||||
else:
|
||||
self.encoder_hid_proj = None
|
||||
|
||||
# class embedding
|
||||
if class_embed_type is None and num_class_embeds is not None:
|
||||
self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim)
|
||||
elif class_embed_type == "timestep":
|
||||
self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim)
|
||||
elif class_embed_type == "identity":
|
||||
self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim)
|
||||
elif class_embed_type == "projection":
|
||||
if projection_class_embeddings_input_dim is None:
|
||||
raise ValueError(
|
||||
"`class_embed_type`: 'projection' requires `projection_class_embeddings_input_dim` be set"
|
||||
)
|
||||
# The projection `class_embed_type` is the same as the timestep `class_embed_type` except
|
||||
# 1. the `class_labels` inputs are not first converted to sinusoidal embeddings
|
||||
# 2. it projects from an arbitrary input dimension.
|
||||
#
|
||||
# Note that `TimestepEmbedding` is quite general, being mainly linear layers and activations.
|
||||
# When used for embedding actual timesteps, the timesteps are first converted to sinusoidal embeddings.
|
||||
# As a result, `TimestepEmbedding` can be passed arbitrary vectors.
|
||||
self.class_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim)
|
||||
else:
|
||||
self.class_embedding = None
|
||||
|
||||
if addition_embed_type == "text":
|
||||
if encoder_hid_dim is not None:
|
||||
text_time_embedding_from_dim = encoder_hid_dim
|
||||
else:
|
||||
text_time_embedding_from_dim = cross_attention_dim
|
||||
|
||||
self.add_embedding = TextTimeEmbedding(
|
||||
text_time_embedding_from_dim, time_embed_dim, num_heads=addition_embed_type_num_heads
|
||||
)
|
||||
elif addition_embed_type == "text_image":
|
||||
# text_embed_dim and image_embed_dim DON'T have to be `cross_attention_dim`. To not clutter the __init__ too much
|
||||
# they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use
|
||||
# case when `addition_embed_type == "text_image"` (Kandinsky 2.1)`
|
||||
self.add_embedding = TextImageTimeEmbedding(
|
||||
text_embed_dim=cross_attention_dim, image_embed_dim=cross_attention_dim, time_embed_dim=time_embed_dim
|
||||
)
|
||||
elif addition_embed_type == "text_time":
|
||||
self.add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift)
|
||||
self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim)
|
||||
|
||||
elif addition_embed_type is not None:
|
||||
raise ValueError(f"addition_embed_type: {addition_embed_type} must be None, 'text' or 'text_image'.")
|
||||
|
||||
# control net conditioning embedding
|
||||
self.ref_conv_in = nn.Conv2d(
|
||||
in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding
|
||||
)
|
||||
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
self.controlnet_down_blocks = nn.ModuleList([])
|
||||
|
||||
if isinstance(only_cross_attention, bool):
|
||||
only_cross_attention = [only_cross_attention] * len(down_block_types)
|
||||
|
||||
if isinstance(attention_head_dim, int):
|
||||
attention_head_dim = (attention_head_dim,) * len(down_block_types)
|
||||
|
||||
if isinstance(num_attention_heads, int):
|
||||
num_attention_heads = (num_attention_heads,) * len(down_block_types)
|
||||
|
||||
# down
|
||||
output_channel = block_out_channels[0]
|
||||
|
||||
# controlnet_block = ZeroConv(output_channel, output_channel)
|
||||
controlnet_block = nn.Sequential(
|
||||
SFT(output_channel, output_channel),
|
||||
zero_module(nn.Conv2d(output_channel, output_channel, kernel_size=1))
|
||||
)
|
||||
self.controlnet_down_blocks.append(controlnet_block)
|
||||
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
down_block = get_down_block(
|
||||
down_block_type,
|
||||
num_layers=layers_per_block,
|
||||
transformer_layers_per_block=transformer_layers_per_block[i],
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
temb_channels=time_embed_dim,
|
||||
add_downsample=not is_final_block,
|
||||
resnet_eps=norm_eps,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
num_attention_heads=num_attention_heads[i],
|
||||
attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel,
|
||||
downsample_padding=downsample_padding,
|
||||
use_linear_projection=use_linear_projection,
|
||||
only_cross_attention=only_cross_attention[i],
|
||||
upcast_attention=upcast_attention,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
for _ in range(layers_per_block):
|
||||
# controlnet_block = ZeroConv(output_channel, output_channel)
|
||||
controlnet_block = nn.Sequential(
|
||||
SFT(output_channel, output_channel),
|
||||
zero_module(nn.Conv2d(output_channel, output_channel, kernel_size=1))
|
||||
)
|
||||
self.controlnet_down_blocks.append(controlnet_block)
|
||||
|
||||
if not is_final_block:
|
||||
# controlnet_block = ZeroConv(output_channel, output_channel)
|
||||
controlnet_block = nn.Sequential(
|
||||
SFT(output_channel, output_channel),
|
||||
zero_module(nn.Conv2d(output_channel, output_channel, kernel_size=1))
|
||||
)
|
||||
self.controlnet_down_blocks.append(controlnet_block)
|
||||
|
||||
# mid
|
||||
mid_block_channel = block_out_channels[-1]
|
||||
|
||||
# controlnet_block = ZeroConv(mid_block_channel, mid_block_channel)
|
||||
controlnet_block = nn.Sequential(
|
||||
SFT(mid_block_channel, mid_block_channel),
|
||||
zero_module(nn.Conv2d(mid_block_channel, mid_block_channel, kernel_size=1))
|
||||
)
|
||||
self.controlnet_mid_block = controlnet_block
|
||||
|
||||
if mid_block_type == "UNetMidBlock2DCrossAttn":
|
||||
self.mid_block = UNetMidBlock2DCrossAttn(
|
||||
transformer_layers_per_block=transformer_layers_per_block[-1],
|
||||
in_channels=mid_block_channel,
|
||||
temb_channels=time_embed_dim,
|
||||
resnet_eps=norm_eps,
|
||||
resnet_act_fn=act_fn,
|
||||
output_scale_factor=mid_block_scale_factor,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
num_attention_heads=num_attention_heads[-1],
|
||||
resnet_groups=norm_num_groups,
|
||||
use_linear_projection=use_linear_projection,
|
||||
upcast_attention=upcast_attention,
|
||||
)
|
||||
elif mid_block_type == "UNetMidBlock2D":
|
||||
self.mid_block = UNetMidBlock2D(
|
||||
in_channels=block_out_channels[-1],
|
||||
temb_channels=time_embed_dim,
|
||||
num_layers=0,
|
||||
resnet_eps=norm_eps,
|
||||
resnet_act_fn=act_fn,
|
||||
output_scale_factor=mid_block_scale_factor,
|
||||
resnet_groups=norm_num_groups,
|
||||
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||
add_attention=False,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"unknown mid_block_type : {mid_block_type}")
|
||||
|
||||
@classmethod
|
||||
def from_unet(
|
||||
cls,
|
||||
unet: UNet2DConditionModel,
|
||||
controlnet_conditioning_channel_order: str = "rgb",
|
||||
conditioning_embedding_out_channels: Optional[Tuple[int, ...]] = (16, 32, 96, 256),
|
||||
load_weights_from_unet: bool = True,
|
||||
conditioning_channels: int = 3,
|
||||
):
|
||||
r"""
|
||||
Instantiate a [`ControlNetModel`] from [`UNet2DConditionModel`].
|
||||
|
||||
Parameters:
|
||||
unet (`UNet2DConditionModel`):
|
||||
The UNet model weights to copy to the [`ControlNetModel`]. All configuration options are also copied
|
||||
where applicable.
|
||||
"""
|
||||
transformer_layers_per_block = (
|
||||
unet.config.transformer_layers_per_block if "transformer_layers_per_block" in unet.config else 1
|
||||
)
|
||||
encoder_hid_dim = unet.config.encoder_hid_dim if "encoder_hid_dim" in unet.config else None
|
||||
encoder_hid_dim_type = unet.config.encoder_hid_dim_type if "encoder_hid_dim_type" in unet.config else None
|
||||
addition_embed_type = unet.config.addition_embed_type if "addition_embed_type" in unet.config else None
|
||||
addition_time_embed_dim = (
|
||||
unet.config.addition_time_embed_dim if "addition_time_embed_dim" in unet.config else None
|
||||
)
|
||||
|
||||
controlnet = cls(
|
||||
encoder_hid_dim=encoder_hid_dim,
|
||||
encoder_hid_dim_type=encoder_hid_dim_type,
|
||||
addition_embed_type=addition_embed_type,
|
||||
addition_time_embed_dim=addition_time_embed_dim,
|
||||
transformer_layers_per_block=transformer_layers_per_block,
|
||||
in_channels=unet.config.in_channels,
|
||||
flip_sin_to_cos=unet.config.flip_sin_to_cos,
|
||||
freq_shift=unet.config.freq_shift,
|
||||
down_block_types=unet.config.down_block_types,
|
||||
only_cross_attention=unet.config.only_cross_attention,
|
||||
block_out_channels=unet.config.block_out_channels,
|
||||
layers_per_block=unet.config.layers_per_block,
|
||||
downsample_padding=unet.config.downsample_padding,
|
||||
mid_block_scale_factor=unet.config.mid_block_scale_factor,
|
||||
act_fn=unet.config.act_fn,
|
||||
norm_num_groups=unet.config.norm_num_groups,
|
||||
norm_eps=unet.config.norm_eps,
|
||||
cross_attention_dim=unet.config.cross_attention_dim,
|
||||
attention_head_dim=unet.config.attention_head_dim,
|
||||
num_attention_heads=unet.config.num_attention_heads,
|
||||
use_linear_projection=unet.config.use_linear_projection,
|
||||
class_embed_type=unet.config.class_embed_type,
|
||||
num_class_embeds=unet.config.num_class_embeds,
|
||||
upcast_attention=unet.config.upcast_attention,
|
||||
resnet_time_scale_shift=unet.config.resnet_time_scale_shift,
|
||||
projection_class_embeddings_input_dim=unet.config.projection_class_embeddings_input_dim,
|
||||
mid_block_type=unet.config.mid_block_type,
|
||||
controlnet_conditioning_channel_order=controlnet_conditioning_channel_order,
|
||||
conditioning_embedding_out_channels=conditioning_embedding_out_channels,
|
||||
conditioning_channels=conditioning_channels,
|
||||
)
|
||||
|
||||
if load_weights_from_unet:
|
||||
controlnet.conv_in.load_state_dict(unet.conv_in.state_dict())
|
||||
controlnet.ref_conv_in.load_state_dict(unet.conv_in.state_dict())
|
||||
controlnet.time_proj.load_state_dict(unet.time_proj.state_dict())
|
||||
controlnet.time_embedding.load_state_dict(unet.time_embedding.state_dict())
|
||||
|
||||
if controlnet.class_embedding:
|
||||
controlnet.class_embedding.load_state_dict(unet.class_embedding.state_dict())
|
||||
|
||||
if hasattr(controlnet, "add_embedding"):
|
||||
controlnet.add_embedding.load_state_dict(unet.add_embedding.state_dict())
|
||||
|
||||
controlnet.down_blocks.load_state_dict(unet.down_blocks.state_dict())
|
||||
controlnet.mid_block.load_state_dict(unet.mid_block.state_dict())
|
||||
|
||||
return controlnet
|
||||
|
||||
@property
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
return processors
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
|
||||
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
if isinstance(processor, dict) and len(processor) != count:
|
||||
raise ValueError(
|
||||
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
|
||||
)
|
||||
|
||||
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
|
||||
if hasattr(module, "set_processor"):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor)
|
||||
else:
|
||||
module.set_processor(processor.pop(f"{name}.processor"))
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
|
||||
def set_default_attn_processor(self):
|
||||
"""
|
||||
Disables custom attention processors and sets the default attention implementation.
|
||||
"""
|
||||
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
processor = AttnAddedKVProcessor()
|
||||
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
|
||||
)
|
||||
|
||||
self.set_attn_processor(processor)
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attention_slice
|
||||
def set_attention_slice(self, slice_size: Union[str, int, List[int]]) -> None:
|
||||
r"""
|
||||
Enable sliced attention computation.
|
||||
|
||||
When this option is enabled, the attention module splits the input tensor in slices to compute attention in
|
||||
several steps. This is useful for saving some memory in exchange for a small decrease in speed.
|
||||
|
||||
Args:
|
||||
slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`):
|
||||
When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If
|
||||
`"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is
|
||||
provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim`
|
||||
must be a multiple of `slice_size`.
|
||||
"""
|
||||
sliceable_head_dims = []
|
||||
|
||||
def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module):
|
||||
if hasattr(module, "set_attention_slice"):
|
||||
sliceable_head_dims.append(module.sliceable_head_dim)
|
||||
|
||||
for child in module.children():
|
||||
fn_recursive_retrieve_sliceable_dims(child)
|
||||
|
||||
# retrieve number of attention layers
|
||||
for module in self.children():
|
||||
fn_recursive_retrieve_sliceable_dims(module)
|
||||
|
||||
num_sliceable_layers = len(sliceable_head_dims)
|
||||
|
||||
if slice_size == "auto":
|
||||
# half the attention head size is usually a good trade-off between
|
||||
# speed and memory
|
||||
slice_size = [dim // 2 for dim in sliceable_head_dims]
|
||||
elif slice_size == "max":
|
||||
# make smallest slice possible
|
||||
slice_size = num_sliceable_layers * [1]
|
||||
|
||||
slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size
|
||||
|
||||
if len(slice_size) != len(sliceable_head_dims):
|
||||
raise ValueError(
|
||||
f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different"
|
||||
f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}."
|
||||
)
|
||||
|
||||
for i in range(len(slice_size)):
|
||||
size = slice_size[i]
|
||||
dim = sliceable_head_dims[i]
|
||||
if size is not None and size > dim:
|
||||
raise ValueError(f"size {size} has to be smaller or equal to {dim}.")
|
||||
|
||||
# Recursively walk through all the children.
|
||||
# Any children which exposes the set_attention_slice method
|
||||
# gets the message
|
||||
def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: List[int]):
|
||||
if hasattr(module, "set_attention_slice"):
|
||||
module.set_attention_slice(slice_size.pop())
|
||||
|
||||
for child in module.children():
|
||||
fn_recursive_set_attention_slice(child, slice_size)
|
||||
|
||||
reversed_slice_size = list(reversed(slice_size))
|
||||
for module in self.children():
|
||||
fn_recursive_set_attention_slice(module, reversed_slice_size)
|
||||
|
||||
def process_encoder_hidden_states(
|
||||
self, encoder_hidden_states: torch.Tensor, added_cond_kwargs: Dict[str, Any]
|
||||
) -> torch.Tensor:
|
||||
if self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_proj":
|
||||
encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states)
|
||||
elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_image_proj":
|
||||
# Kandinsky 2.1 - style
|
||||
if "image_embeds" not in added_cond_kwargs:
|
||||
raise ValueError(
|
||||
f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'text_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`"
|
||||
)
|
||||
|
||||
image_embeds = added_cond_kwargs.get("image_embeds")
|
||||
encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states, image_embeds)
|
||||
elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "image_proj":
|
||||
# Kandinsky 2.2 - style
|
||||
if "image_embeds" not in added_cond_kwargs:
|
||||
raise ValueError(
|
||||
f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`"
|
||||
)
|
||||
image_embeds = added_cond_kwargs.get("image_embeds")
|
||||
encoder_hidden_states = self.encoder_hid_proj(image_embeds)
|
||||
elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "ip_image_proj":
|
||||
if "image_embeds" not in added_cond_kwargs:
|
||||
raise ValueError(
|
||||
f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'ip_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`"
|
||||
)
|
||||
image_embeds = added_cond_kwargs.get("image_embeds")
|
||||
image_embeds = self.encoder_hid_proj(image_embeds)
|
||||
encoder_hidden_states = (encoder_hidden_states, image_embeds)
|
||||
return encoder_hidden_states
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value: bool = False) -> None:
|
||||
if isinstance(module, (CrossAttnDownBlock2D, DownBlock2D)):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[torch.Tensor, float, int],
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
controlnet_cond: torch.FloatTensor,
|
||||
cat_dim: int = -2,
|
||||
conditioning_scale: float = 1.0,
|
||||
class_labels: Optional[torch.Tensor] = None,
|
||||
timestep_cond: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None,
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[AggregatorOutput, Tuple[Tuple[torch.FloatTensor, ...], torch.FloatTensor]]:
|
||||
"""
|
||||
The [`Aggregator`] forward method.
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The noisy input tensor.
|
||||
timestep (`Union[torch.Tensor, float, int]`):
|
||||
The number of timesteps to denoise an input.
|
||||
encoder_hidden_states (`torch.Tensor`):
|
||||
The encoder hidden states.
|
||||
controlnet_cond (`torch.FloatTensor`):
|
||||
The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`.
|
||||
conditioning_scale (`float`, defaults to `1.0`):
|
||||
The scale factor for ControlNet outputs.
|
||||
class_labels (`torch.Tensor`, *optional*, defaults to `None`):
|
||||
Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings.
|
||||
timestep_cond (`torch.Tensor`, *optional*, defaults to `None`):
|
||||
Additional conditional embeddings for timestep. If provided, the embeddings will be summed with the
|
||||
timestep_embedding passed through the `self.time_embedding` layer to obtain the final timestep
|
||||
embeddings.
|
||||
attention_mask (`torch.Tensor`, *optional*, defaults to `None`):
|
||||
An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask
|
||||
is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large
|
||||
negative values to the attention scores corresponding to "discard" tokens.
|
||||
added_cond_kwargs (`dict`):
|
||||
Additional conditions for the Stable Diffusion XL UNet.
|
||||
cross_attention_kwargs (`dict[str]`, *optional*, defaults to `None`):
|
||||
A kwargs dictionary that if specified is passed along to the `AttnProcessor`.
|
||||
return_dict (`bool`, defaults to `True`):
|
||||
Whether or not to return a [`~models.controlnet.ControlNetOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.controlnet.ControlNetOutput`] **or** `tuple`:
|
||||
If `return_dict` is `True`, a [`~models.controlnet.ControlNetOutput`] is returned, otherwise a tuple is
|
||||
returned where the first element is the sample tensor.
|
||||
"""
|
||||
# check channel order
|
||||
channel_order = self.config.controlnet_conditioning_channel_order
|
||||
|
||||
if channel_order == "rgb":
|
||||
# in rgb order by default
|
||||
...
|
||||
else:
|
||||
raise ValueError(f"unknown `controlnet_conditioning_channel_order`: {channel_order}")
|
||||
|
||||
# prepare attention_mask
|
||||
if attention_mask is not None:
|
||||
attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0
|
||||
attention_mask = attention_mask.unsqueeze(1)
|
||||
|
||||
# 1. time
|
||||
timesteps = timestep
|
||||
if not torch.is_tensor(timesteps):
|
||||
# TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can
|
||||
# This would be a good case for the `match` statement (Python 3.10+)
|
||||
is_mps = sample.device.type == "mps"
|
||||
if isinstance(timestep, float):
|
||||
dtype = torch.float32 if is_mps else torch.float64
|
||||
else:
|
||||
dtype = torch.int32 if is_mps else torch.int64
|
||||
timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device)
|
||||
elif len(timesteps.shape) == 0:
|
||||
timesteps = timesteps[None].to(sample.device)
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timesteps = timesteps.expand(sample.shape[0])
|
||||
|
||||
t_emb = self.time_proj(timesteps)
|
||||
|
||||
# timesteps does not contain any weights and will always return f32 tensors
|
||||
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||
# there might be better ways to encapsulate this.
|
||||
t_emb = t_emb.to(dtype=sample.dtype)
|
||||
|
||||
emb = self.time_embedding(t_emb, timestep_cond)
|
||||
aug_emb = None
|
||||
|
||||
if self.class_embedding is not None:
|
||||
if class_labels is None:
|
||||
raise ValueError("class_labels should be provided when num_class_embeds > 0")
|
||||
|
||||
if self.config.class_embed_type == "timestep":
|
||||
class_labels = self.time_proj(class_labels)
|
||||
|
||||
class_emb = self.class_embedding(class_labels).to(dtype=self.dtype)
|
||||
emb = emb + class_emb
|
||||
|
||||
if self.config.addition_embed_type is not None:
|
||||
if self.config.addition_embed_type == "text":
|
||||
aug_emb = self.add_embedding(encoder_hidden_states)
|
||||
|
||||
elif self.config.addition_embed_type == "text_time":
|
||||
if "text_embeds" not in added_cond_kwargs:
|
||||
raise ValueError(
|
||||
f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`"
|
||||
)
|
||||
text_embeds = added_cond_kwargs.get("text_embeds")
|
||||
if "time_ids" not in added_cond_kwargs:
|
||||
raise ValueError(
|
||||
f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`"
|
||||
)
|
||||
time_ids = added_cond_kwargs.get("time_ids")
|
||||
time_embeds = self.add_time_proj(time_ids.flatten())
|
||||
time_embeds = time_embeds.reshape((text_embeds.shape[0], -1))
|
||||
|
||||
add_embeds = torch.concat([text_embeds, time_embeds], dim=-1)
|
||||
add_embeds = add_embeds.to(emb.dtype)
|
||||
aug_emb = self.add_embedding(add_embeds)
|
||||
|
||||
emb = emb + aug_emb if aug_emb is not None else emb
|
||||
|
||||
encoder_hidden_states = self.process_encoder_hidden_states(
|
||||
encoder_hidden_states=encoder_hidden_states, added_cond_kwargs=added_cond_kwargs
|
||||
)
|
||||
|
||||
# 2. prepare input
|
||||
cond_latent = self.conv_in(sample)
|
||||
ref_latent = self.ref_conv_in(controlnet_cond)
|
||||
batch_size, channel, height, width = cond_latent.shape
|
||||
if self.pad_concat:
|
||||
if cat_dim == -2 or cat_dim == 2:
|
||||
concat_pad = torch.zeros(batch_size, channel, 1, width)
|
||||
elif cat_dim == -1 or cat_dim == 3:
|
||||
concat_pad = torch.zeros(batch_size, channel, height, 1)
|
||||
else:
|
||||
raise ValueError(f"Aggregator shall concat along spatial dimension, but is asked to concat dim: {cat_dim}.")
|
||||
concat_pad = concat_pad.to(cond_latent.device, dtype=cond_latent.dtype)
|
||||
sample = torch.cat([cond_latent, concat_pad, ref_latent], dim=cat_dim)
|
||||
else:
|
||||
sample = torch.cat([cond_latent, ref_latent], dim=cat_dim)
|
||||
|
||||
# 3. down
|
||||
down_block_res_samples = (sample,)
|
||||
for downsample_block in self.down_blocks:
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
)
|
||||
|
||||
# rebuild sample: split and concat
|
||||
if self.pad_concat:
|
||||
batch_size, channel, height, width = sample.shape
|
||||
if cat_dim == -2 or cat_dim == 2:
|
||||
cond_latent = sample[:, :, :height//2, :]
|
||||
ref_latent = sample[:, :, -(height//2):, :]
|
||||
concat_pad = torch.zeros(batch_size, channel, 1, width)
|
||||
elif cat_dim == -1 or cat_dim == 3:
|
||||
cond_latent = sample[:, :, :, :width//2]
|
||||
ref_latent = sample[:, :, :, -(width//2):]
|
||||
concat_pad = torch.zeros(batch_size, channel, height, 1)
|
||||
concat_pad = concat_pad.to(cond_latent.device, dtype=cond_latent.dtype)
|
||||
sample = torch.cat([cond_latent, concat_pad, ref_latent], dim=cat_dim)
|
||||
res_samples = res_samples[:-1] + (sample,)
|
||||
|
||||
down_block_res_samples += res_samples
|
||||
|
||||
# 4. mid
|
||||
if self.mid_block is not None:
|
||||
sample = self.mid_block(
|
||||
sample,
|
||||
emb,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
)
|
||||
|
||||
# 5. split samples and SFT.
|
||||
controlnet_down_block_res_samples = ()
|
||||
for down_block_res_sample, controlnet_block in zip(down_block_res_samples, self.controlnet_down_blocks):
|
||||
batch_size, channel, height, width = down_block_res_sample.shape
|
||||
if cat_dim == -2 or cat_dim == 2:
|
||||
cond_latent = down_block_res_sample[:, :, :height//2, :]
|
||||
ref_latent = down_block_res_sample[:, :, -(height//2):, :]
|
||||
elif cat_dim == -1 or cat_dim == 3:
|
||||
cond_latent = down_block_res_sample[:, :, :, :width//2]
|
||||
ref_latent = down_block_res_sample[:, :, :, -(width//2):]
|
||||
down_block_res_sample = controlnet_block((cond_latent, ref_latent), )
|
||||
controlnet_down_block_res_samples = controlnet_down_block_res_samples + (down_block_res_sample,)
|
||||
|
||||
down_block_res_samples = controlnet_down_block_res_samples
|
||||
|
||||
batch_size, channel, height, width = sample.shape
|
||||
if cat_dim == -2 or cat_dim == 2:
|
||||
cond_latent = sample[:, :, :height//2, :]
|
||||
ref_latent = sample[:, :, -(height//2):, :]
|
||||
elif cat_dim == -1 or cat_dim == 3:
|
||||
cond_latent = sample[:, :, :, :width//2]
|
||||
ref_latent = sample[:, :, :, -(width//2):]
|
||||
mid_block_res_sample = self.controlnet_mid_block((cond_latent, ref_latent), )
|
||||
|
||||
# 6. scaling
|
||||
down_block_res_samples = [sample*conditioning_scale for sample in down_block_res_samples]
|
||||
mid_block_res_sample = mid_block_res_sample*conditioning_scale
|
||||
|
||||
if self.config.global_pool_conditions:
|
||||
down_block_res_samples = [
|
||||
torch.mean(sample, dim=(2, 3), keepdim=True) for sample in down_block_res_samples
|
||||
]
|
||||
mid_block_res_sample = torch.mean(mid_block_res_sample, dim=(2, 3), keepdim=True)
|
||||
|
||||
if not return_dict:
|
||||
return (down_block_res_samples, mid_block_res_sample)
|
||||
|
||||
return AggregatorOutput(
|
||||
down_block_res_samples=down_block_res_samples, mid_block_res_sample=mid_block_res_sample
|
||||
)
|
||||
|
||||
|
||||
def zero_module(module):
|
||||
for p in module.parameters():
|
||||
nn.init.zeros_(p)
|
||||
return module
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,236 @@
|
||||
import os
|
||||
import torch
|
||||
from typing import List
|
||||
from collections import namedtuple, OrderedDict
|
||||
|
||||
def is_torch2_available():
|
||||
return hasattr(torch.nn.functional, "scaled_dot_product_attention")
|
||||
|
||||
if is_torch2_available():
|
||||
from .attention_processor import (
|
||||
AttnProcessor2_0 as AttnProcessor,
|
||||
)
|
||||
from .attention_processor import (
|
||||
CNAttnProcessor2_0 as CNAttnProcessor,
|
||||
)
|
||||
from .attention_processor import (
|
||||
IPAttnProcessor2_0 as IPAttnProcessor,
|
||||
)
|
||||
from .attention_processor import (
|
||||
TA_IPAttnProcessor2_0 as TA_IPAttnProcessor,
|
||||
)
|
||||
else:
|
||||
from .attention_processor import AttnProcessor, CNAttnProcessor, IPAttnProcessor, TA_IPAttnProcessor
|
||||
|
||||
|
||||
class ImageProjModel(torch.nn.Module):
|
||||
"""Projection Model"""
|
||||
|
||||
def __init__(self, cross_attention_dim=2048, clip_embeddings_dim=1280, clip_extra_context_tokens=4):
|
||||
super().__init__()
|
||||
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.clip_extra_context_tokens = clip_extra_context_tokens
|
||||
self.proj = torch.nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
|
||||
self.norm = torch.nn.LayerNorm(cross_attention_dim)
|
||||
|
||||
def forward(self, image_embeds):
|
||||
embeds = image_embeds
|
||||
clip_extra_context_tokens = self.proj(embeds).reshape(
|
||||
-1, self.clip_extra_context_tokens, self.cross_attention_dim
|
||||
)
|
||||
clip_extra_context_tokens = self.norm(clip_extra_context_tokens)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
|
||||
class MLPProjModel(torch.nn.Module):
|
||||
"""SD model with image prompt"""
|
||||
def __init__(self, cross_attention_dim=2048, clip_embeddings_dim=1280):
|
||||
super().__init__()
|
||||
|
||||
self.proj = torch.nn.Sequential(
|
||||
torch.nn.Linear(clip_embeddings_dim, clip_embeddings_dim),
|
||||
torch.nn.GELU(),
|
||||
torch.nn.Linear(clip_embeddings_dim, cross_attention_dim),
|
||||
torch.nn.LayerNorm(cross_attention_dim)
|
||||
)
|
||||
|
||||
def forward(self, image_embeds):
|
||||
clip_extra_context_tokens = self.proj(image_embeds)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
|
||||
class MultiIPAdapterImageProjection(torch.nn.Module):
|
||||
def __init__(self, IPAdapterImageProjectionLayers):
|
||||
super().__init__()
|
||||
self.image_projection_layers = torch.nn.ModuleList(IPAdapterImageProjectionLayers)
|
||||
|
||||
def forward(self, image_embeds: List[torch.FloatTensor]):
|
||||
projected_image_embeds = []
|
||||
|
||||
# currently, we accept `image_embeds` as
|
||||
# 1. a tensor (deprecated) with shape [batch_size, embed_dim] or [batch_size, sequence_length, embed_dim]
|
||||
# 2. list of `n` tensors where `n` is number of ip-adapters, each tensor can hae shape [batch_size, num_images, embed_dim] or [batch_size, num_images, sequence_length, embed_dim]
|
||||
if not isinstance(image_embeds, list):
|
||||
image_embeds = [image_embeds.unsqueeze(1)]
|
||||
|
||||
if len(image_embeds) != len(self.image_projection_layers):
|
||||
raise ValueError(
|
||||
f"image_embeds must have the same length as image_projection_layers, got {len(image_embeds)} and {len(self.image_projection_layers)}"
|
||||
)
|
||||
|
||||
for image_embed, image_projection_layer in zip(image_embeds, self.image_projection_layers):
|
||||
batch_size, num_images = image_embed.shape[0], image_embed.shape[1]
|
||||
image_embed = image_embed.reshape((batch_size * num_images,) + image_embed.shape[2:])
|
||||
image_embed = image_projection_layer(image_embed)
|
||||
# image_embed = image_embed.reshape((batch_size, num_images) + image_embed.shape[1:])
|
||||
|
||||
projected_image_embeds.append(image_embed)
|
||||
|
||||
return projected_image_embeds
|
||||
|
||||
|
||||
class IPAdapter(torch.nn.Module):
|
||||
"""IP-Adapter"""
|
||||
def __init__(self, unet, image_proj_model, adapter_modules, ckpt_path=None):
|
||||
super().__init__()
|
||||
self.unet = unet
|
||||
self.image_proj = image_proj_model
|
||||
self.ip_adapter = adapter_modules
|
||||
|
||||
if ckpt_path is not None:
|
||||
self.load_from_checkpoint(ckpt_path)
|
||||
|
||||
def forward(self, noisy_latents, timesteps, encoder_hidden_states, image_embeds):
|
||||
ip_tokens = self.image_proj(image_embeds)
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1)
|
||||
# Predict the noise residual
|
||||
noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states).sample
|
||||
return noise_pred
|
||||
|
||||
def load_from_checkpoint(self, ckpt_path: str):
|
||||
# Calculate original checksums
|
||||
orig_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()]))
|
||||
orig_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()]))
|
||||
|
||||
state_dict = torch.load(ckpt_path, map_location="cpu")
|
||||
keys = list(state_dict.keys())
|
||||
if keys != ["image_proj", "ip_adapter"]:
|
||||
state_dict = revise_state_dict(state_dict)
|
||||
|
||||
# Load state dict for image_proj_model and adapter_modules
|
||||
self.image_proj.load_state_dict(state_dict["image_proj"], strict=True)
|
||||
self.ip_adapter.load_state_dict(state_dict["ip_adapter"], strict=True)
|
||||
|
||||
# Calculate new checksums
|
||||
new_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()]))
|
||||
new_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()]))
|
||||
|
||||
# Verify if the weights have changed
|
||||
assert orig_ip_proj_sum != new_ip_proj_sum, "Weights of image_proj_model did not change!"
|
||||
assert orig_adapter_sum != new_adapter_sum, "Weights of adapter_modules did not change!"
|
||||
|
||||
|
||||
class IPAdapterPlus(torch.nn.Module):
|
||||
"""IP-Adapter"""
|
||||
def __init__(self, unet, image_proj_model, adapter_modules, ckpt_path=None):
|
||||
super().__init__()
|
||||
self.unet = unet
|
||||
self.image_proj = image_proj_model
|
||||
self.ip_adapter = adapter_modules
|
||||
|
||||
if ckpt_path is not None:
|
||||
self.load_from_checkpoint(ckpt_path)
|
||||
|
||||
def forward(self, noisy_latents, timesteps, encoder_hidden_states, image_embeds):
|
||||
ip_tokens = self.image_proj(image_embeds)
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1)
|
||||
# Predict the noise residual
|
||||
noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states).sample
|
||||
return noise_pred
|
||||
|
||||
def load_from_checkpoint(self, ckpt_path: str):
|
||||
# Calculate original checksums
|
||||
orig_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()]))
|
||||
orig_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()]))
|
||||
org_unet_sum = []
|
||||
for attn_name, attn_proc in self.unet.attn_processors.items():
|
||||
if isinstance(attn_proc, (TA_IPAttnProcessor, IPAttnProcessor)):
|
||||
org_unet_sum.append(torch.sum(torch.stack([torch.sum(p) for p in attn_proc.parameters()])))
|
||||
org_unet_sum = torch.sum(torch.stack(org_unet_sum))
|
||||
|
||||
state_dict = torch.load(ckpt_path, map_location="cpu")
|
||||
keys = list(state_dict.keys())
|
||||
if keys != ["image_proj", "ip_adapter"]:
|
||||
state_dict = revise_state_dict(state_dict)
|
||||
|
||||
# Check if 'latents' exists in both the saved state_dict and the current model's state_dict
|
||||
strict_load_image_proj_model = True
|
||||
if "latents" in state_dict["image_proj"] and "latents" in self.image_proj.state_dict():
|
||||
# Check if the shapes are mismatched
|
||||
if state_dict["image_proj"]["latents"].shape != self.image_proj.state_dict()["latents"].shape:
|
||||
print(f"Shapes of 'image_proj.latents' in checkpoint {ckpt_path} and current model do not match.")
|
||||
print("Removing 'latents' from checkpoint and loading the rest of the weights.")
|
||||
del state_dict["image_proj"]["latents"]
|
||||
strict_load_image_proj_model = False
|
||||
|
||||
# Load state dict for image_proj_model and adapter_modules
|
||||
self.image_proj.load_state_dict(state_dict["image_proj"], strict=strict_load_image_proj_model)
|
||||
missing_key, unexpected_key = self.ip_adapter.load_state_dict(state_dict["ip_adapter"], strict=False)
|
||||
if len(missing_key) > 0:
|
||||
for ms in missing_key:
|
||||
if "ln" not in ms:
|
||||
raise ValueError(f"Missing key in adapter_modules: {len(missing_key)}")
|
||||
if len(unexpected_key) > 0:
|
||||
raise ValueError(f"Unexpected key in adapter_modules: {len(unexpected_key)}")
|
||||
|
||||
# Calculate new checksums
|
||||
new_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()]))
|
||||
new_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()]))
|
||||
|
||||
# Verify if the weights loaded to unet
|
||||
unet_sum = []
|
||||
for attn_name, attn_proc in self.unet.attn_processors.items():
|
||||
if isinstance(attn_proc, (TA_IPAttnProcessor, IPAttnProcessor)):
|
||||
unet_sum.append(torch.sum(torch.stack([torch.sum(p) for p in attn_proc.parameters()])))
|
||||
unet_sum = torch.sum(torch.stack(unet_sum))
|
||||
|
||||
assert org_unet_sum != unet_sum, "Weights of adapter_modules in unet did not change!"
|
||||
assert (unet_sum - new_adapter_sum < 1e-4), "Weights of adapter_modules did not load to unet!"
|
||||
|
||||
# Verify if the weights have changed
|
||||
assert orig_ip_proj_sum != new_ip_proj_sum, "Weights of image_proj_model did not change!"
|
||||
assert orig_adapter_sum != new_adapter_sum, "Weights of adapter_mod`ules did not change!"
|
||||
|
||||
|
||||
class IPAdapterXL(IPAdapter):
|
||||
"""SDXL"""
|
||||
|
||||
def forward(self, noisy_latents, timesteps, encoder_hidden_states, unet_added_cond_kwargs, image_embeds):
|
||||
ip_tokens = self.image_proj(image_embeds)
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1)
|
||||
# Predict the noise residual
|
||||
noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states, added_cond_kwargs=unet_added_cond_kwargs).sample
|
||||
return noise_pred
|
||||
|
||||
|
||||
class IPAdapterPlusXL(IPAdapterPlus):
|
||||
"""IP-Adapter with fine-grained features"""
|
||||
|
||||
def forward(self, noisy_latents, timesteps, encoder_hidden_states, unet_added_cond_kwargs, image_embeds):
|
||||
ip_tokens = self.image_proj(image_embeds)
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1)
|
||||
# Predict the noise residual
|
||||
noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states, added_cond_kwargs=unet_added_cond_kwargs).sample
|
||||
return noise_pred
|
||||
|
||||
|
||||
class IPAdapterFull(IPAdapterPlus):
|
||||
"""IP-Adapter with full features"""
|
||||
|
||||
def init_proj(self):
|
||||
image_proj_model = MLPProjModel(
|
||||
cross_attention_dim=self.pipe.unet.config.cross_attention_dim,
|
||||
clip_embeddings_dim=self.image_encoder.config.hidden_size,
|
||||
).to(self.device, dtype=torch.float16)
|
||||
return image_proj_model
|
||||
@@ -0,0 +1,158 @@
|
||||
# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py
|
||||
# and https://github.com/lucidrains/imagen-pytorch/blob/main/imagen_pytorch/imagen_pytorch.py
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from einops.layers.torch import Rearrange
|
||||
|
||||
|
||||
# FFN
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, width = x.shape
|
||||
# (bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
||||
x = x.view(bs, length, heads, -1)
|
||||
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
||||
x = x.transpose(1, 2)
|
||||
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
||||
super().__init__()
|
||||
self.scale = dim_head**-0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
|
||||
b, l, _ = latents.shape
|
||||
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Resampler(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=1280,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=20,
|
||||
num_queries=64,
|
||||
embedding_dim=768,
|
||||
output_dim=1024,
|
||||
ff_mult=4,
|
||||
max_seq_len: int = 257, # CLIP tokens + CLS token
|
||||
apply_pos_emb: bool = False,
|
||||
num_latents_mean_pooled: int = 0, # number of latents derived from mean pooled representation of the sequence
|
||||
):
|
||||
super().__init__()
|
||||
self.pos_emb = nn.Embedding(max_seq_len, embedding_dim) if apply_pos_emb else None
|
||||
|
||||
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
|
||||
|
||||
self.proj_in = nn.Linear(embedding_dim, dim)
|
||||
|
||||
self.proj_out = nn.Linear(dim, output_dim)
|
||||
self.norm_out = nn.LayerNorm(output_dim)
|
||||
|
||||
self.to_latents_from_mean_pooled_seq = (
|
||||
nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, dim * num_latents_mean_pooled),
|
||||
Rearrange("b (n d) -> b n d", n=num_latents_mean_pooled),
|
||||
)
|
||||
if num_latents_mean_pooled > 0
|
||||
else None
|
||||
)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
if self.pos_emb is not None:
|
||||
n, device = x.shape[1], x.device
|
||||
pos_emb = self.pos_emb(torch.arange(n, device=device))
|
||||
x = x + pos_emb
|
||||
|
||||
latents = self.latents.repeat(x.size(0), 1, 1)
|
||||
|
||||
x = self.proj_in(x)
|
||||
|
||||
if self.to_latents_from_mean_pooled_seq:
|
||||
meanpooled_seq = masked_mean(x, dim=1, mask=torch.ones(x.shape[:2], device=x.device, dtype=torch.bool))
|
||||
meanpooled_latents = self.to_latents_from_mean_pooled_seq(meanpooled_seq)
|
||||
latents = torch.cat((meanpooled_latents, latents), dim=-2)
|
||||
|
||||
for attn, ff in self.layers:
|
||||
latents = attn(x, latents) + latents
|
||||
latents = ff(latents) + latents
|
||||
|
||||
latents = self.proj_out(latents)
|
||||
return self.norm_out(latents)
|
||||
|
||||
|
||||
def masked_mean(t, *, dim, mask=None):
|
||||
if mask is None:
|
||||
return t.mean(dim=dim)
|
||||
|
||||
denom = mask.sum(dim=dim, keepdim=True)
|
||||
mask = rearrange(mask, "b n -> b n 1")
|
||||
masked_t = t.masked_fill(~mask, 0.0)
|
||||
|
||||
return masked_t.sum(dim=dim) / denom.clamp(min=1e-5)
|
||||
@@ -0,0 +1,248 @@
|
||||
import torch
|
||||
from collections import namedtuple, OrderedDict
|
||||
from safetensors import safe_open
|
||||
from .attention_processor import init_attn_proc
|
||||
from .ip_adapter import MultiIPAdapterImageProjection
|
||||
from .resampler import Resampler
|
||||
from transformers import (
|
||||
AutoModel, AutoImageProcessor,
|
||||
CLIPVisionModelWithProjection, CLIPImageProcessor)
|
||||
|
||||
|
||||
def init_adapter_in_unet(
|
||||
unet,
|
||||
image_proj_model=None,
|
||||
pretrained_model_path_or_dict=None,
|
||||
adapter_tokens=64,
|
||||
embedding_dim=None,
|
||||
use_lcm=False,
|
||||
use_adaln=True,
|
||||
):
|
||||
device = unet.device
|
||||
dtype = unet.dtype
|
||||
if image_proj_model is None:
|
||||
assert embedding_dim is not None, "embedding_dim must be provided if image_proj_model is None."
|
||||
image_proj_model = Resampler(
|
||||
embedding_dim=embedding_dim,
|
||||
output_dim=unet.config.cross_attention_dim,
|
||||
num_queries=adapter_tokens,
|
||||
)
|
||||
if pretrained_model_path_or_dict is not None:
|
||||
if not isinstance(pretrained_model_path_or_dict, dict):
|
||||
if pretrained_model_path_or_dict.endswith(".safetensors"):
|
||||
state_dict = {"image_proj": {}, "ip_adapter": {}}
|
||||
with safe_open(pretrained_model_path_or_dict, framework="pt", device=unet.device) as f:
|
||||
for key in f.keys():
|
||||
if key.startswith("image_proj."):
|
||||
state_dict["image_proj"][key.replace("image_proj.", "")] = f.get_tensor(key)
|
||||
elif key.startswith("ip_adapter."):
|
||||
state_dict["ip_adapter"][key.replace("ip_adapter.", "")] = f.get_tensor(key)
|
||||
else:
|
||||
state_dict = torch.load(pretrained_model_path_or_dict, map_location=unet.device)
|
||||
else:
|
||||
state_dict = pretrained_model_path_or_dict
|
||||
keys = list(state_dict.keys())
|
||||
if "image_proj" not in keys and "ip_adapter" not in keys:
|
||||
state_dict = revise_state_dict(state_dict)
|
||||
|
||||
# Creat IP cross-attention in unet.
|
||||
attn_procs = init_attn_proc(unet, adapter_tokens, use_lcm, use_adaln)
|
||||
unet.set_attn_processor(attn_procs)
|
||||
|
||||
# Load pretrinaed model if needed.
|
||||
if pretrained_model_path_or_dict is not None:
|
||||
if "ip_adapter" in state_dict.keys():
|
||||
adapter_modules = torch.nn.ModuleList(unet.attn_processors.values())
|
||||
missing, unexpected = adapter_modules.load_state_dict(state_dict["ip_adapter"], strict=False)
|
||||
for mk in missing:
|
||||
if "ln" not in mk:
|
||||
raise ValueError(f"Missing keys in adapter_modules: {missing}")
|
||||
if "image_proj" in state_dict.keys():
|
||||
image_proj_model.load_state_dict(state_dict["image_proj"])
|
||||
|
||||
# Load image projectors into iterable ModuleList.
|
||||
image_projection_layers = []
|
||||
image_projection_layers.append(image_proj_model)
|
||||
unet.encoder_hid_proj = MultiIPAdapterImageProjection(image_projection_layers)
|
||||
|
||||
# Adjust unet config to handle addtional ip hidden states.
|
||||
unet.config.encoder_hid_dim_type = "ip_image_proj"
|
||||
unet.to(dtype=dtype, device=device)
|
||||
|
||||
|
||||
def load_adapter_to_pipe(
|
||||
pipe,
|
||||
pretrained_model_path_or_dict,
|
||||
image_encoder_or_path=None,
|
||||
feature_extractor_or_path=None,
|
||||
use_clip_encoder=False,
|
||||
adapter_tokens=64,
|
||||
use_lcm=False,
|
||||
use_adaln=True,
|
||||
):
|
||||
|
||||
if not isinstance(pretrained_model_path_or_dict, dict):
|
||||
if pretrained_model_path_or_dict.endswith(".safetensors"):
|
||||
state_dict = {"image_proj": {}, "ip_adapter": {}}
|
||||
with safe_open(pretrained_model_path_or_dict, framework="pt", device=pipe.device) as f:
|
||||
for key in f.keys():
|
||||
if key.startswith("image_proj."):
|
||||
state_dict["image_proj"][key.replace("image_proj.", "")] = f.get_tensor(key)
|
||||
elif key.startswith("ip_adapter."):
|
||||
state_dict["ip_adapter"][key.replace("ip_adapter.", "")] = f.get_tensor(key)
|
||||
else:
|
||||
state_dict = torch.load(pretrained_model_path_or_dict, map_location=pipe.device)
|
||||
else:
|
||||
state_dict = pretrained_model_path_or_dict
|
||||
keys = list(state_dict.keys())
|
||||
if "image_proj" not in keys and "ip_adapter" not in keys:
|
||||
state_dict = revise_state_dict(state_dict)
|
||||
|
||||
# load CLIP image encoder here if it has not been registered to the pipeline yet
|
||||
if image_encoder_or_path is not None:
|
||||
if isinstance(image_encoder_or_path, str):
|
||||
feature_extractor_or_path = image_encoder_or_path if feature_extractor_or_path is None else feature_extractor_or_path
|
||||
|
||||
image_encoder_or_path = (
|
||||
CLIPVisionModelWithProjection.from_pretrained(
|
||||
image_encoder_or_path
|
||||
) if use_clip_encoder else
|
||||
AutoModel.from_pretrained(image_encoder_or_path)
|
||||
)
|
||||
|
||||
if feature_extractor_or_path is not None:
|
||||
if isinstance(feature_extractor_or_path, str):
|
||||
feature_extractor_or_path = (
|
||||
CLIPImageProcessor() if use_clip_encoder else
|
||||
AutoImageProcessor.from_pretrained(feature_extractor_or_path)
|
||||
)
|
||||
|
||||
# create image encoder if it has not been registered to the pipeline yet
|
||||
if hasattr(pipe, "image_encoder") and getattr(pipe, "image_encoder", None) is None:
|
||||
image_encoder = image_encoder_or_path.to(pipe.device, dtype=pipe.dtype)
|
||||
pipe.register_modules(image_encoder=image_encoder)
|
||||
else:
|
||||
image_encoder = pipe.image_encoder
|
||||
|
||||
# create feature extractor if it has not been registered to the pipeline yet
|
||||
if hasattr(pipe, "feature_extractor") and getattr(pipe, "feature_extractor", None) is None:
|
||||
feature_extractor = feature_extractor_or_path
|
||||
pipe.register_modules(feature_extractor=feature_extractor)
|
||||
else:
|
||||
feature_extractor = pipe.feature_extractor
|
||||
|
||||
# load adapter into unet
|
||||
unet = getattr(pipe, pipe.unet_name) if not hasattr(pipe, "unet") else pipe.unet
|
||||
attn_procs = init_attn_proc(unet, adapter_tokens, use_lcm, use_adaln)
|
||||
unet.set_attn_processor(attn_procs)
|
||||
image_proj_model = Resampler(
|
||||
embedding_dim=image_encoder.config.hidden_size,
|
||||
output_dim=unet.config.cross_attention_dim,
|
||||
num_queries=adapter_tokens,
|
||||
)
|
||||
|
||||
# Load pretrinaed model if needed.
|
||||
if "ip_adapter" in state_dict.keys():
|
||||
adapter_modules = torch.nn.ModuleList(unet.attn_processors.values())
|
||||
missing, unexpected = adapter_modules.load_state_dict(state_dict["ip_adapter"], strict=False)
|
||||
for mk in missing:
|
||||
if "ln" not in mk:
|
||||
raise ValueError(f"Missing keys in adapter_modules: {missing}")
|
||||
if "image_proj" in state_dict.keys():
|
||||
image_proj_model.load_state_dict(state_dict["image_proj"])
|
||||
|
||||
# convert IP-Adapter Image Projection layers to diffusers
|
||||
image_projection_layers = []
|
||||
image_projection_layers.append(image_proj_model)
|
||||
unet.encoder_hid_proj = MultiIPAdapterImageProjection(image_projection_layers)
|
||||
|
||||
# Adjust unet config to handle addtional ip hidden states.
|
||||
unet.config.encoder_hid_dim_type = "ip_image_proj"
|
||||
unet.to(dtype=pipe.dtype, device=pipe.device)
|
||||
|
||||
|
||||
def revise_state_dict(old_state_dict_or_path, map_location="cpu"):
|
||||
new_state_dict = OrderedDict()
|
||||
new_state_dict["image_proj"] = OrderedDict()
|
||||
new_state_dict["ip_adapter"] = OrderedDict()
|
||||
if isinstance(old_state_dict_or_path, str):
|
||||
old_state_dict = torch.load(old_state_dict_or_path, map_location=map_location)
|
||||
else:
|
||||
old_state_dict = old_state_dict_or_path
|
||||
for name, weight in old_state_dict.items():
|
||||
if name.startswith("image_proj_model."):
|
||||
new_state_dict["image_proj"][name[len("image_proj_model."):]] = weight
|
||||
elif name.startswith("adapter_modules."):
|
||||
new_state_dict["ip_adapter"][name[len("adapter_modules."):]] = weight
|
||||
return new_state_dict
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.encode_image
|
||||
def encode_image(image_encoder, feature_extractor, image, device, num_images_per_prompt, output_hidden_states=None):
|
||||
dtype = next(image_encoder.parameters()).dtype
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
image = feature_extractor(image, return_tensors="pt").pixel_values
|
||||
|
||||
image = image.to(device=device, dtype=dtype)
|
||||
if output_hidden_states:
|
||||
image_enc_hidden_states = image_encoder(image, output_hidden_states=True).hidden_states[-2]
|
||||
image_enc_hidden_states = image_enc_hidden_states.repeat_interleave(num_images_per_prompt, dim=0)
|
||||
return image_enc_hidden_states
|
||||
else:
|
||||
if isinstance(image_encoder, CLIPVisionModelWithProjection):
|
||||
# CLIP image encoder.
|
||||
image_embeds = image_encoder(image).image_embeds
|
||||
else:
|
||||
# DINO image encoder.
|
||||
image_embeds = image_encoder(image).last_hidden_state
|
||||
image_embeds = image_embeds.repeat_interleave(num_images_per_prompt, dim=0)
|
||||
return image_embeds
|
||||
|
||||
|
||||
def prepare_training_image_embeds(
|
||||
image_encoder, feature_extractor,
|
||||
ip_adapter_image, ip_adapter_image_embeds,
|
||||
device, drop_rate, output_hidden_state, idx_to_replace=None
|
||||
):
|
||||
if ip_adapter_image_embeds is None:
|
||||
if not isinstance(ip_adapter_image, list):
|
||||
ip_adapter_image = [ip_adapter_image]
|
||||
|
||||
# if len(ip_adapter_image) != len(unet.encoder_hid_proj.image_projection_layers):
|
||||
# raise ValueError(
|
||||
# f"`ip_adapter_image` must have same length as the number of IP Adapters. Got {len(ip_adapter_image)} images and {len(unet.encoder_hid_proj.image_projection_layers)} IP Adapters."
|
||||
# )
|
||||
|
||||
image_embeds = []
|
||||
for single_ip_adapter_image in ip_adapter_image:
|
||||
if idx_to_replace is None:
|
||||
idx_to_replace = torch.rand(len(single_ip_adapter_image)) < drop_rate
|
||||
zero_ip_adapter_image = torch.zeros_like(single_ip_adapter_image)
|
||||
single_ip_adapter_image[idx_to_replace] = zero_ip_adapter_image[idx_to_replace]
|
||||
single_image_embeds = encode_image(
|
||||
image_encoder, feature_extractor, single_ip_adapter_image, device, 1, output_hidden_state
|
||||
)
|
||||
single_image_embeds = torch.stack([single_image_embeds], dim=1) # FIXME
|
||||
|
||||
image_embeds.append(single_image_embeds)
|
||||
else:
|
||||
repeat_dims = [1]
|
||||
image_embeds = []
|
||||
for single_image_embeds in ip_adapter_image_embeds:
|
||||
if do_classifier_free_guidance:
|
||||
single_negative_image_embeds, single_image_embeds = single_image_embeds.chunk(2)
|
||||
single_image_embeds = single_image_embeds.repeat(
|
||||
num_images_per_prompt, *(repeat_dims * len(single_image_embeds.shape[1:]))
|
||||
)
|
||||
single_negative_image_embeds = single_negative_image_embeds.repeat(
|
||||
num_images_per_prompt, *(repeat_dims * len(single_negative_image_embeds.shape[1:]))
|
||||
)
|
||||
single_image_embeds = torch.cat([single_negative_image_embeds, single_image_embeds])
|
||||
else:
|
||||
single_image_embeds = single_image_embeds.repeat(
|
||||
num_images_per_prompt, *(repeat_dims * len(single_image_embeds.shape[1:]))
|
||||
)
|
||||
image_embeds.append(single_image_embeds)
|
||||
|
||||
return image_embeds
|
||||
@@ -0,0 +1,537 @@
|
||||
# Copyright 2023 Stanford University Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
# DISCLAIMER: This code is strongly influenced by https://github.com/pesser/pytorch_diffusion
|
||||
# and https://github.com/hojonathanho/diffusion
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class LCMSingleStepSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
pred_original_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
The predicted denoised sample `(x_{0})` based on the model output from the current timestep.
|
||||
`pred_original_sample` can be used to preview progress or for guidance.
|
||||
"""
|
||||
|
||||
denoised: Optional[torch.FloatTensor] = None
|
||||
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.betas_for_alpha_bar
|
||||
def betas_for_alpha_bar(
|
||||
num_diffusion_timesteps,
|
||||
max_beta=0.999,
|
||||
alpha_transform_type="cosine",
|
||||
):
|
||||
"""
|
||||
Create a beta schedule that discretizes the given alpha_t_bar function, which defines the cumulative product of
|
||||
(1-beta) over time from t = [0,1].
|
||||
|
||||
Contains a function alpha_bar that takes an argument t and transforms it to the cumulative product of (1-beta) up
|
||||
to that part of the diffusion process.
|
||||
|
||||
|
||||
Args:
|
||||
num_diffusion_timesteps (`int`): the number of betas to produce.
|
||||
max_beta (`float`): the maximum beta to use; use values lower than 1 to
|
||||
prevent singularities.
|
||||
alpha_transform_type (`str`, *optional*, default to `cosine`): the type of noise schedule for alpha_bar.
|
||||
Choose from `cosine` or `exp`
|
||||
|
||||
Returns:
|
||||
betas (`np.ndarray`): the betas used by the scheduler to step the model outputs
|
||||
"""
|
||||
if alpha_transform_type == "cosine":
|
||||
|
||||
def alpha_bar_fn(t):
|
||||
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
|
||||
|
||||
elif alpha_transform_type == "exp":
|
||||
|
||||
def alpha_bar_fn(t):
|
||||
return math.exp(t * -12.0)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unsupported alpha_tranform_type: {alpha_transform_type}")
|
||||
|
||||
betas = []
|
||||
for i in range(num_diffusion_timesteps):
|
||||
t1 = i / num_diffusion_timesteps
|
||||
t2 = (i + 1) / num_diffusion_timesteps
|
||||
betas.append(min(1 - alpha_bar_fn(t2) / alpha_bar_fn(t1), max_beta))
|
||||
return torch.tensor(betas, dtype=torch.float32)
|
||||
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddim.rescale_zero_terminal_snr
|
||||
def rescale_zero_terminal_snr(betas: torch.FloatTensor) -> torch.FloatTensor:
|
||||
"""
|
||||
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
|
||||
|
||||
|
||||
Args:
|
||||
betas (`torch.FloatTensor`):
|
||||
the betas that the scheduler is being initialized with.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`: rescaled betas with zero terminal SNR
|
||||
"""
|
||||
# Convert betas to alphas_bar_sqrt
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
alphas_bar_sqrt = alphas_cumprod.sqrt()
|
||||
|
||||
# Store old values.
|
||||
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
|
||||
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
|
||||
|
||||
# Shift so the last timestep is zero.
|
||||
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
||||
|
||||
# Scale so the first timestep is back to the old value.
|
||||
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
||||
|
||||
# Convert alphas_bar_sqrt to betas
|
||||
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
|
||||
alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod
|
||||
alphas = torch.cat([alphas_bar[0:1], alphas])
|
||||
betas = 1 - alphas
|
||||
|
||||
return betas
|
||||
|
||||
|
||||
class LCMSingleStepScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
`LCMSingleStepScheduler` extends the denoising procedure introduced in denoising diffusion probabilistic models (DDPMs) with
|
||||
non-Markovian guidance.
|
||||
|
||||
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. [`~ConfigMixin`] takes care of storing all config
|
||||
attributes that are passed in the scheduler's `__init__` function, such as `num_train_timesteps`. They can be
|
||||
accessed via `scheduler.config.num_train_timesteps`. [`SchedulerMixin`] provides general loading and saving
|
||||
functionality via the [`SchedulerMixin.save_pretrained`] and [`~SchedulerMixin.from_pretrained`] functions.
|
||||
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
beta_start (`float`, defaults to 0.0001):
|
||||
The starting `beta` value of inference.
|
||||
beta_end (`float`, defaults to 0.02):
|
||||
The final `beta` value.
|
||||
beta_schedule (`str`, defaults to `"linear"`):
|
||||
The beta schedule, a mapping from a beta range to a sequence of betas for stepping the model. Choose from
|
||||
`linear`, `scaled_linear`, or `squaredcos_cap_v2`.
|
||||
trained_betas (`np.ndarray`, *optional*):
|
||||
Pass an array of betas directly to the constructor to bypass `beta_start` and `beta_end`.
|
||||
original_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The default number of inference steps used to generate a linearly-spaced timestep schedule, from which we
|
||||
will ultimately take `num_inference_steps` evenly spaced timesteps to form the final timestep schedule.
|
||||
clip_sample (`bool`, defaults to `True`):
|
||||
Clip the predicted sample for numerical stability.
|
||||
clip_sample_range (`float`, defaults to 1.0):
|
||||
The maximum magnitude for sample clipping. Valid only when `clip_sample=True`.
|
||||
set_alpha_to_one (`bool`, defaults to `True`):
|
||||
Each diffusion step uses the alphas product value at that step and at the previous one. For the final step
|
||||
there is no previous alpha. When this option is `True` the previous alpha product is fixed to `1`,
|
||||
otherwise it uses the alpha value at step 0.
|
||||
steps_offset (`int`, defaults to 0):
|
||||
An offset added to the inference steps. You can use a combination of `offset=1` and
|
||||
`set_alpha_to_one=False` to make the last step use step 0 for the previous alpha product like in Stable
|
||||
Diffusion.
|
||||
prediction_type (`str`, defaults to `epsilon`, *optional*):
|
||||
Prediction type of the scheduler function; can be `epsilon` (predicts the noise of the diffusion process),
|
||||
`sample` (directly predicts the noisy sample`) or `v_prediction` (see section 2.4 of [Imagen
|
||||
Video](https://imagen.research.google/video/paper.pdf) paper).
|
||||
thresholding (`bool`, defaults to `False`):
|
||||
Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such
|
||||
as Stable Diffusion.
|
||||
dynamic_thresholding_ratio (`float`, defaults to 0.995):
|
||||
The ratio for the dynamic thresholding method. Valid only when `thresholding=True`.
|
||||
sample_max_value (`float`, defaults to 1.0):
|
||||
The threshold value for dynamic thresholding. Valid only when `thresholding=True`.
|
||||
timestep_spacing (`str`, defaults to `"leading"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
timestep_scaling (`float`, defaults to 10.0):
|
||||
The factor the timesteps will be multiplied by when calculating the consistency model boundary conditions
|
||||
`c_skip` and `c_out`. Increasing this will decrease the approximation error (although the approximation
|
||||
error at the default of `10.0` is already pretty small).
|
||||
rescale_betas_zero_snr (`bool`, defaults to `False`):
|
||||
Whether to rescale the betas to have zero terminal SNR. This enables the model to generate very bright and
|
||||
dark samples instead of limiting it to samples with medium brightness. Loosely related to
|
||||
[`--offset_noise`](https://github.com/huggingface/diffusers/blob/74fd735eb073eb1d774b1ab4154a0876eb82f055/examples/dreambooth/train_dreambooth.py#L506).
|
||||
"""
|
||||
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
beta_start: float = 0.00085,
|
||||
beta_end: float = 0.012,
|
||||
beta_schedule: str = "scaled_linear",
|
||||
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
|
||||
original_inference_steps: int = 50,
|
||||
clip_sample: bool = False,
|
||||
clip_sample_range: float = 1.0,
|
||||
set_alpha_to_one: bool = True,
|
||||
steps_offset: int = 0,
|
||||
prediction_type: str = "epsilon",
|
||||
thresholding: bool = False,
|
||||
dynamic_thresholding_ratio: float = 0.995,
|
||||
sample_max_value: float = 1.0,
|
||||
timestep_spacing: str = "leading",
|
||||
timestep_scaling: float = 10.0,
|
||||
rescale_betas_zero_snr: bool = False,
|
||||
):
|
||||
if trained_betas is not None:
|
||||
self.betas = torch.tensor(trained_betas, dtype=torch.float32)
|
||||
elif beta_schedule == "linear":
|
||||
self.betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32)
|
||||
elif beta_schedule == "scaled_linear":
|
||||
# this schedule is very specific to the latent diffusion model.
|
||||
self.betas = (
|
||||
torch.linspace(beta_start**0.5, beta_end**0.5, num_train_timesteps, dtype=torch.float32) ** 2
|
||||
)
|
||||
elif beta_schedule == "squaredcos_cap_v2":
|
||||
# Glide cosine schedule
|
||||
self.betas = betas_for_alpha_bar(num_train_timesteps)
|
||||
else:
|
||||
raise NotImplementedError(f"{beta_schedule} does is not implemented for {self.__class__}")
|
||||
|
||||
# Rescale for zero SNR
|
||||
if rescale_betas_zero_snr:
|
||||
self.betas = rescale_zero_terminal_snr(self.betas)
|
||||
|
||||
self.alphas = 1.0 - self.betas
|
||||
self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)
|
||||
|
||||
# At every step in ddim, we are looking into the previous alphas_cumprod
|
||||
# For the final step, there is no previous alphas_cumprod because we are already at 0
|
||||
# `set_alpha_to_one` decides whether we set this parameter simply to one or
|
||||
# whether we use the final alpha of the "non-previous" one.
|
||||
self.final_alpha_cumprod = torch.tensor(1.0) if set_alpha_to_one else self.alphas_cumprod[0]
|
||||
|
||||
# standard deviation of the initial noise distribution
|
||||
self.init_noise_sigma = 1.0
|
||||
|
||||
# setable values
|
||||
self.num_inference_steps = None
|
||||
self.timesteps = torch.from_numpy(np.arange(0, num_train_timesteps)[::-1].copy().astype(np.int64))
|
||||
|
||||
self._step_index = None
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._init_step_index
|
||||
def _init_step_index(self, timestep):
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
|
||||
index_candidates = (self.timesteps == timestep).nonzero()
|
||||
|
||||
# The sigma index that is taken for the **very** first `step`
|
||||
# is always the second index (or the last index if there is only 1)
|
||||
# This way we can ensure we don't accidentally skip a sigma in
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
if len(index_candidates) > 1:
|
||||
step_index = index_candidates[1]
|
||||
else:
|
||||
step_index = index_candidates[0]
|
||||
|
||||
self._step_index = step_index.item()
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
return self._step_index
|
||||
|
||||
def scale_model_input(self, sample: torch.FloatTensor, timestep: Optional[int] = None) -> torch.FloatTensor:
|
||||
"""
|
||||
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
|
||||
current timestep.
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The input sample.
|
||||
timestep (`int`, *optional*):
|
||||
The current timestep in the diffusion chain.
|
||||
Returns:
|
||||
`torch.FloatTensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
return sample
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample
|
||||
def _threshold_sample(self, sample: torch.FloatTensor) -> torch.FloatTensor:
|
||||
"""
|
||||
"Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the
|
||||
prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by
|
||||
s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing
|
||||
pixels from saturation at each step. We find that dynamic thresholding results in significantly better
|
||||
photorealism as well as better image-text alignment, especially when using very large guidance weights."
|
||||
|
||||
https://arxiv.org/abs/2205.11487
|
||||
"""
|
||||
dtype = sample.dtype
|
||||
batch_size, channels, *remaining_dims = sample.shape
|
||||
|
||||
if dtype not in (torch.float32, torch.float64):
|
||||
sample = sample.float() # upcast for quantile calculation, and clamp not implemented for cpu half
|
||||
|
||||
# Flatten sample for doing quantile calculation along each image
|
||||
sample = sample.reshape(batch_size, channels * np.prod(remaining_dims))
|
||||
|
||||
abs_sample = sample.abs() # "a certain percentile absolute pixel value"
|
||||
|
||||
s = torch.quantile(abs_sample, self.config.dynamic_thresholding_ratio, dim=1)
|
||||
s = torch.clamp(
|
||||
s, min=1, max=self.config.sample_max_value
|
||||
) # When clamped to min=1, equivalent to standard clipping to [-1, 1]
|
||||
s = s.unsqueeze(1) # (batch_size, 1) because clamp will broadcast along dim=0
|
||||
sample = torch.clamp(sample, -s, s) / s # "we threshold xt0 to the range [-s, s] and then divide by s"
|
||||
|
||||
sample = sample.reshape(batch_size, channels, *remaining_dims)
|
||||
sample = sample.to(dtype)
|
||||
|
||||
return sample
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
original_inference_steps: Optional[int] = None,
|
||||
strength: int = 1.0,
|
||||
timesteps: Optional[list] = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
original_inference_steps (`int`, *optional*):
|
||||
The original number of inference steps, which will be used to generate a linearly-spaced timestep
|
||||
schedule (which is different from the standard `diffusers` implementation). We will then take
|
||||
`num_inference_steps` timesteps from this schedule, evenly spaced in terms of indices, and use that as
|
||||
our final timestep schedule. If not set, this will default to the `original_inference_steps` attribute.
|
||||
"""
|
||||
|
||||
if num_inference_steps is not None and timesteps is not None:
|
||||
raise ValueError("Can only pass one of `num_inference_steps` or `custom_timesteps`.")
|
||||
|
||||
if timesteps is not None:
|
||||
for i in range(1, len(timesteps)):
|
||||
if timesteps[i] >= timesteps[i - 1]:
|
||||
raise ValueError("`custom_timesteps` must be in descending order.")
|
||||
|
||||
if timesteps[0] >= self.config.num_train_timesteps:
|
||||
raise ValueError(
|
||||
f"`timesteps` must start before `self.config.train_timesteps`:"
|
||||
f" {self.config.num_train_timesteps}."
|
||||
)
|
||||
|
||||
timesteps = np.array(timesteps, dtype=np.int64)
|
||||
else:
|
||||
if num_inference_steps > self.config.num_train_timesteps:
|
||||
raise ValueError(
|
||||
f"`num_inference_steps`: {num_inference_steps} cannot be larger than `self.config.train_timesteps`:"
|
||||
f" {self.config.num_train_timesteps} as the unet model trained with this scheduler can only handle"
|
||||
f" maximal {self.config.num_train_timesteps} timesteps."
|
||||
)
|
||||
|
||||
self.num_inference_steps = num_inference_steps
|
||||
original_steps = (
|
||||
original_inference_steps if original_inference_steps is not None else self.config.original_inference_steps
|
||||
)
|
||||
|
||||
if original_steps > self.config.num_train_timesteps:
|
||||
raise ValueError(
|
||||
f"`original_steps`: {original_steps} cannot be larger than `self.config.train_timesteps`:"
|
||||
f" {self.config.num_train_timesteps} as the unet model trained with this scheduler can only handle"
|
||||
f" maximal {self.config.num_train_timesteps} timesteps."
|
||||
)
|
||||
|
||||
if num_inference_steps > original_steps:
|
||||
raise ValueError(
|
||||
f"`num_inference_steps`: {num_inference_steps} cannot be larger than `original_inference_steps`:"
|
||||
f" {original_steps} because the final timestep schedule will be a subset of the"
|
||||
f" `original_inference_steps`-sized initial timestep schedule."
|
||||
)
|
||||
|
||||
# LCM Timesteps Setting
|
||||
# Currently, only linear spacing is supported.
|
||||
c = self.config.num_train_timesteps // original_steps
|
||||
# LCM Training Steps Schedule
|
||||
lcm_origin_timesteps = np.asarray(list(range(1, int(original_steps * strength) + 1))) * c - 1
|
||||
skipping_step = len(lcm_origin_timesteps) // num_inference_steps
|
||||
# LCM Inference Steps Schedule
|
||||
timesteps = lcm_origin_timesteps[::-skipping_step][:num_inference_steps]
|
||||
|
||||
self.timesteps = torch.from_numpy(timesteps.copy()).to(device=device, dtype=torch.long)
|
||||
|
||||
self._step_index = None
|
||||
|
||||
def get_scalings_for_boundary_condition_discrete(self, timestep):
|
||||
self.sigma_data = 0.5 # Default: 0.5
|
||||
scaled_timestep = timestep * self.config.timestep_scaling
|
||||
|
||||
c_skip = self.sigma_data**2 / (scaled_timestep**2 + self.sigma_data**2)
|
||||
c_out = scaled_timestep / (scaled_timestep**2 + self.sigma_data**2) ** 0.5
|
||||
return c_skip, c_out
|
||||
|
||||
def append_dims(self, x, target_dims):
|
||||
"""Appends dimensions to the end of a tensor until it has target_dims dimensions."""
|
||||
dims_to_append = target_dims - x.ndim
|
||||
if dims_to_append < 0:
|
||||
raise ValueError(f"input has {x.ndim} dims but target_dims is {target_dims}, which is less")
|
||||
return x[(...,) + (None,) * dims_to_append]
|
||||
|
||||
def extract_into_tensor(self, a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: torch.Tensor,
|
||||
sample: torch.FloatTensor,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[LCMSingleStepSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
|
||||
Args:
|
||||
model_output (`torch.FloatTensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`float`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~schedulers.scheduling_lcm.LCMSchedulerOutput`] or `tuple`.
|
||||
Returns:
|
||||
[`~schedulers.scheduling_utils.LCMSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_lcm.LCMSchedulerOutput`] is returned, otherwise a
|
||||
tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
# 0. make sure everything is on the same device
|
||||
alphas_cumprod = self.alphas_cumprod.to(sample.device)
|
||||
|
||||
# 1. compute alphas, betas
|
||||
if timestep.ndim == 0:
|
||||
timestep = timestep.unsqueeze(0)
|
||||
alpha_prod_t = self.extract_into_tensor(alphas_cumprod, timestep, sample.shape)
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
# 2. Get scalings for boundary conditions
|
||||
c_skip, c_out = self.get_scalings_for_boundary_condition_discrete(timestep)
|
||||
c_skip, c_out = [self.append_dims(x, sample.ndim) for x in [c_skip, c_out]]
|
||||
|
||||
# 3. Compute the predicted original sample x_0 based on the model parameterization
|
||||
if self.config.prediction_type == "epsilon": # noise-prediction
|
||||
predicted_original_sample = (sample - torch.sqrt(beta_prod_t) * model_output) / torch.sqrt(alpha_prod_t)
|
||||
elif self.config.prediction_type == "sample": # x-prediction
|
||||
predicted_original_sample = model_output
|
||||
elif self.config.prediction_type == "v_prediction": # v-prediction
|
||||
predicted_original_sample = torch.sqrt(alpha_prod_t) * sample - torch.sqrt(beta_prod_t) * model_output
|
||||
else:
|
||||
raise ValueError(
|
||||
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample` or"
|
||||
" `v_prediction` for `LCMScheduler`."
|
||||
)
|
||||
|
||||
# 4. Clip or threshold "predicted x_0"
|
||||
if self.config.thresholding:
|
||||
predicted_original_sample = self._threshold_sample(predicted_original_sample)
|
||||
elif self.config.clip_sample:
|
||||
predicted_original_sample = predicted_original_sample.clamp(
|
||||
-self.config.clip_sample_range, self.config.clip_sample_range
|
||||
)
|
||||
|
||||
# 5. Denoise model output using boundary conditions
|
||||
denoised = c_out * predicted_original_sample + c_skip * sample
|
||||
|
||||
if not return_dict:
|
||||
return (denoised, )
|
||||
|
||||
return LCMSingleStepSchedulerOutput(denoised=denoised)
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler.add_noise
|
||||
def add_noise(
|
||||
self,
|
||||
original_samples: torch.FloatTensor,
|
||||
noise: torch.FloatTensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.FloatTensor:
|
||||
# Make sure alphas_cumprod and timestep have same device and dtype as original_samples
|
||||
alphas_cumprod = self.alphas_cumprod.to(device=original_samples.device, dtype=original_samples.dtype)
|
||||
timesteps = timesteps.to(original_samples.device)
|
||||
|
||||
sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.flatten()
|
||||
while len(sqrt_alpha_prod.shape) < len(original_samples.shape):
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1)
|
||||
|
||||
sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.flatten()
|
||||
while len(sqrt_one_minus_alpha_prod.shape) < len(original_samples.shape):
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1)
|
||||
|
||||
noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise
|
||||
return noisy_samples
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler.get_velocity
|
||||
def get_velocity(
|
||||
self, sample: torch.FloatTensor, noise: torch.FloatTensor, timesteps: torch.IntTensor
|
||||
) -> torch.FloatTensor:
|
||||
# Make sure alphas_cumprod and timestep have same device and dtype as sample
|
||||
alphas_cumprod = self.alphas_cumprod.to(device=sample.device, dtype=sample.dtype)
|
||||
timesteps = timesteps.to(sample.device)
|
||||
|
||||
sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.flatten()
|
||||
while len(sqrt_alpha_prod.shape) < len(sample.shape):
|
||||
sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1)
|
||||
|
||||
sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.flatten()
|
||||
while len(sqrt_one_minus_alpha_prod.shape) < len(sample.shape):
|
||||
sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1)
|
||||
|
||||
velocity = sqrt_alpha_prod * noise - sqrt_one_minus_alpha_prod * sample
|
||||
return velocity
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
File diff suppressed because it is too large
Load Diff
+2
-1
@@ -126,4 +126,5 @@ except ImportError:
|
||||
except ImportError:
|
||||
pass # shrug...
|
||||
|
||||
errors.log.info(f'System packages: {get_packages()}')
|
||||
errors.log.info(f'Torch: torch=={torch.__version__} torchvision=={torchvision.__version__}')
|
||||
errors.log.info(f'Packages: diffusers=={diffusers.__version__} transformers=={transformers.__version__} accelerate=={accelerate.__version__} gradio=={gradio.__version__}')
|
||||
|
||||
@@ -32,14 +32,14 @@ def load_bnb(msg='', silent=False):
|
||||
global bnb # pylint: disable=global-statement
|
||||
if bnb is not None:
|
||||
return bnb
|
||||
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
|
||||
log.debug(f'Quantization: type=bitsandbytes fn={fn}') # pylint: disable=protected-access
|
||||
install('bitsandbytes', quiet=True)
|
||||
try:
|
||||
import bitsandbytes
|
||||
bnb = bitsandbytes
|
||||
diffusers.utils.import_utils._bitsandbytes_available = True # pylint: disable=protected-access
|
||||
diffusers.utils.import_utils._bitsandbytes_version = '0.43.3' # pylint: disable=protected-access
|
||||
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
|
||||
log.debug(f'Quantization: type=bitsandbytes version={bnb.__version__} fn={fn}') # pylint: disable=protected-access
|
||||
return bnb
|
||||
except Exception as e:
|
||||
if len(msg) > 0:
|
||||
@@ -54,12 +54,12 @@ def load_quanto(msg='', silent=False):
|
||||
global quanto # pylint: disable=global-statement
|
||||
if quanto is not None:
|
||||
return quanto
|
||||
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
|
||||
log.debug(f'Quantization: type=quanto fn={fn}') # pylint: disable=protected-access
|
||||
install('optimum-quanto', quiet=True)
|
||||
try:
|
||||
from optimum import quanto as optimum_quanto # pylint: disable=no-name-in-module
|
||||
quanto = optimum_quanto
|
||||
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
|
||||
log.debug(f'Quantization: type=quanto version={quanto.__version__} fn={fn}') # pylint: disable=protected-access
|
||||
return quanto
|
||||
except Exception as e:
|
||||
if len(msg) > 0:
|
||||
|
||||
+24
-2
@@ -5,13 +5,35 @@ import safetensors.torch
|
||||
from modules import shared, devices, model_quant
|
||||
|
||||
|
||||
def remove_entries_after_depth(d, depth, current_depth=0):
|
||||
if current_depth >= depth:
|
||||
return None
|
||||
if isinstance(d, dict):
|
||||
return {k: remove_entries_after_depth(v, depth, current_depth + 1) for k, v in d.items() if remove_entries_after_depth(v, depth, current_depth + 1) is not None}
|
||||
return d
|
||||
|
||||
|
||||
def list_to_dict(flat_list):
|
||||
result_dict = {}
|
||||
try:
|
||||
for item in flat_list:
|
||||
keys = item.split('.')
|
||||
d = result_dict
|
||||
for key in keys[:-1]:
|
||||
d = d.setdefault(key, {})
|
||||
d[keys[-1]] = None
|
||||
except Exception:
|
||||
pass
|
||||
return result_dict
|
||||
|
||||
|
||||
def get_safetensor_keys(filename):
|
||||
keys = []
|
||||
try:
|
||||
with safetensors.torch.safe_open(filename, framework="pt", device="cpu") as f:
|
||||
keys = f.keys()
|
||||
except Exception as e:
|
||||
shared.log.error(f'Load dict: path="{filename}" {e}')
|
||||
except Exception:
|
||||
pass
|
||||
return keys
|
||||
|
||||
|
||||
|
||||
+18
-11
@@ -18,7 +18,7 @@ PREDEFINED = [ # <https://huggingface.co/vladmandic/yolo-detailers/tree/main>
|
||||
|
||||
|
||||
class YoloResult:
|
||||
def __init__(self, cls: int, label: str, score: float, box: list[int], mask: Image.Image = None, item: Image.Image = None, size: float = 0, width = 0, height = 0, args = {}):
|
||||
def __init__(self, cls: int, label: str, score: float, box: list[int], mask: Image.Image = None, item: Image.Image = None, width = 0, height = 0, args = {}):
|
||||
self.cls = cls
|
||||
self.label = label
|
||||
self.score = score
|
||||
@@ -29,6 +29,9 @@ class YoloResult:
|
||||
self.height = height
|
||||
self.args = args
|
||||
|
||||
def __str__(self):
|
||||
return f'cls={self.cls} label={self.label} score={self.score} box={self.box} mask={self.mask} item={self.item} size={self.width}x{self.height} args={self.args}'
|
||||
|
||||
|
||||
class YoloRestorer(Detailer):
|
||||
def __init__(self):
|
||||
@@ -76,11 +79,15 @@ class YoloRestorer(Detailer):
|
||||
offload: bool = shared.opts.detailer_unload,
|
||||
) -> list[YoloResult]:
|
||||
|
||||
if model is None or (isinstance(model, str) and len(model) == 0):
|
||||
model = 'yolo11m'
|
||||
result = []
|
||||
if isinstance(model, str):
|
||||
model = self.models.get(model, None)
|
||||
if model is None:
|
||||
cached = self.models.get(model, None)
|
||||
if cached is None:
|
||||
_, model = self.load(model)
|
||||
else:
|
||||
model = cached
|
||||
if model is None:
|
||||
return result
|
||||
args = {
|
||||
@@ -136,7 +143,8 @@ class YoloRestorer(Detailer):
|
||||
draw = ImageDraw.Draw(mask_image)
|
||||
draw.rectangle(box, fill="white", outline=None, width=0)
|
||||
cropped = image.crop(box)
|
||||
result.append(YoloResult(cls=cls, label=label, score=round(score, 2), box=box, mask=mask_image, item=cropped, width=w, height=h, args=args))
|
||||
res = YoloResult(cls=cls, label=label, score=round(score, 2), box=box, mask=mask_image, item=cropped, width=w, height=h, args=args)
|
||||
result.append(res)
|
||||
if len(result) >= shared.opts.detailer_max:
|
||||
break
|
||||
return result
|
||||
@@ -156,10 +164,10 @@ class YoloRestorer(Detailer):
|
||||
try:
|
||||
model_file = modelloader.load_file_from_url(url=model_url, model_dir=shared.opts.yolo_dir, file_name=file_name)
|
||||
if model_file is not None:
|
||||
from ultralytics import YOLO # pylint: disable=import-outside-toplevel
|
||||
model = YOLO(model_file)
|
||||
import ultralytics
|
||||
model = ultralytics.YOLO(model_file)
|
||||
classes = list(model.names.values())
|
||||
shared.log.info(f'Load: type=Detailer name="{model_name}" model="{model_file}" classes={classes}')
|
||||
shared.log.info(f'Load: type=Detailer name="{model_name}" model="{model_file}" ultralytics={ultralytics.__version__} classes={classes}')
|
||||
self.models[model_name] = model
|
||||
return model_name, model
|
||||
except Exception as e:
|
||||
@@ -194,7 +202,6 @@ class YoloRestorer(Detailer):
|
||||
shared.log.info(f'Detailer: model="{name}" no items detected')
|
||||
continue
|
||||
|
||||
pp = None
|
||||
shared.opts.data['mask_apply_overlay'] = True
|
||||
resolution = 512 if shared.sd_model_type in ['none', 'sd', 'lcm', 'unknown'] else 1024
|
||||
orig_prompt: str = orig_p.get('all_prompts', [''])[0]
|
||||
@@ -330,9 +337,9 @@ class YoloRestorer(Detailer):
|
||||
iou = gr.Slider(label="Max overlap", elem_id=f"{tab}_detailer_iou", value=shared.opts.detailer_iou, minimum=0, maximum=1.0, step=0.05)
|
||||
with gr.Row():
|
||||
min_size = shared.opts.detailer_min_size if shared.opts.detailer_min_size < 1 else 0.0
|
||||
min_size = gr.Slider(label="Min size", elem_id=f"{tab}_detailer_min_size", value=min_size, minimum=0.1, maximum=1.0, step=0.05)
|
||||
max_size = shared.opts.detailer_min_size if shared.opts.detailer_min_size < 1 and shared.opts.detailer_min_size > 0 else 1.0
|
||||
max_size = gr.Slider(label="Max size", elem_id=f"{tab}_detailer_max_size", value=max_size, minimum=0.1, maximum=1.0, step=0.05)
|
||||
min_size = gr.Slider(label="Min size", elem_id=f"{tab}_detailer_min_size", value=min_size, minimum=0.0, maximum=1.0, step=0.05)
|
||||
max_size = shared.opts.detailer_max_size if shared.opts.detailer_max_size < 1 and shared.opts.detailer_max_size > 0 else 1.0
|
||||
max_size = gr.Slider(label="Max size", elem_id=f"{tab}_detailer_max_size", value=max_size, minimum=0.0, maximum=1.0, step=0.05)
|
||||
detailers.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[])
|
||||
classes.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[])
|
||||
strength.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[])
|
||||
|
||||
@@ -34,14 +34,14 @@ images_tensor_to_samples = processing_helpers.images_tensor_to_samples
|
||||
|
||||
|
||||
class Processed:
|
||||
def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info="", subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments=""):
|
||||
def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info=None, subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments=""):
|
||||
self.images = images_list
|
||||
self.prompt = p.prompt or ''
|
||||
self.negative_prompt = p.negative_prompt or ''
|
||||
self.seed = seed if seed != -1 else p.seed
|
||||
self.subseed = subseed
|
||||
self.subseed_strength = p.subseed_strength
|
||||
self.info = info
|
||||
self.info = info or create_infotext(p)
|
||||
self.comments = comments or ''
|
||||
self.width = p.width if hasattr(p, 'width') else (self.images[0].width if len(self.images) > 0 else 0)
|
||||
self.height = p.height if hasattr(p, 'height') else (self.images[0].height if len(self.images) > 0 else 0)
|
||||
@@ -80,7 +80,7 @@ class Processed:
|
||||
self.all_negative_prompts = all_negative_prompts or p.all_negative_prompts or [self.negative_prompt]
|
||||
self.all_seeds = all_seeds or p.all_seeds or [self.seed]
|
||||
self.all_subseeds = all_subseeds or p.all_subseeds or [self.subseed]
|
||||
self.infotexts = infotexts or [info]
|
||||
self.infotexts = infotexts or [self.info]
|
||||
|
||||
def js(self):
|
||||
obj = {
|
||||
|
||||
+33
-26
@@ -9,6 +9,7 @@ import numpy as np
|
||||
from modules import shared, errors, sd_models, processing, processing_vae, processing_helpers, sd_hijack_hypertile, prompt_parser_diffusers, timer
|
||||
from modules.processing_callbacks import diffusers_callback_legacy, diffusers_callback, set_callbacks_p
|
||||
from modules.processing_helpers import resize_hires, fix_prompts, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, get_generator, set_latents, apply_circular # pylint: disable=unused-import
|
||||
from modules.api import helpers
|
||||
|
||||
|
||||
debug = shared.log.trace if os.environ.get('SD_DIFFUSERS_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
@@ -18,7 +19,9 @@ def task_specific_kwargs(p, model):
|
||||
task_args = {}
|
||||
is_img2img_model = bool('Zero123' in shared.sd_model.__class__.__name__)
|
||||
if len(getattr(p, 'init_images', [])) > 0:
|
||||
p.init_images = [p.convert('RGB') for p in p.init_images]
|
||||
if isinstance(p.init_images[0], str):
|
||||
p.init_images = [helpers.decode_base64_to_image(i, quiet=True) for i in p.init_images]
|
||||
p.init_images = [i.convert('RGB') if i.mode != 'RGB' else i for i in p.init_images]
|
||||
if sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.TEXT_2_IMAGE or len(getattr(p, 'init_images', [])) == 0 and not is_img2img_model:
|
||||
p.ops.append('txt2img')
|
||||
if hasattr(p, 'width') and hasattr(p, 'height'):
|
||||
@@ -27,7 +30,7 @@ def task_specific_kwargs(p, model):
|
||||
'height': 8 * math.ceil(p.height / 8),
|
||||
}
|
||||
elif (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.IMAGE_2_IMAGE or is_img2img_model) and len(getattr(p, 'init_images', [])) > 0:
|
||||
if shared.sd_model_type == 'sdxl':
|
||||
if shared.sd_model_type == 'sdxl' and hasattr(model, 'register_to_config'):
|
||||
model.register_to_config(requires_aesthetics_score = False)
|
||||
p.ops.append('img2img')
|
||||
task_args = {
|
||||
@@ -55,7 +58,7 @@ def task_specific_kwargs(p, model):
|
||||
'strength': p.denoising_strength,
|
||||
}
|
||||
elif (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.INPAINTING or is_img2img_model) and len(getattr(p, 'init_images', [])) > 0:
|
||||
if shared.sd_model_type == 'sdxl':
|
||||
if shared.sd_model_type == 'sdxl' and hasattr(model, 'register_to_config'):
|
||||
model.register_to_config(requires_aesthetics_score = False)
|
||||
if p.detailer:
|
||||
p.ops.append('detailer')
|
||||
@@ -100,62 +103,64 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2
|
||||
if hasattr(model, "set_progress_bar_config"):
|
||||
model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba')
|
||||
args = {}
|
||||
if hasattr(model, 'pipe'): # recurse
|
||||
if hasattr(model, 'pipe') and not hasattr(model, 'no_recurse'): # recurse
|
||||
model = model.pipe
|
||||
signature = inspect.signature(type(model).__call__, follow_wrapped=True)
|
||||
possible = list(signature.parameters)
|
||||
|
||||
debug(f'Diffusers pipeline possible: {possible}')
|
||||
prompts, negative_prompts, prompts_2, negative_prompts_2 = fix_prompts(prompts, negative_prompts, prompts_2, negative_prompts_2)
|
||||
parser = 'Fixed attention'
|
||||
steps = kwargs.get("num_inference_steps", None) or len(getattr(p, 'timesteps', ['1']))
|
||||
clip_skip = kwargs.pop("clip_skip", 1)
|
||||
|
||||
# prompt_parser_diffusers.fix_position_ids(model)
|
||||
if shared.opts.prompt_attention != 'Fixed attention' and 'Onnx' not in model.__class__.__name__ and (
|
||||
parser = 'fixed'
|
||||
if shared.opts.prompt_attention != 'fixed' and 'Onnx' not in model.__class__.__name__ and (
|
||||
'StableDiffusion' in model.__class__.__name__ or
|
||||
'StableCascade' in model.__class__.__name__ or
|
||||
'Flux' in model.__class__.__name__
|
||||
):
|
||||
try:
|
||||
prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, steps=steps, clip_skip=clip_skip)
|
||||
prompt_parser_diffusers.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, steps, clip_skip, p)
|
||||
parser = shared.opts.prompt_attention
|
||||
except Exception as e:
|
||||
shared.log.error(f'Prompt parser encode: {e}')
|
||||
if os.environ.get('SD_PROMPT_DEBUG', None) is not None:
|
||||
errors.display(e, 'Prompt parser encode')
|
||||
timer.process.record('encode', reset=False)
|
||||
else:
|
||||
prompt_parser_diffusers.embedder = None
|
||||
|
||||
if 'prompt' in possible:
|
||||
if 'OmniGen' in model.__class__.__name__:
|
||||
prompts = [p.replace('|image|', '<|image_1|>') for p in prompts]
|
||||
if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and len(p.prompt_embeds) > 0 and p.prompt_embeds[0] is not None:
|
||||
args['prompt_embeds'] = p.prompt_embeds[0]
|
||||
if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
|
||||
args['prompt_embeds'] = prompt_parser_diffusers.embedder('prompt_embeds')
|
||||
if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0:
|
||||
args['prompt_embeds_pooled'] = p.positive_pooleds[0].unsqueeze(0)
|
||||
elif 'XL' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0:
|
||||
args['pooled_prompt_embeds'] = p.positive_pooleds[0]
|
||||
elif 'StableDiffusion3' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0:
|
||||
args['pooled_prompt_embeds'] = p.positive_pooleds[0]
|
||||
elif 'Flux' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0:
|
||||
args['pooled_prompt_embeds'] = p.positive_pooleds[0]
|
||||
args['prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('positive_pooleds').unsqueeze(0)
|
||||
elif 'XL' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
|
||||
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
|
||||
elif 'StableDiffusion3' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
|
||||
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
|
||||
elif 'Flux' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
|
||||
args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds')
|
||||
else:
|
||||
args['prompt'] = prompts
|
||||
if 'negative_prompt' in possible:
|
||||
if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and len(p.negative_embeds) > 0 and p.negative_embeds[0] is not None:
|
||||
args['negative_prompt_embeds'] = p.negative_embeds[0]
|
||||
if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0:
|
||||
args['negative_prompt_embeds_pooled'] = p.negative_pooleds[0].unsqueeze(0)
|
||||
if 'XL' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0:
|
||||
args['negative_pooled_prompt_embeds'] = p.negative_pooleds[0]
|
||||
if 'StableDiffusion3' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0:
|
||||
args['negative_pooled_prompt_embeds'] = p.negative_pooleds[0]
|
||||
if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
|
||||
args['negative_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_prompt_embeds')
|
||||
if 'StableCascade' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
|
||||
args['negative_prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('negative_pooleds').unsqueeze(0)
|
||||
if 'XL' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
|
||||
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
|
||||
if 'StableDiffusion3' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
|
||||
args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds')
|
||||
else:
|
||||
if 'PixArtSigmaPipeline' in model.__class__.__name__: # pixart-sigma pipeline throws list-of-list for negative prompt
|
||||
args['negative_prompt'] = negative_prompts[0]
|
||||
else:
|
||||
args['negative_prompt'] = negative_prompts
|
||||
|
||||
if 'clip_skip' in possible and parser == 'Fixed attention':
|
||||
if 'clip_skip' in possible and parser == 'fixed':
|
||||
if clip_skip == 1:
|
||||
pass # clip_skip = None
|
||||
else:
|
||||
@@ -180,6 +185,8 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2
|
||||
if hasattr(model, 'scheduler') and hasattr(model.scheduler, 'noise_sampler_seed') and hasattr(model.scheduler, 'noise_sampler'):
|
||||
model.scheduler.noise_sampler = None # noise needs to be reset instead of using cached values
|
||||
model.scheduler.noise_sampler_seed = p.seeds # some schedulers have internal noise generator and do not use pipeline generator
|
||||
if 'seed' in possible:
|
||||
args['seed'] = p.seed
|
||||
if 'noise_sampler_seed' in possible:
|
||||
args['noise_sampler_seed'] = p.seeds
|
||||
if 'guidance_scale' in possible:
|
||||
|
||||
@@ -3,8 +3,7 @@ import os
|
||||
import time
|
||||
import torch
|
||||
import numpy as np
|
||||
from modules import shared, processing_correction, extra_networks, timer
|
||||
|
||||
from modules import shared, processing_correction, extra_networks, timer, prompt_parser_diffusers
|
||||
|
||||
p = None
|
||||
debug_callback = shared.log.trace if os.environ.get('SD_CALLBACK_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
@@ -14,6 +13,19 @@ def set_callbacks_p(processing):
|
||||
global p # pylint: disable=global-statement
|
||||
p = processing
|
||||
|
||||
def prompt_callback(step, kwargs):
|
||||
if prompt_parser_diffusers.embedder is None or 'prompt_embeds' not in kwargs:
|
||||
return kwargs
|
||||
try:
|
||||
prompt_embeds = prompt_parser_diffusers.embedder('prompt_embeds', step + 1)
|
||||
negative_prompt_embeds = prompt_parser_diffusers.embedder('negative_prompt_embeds', step + 1)
|
||||
if p.cfg_scale > 1: # Perform guidance
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) # Combined embeds
|
||||
assert prompt_embeds.shape == kwargs['prompt_embeds'].shape, f"prompt_embed shape mismatch {kwargs['prompt_embeds'].shape} {prompt_embeds.shape}"
|
||||
kwargs['prompt_embeds'] = prompt_embeds
|
||||
except Exception as e:
|
||||
debug_callback(f"Callback: {e}")
|
||||
return kwargs
|
||||
|
||||
def diffusers_callback_legacy(step: int, timestep: int, latents: typing.Union[torch.FloatTensor, np.ndarray]):
|
||||
if p is None:
|
||||
@@ -33,7 +45,7 @@ def diffusers_callback_legacy(step: int, timestep: int, latents: typing.Union[to
|
||||
time.sleep(0.1)
|
||||
|
||||
|
||||
def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict):
|
||||
def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict = {}):
|
||||
t0 = time.time()
|
||||
if p is None:
|
||||
return kwargs
|
||||
@@ -49,7 +61,7 @@ def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict):
|
||||
if shared.state.interrupted or shared.state.skipped:
|
||||
raise AssertionError('Interrupted...')
|
||||
time.sleep(0.1)
|
||||
if hasattr(p, "extra_network_data"):
|
||||
if hasattr(p, "stepwise_lora"):
|
||||
extra_networks.activate(p, p.extra_network_data, step=step)
|
||||
if latents is None:
|
||||
return kwargs
|
||||
@@ -67,14 +79,7 @@ def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict):
|
||||
pipe.set_ip_adapter_scale(ip_adapter_scales)
|
||||
if step != getattr(pipe, 'num_timesteps', 0):
|
||||
kwargs = processing_correction.correction_callback(p, timestep, kwargs)
|
||||
if p.scheduled_prompt and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs:
|
||||
try:
|
||||
i = (step + 1) % len(p.prompt_embeds)
|
||||
kwargs["prompt_embeds"] = p.prompt_embeds[i][0:1].expand(kwargs["prompt_embeds"].shape)
|
||||
j = (step + 1) % len(p.negative_embeds)
|
||||
kwargs["negative_prompt_embeds"] = p.negative_embeds[j][0:1].expand(kwargs["negative_prompt_embeds"].shape)
|
||||
except Exception as e:
|
||||
shared.log.debug(f"Callback: {e}")
|
||||
kwargs = prompt_callback(step, kwargs) # monkey patch for diffusers callback issues
|
||||
if step == int(getattr(pipe, 'num_timesteps', 100) * p.cfg_end) and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs:
|
||||
if "PAG" in shared.sd_model.__class__.__name__:
|
||||
pipe._guidance_scale = 1.001 if pipe._guidance_scale > 1 else pipe._guidance_scale # pylint: disable=protected-access
|
||||
|
||||
+206
-237
@@ -17,53 +17,44 @@ debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None
|
||||
|
||||
@dataclass(repr=False)
|
||||
class StableDiffusionProcessing:
|
||||
"""
|
||||
The first set of paramaters: sd_models -> do_not_reload_embeddings represent the minimum required to create a StableDiffusionProcessing
|
||||
"""
|
||||
def __init__(self,
|
||||
sd_model=None,
|
||||
outpath_samples=None,
|
||||
outpath_grids=None,
|
||||
sd_model=None, # pylint: disable=unused-argument # local instance of sd_model
|
||||
# base params
|
||||
prompt: str = "",
|
||||
styles: List[str] = None,
|
||||
negative_prompt: str = "",
|
||||
seed: int = -1,
|
||||
subseed: int = -1,
|
||||
subseed_strength: float = 0,
|
||||
seed_resize_from_h: int = -1,
|
||||
seed_resize_from_w: int = -1,
|
||||
seed_enable_extras: bool = True,
|
||||
sampler_name: str = None,
|
||||
hr_sampler_name: str = None,
|
||||
batch_size: int = 1,
|
||||
n_iter: int = 1,
|
||||
steps: int = 50,
|
||||
cfg_scale: float = 7.0,
|
||||
image_cfg_scale: float = None,
|
||||
clip_skip: int = 1,
|
||||
width: int = 512,
|
||||
height: int = 512,
|
||||
full_quality: bool = True,
|
||||
detailer: bool = False,
|
||||
restore_faces: bool = False,
|
||||
tiling: bool = False,
|
||||
hidiffusion: bool = False,
|
||||
do_not_save_samples: bool = False,
|
||||
do_not_save_grid: bool = False,
|
||||
extra_generation_params: Dict[Any, Any] = None,
|
||||
overlay_images: Any = None,
|
||||
negative_prompt: str = None,
|
||||
# samplers
|
||||
sampler_index: int = None, # pylint: disable=unused-argument # used only to set sampler_name
|
||||
sampler_name: str = None,
|
||||
hr_sampler_name: str = None,
|
||||
eta: float = None,
|
||||
do_not_reload_embeddings: bool = False,
|
||||
denoising_strength: float = 0,
|
||||
# guidance
|
||||
cfg_scale: float = 7.0,
|
||||
cfg_end: float = 1,
|
||||
diffusers_guidance_rescale: float = 0.7,
|
||||
pag_scale: float = 0.0,
|
||||
pag_adaptive: float = 0.5,
|
||||
cfg_end: float = 1,
|
||||
resize_mode: int = 0,
|
||||
resize_name: str = 'None',
|
||||
resize_context: str = 'None',
|
||||
scale_by: float = 0,
|
||||
selected_scale_tab: int = 0,
|
||||
# styles
|
||||
styles: List[str] = [],
|
||||
# vae
|
||||
tiling: bool = False,
|
||||
full_quality: bool = True,
|
||||
# other
|
||||
hidiffusion: bool = False,
|
||||
do_not_reload_embeddings: bool = False,
|
||||
detailer: bool = False,
|
||||
restore_faces: bool = False,
|
||||
# hdr corrections
|
||||
hdr_mode: int = 0,
|
||||
hdr_brightness: float = 0,
|
||||
hdr_color: float = 0,
|
||||
@@ -76,92 +67,196 @@ class StableDiffusionProcessing:
|
||||
hdr_max_boundry: float = 1.0,
|
||||
hdr_color_picker: str = None,
|
||||
hdr_tint_ratio: float = 0,
|
||||
override_settings: Dict[str, Any] = None,
|
||||
# img2img
|
||||
init_images: list = None,
|
||||
resize_mode: int = 0,
|
||||
resize_name: str = 'None',
|
||||
resize_context: str = 'None',
|
||||
denoising_strength: float = 0.3,
|
||||
image_cfg_scale: float = None,
|
||||
initial_noise_multiplier: float = None, # pylint: disable=unused-argument # a1111 compatibility
|
||||
scale_by: float = 1,
|
||||
selected_scale_tab: int = 0, # pylint: disable=unused-argument # a1111 compatibility
|
||||
# inpaint
|
||||
mask: Any = None,
|
||||
latent_mask: Any = None,
|
||||
mask_for_overlay: Any = None,
|
||||
mask_blur: int = 4,
|
||||
paste_to: Any = None,
|
||||
inpainting_fill: int = 0,
|
||||
inpaint_full_res: bool = False,
|
||||
inpaint_full_res_padding: int = 0,
|
||||
inpainting_mask_invert: int = 0,
|
||||
overlay_images: Any = None,
|
||||
# refiner
|
||||
enable_hr: bool = False,
|
||||
firstphase_width: int = 0,
|
||||
firstphase_height: int = 0,
|
||||
hr_scale: float = 2.0,
|
||||
hr_force: bool = False,
|
||||
hr_resize_mode: int = 0,
|
||||
hr_resize_context: str = 'None',
|
||||
hr_upscaler: str = None,
|
||||
hr_second_pass_steps: int = 0,
|
||||
hr_resize_x: int = 0,
|
||||
hr_resize_y: int = 0,
|
||||
hr_denoising_strength: float = 0.0,
|
||||
refiner_steps: int = 5,
|
||||
refiner_start: float = 0,
|
||||
refiner_prompt: str = '',
|
||||
refiner_negative: str = '',
|
||||
hr_refiner_start: float = 0,
|
||||
# save options
|
||||
outpath_samples=None,
|
||||
outpath_grids=None,
|
||||
do_not_save_samples: bool = False,
|
||||
do_not_save_grid: bool = False,
|
||||
# scripts
|
||||
script_args: list = [],
|
||||
# overrides
|
||||
override_settings: Dict[str, Any] = {},
|
||||
override_settings_restore_afterwards: bool = True,
|
||||
sampler_index: int = None,
|
||||
script_args: list = None
|
||||
): # pylint: disable=unused-argument
|
||||
# metadata
|
||||
extra_generation_params: Dict[Any, Any] = {},
|
||||
):
|
||||
|
||||
# extra args set by processing loop
|
||||
self.task_args = {}
|
||||
|
||||
# state items
|
||||
self.state: str = ''
|
||||
self.ops = []
|
||||
self.skip = []
|
||||
self.outpath_samples: str = outpath_samples
|
||||
self.outpath_grids: str = outpath_grids
|
||||
self.prompt: str = prompt
|
||||
self.prompt_for_display: str = None
|
||||
self.negative_prompt: str = (negative_prompt or "")
|
||||
self.styles: list = styles or []
|
||||
self.seed: int = seed
|
||||
self.subseed: int = subseed
|
||||
self.subseed_strength: float = subseed_strength
|
||||
self.seed_resize_from_h: int = seed_resize_from_h
|
||||
self.seed_resize_from_w: int = seed_resize_from_w
|
||||
self.sampler_name: str = sampler_name
|
||||
self.hr_sampler_name: str = hr_sampler_name if hr_sampler_name != 'Same as primary' else sampler_name
|
||||
self.batch_size: int = batch_size
|
||||
self.n_iter: int = n_iter
|
||||
self.steps: int = steps
|
||||
self.hr_second_pass_steps = 0
|
||||
self.cfg_scale: float = cfg_scale
|
||||
self.scale_by: float = scale_by
|
||||
self.color_corrections = []
|
||||
self.is_control = False
|
||||
self.is_hr_pass = False
|
||||
self.is_refiner_pass = False
|
||||
self.is_api = False
|
||||
self.scheduled_prompt = False
|
||||
self.prompt_embeds = []
|
||||
self.positive_pooleds = []
|
||||
self.negative_embeds = []
|
||||
self.negative_pooleds = []
|
||||
self.disable_extra_networks = False
|
||||
self.iteration = 0
|
||||
|
||||
# initializers
|
||||
self.prompt = prompt
|
||||
self.seed = seed
|
||||
self.subseed = subseed
|
||||
self.subseed_strength = subseed_strength
|
||||
self.seed_resize_from_h = seed_resize_from_h
|
||||
self.seed_resize_from_w = seed_resize_from_w
|
||||
self.batch_size = batch_size
|
||||
self.n_iter = n_iter
|
||||
self.steps = steps
|
||||
self.clip_skip = clip_skip
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.negative_prompt = negative_prompt
|
||||
self.styles = styles
|
||||
self.tiling = tiling
|
||||
self.full_quality = full_quality
|
||||
self.hidiffusion = hidiffusion
|
||||
self.do_not_reload_embeddings = do_not_reload_embeddings
|
||||
self.detailer = detailer
|
||||
self.restore_faces = restore_faces
|
||||
self.init_images = init_images
|
||||
self.resize_mode = resize_mode
|
||||
self.resize_name = resize_name
|
||||
self.resize_context = resize_context
|
||||
self.denoising_strength = denoising_strength
|
||||
self.image_cfg_scale = image_cfg_scale
|
||||
self.scale_by = scale_by
|
||||
self.mask = mask
|
||||
self.image_mask = mask # TODO duplciate mask params
|
||||
self.latent_mask = latent_mask
|
||||
self.mask_blur = mask_blur
|
||||
self.inpainting_fill = inpainting_fill
|
||||
self.inpaint_full_res_padding = inpaint_full_res_padding
|
||||
self.inpainting_mask_invert = inpainting_mask_invert
|
||||
self.overlay_images = overlay_images
|
||||
self.enable_hr = enable_hr
|
||||
self.firstphase_width = firstphase_width
|
||||
self.firstphase_height = firstphase_height
|
||||
self.hr_scale = hr_scale
|
||||
self.hr_force = hr_force
|
||||
self.hr_resize_mode = hr_resize_mode
|
||||
self.hr_resize_context = hr_resize_context
|
||||
self.hr_upscaler = hr_upscaler
|
||||
self.hr_second_pass_steps = hr_second_pass_steps
|
||||
self.hr_resize_x = hr_resize_x
|
||||
self.hr_resize_y = hr_resize_y
|
||||
self.hr_upscale_to_x = hr_resize_x
|
||||
self.hr_upscale_to_y = hr_resize_y
|
||||
self.hr_denoising_strength = hr_denoising_strength
|
||||
self.refiner_steps = refiner_steps
|
||||
self.refiner_start = refiner_start
|
||||
self.refiner_prompt = refiner_prompt
|
||||
self.refiner_negative = refiner_negative
|
||||
self.hr_refiner_start = hr_refiner_start
|
||||
self.outpath_samples = outpath_samples
|
||||
self.outpath_grids = outpath_grids
|
||||
self.do_not_save_samples = do_not_save_samples
|
||||
self.do_not_save_grid = do_not_save_grid
|
||||
self.override_settings_restore_afterwards = override_settings_restore_afterwards
|
||||
self.extra_generation_params = extra_generation_params
|
||||
self.eta = eta
|
||||
self.cfg_scale = cfg_scale
|
||||
self.cfg_end = cfg_end
|
||||
self.diffusers_guidance_rescale = diffusers_guidance_rescale
|
||||
self.pag_scale = pag_scale
|
||||
self.pag_adaptive = pag_adaptive
|
||||
self.cfg_end = cfg_end
|
||||
self.width: int = width
|
||||
self.height: int = height
|
||||
self.full_quality: bool = full_quality
|
||||
self.detailer: bool = detailer
|
||||
self.restore_faces: bool = restore_faces
|
||||
self.tiling: bool = tiling
|
||||
self.hidiffusion: bool = hidiffusion
|
||||
self.do_not_save_samples: bool = do_not_save_samples
|
||||
self.do_not_save_grid: bool = do_not_save_grid
|
||||
self.extra_generation_params: dict = extra_generation_params or {}
|
||||
self.overlay_images = overlay_images
|
||||
self.eta = eta
|
||||
self.do_not_reload_embeddings = do_not_reload_embeddings
|
||||
self.paste_to = None
|
||||
self.color_corrections = None
|
||||
self.denoising_strength: float = denoising_strength
|
||||
self.selected_scale_tab = selected_scale_tab
|
||||
self.mask_for_overlay = mask_for_overlay
|
||||
self.paste_to = paste_to
|
||||
self.init_latent = None
|
||||
|
||||
# special handled items
|
||||
if firstphase_width != 0 or firstphase_height != 0:
|
||||
self.hr_upscale_to_x = self.width
|
||||
self.hr_upscale_to_y = self.height
|
||||
self.width = firstphase_width
|
||||
self.height = firstphase_height
|
||||
self.sampler_name = sampler_name or processing_helpers.get_sampler_name(sampler_index, img=True)
|
||||
self.hr_sampler_name: str = hr_sampler_name if hr_sampler_name != 'Same as primary' else self.sampler_name
|
||||
self.override_settings = {k: v for k, v in (override_settings or {}).items() if k not in shared.restricted_opts}
|
||||
self.override_settings_restore_afterwards = override_settings_restore_afterwards
|
||||
self.is_using_inpainting_conditioning = False # a111 compatibility
|
||||
self.disable_extra_networks = False
|
||||
# self.scripts = scripts.ScriptRunner() # set via property
|
||||
# self.script_args = script_args or [] # set via property
|
||||
self.per_script_args = {}
|
||||
self.inpaint_full_res = inpaint_full_res if isinstance(inpaint_full_res, bool) else self.inpaint_full_res
|
||||
self.inpaint_full_res = inpaint_full_res != 0 if isinstance(inpaint_full_res, int) else self.inpaint_full_res
|
||||
|
||||
# null items initialized later
|
||||
self.all_prompts = None
|
||||
self.all_negative_prompts = None
|
||||
self.all_seeds = None
|
||||
self.all_subseeds = None
|
||||
self.clip_skip = clip_skip
|
||||
|
||||
# a1111 compatibility items
|
||||
shared.opts.data['clip_skip'] = int(self.clip_skip) # for compatibility with a1111 sd_hijack_clip
|
||||
self.iteration = 0
|
||||
self.is_control = False
|
||||
self.is_hr_pass = False
|
||||
self.is_refiner_pass = False
|
||||
self.hr_force = False
|
||||
self.enable_hr = None
|
||||
self.hr_scale = None
|
||||
self.hr_upscaler = None
|
||||
self.hr_resize_mode = 0
|
||||
self.hr_resize_context = 'None'
|
||||
self.hr_resize_x = 0
|
||||
self.hr_resize_y = 0
|
||||
self.hr_upscale_to_x = 0
|
||||
self.hr_upscale_to_y = 0
|
||||
self.seed_enable_extras: bool = True
|
||||
self.is_using_inpainting_conditioning = False # a111 compatibility
|
||||
self.batch_index = 0
|
||||
self.refiner_switch_at = 0
|
||||
self.hr_prompt = ''
|
||||
self.all_hr_prompts = []
|
||||
self.hr_negative_prompt = ''
|
||||
self.all_hr_negative_prompts = []
|
||||
self.truncate_x = 0
|
||||
self.truncate_y = 0
|
||||
self.applied_old_hires_behavior_to = None
|
||||
self.refiner_steps = 5
|
||||
self.refiner_start = 0
|
||||
self.refiner_prompt = ''
|
||||
self.refiner_negative = ''
|
||||
self.ops = []
|
||||
self.resize_mode: int = resize_mode
|
||||
self.resize_name: str = resize_name
|
||||
self.resize_context: str = resize_context
|
||||
self.comments = {}
|
||||
self.sampler = None
|
||||
self.nmask = None
|
||||
self.initial_noise_multiplier = initial_noise_multiplier or shared.opts.initial_noise_multiplier
|
||||
self.image_conditioning = None
|
||||
self.prompt_for_display: str = None
|
||||
|
||||
# scripts
|
||||
self.scripts_value: scripts.ScriptRunner = field(default=None, init=False)
|
||||
self.script_args_value: list = field(default=None, init=False)
|
||||
self.scripts_setup_complete: bool = field(default=False, init=False)
|
||||
self.script_args = script_args
|
||||
self.per_script_args = {}
|
||||
|
||||
# settings to processing
|
||||
self.ddim_discretize = shared.opts.ddim_discretize
|
||||
self.s_min_uncond = shared.opts.s_min_uncond
|
||||
self.s_churn = shared.opts.s_churn
|
||||
@@ -171,18 +266,7 @@ class StableDiffusionProcessing:
|
||||
self.s_tmin = shared.opts.s_tmin
|
||||
self.s_tmax = float('inf') # not representable as a standard ui option
|
||||
self.task_args = {}
|
||||
# a1111 compatibility items
|
||||
self.batch_index = 0
|
||||
self.refiner_switch_at = 0
|
||||
self.hr_prompt = ''
|
||||
self.all_hr_prompts = []
|
||||
self.hr_negative_prompt = ''
|
||||
self.all_hr_negative_prompts = []
|
||||
self.comments = {}
|
||||
self.is_api = False
|
||||
self.scripts_value: scripts.ScriptRunner = field(default=None, init=False)
|
||||
self.script_args_value: list = field(default=None, init=False)
|
||||
self.scripts_setup_complete: bool = field(default=False, init=False)
|
||||
|
||||
# ip adapter
|
||||
self.ip_adapter_names = []
|
||||
self.ip_adapter_scales = [0.0]
|
||||
@@ -190,6 +274,7 @@ class StableDiffusionProcessing:
|
||||
self.ip_adapter_starts = [0.0]
|
||||
self.ip_adapter_ends = [1.0]
|
||||
self.ip_adapter_crops = []
|
||||
|
||||
# hdr
|
||||
self.hdr_mode=hdr_mode
|
||||
self.hdr_brightness=hdr_brightness
|
||||
@@ -203,7 +288,10 @@ class StableDiffusionProcessing:
|
||||
self.hdr_max_boundry=hdr_max_boundry
|
||||
self.hdr_color_picker=hdr_color_picker
|
||||
self.hdr_tint_ratio=hdr_tint_ratio
|
||||
|
||||
# globals
|
||||
self.embedder = None
|
||||
self.override = None
|
||||
self.scheduled_prompt: bool = False
|
||||
self.prompt_embeds = []
|
||||
self.positive_pooleds = []
|
||||
@@ -252,57 +340,9 @@ class StableDiffusionProcessing:
|
||||
|
||||
|
||||
class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
|
||||
|
||||
def __init__(self,
|
||||
enable_hr: bool = False,
|
||||
denoising_strength: float = 0.75,
|
||||
firstphase_width: int = 0,
|
||||
firstphase_height: int = 0,
|
||||
hr_scale: float = 2.0,
|
||||
hr_force: bool = False,
|
||||
hr_resize_mode: int = 0,
|
||||
hr_resize_context: str = 'None',
|
||||
hr_upscaler: str = None,
|
||||
hr_second_pass_steps: int = 0,
|
||||
hr_resize_x: int = 0,
|
||||
hr_resize_y: int = 0,
|
||||
refiner_steps: int = 5,
|
||||
refiner_start: float = 0,
|
||||
refiner_prompt: str = '',
|
||||
refiner_negative: str = '',
|
||||
**kwargs
|
||||
):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access
|
||||
super().__init__(**kwargs)
|
||||
self.reprocess = {}
|
||||
self.enable_hr = enable_hr
|
||||
self.denoising_strength = denoising_strength
|
||||
self.hr_scale = hr_scale
|
||||
self.hr_upscaler = hr_upscaler
|
||||
self.hr_resize_mode = hr_resize_mode
|
||||
self.hr_resize_context = hr_resize_context
|
||||
self.hr_force = hr_force
|
||||
self.hr_second_pass_steps = hr_second_pass_steps
|
||||
self.hr_resize_x = hr_resize_x
|
||||
self.hr_resize_y = hr_resize_y
|
||||
self.hr_upscale_to_x = hr_resize_x
|
||||
self.hr_upscale_to_y = hr_resize_y
|
||||
if firstphase_width != 0 or firstphase_height != 0:
|
||||
self.hr_upscale_to_x = self.width
|
||||
self.hr_upscale_to_y = self.height
|
||||
self.width = firstphase_width
|
||||
self.height = firstphase_height
|
||||
self.truncate_x = 0
|
||||
self.truncate_y = 0
|
||||
self.applied_old_hires_behavior_to = None
|
||||
self.refiner_steps = refiner_steps
|
||||
self.refiner_start = refiner_start
|
||||
self.refiner_prompt = refiner_prompt
|
||||
self.refiner_negative = refiner_negative
|
||||
self.sampler = None
|
||||
self.scripts = None
|
||||
self.script_args = []
|
||||
|
||||
|
||||
def init(self, all_prompts=None, all_seeds=None, all_subseeds=None):
|
||||
if shared.native:
|
||||
@@ -360,41 +400,9 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
|
||||
|
||||
|
||||
class StableDiffusionProcessingImg2Img(StableDiffusionProcessing):
|
||||
|
||||
def __init__(self, init_images: list = None, resize_mode: int = 0, resize_name: str = 'None', resize_context: str = 'None', denoising_strength: float = 0.3, image_cfg_scale: float = None, mask: Any = None, mask_blur: int = 4, inpainting_fill: int = 0, inpaint_full_res: bool = False, inpaint_full_res_padding: int = 0, inpainting_mask_invert: int = 0, initial_noise_multiplier: float = None, scale_by: float = 1, refiner_steps: int = 5, refiner_start: float = 0, refiner_prompt: str = '', refiner_negative: str = '', **kwargs):
|
||||
def __init__(self, **kwargs):
|
||||
debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access
|
||||
super().__init__(**kwargs)
|
||||
self.init_images = init_images
|
||||
self.resize_mode: int = resize_mode
|
||||
self.resize_name: str = resize_name
|
||||
self.resize_context: str = resize_context
|
||||
self.denoising_strength: float = denoising_strength
|
||||
self.hr_denoising_strength: float = denoising_strength
|
||||
self.image_cfg_scale: float = image_cfg_scale
|
||||
self.init_latent = None
|
||||
self.image_mask = mask
|
||||
self.latent_mask = None
|
||||
self.mask_for_overlay = None
|
||||
self.mask_blur_x = mask_blur # a1111 compatibility item
|
||||
self.mask_blur_y = mask_blur # a1111 compatibility item
|
||||
self.mask_blur = mask_blur
|
||||
self.inpainting_fill = inpainting_fill
|
||||
self.inpaint_full_res = inpaint_full_res
|
||||
self.inpaint_full_res_padding = inpaint_full_res_padding
|
||||
self.inpainting_mask_invert = inpainting_mask_invert
|
||||
self.initial_noise_multiplier = shared.opts.initial_noise_multiplier if initial_noise_multiplier is None else initial_noise_multiplier
|
||||
self.mask = None
|
||||
self.nmask = None
|
||||
self.image_conditioning = None
|
||||
self.refiner_steps = refiner_steps
|
||||
self.refiner_start = refiner_start
|
||||
self.refiner_prompt = refiner_prompt
|
||||
self.refiner_negative = refiner_negative
|
||||
self.enable_hr = None
|
||||
self.is_batch = False
|
||||
self.scale_by = scale_by
|
||||
self.sampler = None
|
||||
self.scripts = None
|
||||
self.script_args = []
|
||||
|
||||
def init(self, all_prompts=None, all_seeds=None, all_subseeds=None):
|
||||
if hasattr(self, 'init_images') and self.init_images is not None and len(self.init_images) > 0:
|
||||
@@ -485,7 +493,7 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing):
|
||||
image = images.resize_image(self.resize_mode, image, self.width, self.height, upscaler_name=self.resize_name, context=self.resize_context)
|
||||
self.width = image.width
|
||||
self.height = image.height
|
||||
if self.image_mask is not None and shared.opts.mask_apply_overlay:
|
||||
if self.image_mask is not None and shared.opts.mask_apply_overlay and not hasattr(self, 'xyz'):
|
||||
image_masked = Image.new('RGBa', (image.width, image.height))
|
||||
image_to_paste = image.convert("RGBA").convert("RGBa")
|
||||
image_to_mask = ImageOps.invert(self.mask_for_overlay.convert('L')) if self.mask_for_overlay is not None else None
|
||||
@@ -544,47 +552,8 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing):
|
||||
|
||||
class StableDiffusionProcessingControl(StableDiffusionProcessingImg2Img):
|
||||
def __init__(self, **kwargs):
|
||||
debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access
|
||||
super().__init__(**kwargs)
|
||||
self.strength = None
|
||||
self.adapter_conditioning_scale = None
|
||||
self.adapter_conditioning_factor = None
|
||||
self.guess_mode = None
|
||||
self.controlnet_conditioning_scale = None
|
||||
self.control_guidance_start = None
|
||||
self.control_guidance_end = None
|
||||
self.control_mode = None
|
||||
self.reference_attn = None
|
||||
self.reference_adain = None
|
||||
self.attention_auto_machine_weight = None
|
||||
self.gn_auto_machine_weight = None
|
||||
self.style_fidelity = None
|
||||
self.ref_image = None
|
||||
self.image = None
|
||||
self.query_weight = None
|
||||
self.adain_weight = None
|
||||
self.adapter_conditioning_factor = 1.0
|
||||
self.attention = 'Attention'
|
||||
self.fidelity = 0.5
|
||||
self.mask_image = None
|
||||
self.override = None
|
||||
self.resize_mode_before = None
|
||||
self.resize_name_before = None
|
||||
self.width_before = None
|
||||
self.height_before = None
|
||||
self.scale_by_before = None
|
||||
self.selected_scale_tab_before = None
|
||||
self.resize_mode_after = None
|
||||
self.resize_name_after = None
|
||||
self.width_after = None
|
||||
self.height_after = None
|
||||
self.scale_by_after = None
|
||||
self.selected_scale_tab_after = None
|
||||
self.resize_mode_mask = None
|
||||
self.resize_name_mask = None
|
||||
self.width_mask = None
|
||||
self.height_mask = None
|
||||
self.scale_by_mask = None
|
||||
self.selected_scale_tab_mask = None
|
||||
|
||||
def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): # abstract
|
||||
pass
|
||||
|
||||
@@ -71,7 +71,7 @@ def process_base(p: processing.StableDiffusionProcessing):
|
||||
guidance_rescale=p.diffusers_guidance_rescale,
|
||||
denoising_start=0 if use_refiner_start else p.refiner_start if use_denoise_start else None,
|
||||
denoising_end=p.refiner_start if use_refiner_start else 1 if use_denoise_start else None,
|
||||
output_type='latent' if hasattr(shared.sd_model, 'vae') else 'np',
|
||||
output_type='latent',
|
||||
clip_skip=p.clip_skip,
|
||||
desc='Base',
|
||||
)
|
||||
@@ -101,7 +101,8 @@ def process_base(p: processing.StableDiffusionProcessing):
|
||||
output = SimpleNamespace(**output)
|
||||
if isinstance(output, list):
|
||||
output = SimpleNamespace(images=output)
|
||||
shared.history.add(output.images, info=processing.create_infotext(p), ops=p.ops)
|
||||
if hasattr(output, 'images'):
|
||||
shared.history.add(output.images, info=processing.create_infotext(p), ops=p.ops)
|
||||
timer.process.record('pipeline')
|
||||
hidiffusion.unapply()
|
||||
sd_models_compile.openvino_post_compile(op="base") # only executes on compiled vino models
|
||||
@@ -151,13 +152,14 @@ def process_hires(p: processing.StableDiffusionProcessing, output):
|
||||
p.is_hr_pass = True
|
||||
if hasattr(p, 'init_hr'):
|
||||
p.init_hr(p.hr_scale, p.hr_upscaler, force=p.hr_force)
|
||||
else: # fake hires for img2img
|
||||
p.hr_scale = p.scale_by
|
||||
p.hr_upscaler = p.resize_name
|
||||
p.hr_resize_mode = p.resize_mode
|
||||
p.hr_resize_context = p.resize_context
|
||||
p.hr_upscale_to_x = p.width
|
||||
p.hr_upscale_to_y = p.height
|
||||
else:
|
||||
if not p.is_hr_pass: # fake hires for img2img if not actual hr pass
|
||||
p.hr_scale = p.scale_by
|
||||
p.hr_upscaler = p.resize_name
|
||||
p.hr_resize_mode = p.resize_mode
|
||||
p.hr_resize_context = p.resize_context
|
||||
p.hr_upscale_to_x = p.width * p.hr_scale if p.hr_resize_x == 0 else p.hr_resize_x
|
||||
p.hr_upscale_to_y = p.height * p.hr_scale if p.hr_resize_y == 0 else p.hr_resize_y
|
||||
prev_job = shared.state.job
|
||||
|
||||
# hires runs on original pipeline
|
||||
@@ -175,7 +177,8 @@ def process_hires(p: processing.StableDiffusionProcessing, output):
|
||||
sd_hijack_hypertile.hypertile_set(p, hr=True)
|
||||
|
||||
latent_upscale = shared.latent_upscale_modes.get(p.hr_upscaler, None)
|
||||
if (latent_upscale is not None or p.hr_force) and getattr(p, 'hr_denoising_strength', p.denoising_strength) > 0:
|
||||
strength = p.hr_denoising_strength if p.hr_denoising_strength > 0 else p.denoising_strength
|
||||
if (latent_upscale is not None or p.hr_force) and strength > 0:
|
||||
p.ops.append('hires')
|
||||
sd_models_compile.openvino_recompile_model(p, hires=True, refiner=False)
|
||||
if shared.sd_model.__class__.__name__ == "OnnxRawPipeline":
|
||||
@@ -183,8 +186,7 @@ def process_hires(p: processing.StableDiffusionProcessing, output):
|
||||
p.hr_force = True
|
||||
|
||||
# hires
|
||||
p.denoising_strength = getattr(p, 'hr_denoising_strength', p.denoising_strength)
|
||||
if p.hr_force and p.denoising_strength == 0:
|
||||
if p.hr_force and strength == 0:
|
||||
shared.log.warning('HiRes skip: denoising=0')
|
||||
p.hr_force = False
|
||||
if p.hr_force:
|
||||
@@ -202,9 +204,9 @@ def process_hires(p: processing.StableDiffusionProcessing, output):
|
||||
sd_models.move_model(shared.sd_model.unet, devices.device)
|
||||
if hasattr(shared.sd_model, 'transformer'):
|
||||
sd_models.move_model(shared.sd_model.transformer, devices.device)
|
||||
orig_denoise = p.denoising_strength
|
||||
p.denoising_strength = getattr(p, 'hr_denoising_strength', p.denoising_strength)
|
||||
update_sampler(p, shared.sd_model, second_pass=True)
|
||||
orig_denoise = p.denoising_strength
|
||||
p.denoising_strength = strength
|
||||
hires_args = set_pipeline_args(
|
||||
p=p,
|
||||
model=shared.sd_model,
|
||||
@@ -216,10 +218,10 @@ def process_hires(p: processing.StableDiffusionProcessing, output):
|
||||
eta=shared.opts.scheduler_eta,
|
||||
guidance_scale=p.image_cfg_scale if p.image_cfg_scale is not None else p.cfg_scale,
|
||||
guidance_rescale=p.diffusers_guidance_rescale,
|
||||
output_type='latent' if hasattr(shared.sd_model, 'vae') else 'np',
|
||||
output_type='latent',
|
||||
clip_skip=p.clip_skip,
|
||||
image=output.images,
|
||||
strength=p.denoising_strength,
|
||||
strength=strength,
|
||||
desc='Hires',
|
||||
)
|
||||
shared.state.job = 'HiRes'
|
||||
@@ -277,7 +279,7 @@ def process_refine(p: processing.StableDiffusionProcessing, output):
|
||||
for i in range(len(output.images)):
|
||||
image = output.images[i]
|
||||
noise_level = round(350 * p.denoising_strength)
|
||||
output_type='latent' if hasattr(shared.sd_refiner, 'vae') else 'np'
|
||||
output_type='latent'
|
||||
if 'Upscale' in shared.sd_refiner.__class__.__name__ or 'Flux' in shared.sd_refiner.__class__.__name__:
|
||||
image = processing_vae.vae_decode(latents=image, model=shared.sd_model, full_quality=p.full_quality, output_type='pil', width=p.width, height=p.height)
|
||||
p.extra_generation_params['Noise level'] = noise_level
|
||||
@@ -345,7 +347,11 @@ def process_decode(p: processing.StableDiffusionProcessing, output):
|
||||
if not hasattr(output, 'images') and hasattr(output, 'frames'):
|
||||
shared.log.debug(f'Generated: frames={len(output.frames[0])}')
|
||||
output.images = output.frames[0]
|
||||
if hasattr(shared.sd_model, "vae") and output.images is not None and len(output.images) > 0:
|
||||
model = shared.sd_model if not is_refiner_enabled(p) else shared.sd_refiner
|
||||
if not hasattr(model, 'vae'):
|
||||
if hasattr(model, 'pipe') and hasattr(model.pipe, 'vae'):
|
||||
model = model.pipe
|
||||
if hasattr(model, "vae") and output.images is not None and len(output.images) > 0:
|
||||
if p.hr_resize_mode > 0 and (p.hr_upscaler != 'None' or p.hr_resize_mode == 5):
|
||||
width = max(getattr(p, 'width', 0), getattr(p, 'hr_upscale_to_x', 0))
|
||||
height = max(getattr(p, 'height', 0), getattr(p, 'hr_upscale_to_y', 0))
|
||||
@@ -354,7 +360,7 @@ def process_decode(p: processing.StableDiffusionProcessing, output):
|
||||
height = getattr(p, 'height', 0)
|
||||
results = processing_vae.vae_decode(
|
||||
latents = output.images,
|
||||
model = shared.sd_model if not is_refiner_enabled(p) else shared.sd_refiner,
|
||||
model = model,
|
||||
full_quality = p.full_quality,
|
||||
width = width,
|
||||
height = height,
|
||||
|
||||
@@ -47,16 +47,23 @@ def apply_overlay(image: Image, paste_loc, index, overlays):
|
||||
return image
|
||||
debug(f'Apply overlay: image={image} loc={paste_loc} index={index} overlays={overlays}')
|
||||
overlay = overlays[index]
|
||||
if paste_loc is not None:
|
||||
x, y, w, h = paste_loc
|
||||
if image.width != w or image.height != h or x != 0 or y != 0:
|
||||
base_image = Image.new('RGBA', (overlay.width, overlay.height))
|
||||
image = images.resize_image(2, image, w, h)
|
||||
base_image.paste(image, (x, y))
|
||||
image = base_image
|
||||
image = image.convert('RGBA')
|
||||
image.alpha_composite(overlay)
|
||||
image = image.convert('RGB')
|
||||
if not isinstance(image, Image.Image) or not isinstance(overlay, Image.Image):
|
||||
return image
|
||||
try:
|
||||
if paste_loc is not None and (isinstance(paste_loc, tuple) or isinstance(paste_loc, list)):
|
||||
x, y, w, h = paste_loc
|
||||
if x is None or y is None or w is None or h is None:
|
||||
return image
|
||||
if image.width != w or image.height != h or x != 0 or y != 0:
|
||||
base_image = Image.new('RGBA', (overlay.width, overlay.height))
|
||||
image = images.resize_image(2, image, w, h)
|
||||
base_image.paste(image, (x, y))
|
||||
image = base_image
|
||||
image = image.convert('RGBA')
|
||||
image.alpha_composite(overlay)
|
||||
image = image.convert('RGB')
|
||||
except Exception as e:
|
||||
shared.log.error(f'Apply overlay: {e}')
|
||||
return image
|
||||
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ from modules import shared, sd_samplers_common, sd_vae, generation_parameters_co
|
||||
from modules.processing_class import StableDiffusionProcessing
|
||||
|
||||
|
||||
debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
if not shared.native:
|
||||
from modules import sd_hijack
|
||||
else:
|
||||
@@ -39,30 +40,34 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No
|
||||
ops.reverse()
|
||||
args = {
|
||||
# basic
|
||||
"Size": f"{p.width}x{p.height}" if hasattr(p, 'width') and hasattr(p, 'height') else None,
|
||||
"Sampler": p.sampler_name if p.sampler_name != 'Default' else None,
|
||||
"Steps": p.steps,
|
||||
"Seed": all_seeds[index],
|
||||
"Sampler": p.sampler_name if p.sampler_name != 'Default' else None,
|
||||
"Seed resize from": None if p.seed_resize_from_w == 0 or p.seed_resize_from_h == 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}",
|
||||
"CFG scale": p.cfg_scale if p.cfg_scale > 1.0 else None,
|
||||
"CFG end": p.cfg_end if p.cfg_end < 1.0 else None,
|
||||
"Size": f"{p.width}x{p.height}" if hasattr(p, 'width') and hasattr(p, 'height') else None,
|
||||
"Clip skip": p.clip_skip if p.clip_skip > 1 else None,
|
||||
"Batch": f'{p.n_iter}x{p.batch_size}' if p.n_iter > 1 or p.batch_size > 1 else None,
|
||||
"Parser": shared.opts.prompt_attention.split()[0],
|
||||
"Model": None if (not shared.opts.add_model_name_to_info) or (not shared.sd_model.sd_checkpoint_info.model_name) else shared.sd_model.sd_checkpoint_info.model_name.replace(',', '').replace(':', ''),
|
||||
"Model hash": getattr(p, 'sd_model_hash', None if (not shared.opts.add_model_hash_to_info) or (not shared.sd_model.sd_model_hash) else shared.sd_model.sd_model_hash),
|
||||
"VAE": (None if not shared.opts.add_model_name_to_info or sd_vae.loaded_vae_file is None else os.path.splitext(os.path.basename(sd_vae.loaded_vae_file))[0]) if p.full_quality else 'TAESD',
|
||||
"Seed resize from": None if p.seed_resize_from_w == 0 or p.seed_resize_from_h == 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}",
|
||||
"Clip skip": p.clip_skip if p.clip_skip > 1 else None,
|
||||
"Prompt2": p.refiner_prompt if len(p.refiner_prompt) > 0 else None,
|
||||
"Negative2": p.refiner_negative if len(p.refiner_negative) > 0 else None,
|
||||
"Styles": "; ".join(p.styles) if p.styles is not None and len(p.styles) > 0 else None,
|
||||
"Tiling": p.tiling if p.tiling else None,
|
||||
# sdnext
|
||||
"Backend": 'Diffusers' if shared.native else 'Original',
|
||||
"App": 'SD.Next',
|
||||
"Version": git_commit,
|
||||
"Backend": 'Diffusers' if shared.native else 'Original',
|
||||
"Pipeline": 'LDM',
|
||||
"Parser": shared.opts.prompt_attention.split()[0],
|
||||
"Comment": comment,
|
||||
"Operations": '; '.join(ops).replace('"', '') if len(p.ops) > 0 else 'none',
|
||||
}
|
||||
if shared.opts.add_model_name_to_info and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None:
|
||||
args["Model"] = shared.sd_model.sd_checkpoint_info.model_name.replace(',', '').replace(':', '')
|
||||
if shared.opts.add_model_hash_to_info and getattr(shared.sd_model, 'sd_model_hash', None) is not None:
|
||||
args["Model hash"] = shared.sd_model.sd_model_hash
|
||||
# native
|
||||
if grid is None and (p.n_iter > 1 or p.batch_size > 1) and index >= 0:
|
||||
args['Index'] = f'{p.iteration + 1}x{index + 1}'
|
||||
@@ -165,7 +170,9 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No
|
||||
if isinstance(v, str):
|
||||
if len(v) == 0 or v == '0x0':
|
||||
del args[k]
|
||||
debug(f'Infotext: args={args}')
|
||||
params_text = ", ".join([k if k == v else f'{k}: {generation_parameters_copypaste.quote(v)}' for k, v in args.items()])
|
||||
negative_prompt_text = f"\nNegative prompt: {all_negative_prompts[index]}" if all_negative_prompts[index] else ""
|
||||
infotext = f"{all_prompts[index]}{negative_prompt_text}\n{params_text}".strip()
|
||||
debug(f'Infotext: "{infotext}"')
|
||||
return infotext
|
||||
|
||||
@@ -35,7 +35,9 @@ def create_latents(image, p, dtype=None, device=None):
|
||||
|
||||
def full_vae_decode(latents, model):
|
||||
t0 = time.time()
|
||||
if not hasattr(model, 'vae'):
|
||||
if not hasattr(model, 'vae') and hasattr(model, 'pipe'):
|
||||
model = model.pipe
|
||||
if model is None or not hasattr(model, 'vae'):
|
||||
shared.log.error('VAE not found in model')
|
||||
return []
|
||||
if debug:
|
||||
@@ -147,6 +149,9 @@ def taesd_vae_encode(image):
|
||||
|
||||
def vae_decode(latents, model, output_type='np', full_quality=True, width=None, height=None):
|
||||
t0 = time.time()
|
||||
model = model or shared.sd_model
|
||||
if not hasattr(model, 'vae') and hasattr(model, 'pipe'):
|
||||
model = model.pipe
|
||||
if latents is None or not torch.is_tensor(latents): # already decoded
|
||||
return latents
|
||||
prev_job = shared.state.job
|
||||
@@ -169,8 +174,8 @@ def vae_decode(latents, model, output_type='np', full_quality=True, width=None,
|
||||
|
||||
if latents.shape[-1] <= 4: # not a latent, likely an image
|
||||
decoded = latents.float().cpu().numpy()
|
||||
elif full_quality and hasattr(shared.sd_model, "vae"):
|
||||
decoded = full_vae_decode(latents=latents, model=shared.sd_model)
|
||||
elif full_quality and hasattr(model, "vae"):
|
||||
decoded = full_vae_decode(latents=latents, model=model)
|
||||
else:
|
||||
decoded = taesd_vae_decode(latents=latents)
|
||||
|
||||
@@ -195,6 +200,8 @@ def vae_decode(latents, model, output_type='np', full_quality=True, width=None,
|
||||
def vae_encode(image, model, full_quality=True): # pylint: disable=unused-variable
|
||||
if shared.state.interrupted or shared.state.skipped:
|
||||
return []
|
||||
if not hasattr(model, 'vae') and hasattr(model, 'pipe'):
|
||||
model = model.pipe
|
||||
if not hasattr(model, 'vae'):
|
||||
shared.log.error('VAE not found in model')
|
||||
return []
|
||||
|
||||
@@ -73,7 +73,6 @@ def progressapi(req: ProgressRequest):
|
||||
elapsed = time.time() - shared.state.time_start if shared.state.time_start is not None else 0
|
||||
predicted = elapsed / progress if progress > 0 else None
|
||||
eta = predicted - elapsed if predicted is not None else None
|
||||
# shared.log.debug(f'Progress: step={step_x}:{step_y} batch={batch_x}:{batch_y} current={current} total={total} progress={progress} elapsed={elapsed} eta={eta}')
|
||||
id_live_preview = req.id_live_preview
|
||||
live_preview = None
|
||||
shared.state.set_current_image()
|
||||
|
||||
@@ -308,11 +308,11 @@ def parse_prompt_attention(text):
|
||||
res = []
|
||||
round_brackets = []
|
||||
square_brackets = []
|
||||
if opts.prompt_attention == 'Fixed attention':
|
||||
if opts.prompt_attention == 'fixed':
|
||||
res = [[text, 1.0]]
|
||||
debug(f'Prompt: parser="{opts.prompt_attention}" {res}')
|
||||
return res
|
||||
elif opts.prompt_attention == 'Compel parser':
|
||||
elif opts.prompt_attention == 'compel':
|
||||
conjunction = Compel.parse_prompt_string(text)
|
||||
if conjunction is None or conjunction.prompts is None or conjunction.prompts is None or len(conjunction.prompts[0].children) == 0:
|
||||
return [["", 1.0]]
|
||||
@@ -321,7 +321,7 @@ def parse_prompt_attention(text):
|
||||
res.append([frag.text, frag.weight])
|
||||
debug(f'Prompt: parser="{opts.prompt_attention}" {res}')
|
||||
return res
|
||||
elif opts.prompt_attention == 'A1111 parser':
|
||||
elif opts.prompt_attention == 'a1111':
|
||||
re_attention = re_attention_v1
|
||||
whitespace = ''
|
||||
else:
|
||||
@@ -360,7 +360,7 @@ def parse_prompt_attention(text):
|
||||
for i, part in enumerate(parts):
|
||||
if i > 0:
|
||||
res.append(["BREAK", -1])
|
||||
if opts.prompt_attention == 'Full parser':
|
||||
if opts.prompt_attention == 'native':
|
||||
part = re_clean.sub("", part)
|
||||
part = re_whitespace.sub(" ", part).strip()
|
||||
if len(part) == 0:
|
||||
@@ -392,15 +392,15 @@ if __name__ == "__main__":
|
||||
log.info(f'Schedules: {all_schedules}')
|
||||
for schedule in all_schedules:
|
||||
log.info(f'Schedule: {schedule[0]}')
|
||||
opts.data['prompt_attention'] = 'Fixed attention'
|
||||
opts.data['prompt_attention'] = 'fixed'
|
||||
output_list = parse_prompt_attention(schedule[1])
|
||||
log.info(f' Fixed: {output_list}')
|
||||
opts.data['prompt_attention'] = 'Compel parser'
|
||||
opts.data['prompt_attention'] = 'compel'
|
||||
output_list = parse_prompt_attention(schedule[1])
|
||||
log.info(f' Compel: {output_list}')
|
||||
opts.data['prompt_attention'] = 'A1111 parser'
|
||||
opts.data['prompt_attention'] = 'a1111'
|
||||
output_list = parse_prompt_attention(schedule[1])
|
||||
log.info(f' A1111: {output_list}')
|
||||
opts.data['prompt_attention'] = 'Full parser'
|
||||
opts.data['prompt_attention'] = 'native'
|
||||
log.info = parse_prompt_attention(schedule[1])
|
||||
log.info(f' Full: {output_list}')
|
||||
|
||||
+211
-145
@@ -2,6 +2,7 @@ import os
|
||||
import math
|
||||
import time
|
||||
import typing
|
||||
from collections import OrderedDict
|
||||
import torch
|
||||
from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider
|
||||
from transformers import PreTrainedTokenizer
|
||||
@@ -14,7 +15,182 @@ debug('Trace: PROMPT')
|
||||
orig_encode_token_ids_to_embeddings = EmbeddingsProvider._encode_token_ids_to_embeddings # pylint: disable=protected-access
|
||||
token_dict = None # used by helper get_tokens
|
||||
token_type = None # used by helper get_tokens
|
||||
cache = {}
|
||||
cache = OrderedDict()
|
||||
embedder = None
|
||||
|
||||
|
||||
def prompt_compatible(pipe = None):
|
||||
pipe = pipe or shared.sd_model
|
||||
if (
|
||||
'StableDiffusion' not in pipe.__class__.__name__ and
|
||||
'DemoFusion' not in pipe.__class__.__name__ and
|
||||
'StableCascade' not in pipe.__class__.__name__ and
|
||||
'Flux' not in pipe.__class__.__name__
|
||||
):
|
||||
shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}")
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def prepare_model(pipe = None):
|
||||
pipe = pipe or shared.sd_model
|
||||
if not hasattr(pipe, "text_encoder") and hasattr(shared.sd_model, "pipe"):
|
||||
pipe = pipe.pipe
|
||||
if not hasattr(pipe, "text_encoder"):
|
||||
return None
|
||||
if shared.opts.diffusers_offload_mode == "balanced":
|
||||
pipe = sd_models.apply_balanced_offload(pipe)
|
||||
elif hasattr(pipe, "maybe_free_model_hooks"):
|
||||
pipe.maybe_free_model_hooks()
|
||||
devices.torch_gc()
|
||||
return pipe
|
||||
|
||||
|
||||
class PromptEmbedder:
|
||||
def __init__(self, prompts, negative_prompts, steps, clip_skip, p):
|
||||
t0 = time.time()
|
||||
self.prompts = prompts
|
||||
self.negative_prompts = negative_prompts
|
||||
self.batchsize = len(self.prompts)
|
||||
self.attention = None
|
||||
self.allsame = self.compare_prompts() # collapses batched prompts to single prompt if possible
|
||||
self.steps = steps
|
||||
self.clip_skip = clip_skip
|
||||
# All embeds are nested lists, outer list batch length, inner schedule length
|
||||
self.prompt_embeds = [[]] * self.batchsize
|
||||
self.positive_pooleds = [[]] * self.batchsize
|
||||
self.negative_prompt_embeds = [[]] * self.batchsize
|
||||
self.negative_pooleds = [[]] * self.batchsize
|
||||
self.positive_schedule = None
|
||||
self.negative_schedule = None
|
||||
self.scheduled_prompt = False
|
||||
earlyout = self.checkcache(p)
|
||||
if earlyout:
|
||||
return
|
||||
pipe = prepare_model(p.sd_model)
|
||||
if pipe is None:
|
||||
shared.log.error("Prompt encode: cannot find text encoder in model")
|
||||
return
|
||||
# per prompt in batch
|
||||
for batchidx, (prompt, negative_prompt) in enumerate(zip(self.prompts, self.negative_prompts)):
|
||||
self.prepare_schedule(prompt, negative_prompt)
|
||||
if self.scheduled_prompt:
|
||||
self.scheduled_encode(pipe, batchidx)
|
||||
else:
|
||||
self.encode(pipe, prompt, negative_prompt, batchidx)
|
||||
self.checkcache(p)
|
||||
debug(f"Prompt encode: time={(time.time() - t0):.3f}")
|
||||
|
||||
def checkcache(self, p):
|
||||
if shared.opts.sd_textencoder_cache_size == 0:
|
||||
return False
|
||||
if self.attention != shared.opts.prompt_attention:
|
||||
debug(f"Prompt change: parser={shared.opts.prompt_attention}")
|
||||
cache.clear()
|
||||
return False
|
||||
|
||||
def flatten(xss):
|
||||
return [x for xs in xss for x in xs]
|
||||
|
||||
# unpack EN data in case of TE LoRA
|
||||
en_data = p.extra_network_data
|
||||
en_data = [idx.items for item in en_data.values() for idx in item]
|
||||
effective_batch = 1 if self.allsame else self.batchsize
|
||||
key = str([self.prompts, self.negative_prompts, effective_batch, self.clip_skip, self.steps, en_data])
|
||||
item = cache.get(key)
|
||||
if not item:
|
||||
if not any(flatten(emb) for emb in [self.prompt_embeds,
|
||||
self.negative_prompt_embeds,
|
||||
self.positive_pooleds,
|
||||
self.negative_pooleds]):
|
||||
return False
|
||||
else:
|
||||
cache[key] = {'prompt_embeds': self.prompt_embeds,
|
||||
'negative_prompt_embeds': self.negative_prompt_embeds,
|
||||
'positive_pooleds': self.positive_pooleds,
|
||||
'negative_pooleds': self.negative_pooleds,
|
||||
}
|
||||
debug(f"Prompt cache: add={key}")
|
||||
while len(cache) > int(shared.opts.sd_textencoder_cache_size):
|
||||
cache.popitem(last=False)
|
||||
if item:
|
||||
self.__dict__.update(cache[key])
|
||||
cache.move_to_end(key)
|
||||
if self.allsame and len(self.prompt_embeds) < self.batchsize:
|
||||
self.prompt_embeds = [self.prompt_embeds[0]] * self.batchsize
|
||||
self.positive_pooleds = [self.positive_pooleds[0]] * self.batchsize
|
||||
self.negative_prompt_embeds = [self.negative_prompt_embeds[0]] * self.batchsize
|
||||
self.negative_pooleds = [self.negative_pooleds[0]] * self.batchsize
|
||||
debug(f"Prompt cache: get={key}")
|
||||
return True
|
||||
|
||||
def compare_prompts(self):
|
||||
same = (self.prompts == [self.prompts[0]] * len(self.prompts) and self.negative_prompts == [self.negative_prompts[0]] * len(self.negative_prompts))
|
||||
if same:
|
||||
self.prompts = [self.prompts[0]]
|
||||
self.negative_prompts = [self.negative_prompts[0]]
|
||||
return same
|
||||
|
||||
def prepare_schedule(self, prompt, negative_prompt):
|
||||
self.positive_schedule, scheduled = get_prompt_schedule(prompt, self.steps)
|
||||
self.negative_schedule, neg_scheduled = get_prompt_schedule(negative_prompt, self.steps)
|
||||
self.scheduled_prompt = scheduled or neg_scheduled
|
||||
debug(f"Prompt schedule: positive={self.positive_schedule} negative={self.negative_schedule} scheduled={scheduled}")
|
||||
|
||||
def scheduled_encode(self, pipe, batchidx):
|
||||
prompt_dict = {} # index cache
|
||||
for i in range(max(len(self.positive_schedule), len(self.negative_schedule))):
|
||||
positive_prompt = self.positive_schedule[i % len(self.positive_schedule)]
|
||||
negative_prompt = self.negative_schedule[i % len(self.negative_schedule)]
|
||||
# skip repeated scheduled subprompts
|
||||
idx = prompt_dict.get(positive_prompt+negative_prompt)
|
||||
if idx is not None:
|
||||
self.extend_embeds(batchidx, idx)
|
||||
continue
|
||||
self.encode(pipe, positive_prompt, negative_prompt, batchidx)
|
||||
prompt_dict[positive_prompt+negative_prompt] = i
|
||||
|
||||
def extend_embeds(self, batchidx, idx): # Extends scheduled prompt via index
|
||||
if len(self.prompt_embeds[batchidx]) > 0:
|
||||
self.prompt_embeds[batchidx].append(self.prompt_embeds[batchidx][idx])
|
||||
if len(self.negative_prompt_embeds[batchidx]) > 0:
|
||||
self.negative_prompt_embeds[batchidx].append(self.negative_prompt_embeds[batchidx][idx])
|
||||
if len(self.positive_pooleds[batchidx]) > 0:
|
||||
self.positive_pooleds[batchidx].append(self.positive_pooleds[batchidx][idx])
|
||||
if len(self.negative_pooleds[batchidx]) > 0:
|
||||
self.negative_pooleds[batchidx].append(self.negative_pooleds[batchidx][idx])
|
||||
|
||||
def encode(self, pipe, positive_prompt, negative_prompt, batchidx):
|
||||
self.attention = shared.opts.prompt_attention
|
||||
if self.attention == "xhinker" or 'Flux' in pipe.__class__.__name__:
|
||||
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip)
|
||||
else:
|
||||
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip)
|
||||
if prompt_embed is not None:
|
||||
self.prompt_embeds[batchidx].append(prompt_embed)
|
||||
if negative_embed is not None:
|
||||
self.negative_prompt_embeds[batchidx].append(negative_embed)
|
||||
if positive_pooled is not None:
|
||||
self.positive_pooleds[batchidx].append(positive_pooled)
|
||||
if negative_pooled is not None:
|
||||
self.negative_pooleds[batchidx].append(negative_pooled)
|
||||
|
||||
if debug_enabled:
|
||||
get_tokens(pipe, 'positive', positive_prompt)
|
||||
get_tokens(pipe, 'negative', negative_prompt)
|
||||
pipe = prepare_model()
|
||||
|
||||
def __call__(self, key, step=0):
|
||||
batch = getattr(self, key)
|
||||
res = []
|
||||
for i in range(self.batchsize):
|
||||
if len(batch[i]) == 0: # if asking for a null key, ie pooled on SD1.5
|
||||
return None
|
||||
try:
|
||||
res.append(batch[i][step])
|
||||
except IndexError:
|
||||
res.append(batch[i][0]) # if not scheduled, return default
|
||||
return torch.cat(res)
|
||||
|
||||
|
||||
def compel_hijack(self, token_ids: torch.Tensor, attention_mask: typing.Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
@@ -59,9 +235,9 @@ def insert_parser_highjack(pipename):
|
||||
debug("Load Standard Parser hijack")
|
||||
|
||||
|
||||
|
||||
insert_parser_highjack("Initialize")
|
||||
|
||||
|
||||
# from https://github.com/damian0815/compel/blob/main/src/compel/diffusers_textual_inversion_manager.py
|
||||
class DiffusersTextualInversionManager(BaseTextualInversionManager):
|
||||
def __init__(self, pipe, tokenizer):
|
||||
@@ -108,12 +284,6 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager):
|
||||
|
||||
def get_prompt_schedule(prompt, steps):
|
||||
t0 = time.time()
|
||||
if shared.native:
|
||||
# TODO prompt scheduling
|
||||
# prompt schedule returns array of prompts which would require that each prompt is fed to the model per-step
|
||||
# prompt scheduling should instead interpolate between each prompt in schedule
|
||||
# this temporarily disables prompt scheduling
|
||||
return [prompt], False
|
||||
temp = []
|
||||
schedule = prompt_parser.get_learned_conditioning_prompt_schedules([prompt], steps)[0]
|
||||
if all(x == schedule[0] for x in schedule):
|
||||
@@ -126,25 +296,25 @@ def get_prompt_schedule(prompt, steps):
|
||||
return temp, len(schedule) > 1
|
||||
|
||||
|
||||
def get_tokens(msg, prompt):
|
||||
def get_tokens(pipe, msg, prompt):
|
||||
global token_dict, token_type # pylint: disable=global-statement
|
||||
if not shared.native:
|
||||
return
|
||||
if shared.sd_loaded and hasattr(shared.sd_model, 'tokenizer') and shared.sd_model.tokenizer is not None:
|
||||
return 0
|
||||
if shared.sd_loaded and hasattr(pipe, 'tokenizer') and pipe.tokenizer is not None:
|
||||
if token_dict is None or token_type != shared.sd_model_type:
|
||||
token_type = shared.sd_model_type
|
||||
fn = shared.sd_model.tokenizer.name_or_path
|
||||
fn = pipe.tokenizer.name_or_path
|
||||
if fn.endswith('tokenizer'):
|
||||
fn = os.path.join(shared.sd_model.tokenizer.name_or_path, 'vocab.json')
|
||||
fn = os.path.join(pipe.tokenizer.name_or_path, 'vocab.json')
|
||||
else:
|
||||
fn = os.path.join(shared.sd_model.tokenizer.name_or_path, 'tokenizer', 'vocab.json')
|
||||
fn = os.path.join(pipe.tokenizer.name_or_path, 'tokenizer', 'vocab.json')
|
||||
token_dict = shared.readfile(fn, silent=True)
|
||||
for k, v in shared.sd_model.tokenizer.added_tokens_decoder.items():
|
||||
for k, v in pipe.tokenizer.added_tokens_decoder.items():
|
||||
token_dict[str(v)] = k
|
||||
shared.log.debug(f'Tokenizer: words={len(token_dict)} file="{fn}"')
|
||||
has_bos_token = shared.sd_model.tokenizer.bos_token_id is not None
|
||||
has_eos_token = shared.sd_model.tokenizer.eos_token_id is not None
|
||||
ids = shared.sd_model.tokenizer(prompt)
|
||||
has_bos_token = pipe.tokenizer.bos_token_id is not None
|
||||
has_eos_token = pipe.tokenizer.eos_token_id is not None
|
||||
ids = pipe.tokenizer(prompt)
|
||||
ids = getattr(ids, 'input_ids', [])
|
||||
tokens = []
|
||||
for i in ids:
|
||||
@@ -155,118 +325,7 @@ def get_tokens(msg, prompt):
|
||||
tokens.append(f'UNK_{i}')
|
||||
token_count = len(ids) - int(has_bos_token) - int(has_eos_token)
|
||||
debug(f'Prompt tokenizer: type={msg} tokens={token_count} {tokens}')
|
||||
|
||||
|
||||
def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, clip_skip: typing.Optional[int] = None):
|
||||
params_match = prompts == cache.get('prompts', None) and negative_prompts == cache.get('negative_prompts', None) and clip_skip == cache.get('clip_skip', None) and steps == cache.get('steps', None)
|
||||
if (
|
||||
'StableDiffusion' not in pipe.__class__.__name__ and
|
||||
'DemoFusion' not in pipe.__class__.__name__ and
|
||||
'StableCascade' not in pipe.__class__.__name__ and
|
||||
'Flux' not in pipe.__class__.__name__
|
||||
):
|
||||
shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}")
|
||||
return
|
||||
elif shared.opts.sd_textencoder_cache and cache.get('model_type', None) == shared.sd_model_type and params_match:
|
||||
p.prompt_embeds = cache.get('prompt_embeds', None)
|
||||
p.positive_pooleds = cache.get('positive_pooleds', None)
|
||||
p.negative_embeds = cache.get('negative_embeds', None)
|
||||
p.negative_pooleds = cache.get('negative_pooleds', None)
|
||||
p.scheduled_prompt = cache.get('scheduled_prompt', None)
|
||||
debug("Prompt encode: cached")
|
||||
return
|
||||
else:
|
||||
t0 = time.time()
|
||||
if shared.opts.diffusers_offload_mode == "balanced":
|
||||
pipe = sd_models.apply_balanced_offload(pipe)
|
||||
elif hasattr(pipe, "maybe_free_model_hooks"):
|
||||
pipe.maybe_free_model_hooks()
|
||||
devices.torch_gc()
|
||||
|
||||
prompt_embeds, positive_pooleds, negative_embeds, negative_pooleds = [], [], [], []
|
||||
last_prompt, last_negative = None, None
|
||||
for prompt, negative in zip(prompts, negative_prompts):
|
||||
prompt_embed, positive_pooled, negative_embed, negative_pooled = None, None, None, None
|
||||
if last_prompt == prompt and last_negative == negative:
|
||||
prompt_embeds.append(prompt_embeds[-1])
|
||||
negative_embeds.append(negative_embeds[-1])
|
||||
if len(positive_pooleds) > 0:
|
||||
positive_pooleds.append(positive_pooleds[-1])
|
||||
if len(negative_pooleds) > 0:
|
||||
negative_pooleds.append(negative_pooleds[-1])
|
||||
continue
|
||||
positive_schedule, scheduled = get_prompt_schedule(prompt, steps)
|
||||
negative_schedule, neg_scheduled = get_prompt_schedule(negative, steps)
|
||||
p.scheduled_prompt = scheduled or neg_scheduled
|
||||
p.prompt_embeds = []
|
||||
p.positive_pooleds = []
|
||||
p.negative_embeds = []
|
||||
p.negative_pooleds = []
|
||||
|
||||
for i in range(max(len(positive_schedule), len(negative_schedule))):
|
||||
positive_prompt = positive_schedule[i % len(positive_schedule)]
|
||||
negative_prompt = negative_schedule[i % len(negative_schedule)]
|
||||
if shared.opts.prompt_attention == "xhinker parser" or 'Flux' in pipe.__class__.__name__:
|
||||
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip)
|
||||
else:
|
||||
prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip)
|
||||
if prompt_embed is not None:
|
||||
prompt_embeds.append(prompt_embed)
|
||||
if negative_embed is not None:
|
||||
negative_embeds.append(negative_embed)
|
||||
if positive_pooled is not None:
|
||||
positive_pooleds.append(positive_pooled)
|
||||
if negative_pooled is not None:
|
||||
negative_pooleds.append(negative_pooled)
|
||||
last_prompt, last_negative = prompt, negative
|
||||
# TODO prompt scheduling
|
||||
# interpolation should happen here and then we can re-enable prompt scheduling
|
||||
# ive tried simple torch.mean and its not good-enough
|
||||
|
||||
def fix_length(embeds):
|
||||
max_len = max([e.shape[1] for e in embeds if e is not None])
|
||||
for i, e in enumerate(embeds):
|
||||
if e is not None and e.shape[1] < max_len:
|
||||
expanded = torch.zeros((e.shape[0], max_len, e.shape[2]), device=e.device, dtype=e.dtype)
|
||||
expanded[:, :e.shape[1], :] = e
|
||||
embeds[i] = expanded
|
||||
return torch.cat(embeds, dim=0).to(devices.device, dtype=devices.dtype)
|
||||
|
||||
if len(prompt_embeds) > 0:
|
||||
p.prompt_embeds.append(fix_length(prompt_embeds))
|
||||
if len(negative_embeds) > 0:
|
||||
p.negative_embeds.append(fix_length(negative_embeds))
|
||||
if len(positive_pooleds) > 0:
|
||||
p.positive_pooleds.append(fix_length(positive_pooleds))
|
||||
if len(negative_pooleds) > 0:
|
||||
p.negative_pooleds.append(fix_length(negative_pooleds))
|
||||
|
||||
if shared.opts.sd_textencoder_cache and p.batch_size == 1:
|
||||
cache.update({
|
||||
'prompt_embeds': p.prompt_embeds,
|
||||
'negative_embeds': p.negative_embeds,
|
||||
'positive_pooleds': p.positive_pooleds,
|
||||
'negative_pooleds': p.negative_pooleds,
|
||||
'scheduled_prompt': p.scheduled_prompt,
|
||||
'prompts': prompts,
|
||||
'negative_prompts': negative_prompts,
|
||||
'clip_skip': clip_skip,
|
||||
'steps': steps,
|
||||
'model_type': shared.sd_model_type
|
||||
})
|
||||
else:
|
||||
cache.clear()
|
||||
if debug_enabled:
|
||||
get_tokens('positive', prompts[0])
|
||||
get_tokens('negative', negative_prompts[0])
|
||||
if shared.opts.diffusers_offload_mode == "balanced":
|
||||
pipe = sd_models.apply_balanced_offload(pipe)
|
||||
elif hasattr(pipe, "maybe_free_model_hooks"):
|
||||
# text encoder will stay in the vram and cause oom, send everything back to cpu before continuing
|
||||
pipe.maybe_free_model_hooks()
|
||||
debug(f"Prompt encode: time={(time.time() - t0):.3f}")
|
||||
devices.torch_gc()
|
||||
return
|
||||
return token_count
|
||||
|
||||
|
||||
def normalize_prompt(pairs: list):
|
||||
@@ -286,14 +345,20 @@ def normalize_prompt(pairs: list):
|
||||
return pairs
|
||||
|
||||
|
||||
def get_prompts_with_weights(prompt: str):
|
||||
def get_prompts_with_weights(pipe, prompt: str):
|
||||
t0 = time.time()
|
||||
manager = DiffusersTextualInversionManager(shared.sd_model, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2)
|
||||
prompt = manager.maybe_convert_prompt(prompt, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2)
|
||||
manager = DiffusersTextualInversionManager(pipe, pipe.tokenizer or pipe.tokenizer_2)
|
||||
prompt = manager.maybe_convert_prompt(prompt, pipe.tokenizer or pipe.tokenizer_2)
|
||||
texts_and_weights = prompt_parser.parse_prompt_attention(prompt)
|
||||
if shared.opts.prompt_mean_norm:
|
||||
texts_and_weights = normalize_prompt(texts_and_weights)
|
||||
texts, text_weights = zip(*texts_and_weights)
|
||||
if debug_enabled:
|
||||
all_tokens = 0
|
||||
for text in texts:
|
||||
tokens = get_tokens(pipe, 'section', text)
|
||||
all_tokens += tokens
|
||||
debug(f'Prompt tokenizer: parser={shared.opts.prompt_attention} tokens={all_tokens}')
|
||||
debug(f'Prompt: weights={texts_and_weights} time={(time.time() - t0):.3f}')
|
||||
return texts, text_weights
|
||||
|
||||
@@ -354,7 +419,8 @@ def pad_to_same_length(pipe, embeds, empty_embedding_providers=None):
|
||||
embeds[i] = embed
|
||||
return embeds
|
||||
|
||||
def split_prompts(prompt, SD3 = False):
|
||||
|
||||
def split_prompts(pipe, prompt, SD3 = False):
|
||||
if prompt.find("TE2:") != -1:
|
||||
prompt, prompt2 = prompt.split("TE2:")
|
||||
else:
|
||||
@@ -372,7 +438,7 @@ def split_prompts(prompt, SD3 = False):
|
||||
prompt3 = " " if prompt3.strip() == "" else prompt3.strip()
|
||||
|
||||
if SD3 and prompt3 != " ":
|
||||
ps, _ws = get_prompts_with_weights(prompt3)
|
||||
ps, _ws = get_prompts_with_weights(pipe, prompt3)
|
||||
prompt3 = " ".join(ps)
|
||||
return prompt, prompt2, prompt3
|
||||
|
||||
@@ -380,15 +446,15 @@ def split_prompts(prompt, SD3 = False):
|
||||
def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None):
|
||||
device = devices.device
|
||||
SD3 = hasattr(pipe, 'text_encoder_3')
|
||||
prompt, prompt_2, prompt_3 = split_prompts(prompt, SD3)
|
||||
neg_prompt, neg_prompt_2, neg_prompt_3 = split_prompts(neg_prompt, SD3)
|
||||
prompt, prompt_2, prompt_3 = split_prompts(pipe, prompt, SD3)
|
||||
neg_prompt, neg_prompt_2, neg_prompt_3 = split_prompts(pipe, neg_prompt, SD3)
|
||||
|
||||
if prompt != prompt_2:
|
||||
ps = [get_prompts_with_weights(p) for p in [prompt, prompt_2]]
|
||||
ns = [get_prompts_with_weights(p) for p in [neg_prompt, neg_prompt_2]]
|
||||
ps = [get_prompts_with_weights(pipe, p) for p in [prompt, prompt_2]]
|
||||
ns = [get_prompts_with_weights(pipe, p) for p in [neg_prompt, neg_prompt_2]]
|
||||
else:
|
||||
ps = 2 * [get_prompts_with_weights(prompt)]
|
||||
ns = 2 * [get_prompts_with_weights(neg_prompt)]
|
||||
ps = 2 * [get_prompts_with_weights(pipe, prompt)]
|
||||
ns = 2 * [get_prompts_with_weights(pipe, neg_prompt)]
|
||||
|
||||
positives, positive_weights = zip(*ps)
|
||||
negatives, negative_weights = zip(*ns)
|
||||
@@ -434,7 +500,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
|
||||
# negative prompt has no keywords
|
||||
embed, ntokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[negatives[i]], fragment_weights_batch=[negative_weights[i]], device=device, should_return_tokens=True)
|
||||
negative_prompt_embeds.append(embed)
|
||||
debug(f'Prompt: unpadded shape={prompt_embeds[0].shape} TE{i+1} ptokens={torch.count_nonzero(ptokens)} ntokens={torch.count_nonzero(ntokens)} time={(time.time() - t0):.3f}')
|
||||
debug(f'Prompt: unpadded={prompt_embeds[0].shape} TE{i+1} ptokens={torch.count_nonzero(ptokens)} ntokens={torch.count_nonzero(ntokens)} time={(time.time() - t0):.3f}')
|
||||
if SD3:
|
||||
t0 = time.time()
|
||||
pooled_prompt_embeds.append(embedding_providers[0].get_pooled_embeddings(texts=positives[0] if len(positives[0]) == 1 else [" ".join(positives[0])], device=device))
|
||||
@@ -443,7 +509,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
|
||||
negative_pooled_prompt_embeds.append(embedding_providers[1].get_pooled_embeddings(texts=negatives[-1] if len(negatives[-1]) == 1 else [" ".join(negatives[-1])], device=device))
|
||||
pooled_prompt_embeds = torch.cat(pooled_prompt_embeds, dim=-1)
|
||||
negative_pooled_prompt_embeds = torch.cat(negative_pooled_prompt_embeds, dim=-1)
|
||||
debug(f'Prompt: pooled shape={pooled_prompt_embeds[0].shape} time={(time.time() - t0):.3f}')
|
||||
debug(f'Prompt: pooled={pooled_prompt_embeds[0].shape} time={(time.time() - t0):.3f}')
|
||||
elif prompt_embeds[-1].shape[-1] > 768:
|
||||
t0 = time.time()
|
||||
if shared.opts.diffusers_pooled == "weighted":
|
||||
@@ -503,8 +569,8 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c
|
||||
|
||||
def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None):
|
||||
is_sd3 = hasattr(pipe, 'text_encoder_3')
|
||||
prompt, prompt_2, _prompt_3 = split_prompts(prompt, is_sd3)
|
||||
neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(neg_prompt, is_sd3)
|
||||
prompt, prompt_2, _prompt_3 = split_prompts(pipe, prompt, is_sd3)
|
||||
neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(pipe, neg_prompt, is_sd3)
|
||||
try:
|
||||
prompt = pipe.maybe_convert_prompt(prompt, pipe.tokenizer)
|
||||
neg_prompt = pipe.maybe_convert_prompt(neg_prompt, pipe.tokenizer)
|
||||
|
||||
@@ -1305,7 +1305,7 @@ def get_weighted_text_embeddings_sd3(
|
||||
# ---------------------- get neg t5 embeddings -------------------------
|
||||
neg_prompt_tokens_3 = torch.tensor([neg_prompt_tokens_3], dtype=torch.long)
|
||||
|
||||
t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.pipe.text_encoder_3.device))[0].squeeze(0)
|
||||
t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.text_encoder_3.device))[0].squeeze(0)
|
||||
t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(device=pipe.text_encoder_3.device)
|
||||
|
||||
# add weight to neg t5 embeddings
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
"""
|
||||
Credit and original implementation: <https://github.com/ToTheBeginning/PuLID>
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
sys.path.append(os.path.dirname(__file__))
|
||||
from pulid_sdxl import StableDiffusionXLPuLIDPipeline, StableDiffusionXLPuLIDPipelineImage, StableDiffusionXLPuLIDPipelineInpaint
|
||||
from pulid_utils import resize_numpy_image_long as resize
|
||||
import attention_processor as attention
|
||||
import pulid_sampling as sampling
|
||||
@@ -0,0 +1,418 @@
|
||||
# modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
NUM_ZERO = 0
|
||||
ORTHO = False
|
||||
ORTHO_v2 = False
|
||||
|
||||
|
||||
class AttnProcessor(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
id_embedding=None,
|
||||
id_scale=1.0,
|
||||
):
|
||||
residual = hidden_states
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class IDAttnProcessor(nn.Module):
|
||||
r"""
|
||||
Attention processor for ID-Adapater.
|
||||
Args:
|
||||
hidden_size (`int`):
|
||||
The hidden size of the attention layer.
|
||||
cross_attention_dim (`int`):
|
||||
The number of channels in the `encoder_hidden_states`.
|
||||
scale (`float`, defaults to 1.0):
|
||||
the weight scale of image prompt.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, cross_attention_dim=None):
|
||||
super().__init__()
|
||||
self.id_to_k = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
self.id_to_v = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
id_embedding=None,
|
||||
id_scale=1.0,
|
||||
):
|
||||
residual = hidden_states
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query = attn.head_to_batch_dim(query)
|
||||
key = attn.head_to_batch_dim(key)
|
||||
value = attn.head_to_batch_dim(value)
|
||||
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# for id-adapter
|
||||
if id_embedding is not None:
|
||||
if NUM_ZERO == 0:
|
||||
id_key = self.id_to_k(id_embedding)
|
||||
id_value = self.id_to_v(id_embedding)
|
||||
else:
|
||||
zero_tensor = torch.zeros(
|
||||
(id_embedding.size(0), NUM_ZERO, id_embedding.size(-1)),
|
||||
dtype=id_embedding.dtype,
|
||||
device=id_embedding.device,
|
||||
)
|
||||
id_key = self.id_to_k(torch.cat((id_embedding, zero_tensor), dim=1))
|
||||
id_value = self.id_to_v(torch.cat((id_embedding, zero_tensor), dim=1))
|
||||
|
||||
id_key = attn.head_to_batch_dim(id_key).to(query.dtype)
|
||||
id_value = attn.head_to_batch_dim(id_value).to(query.dtype)
|
||||
|
||||
id_attention_probs = attn.get_attention_scores(query, id_key, None)
|
||||
id_hidden_states = torch.bmm(id_attention_probs, id_value)
|
||||
id_hidden_states = attn.batch_to_head_dim(id_hidden_states)
|
||||
|
||||
if not ORTHO:
|
||||
hidden_states = hidden_states + id_scale * id_hidden_states
|
||||
else:
|
||||
projection = (
|
||||
torch.sum((hidden_states * id_hidden_states), dim=-2, keepdim=True)
|
||||
/ torch.sum((hidden_states * hidden_states), dim=-2, keepdim=True)
|
||||
* hidden_states
|
||||
)
|
||||
orthogonal = id_hidden_states - projection
|
||||
hidden_states = hidden_states + id_scale * orthogonal
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AttnProcessor2_0(nn.Module):
|
||||
r"""
|
||||
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
id_embedding=None,
|
||||
id_scale=1.0,
|
||||
):
|
||||
residual = hidden_states
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class IDAttnProcessor2_0(torch.nn.Module):
|
||||
r"""
|
||||
Attention processor for ID-Adapater for PyTorch 2.0.
|
||||
Args:
|
||||
hidden_size (`int`):
|
||||
The hidden size of the attention layer.
|
||||
cross_attention_dim (`int`):
|
||||
The number of channels in the `encoder_hidden_states`.
|
||||
"""
|
||||
|
||||
def __init__(self, hidden_size, cross_attention_dim=None):
|
||||
super().__init__()
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
self.id_to_k = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
self.id_to_v = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
id_embedding=None,
|
||||
id_scale=1.0,
|
||||
):
|
||||
residual = hidden_states
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
# scaled_dot_product_attention expects attention_mask shape to be
|
||||
# (batch, heads, source_length, target_length)
|
||||
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# for id embedding
|
||||
if id_embedding is not None:
|
||||
if NUM_ZERO == 0:
|
||||
id_key = self.id_to_k(id_embedding).to(query.dtype)
|
||||
id_value = self.id_to_v(id_embedding).to(query.dtype)
|
||||
else:
|
||||
zero_tensor = torch.zeros(
|
||||
(id_embedding.size(0), NUM_ZERO, id_embedding.size(-1)),
|
||||
dtype=id_embedding.dtype,
|
||||
device=id_embedding.device,
|
||||
)
|
||||
id_cat = torch.cat((id_embedding, zero_tensor), dim=1)
|
||||
id_key = self.id_to_k(id_cat).to(query.dtype)
|
||||
id_value = self.id_to_v(id_cat).to(query.dtype)
|
||||
|
||||
id_key = id_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
id_value = id_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
# the output of sdp = (batch, num_heads, seq_len, head_dim)
|
||||
id_hidden_states = F.scaled_dot_product_attention(query, id_key, id_value, attn_mask=None, dropout_p=0.0, is_causal=False)
|
||||
id_hidden_states = id_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
id_hidden_states = id_hidden_states.to(query.dtype)
|
||||
|
||||
if not ORTHO and not ORTHO_v2:
|
||||
hidden_states = hidden_states + id_scale * id_hidden_states
|
||||
elif ORTHO_v2:
|
||||
orig_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
id_hidden_states = id_hidden_states.to(torch.float32)
|
||||
attn_map = query @ id_key.transpose(-2, -1)
|
||||
attn_mean = attn_map.softmax(dim=-1).mean(dim=1)
|
||||
attn_mean = attn_mean[:, :, :5].sum(dim=-1, keepdim=True)
|
||||
projection = (
|
||||
torch.sum((hidden_states * id_hidden_states), dim=-2, keepdim=True)
|
||||
/ torch.sum((hidden_states * hidden_states), dim=-2, keepdim=True)
|
||||
* hidden_states
|
||||
)
|
||||
orthogonal = id_hidden_states + (attn_mean - 1) * projection
|
||||
hidden_states = hidden_states + id_scale * orthogonal
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
else:
|
||||
orig_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
id_hidden_states = id_hidden_states.to(torch.float32)
|
||||
projection = (
|
||||
torch.sum((hidden_states * id_hidden_states), dim=-2, keepdim=True)
|
||||
/ torch.sum((hidden_states * hidden_states), dim=-2, keepdim=True)
|
||||
* hidden_states
|
||||
)
|
||||
orthogonal = id_hidden_states - projection
|
||||
hidden_states = hidden_states + id_scale * orthogonal
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,250 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
# FFN
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, _width = x.shape
|
||||
# (bs, length, width) --> (bs, length, n_heads, dim_per_head)
|
||||
x = x.view(bs, length, heads, -1)
|
||||
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
|
||||
x = x.transpose(1, 2)
|
||||
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
|
||||
class PerceiverAttentionCA(nn.Module):
|
||||
def __init__(self, *, dim=3072, dim_head=128, heads=16, kv_dim=2048):
|
||||
super().__init__()
|
||||
self.scale = dim_head ** -0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
self.norm1 = nn.LayerNorm(dim if kv_dim is None else kv_dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim if kv_dim is None else kv_dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
b, seq_len, _ = latents.shape
|
||||
q = self.to_q(latents)
|
||||
k, v = self.to_kv(x).chunk(2, dim=-1)
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, seq_len, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8, kv_dim=None):
|
||||
super().__init__()
|
||||
self.scale = dim_head ** -0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
self.norm1 = nn.LayerNorm(dim if kv_dim is None else kv_dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim if kv_dim is None else kv_dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
b, seq_len, _ = latents.shape
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, seq_len, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class IDFormer(nn.Module):
|
||||
"""
|
||||
- perceiver resampler like arch (compared with previous MLP-like arch)
|
||||
- we concat id embedding (generated by arcface) and query tokens as latents
|
||||
- latents will attend each other and interact with vit features through cross-attention
|
||||
- vit features are multi-scaled and inserted into IDFormer in order, currently, each scale corresponds to two
|
||||
IDFormer layers
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
dim=1024,
|
||||
depth=10,
|
||||
dim_head=64,
|
||||
heads=16,
|
||||
num_id_token=5,
|
||||
num_queries=32,
|
||||
output_dim=2048,
|
||||
ff_mult=4,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.num_id_token = num_id_token
|
||||
self.dim = dim
|
||||
self.num_queries = num_queries
|
||||
assert depth % 5 == 0
|
||||
self.depth = depth // 5
|
||||
scale = dim ** -0.5
|
||||
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) * scale)
|
||||
self.proj_out = nn.Parameter(scale * torch.randn(dim, output_dim))
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
for i in range(5):
|
||||
setattr(
|
||||
self,
|
||||
f'mapping_{i}',
|
||||
nn.Sequential(
|
||||
nn.Linear(1024, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, dim),
|
||||
),
|
||||
)
|
||||
|
||||
self.id_embedding_mapping = nn.Sequential(
|
||||
nn.Linear(1280, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, dim * num_id_token),
|
||||
)
|
||||
|
||||
def forward(self, x, y):
|
||||
latents = self.latents.repeat(x.size(0), 1, 1)
|
||||
num_duotu = x.shape[1] if x.ndim == 3 else 1
|
||||
x = self.id_embedding_mapping(x)
|
||||
x = x.reshape(-1, self.num_id_token * num_duotu, self.dim)
|
||||
latents = torch.cat((latents, x), dim=1)
|
||||
for i in range(5):
|
||||
vit_feature = getattr(self, f'mapping_{i}')(y[i])
|
||||
ctx_feature = torch.cat((x, vit_feature), dim=1)
|
||||
for attn, ff in self.layers[i * self.depth: (i + 1) * self.depth]:
|
||||
latents = attn(ctx_feature, latents) + latents
|
||||
latents = ff(latents) + latents
|
||||
latents = latents[:, :self.num_queries]
|
||||
latents = latents @ self.proj_out
|
||||
return latents
|
||||
|
||||
|
||||
class IDEncoder(nn.Module):
|
||||
def __init__(self, width=1280, context_dim=2048, num_token=5):
|
||||
super().__init__()
|
||||
self.num_token = num_token
|
||||
self.context_dim = context_dim
|
||||
h1 = min((context_dim * num_token) // 4, 1024)
|
||||
h2 = min((context_dim * num_token) // 2, 1024)
|
||||
self.body = nn.Sequential(
|
||||
nn.Linear(width, h1),
|
||||
nn.LayerNorm(h1),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(h1, h2),
|
||||
nn.LayerNorm(h2),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(h2, context_dim * num_token),
|
||||
)
|
||||
|
||||
for i in range(5):
|
||||
setattr(
|
||||
self,
|
||||
f'mapping_{i}',
|
||||
nn.Sequential(
|
||||
nn.Linear(1024, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, context_dim),
|
||||
),
|
||||
)
|
||||
setattr(
|
||||
self,
|
||||
f'mapping_patch_{i}',
|
||||
nn.Sequential(
|
||||
nn.Linear(1024, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, 1024),
|
||||
nn.LayerNorm(1024),
|
||||
nn.LeakyReLU(),
|
||||
nn.Linear(1024, context_dim),
|
||||
),
|
||||
)
|
||||
|
||||
def forward(self, x, y):
|
||||
# x shape [N, C]
|
||||
x = self.body(x)
|
||||
x = x.reshape(-1, self.num_token, self.context_dim)
|
||||
|
||||
hidden_states = ()
|
||||
for i, emb in enumerate(y):
|
||||
hidden_state = getattr(self, f'mapping_{i}')(emb[:, :1]) + getattr(self, f'mapping_patch_{i}')(
|
||||
emb[:, 1:]
|
||||
).mean(dim=1, keepdim=True)
|
||||
hidden_states += (hidden_state,)
|
||||
hidden_states = torch.cat(hidden_states, dim=1)
|
||||
|
||||
return torch.cat([x, hidden_states], dim=1)
|
||||
@@ -0,0 +1,11 @@
|
||||
from .constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD
|
||||
from .factory import create_model, create_model_and_transforms, create_model_from_pretrained, get_tokenizer, create_transforms
|
||||
from .factory import list_models, add_model_config, get_model_config, load_checkpoint
|
||||
from .loss import ClipLoss
|
||||
from .model import CLIP, CustomCLIP, CLIPTextCfg, CLIPVisionCfg,\
|
||||
convert_weights_to_lp, convert_weights_to_fp16, trace_model, get_cast_dtype
|
||||
from .openai import load_openai_model, list_openai_models
|
||||
from .pretrained import list_pretrained, list_pretrained_models_by_tag, list_pretrained_tags_by_model,\
|
||||
get_pretrained_url, download_pretrained_from_url, is_pretrained_cfg, get_pretrained_cfg, download_pretrained
|
||||
from .tokenizer import SimpleTokenizer, tokenize
|
||||
from .transform import image_transform
|
||||
Binary file not shown.
@@ -0,0 +1,2 @@
|
||||
OPENAI_DATASET_MEAN = (0.48145466, 0.4578275, 0.40821073)
|
||||
OPENAI_DATASET_STD = (0.26862954, 0.26130258, 0.27577711)
|
||||
@@ -0,0 +1,548 @@
|
||||
# --------------------------------------------------------
|
||||
# Adapted from https://github.com/microsoft/unilm/tree/master/beit
|
||||
# --------------------------------------------------------
|
||||
import math
|
||||
import os
|
||||
from functools import partial
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
try:
|
||||
from timm.models.layers import drop_path, to_2tuple, trunc_normal_
|
||||
except:
|
||||
from timm.layers import drop_path, to_2tuple, trunc_normal_
|
||||
|
||||
from .transformer import PatchDropout
|
||||
from .rope import VisionRotaryEmbedding, VisionRotaryEmbeddingFast
|
||||
|
||||
if os.getenv('ENV_TYPE') == 'deepspeed':
|
||||
try:
|
||||
from deepspeed.runtime.activation_checkpointing.checkpointing import checkpoint
|
||||
except:
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
else:
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops as xops
|
||||
XFORMERS_IS_AVAILBLE = True
|
||||
except:
|
||||
XFORMERS_IS_AVAILBLE = False
|
||||
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
"""
|
||||
def __init__(self, drop_prob=None):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training)
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
return 'p={}'.format(self.drop_prob)
|
||||
|
||||
|
||||
class Mlp(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
act_layer=nn.GELU,
|
||||
norm_layer=nn.LayerNorm,
|
||||
drop=0.,
|
||||
subln=False,
|
||||
|
||||
):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = act_layer()
|
||||
|
||||
self.ffn_ln = norm_layer(hidden_features) if subln else nn.Identity()
|
||||
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
# x = self.drop(x)
|
||||
# commit this for the orignal BERT implement
|
||||
x = self.ffn_ln(x)
|
||||
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
class SwiGLU(nn.Module):
|
||||
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.SiLU, drop=0.,
|
||||
norm_layer=nn.LayerNorm, subln=False):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
|
||||
self.w1 = nn.Linear(in_features, hidden_features)
|
||||
self.w2 = nn.Linear(in_features, hidden_features)
|
||||
|
||||
self.act = act_layer()
|
||||
self.ffn_ln = norm_layer(hidden_features) if subln else nn.Identity()
|
||||
self.w3 = nn.Linear(hidden_features, out_features)
|
||||
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.w1(x)
|
||||
x2 = self.w2(x)
|
||||
hidden = self.act(x1) * x2
|
||||
x = self.ffn_ln(hidden)
|
||||
x = self.w3(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0.,
|
||||
proj_drop=0., window_size=None, attn_head_dim=None, xattn=False, rope=None, subln=False, norm_layer=nn.LayerNorm):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
if attn_head_dim is not None:
|
||||
head_dim = attn_head_dim
|
||||
all_head_dim = head_dim * self.num_heads
|
||||
self.scale = qk_scale or head_dim ** -0.5
|
||||
|
||||
self.subln = subln
|
||||
if self.subln:
|
||||
self.q_proj = nn.Linear(dim, all_head_dim, bias=False)
|
||||
self.k_proj = nn.Linear(dim, all_head_dim, bias=False)
|
||||
self.v_proj = nn.Linear(dim, all_head_dim, bias=False)
|
||||
else:
|
||||
self.qkv = nn.Linear(dim, all_head_dim * 3, bias=False)
|
||||
|
||||
if qkv_bias:
|
||||
self.q_bias = nn.Parameter(torch.zeros(all_head_dim))
|
||||
self.v_bias = nn.Parameter(torch.zeros(all_head_dim))
|
||||
else:
|
||||
self.q_bias = None
|
||||
self.v_bias = None
|
||||
|
||||
if window_size:
|
||||
self.window_size = window_size
|
||||
self.num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) + 3
|
||||
self.relative_position_bias_table = nn.Parameter(
|
||||
torch.zeros(self.num_relative_distance, num_heads)) # 2*Wh-1 * 2*Ww-1, nH
|
||||
# cls to token & token 2 cls & cls to cls
|
||||
|
||||
# get pair-wise relative position index for each token inside the window
|
||||
coords_h = torch.arange(window_size[0])
|
||||
coords_w = torch.arange(window_size[1])
|
||||
coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww
|
||||
coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
|
||||
relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
|
||||
relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
|
||||
relative_coords[:, :, 0] += window_size[0] - 1 # shift to start from 0
|
||||
relative_coords[:, :, 1] += window_size[1] - 1
|
||||
relative_coords[:, :, 0] *= 2 * window_size[1] - 1
|
||||
relative_position_index = \
|
||||
torch.zeros(size=(window_size[0] * window_size[1] + 1, ) * 2, dtype=relative_coords.dtype)
|
||||
relative_position_index[1:, 1:] = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
|
||||
relative_position_index[0, 0:] = self.num_relative_distance - 3
|
||||
relative_position_index[0:, 0] = self.num_relative_distance - 2
|
||||
relative_position_index[0, 0] = self.num_relative_distance - 1
|
||||
|
||||
self.register_buffer("relative_position_index", relative_position_index)
|
||||
else:
|
||||
self.window_size = None
|
||||
self.relative_position_bias_table = None
|
||||
self.relative_position_index = None
|
||||
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.inner_attn_ln = norm_layer(all_head_dim) if subln else nn.Identity()
|
||||
# self.proj = nn.Linear(all_head_dim, all_head_dim)
|
||||
self.proj = nn.Linear(all_head_dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
self.xattn = xattn
|
||||
self.xattn_drop = attn_drop
|
||||
|
||||
self.rope = rope
|
||||
|
||||
def forward(self, x, rel_pos_bias=None, attn_mask=None):
|
||||
B, N, C = x.shape
|
||||
if self.subln:
|
||||
q = F.linear(input=x, weight=self.q_proj.weight, bias=self.q_bias)
|
||||
k = F.linear(input=x, weight=self.k_proj.weight, bias=None)
|
||||
v = F.linear(input=x, weight=self.v_proj.weight, bias=self.v_bias)
|
||||
|
||||
q = q.reshape(B, N, self.num_heads, -1).permute(0, 2, 1, 3) # B, num_heads, N, C
|
||||
k = k.reshape(B, N, self.num_heads, -1).permute(0, 2, 1, 3)
|
||||
v = v.reshape(B, N, self.num_heads, -1).permute(0, 2, 1, 3)
|
||||
else:
|
||||
|
||||
qkv_bias = None
|
||||
if self.q_bias is not None:
|
||||
qkv_bias = torch.cat((self.q_bias, torch.zeros_like(self.v_bias, requires_grad=False), self.v_bias))
|
||||
|
||||
qkv = F.linear(input=x, weight=self.qkv.weight, bias=qkv_bias)
|
||||
qkv = qkv.reshape(B, N, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4) # 3, B, num_heads, N, C
|
||||
q, k, v = qkv[0], qkv[1], qkv[2]
|
||||
|
||||
if self.rope:
|
||||
# slightly fast impl
|
||||
q_t = q[:, :, 1:, :]
|
||||
ro_q_t = self.rope(q_t)
|
||||
q = torch.cat((q[:, :, :1, :], ro_q_t), -2).type_as(v)
|
||||
|
||||
k_t = k[:, :, 1:, :]
|
||||
ro_k_t = self.rope(k_t)
|
||||
k = torch.cat((k[:, :, :1, :], ro_k_t), -2).type_as(v)
|
||||
|
||||
if self.xattn:
|
||||
q = q.permute(0, 2, 1, 3) # B, num_heads, N, C -> B, N, num_heads, C
|
||||
k = k.permute(0, 2, 1, 3)
|
||||
v = v.permute(0, 2, 1, 3)
|
||||
|
||||
x = xops.memory_efficient_attention(
|
||||
q, k, v,
|
||||
p=self.xattn_drop,
|
||||
scale=self.scale,
|
||||
)
|
||||
x = x.reshape(B, N, -1)
|
||||
x = self.inner_attn_ln(x)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
else:
|
||||
q = q * self.scale
|
||||
attn = (q @ k.transpose(-2, -1))
|
||||
|
||||
if self.relative_position_bias_table is not None:
|
||||
relative_position_bias = \
|
||||
self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
|
||||
self.window_size[0] * self.window_size[1] + 1,
|
||||
self.window_size[0] * self.window_size[1] + 1, -1) # Wh*Ww,Wh*Ww,nH
|
||||
relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
|
||||
attn = attn + relative_position_bias.unsqueeze(0).type_as(attn)
|
||||
|
||||
if rel_pos_bias is not None:
|
||||
attn = attn + rel_pos_bias.type_as(attn)
|
||||
|
||||
if attn_mask is not None:
|
||||
attn_mask = attn_mask.bool()
|
||||
attn = attn.masked_fill(~attn_mask[:, None, None, :], float("-inf"))
|
||||
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn @ v).transpose(1, 2).reshape(B, N, -1)
|
||||
x = self.inner_attn_ln(x)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
|
||||
def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
|
||||
drop_path=0., init_values=None, act_layer=nn.GELU, norm_layer=nn.LayerNorm,
|
||||
window_size=None, attn_head_dim=None, xattn=False, rope=None, postnorm=False,
|
||||
subln=False, naiveswiglu=False):
|
||||
super().__init__()
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = Attention(
|
||||
dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
attn_drop=attn_drop, proj_drop=drop, window_size=window_size, attn_head_dim=attn_head_dim,
|
||||
xattn=xattn, rope=rope, subln=subln, norm_layer=norm_layer)
|
||||
# NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
|
||||
if naiveswiglu:
|
||||
self.mlp = SwiGLU(
|
||||
in_features=dim,
|
||||
hidden_features=mlp_hidden_dim,
|
||||
subln=subln,
|
||||
norm_layer=norm_layer,
|
||||
)
|
||||
else:
|
||||
self.mlp = Mlp(
|
||||
in_features=dim,
|
||||
hidden_features=mlp_hidden_dim,
|
||||
act_layer=act_layer,
|
||||
subln=subln,
|
||||
drop=drop
|
||||
)
|
||||
|
||||
if init_values is not None and init_values > 0:
|
||||
self.gamma_1 = nn.Parameter(init_values * torch.ones((dim)),requires_grad=True)
|
||||
self.gamma_2 = nn.Parameter(init_values * torch.ones((dim)),requires_grad=True)
|
||||
else:
|
||||
self.gamma_1, self.gamma_2 = None, None
|
||||
|
||||
self.postnorm = postnorm
|
||||
|
||||
def forward(self, x, rel_pos_bias=None, attn_mask=None):
|
||||
if self.gamma_1 is None:
|
||||
if self.postnorm:
|
||||
x = x + self.drop_path(self.norm1(self.attn(x, rel_pos_bias=rel_pos_bias, attn_mask=attn_mask)))
|
||||
x = x + self.drop_path(self.norm2(self.mlp(x)))
|
||||
else:
|
||||
x = x + self.drop_path(self.attn(self.norm1(x), rel_pos_bias=rel_pos_bias, attn_mask=attn_mask))
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x)))
|
||||
else:
|
||||
if self.postnorm:
|
||||
x = x + self.drop_path(self.gamma_1 * self.norm1(self.attn(x, rel_pos_bias=rel_pos_bias, attn_mask=attn_mask)))
|
||||
x = x + self.drop_path(self.gamma_2 * self.norm2(self.mlp(x)))
|
||||
else:
|
||||
x = x + self.drop_path(self.gamma_1 * self.attn(self.norm1(x), rel_pos_bias=rel_pos_bias, attn_mask=attn_mask))
|
||||
x = x + self.drop_path(self.gamma_2 * self.mlp(self.norm2(x)))
|
||||
return x
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
""" Image to Patch Embedding
|
||||
"""
|
||||
def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768):
|
||||
super().__init__()
|
||||
img_size = to_2tuple(img_size)
|
||||
patch_size = to_2tuple(patch_size)
|
||||
num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0])
|
||||
self.patch_shape = (img_size[0] // patch_size[0], img_size[1] // patch_size[1])
|
||||
self.img_size = img_size
|
||||
self.patch_size = patch_size
|
||||
self.num_patches = num_patches
|
||||
|
||||
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
|
||||
|
||||
def forward(self, x, **kwargs):
|
||||
B, C, H, W = x.shape
|
||||
# FIXME look at relaxing size constraints
|
||||
assert H == self.img_size[0] and W == self.img_size[1], \
|
||||
f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})."
|
||||
x = self.proj(x).flatten(2).transpose(1, 2)
|
||||
return x
|
||||
|
||||
|
||||
class RelativePositionBias(nn.Module):
|
||||
|
||||
def __init__(self, window_size, num_heads):
|
||||
super().__init__()
|
||||
self.window_size = window_size
|
||||
self.num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) + 3
|
||||
self.relative_position_bias_table = nn.Parameter(
|
||||
torch.zeros(self.num_relative_distance, num_heads)) # 2*Wh-1 * 2*Ww-1, nH
|
||||
# cls to token & token 2 cls & cls to cls
|
||||
|
||||
# get pair-wise relative position index for each token inside the window
|
||||
coords_h = torch.arange(window_size[0])
|
||||
coords_w = torch.arange(window_size[1])
|
||||
coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww
|
||||
coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
|
||||
relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
|
||||
relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
|
||||
relative_coords[:, :, 0] += window_size[0] - 1 # shift to start from 0
|
||||
relative_coords[:, :, 1] += window_size[1] - 1
|
||||
relative_coords[:, :, 0] *= 2 * window_size[1] - 1
|
||||
relative_position_index = \
|
||||
torch.zeros(size=(window_size[0] * window_size[1] + 1,) * 2, dtype=relative_coords.dtype)
|
||||
relative_position_index[1:, 1:] = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
|
||||
relative_position_index[0, 0:] = self.num_relative_distance - 3
|
||||
relative_position_index[0:, 0] = self.num_relative_distance - 2
|
||||
relative_position_index[0, 0] = self.num_relative_distance - 1
|
||||
|
||||
self.register_buffer("relative_position_index", relative_position_index)
|
||||
|
||||
def forward(self):
|
||||
relative_position_bias = \
|
||||
self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
|
||||
self.window_size[0] * self.window_size[1] + 1,
|
||||
self.window_size[0] * self.window_size[1] + 1, -1) # Wh*Ww,Wh*Ww,nH
|
||||
return relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
|
||||
|
||||
|
||||
class EVAVisionTransformer(nn.Module):
|
||||
""" Vision Transformer with support for patch or hybrid CNN input stage
|
||||
"""
|
||||
def __init__(self, img_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12,
|
||||
num_heads=12, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop_rate=0., attn_drop_rate=0.,
|
||||
drop_path_rate=0., norm_layer=nn.LayerNorm, init_values=None, patch_dropout=0.,
|
||||
use_abs_pos_emb=True, use_rel_pos_bias=False, use_shared_rel_pos_bias=False, rope=False,
|
||||
use_mean_pooling=True, init_scale=0.001, grad_checkpointing=False, xattn=False, postnorm=False,
|
||||
pt_hw_seq_len=16, intp_freq=False, naiveswiglu=False, subln=False):
|
||||
super().__init__()
|
||||
|
||||
if not XFORMERS_IS_AVAILBLE:
|
||||
xattn = False
|
||||
|
||||
self.image_size = img_size
|
||||
self.num_classes = num_classes
|
||||
self.num_features = self.embed_dim = embed_dim # num_features for consistency with other models
|
||||
|
||||
self.patch_embed = PatchEmbed(
|
||||
img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim)
|
||||
num_patches = self.patch_embed.num_patches
|
||||
|
||||
self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
|
||||
# self.mask_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
|
||||
if use_abs_pos_emb:
|
||||
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
|
||||
else:
|
||||
self.pos_embed = None
|
||||
self.pos_drop = nn.Dropout(p=drop_rate)
|
||||
|
||||
if use_shared_rel_pos_bias:
|
||||
self.rel_pos_bias = RelativePositionBias(window_size=self.patch_embed.patch_shape, num_heads=num_heads)
|
||||
else:
|
||||
self.rel_pos_bias = None
|
||||
|
||||
if rope:
|
||||
half_head_dim = embed_dim // num_heads // 2
|
||||
hw_seq_len = img_size // patch_size
|
||||
self.rope = VisionRotaryEmbeddingFast(
|
||||
dim=half_head_dim,
|
||||
pt_seq_len=pt_hw_seq_len,
|
||||
ft_seq_len=hw_seq_len if intp_freq else None,
|
||||
# patch_dropout=patch_dropout
|
||||
)
|
||||
else:
|
||||
self.rope = None
|
||||
|
||||
self.naiveswiglu = naiveswiglu
|
||||
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule
|
||||
self.use_rel_pos_bias = use_rel_pos_bias
|
||||
self.blocks = nn.ModuleList([
|
||||
Block(
|
||||
dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer,
|
||||
init_values=init_values, window_size=self.patch_embed.patch_shape if use_rel_pos_bias else None,
|
||||
xattn=xattn, rope=self.rope, postnorm=postnorm, subln=subln, naiveswiglu=naiveswiglu)
|
||||
for i in range(depth)])
|
||||
self.norm = nn.Identity() if use_mean_pooling else norm_layer(embed_dim)
|
||||
self.fc_norm = norm_layer(embed_dim) if use_mean_pooling else None
|
||||
self.head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
if self.pos_embed is not None:
|
||||
trunc_normal_(self.pos_embed, std=.02)
|
||||
|
||||
trunc_normal_(self.cls_token, std=.02)
|
||||
# trunc_normal_(self.mask_token, std=.02)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
self.fix_init_weight()
|
||||
|
||||
if isinstance(self.head, nn.Linear):
|
||||
trunc_normal_(self.head.weight, std=.02)
|
||||
self.head.weight.data.mul_(init_scale)
|
||||
self.head.bias.data.mul_(init_scale)
|
||||
|
||||
# setting a patch_dropout of 0. would mean it is disabled and this function would be the identity fn
|
||||
self.patch_dropout = PatchDropout(patch_dropout) if patch_dropout > 0. else nn.Identity()
|
||||
|
||||
self.grad_checkpointing = grad_checkpointing
|
||||
|
||||
def fix_init_weight(self):
|
||||
def rescale(param, layer_id):
|
||||
param.div_(math.sqrt(2.0 * layer_id))
|
||||
|
||||
for layer_id, layer in enumerate(self.blocks):
|
||||
rescale(layer.attn.proj.weight.data, layer_id + 1)
|
||||
if self.naiveswiglu:
|
||||
rescale(layer.mlp.w3.weight.data, layer_id + 1)
|
||||
else:
|
||||
rescale(layer.mlp.fc2.weight.data, layer_id + 1)
|
||||
|
||||
def get_cast_dtype(self) -> torch.dtype:
|
||||
return self.blocks[0].mlp.fc2.weight.dtype
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.blocks)
|
||||
|
||||
def lock(self, unlocked_groups=0, freeze_bn_stats=False):
|
||||
assert unlocked_groups == 0, 'partial locking not currently supported for this model'
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
self.grad_checkpointing = enable
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed', 'cls_token'}
|
||||
|
||||
def get_classifier(self):
|
||||
return self.head
|
||||
|
||||
def reset_classifier(self, num_classes, global_pool=''):
|
||||
self.num_classes = num_classes
|
||||
self.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
def forward_features(self, x, return_all_features=False, return_hidden=False, shuffle=False):
|
||||
|
||||
x = self.patch_embed(x)
|
||||
batch_size, seq_len, _ = x.size()
|
||||
|
||||
if shuffle:
|
||||
idx = torch.randperm(x.shape[1]) + 1
|
||||
zero = torch.LongTensor([0, ])
|
||||
idx = torch.cat([zero, idx])
|
||||
pos_embed = self.pos_embed[:, idx]
|
||||
|
||||
cls_tokens = self.cls_token.expand(batch_size, -1, -1) # stole cls_tokens impl from Phil Wang, thanks
|
||||
x = torch.cat((cls_tokens, x), dim=1)
|
||||
if shuffle:
|
||||
x = x + pos_embed
|
||||
elif self.pos_embed is not None:
|
||||
x = x + self.pos_embed
|
||||
x = self.pos_drop(x)
|
||||
|
||||
# a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in
|
||||
if os.getenv('RoPE') == '1':
|
||||
if self.training and not isinstance(self.patch_dropout, nn.Identity):
|
||||
x, patch_indices_keep = self.patch_dropout(x)
|
||||
self.rope.forward = partial(self.rope.forward, patch_indices_keep=patch_indices_keep)
|
||||
else:
|
||||
self.rope.forward = partial(self.rope.forward, patch_indices_keep=None)
|
||||
x = self.patch_dropout(x)
|
||||
else:
|
||||
x = self.patch_dropout(x)
|
||||
|
||||
rel_pos_bias = self.rel_pos_bias() if self.rel_pos_bias is not None else None
|
||||
hidden_states = []
|
||||
for idx, blk in enumerate(self.blocks):
|
||||
if (0 < idx <= 20) and (idx % 4 == 0) and return_hidden:
|
||||
hidden_states.append(x)
|
||||
if self.grad_checkpointing:
|
||||
x = checkpoint(blk, x, (rel_pos_bias,))
|
||||
else:
|
||||
x = blk(x, rel_pos_bias=rel_pos_bias)
|
||||
|
||||
if not return_all_features:
|
||||
x = self.norm(x)
|
||||
if self.fc_norm is not None:
|
||||
return self.fc_norm(x.mean(1)), hidden_states
|
||||
else:
|
||||
return x[:, 0], hidden_states
|
||||
return x
|
||||
|
||||
def forward(self, x, return_all_features=False, return_hidden=False, shuffle=False):
|
||||
if return_all_features:
|
||||
return self.forward_features(x, return_all_features, return_hidden, shuffle)
|
||||
x, hidden_states = self.forward_features(x, return_all_features, return_hidden, shuffle)
|
||||
x = self.head(x)
|
||||
if return_hidden:
|
||||
return x, hidden_states
|
||||
return x
|
||||
@@ -0,0 +1,517 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import pathlib
|
||||
import re
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from typing import Optional, Tuple, Union, Dict, Any
|
||||
import torch
|
||||
|
||||
from .constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD
|
||||
from .model import CLIP, CustomCLIP, convert_weights_to_lp, convert_to_custom_text_state_dict,\
|
||||
get_cast_dtype
|
||||
from .openai import load_openai_model
|
||||
from .pretrained import is_pretrained_cfg, get_pretrained_cfg, download_pretrained, list_pretrained_tags_by_model
|
||||
from .transform import image_transform
|
||||
from .tokenizer import HFTokenizer, tokenize
|
||||
from .utils import resize_clip_pos_embed, resize_evaclip_pos_embed, resize_visual_pos_embed, resize_eva_pos_embed
|
||||
|
||||
|
||||
_MODEL_CONFIG_PATHS = [Path(__file__).parent / f"model_configs/"]
|
||||
_MODEL_CONFIGS = {} # directory (model_name: config) of model architecture configs
|
||||
|
||||
|
||||
def _natural_key(string_):
|
||||
return [int(s) if s.isdigit() else s for s in re.split(r'(\d+)', string_.lower())]
|
||||
|
||||
|
||||
def _rescan_model_configs():
|
||||
global _MODEL_CONFIGS
|
||||
|
||||
config_ext = ('.json',)
|
||||
config_files = []
|
||||
for config_path in _MODEL_CONFIG_PATHS:
|
||||
if config_path.is_file() and config_path.suffix in config_ext:
|
||||
config_files.append(config_path)
|
||||
elif config_path.is_dir():
|
||||
for ext in config_ext:
|
||||
config_files.extend(config_path.glob(f'*{ext}'))
|
||||
|
||||
for cf in config_files:
|
||||
with open(cf, "r", encoding="utf8") as f:
|
||||
model_cfg = json.load(f)
|
||||
if all(a in model_cfg for a in ('embed_dim', 'vision_cfg', 'text_cfg')):
|
||||
_MODEL_CONFIGS[cf.stem] = model_cfg
|
||||
|
||||
_MODEL_CONFIGS = dict(sorted(_MODEL_CONFIGS.items(), key=lambda x: _natural_key(x[0])))
|
||||
|
||||
|
||||
_rescan_model_configs() # initial populate of model config registry
|
||||
|
||||
|
||||
def list_models():
|
||||
""" enumerate available model architectures based on config files """
|
||||
return list(_MODEL_CONFIGS.keys())
|
||||
|
||||
|
||||
def add_model_config(path):
|
||||
""" add model config path or file and update registry """
|
||||
if not isinstance(path, Path):
|
||||
path = Path(path)
|
||||
_MODEL_CONFIG_PATHS.append(path)
|
||||
_rescan_model_configs()
|
||||
|
||||
|
||||
def get_model_config(model_name):
|
||||
if model_name in _MODEL_CONFIGS:
|
||||
return deepcopy(_MODEL_CONFIGS[model_name])
|
||||
else:
|
||||
return None
|
||||
|
||||
|
||||
def get_tokenizer(model_name):
|
||||
config = get_model_config(model_name)
|
||||
tokenizer = HFTokenizer(config['text_cfg']['hf_tokenizer_name']) if 'hf_tokenizer_name' in config['text_cfg'] else tokenize
|
||||
return tokenizer
|
||||
|
||||
|
||||
# loading openai CLIP weights when is_openai=True for training
|
||||
def load_state_dict(checkpoint_path: str, map_location: str='cpu', model_key: str='model|module|state_dict', is_openai: bool=False, skip_list: list=[]):
|
||||
if is_openai:
|
||||
model = torch.jit.load(checkpoint_path, map_location="cpu").eval()
|
||||
state_dict = model.state_dict()
|
||||
for key in ["input_resolution", "context_length", "vocab_size"]:
|
||||
state_dict.pop(key, None)
|
||||
else:
|
||||
checkpoint = torch.load(checkpoint_path, map_location=map_location)
|
||||
for mk in model_key.split('|'):
|
||||
if isinstance(checkpoint, dict) and mk in checkpoint:
|
||||
state_dict = checkpoint[mk]
|
||||
break
|
||||
else:
|
||||
state_dict = checkpoint
|
||||
if next(iter(state_dict.items()))[0].startswith('module'):
|
||||
state_dict = {k[7:]: v for k, v in state_dict.items()}
|
||||
|
||||
for k in skip_list:
|
||||
if k in list(state_dict.keys()):
|
||||
logging.info(f"Removing key {k} from pretrained checkpoint")
|
||||
del state_dict[k]
|
||||
|
||||
if os.getenv('RoPE') == '1':
|
||||
for k in list(state_dict.keys()):
|
||||
if 'freqs_cos' in k or 'freqs_sin' in k:
|
||||
del state_dict[k]
|
||||
return state_dict
|
||||
|
||||
|
||||
|
||||
def load_checkpoint(model, checkpoint_path, model_key="model|module|state_dict", strict=True):
|
||||
state_dict = load_state_dict(checkpoint_path, model_key=model_key, is_openai=False)
|
||||
# detect old format and make compatible with new format
|
||||
if 'positional_embedding' in state_dict and not hasattr(model, 'positional_embedding'):
|
||||
state_dict = convert_to_custom_text_state_dict(state_dict)
|
||||
if 'text.logit_scale' in state_dict and hasattr(model, 'logit_scale'):
|
||||
state_dict['logit_scale'] = state_dict['text.logit_scale']
|
||||
del state_dict['text.logit_scale']
|
||||
|
||||
# resize_clip_pos_embed for CLIP and open CLIP
|
||||
if 'visual.positional_embedding' in state_dict:
|
||||
resize_clip_pos_embed(state_dict, model)
|
||||
# specified to eva_vit_model
|
||||
elif 'visual.pos_embed' in state_dict:
|
||||
resize_evaclip_pos_embed(state_dict, model)
|
||||
|
||||
# resize_clip_pos_embed(state_dict, model)
|
||||
incompatible_keys = model.load_state_dict(state_dict, strict=strict)
|
||||
logging.info(f"incompatible_keys.missing_keys: {incompatible_keys.missing_keys}")
|
||||
return incompatible_keys
|
||||
|
||||
def load_clip_visual_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=[]):
|
||||
state_dict = load_state_dict(checkpoint_path, map_location=map_location, is_openai=is_openai, skip_list=skip_list)
|
||||
|
||||
for k in list(state_dict.keys()):
|
||||
if not k.startswith('visual.'):
|
||||
del state_dict[k]
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith('visual.'):
|
||||
new_k = k[7:]
|
||||
state_dict[new_k] = state_dict[k]
|
||||
del state_dict[k]
|
||||
return state_dict
|
||||
|
||||
def load_clip_text_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=[]):
|
||||
state_dict = load_state_dict(checkpoint_path, map_location=map_location, is_openai=is_openai, skip_list=skip_list)
|
||||
|
||||
for k in list(state_dict.keys()):
|
||||
if k.startswith('visual.'):
|
||||
del state_dict[k]
|
||||
return state_dict
|
||||
|
||||
def get_pretrained_tag(pretrained_model):
|
||||
pretrained_model = pretrained_model.lower()
|
||||
if "laion" in pretrained_model or "open_clip" in pretrained_model:
|
||||
return "open_clip"
|
||||
elif "openai" in pretrained_model:
|
||||
return "clip"
|
||||
elif "eva" in pretrained_model and "clip" in pretrained_model:
|
||||
return "eva_clip"
|
||||
else:
|
||||
return "other"
|
||||
|
||||
def load_pretrained_checkpoint(
|
||||
model,
|
||||
visual_checkpoint_path,
|
||||
text_checkpoint_path,
|
||||
strict=True,
|
||||
visual_model=None,
|
||||
text_model=None,
|
||||
model_key="model|module|state_dict",
|
||||
skip_list=[]):
|
||||
visual_tag = get_pretrained_tag(visual_model)
|
||||
text_tag = get_pretrained_tag(text_model)
|
||||
|
||||
logging.info(f"num of model state_dict keys: {len(model.state_dict().keys())}")
|
||||
visual_incompatible_keys, text_incompatible_keys = None, None
|
||||
if visual_checkpoint_path:
|
||||
if visual_tag == "eva_clip" or visual_tag == "open_clip":
|
||||
visual_state_dict = load_clip_visual_state_dict(visual_checkpoint_path, is_openai=False, skip_list=skip_list)
|
||||
elif visual_tag == "clip":
|
||||
visual_state_dict = load_clip_visual_state_dict(visual_checkpoint_path, is_openai=True, skip_list=skip_list)
|
||||
else:
|
||||
visual_state_dict = load_state_dict(visual_checkpoint_path, model_key=model_key, is_openai=False, skip_list=skip_list)
|
||||
|
||||
# resize_clip_pos_embed for CLIP and open CLIP
|
||||
if 'positional_embedding' in visual_state_dict:
|
||||
resize_visual_pos_embed(visual_state_dict, model)
|
||||
# specified to EVA model
|
||||
elif 'pos_embed' in visual_state_dict:
|
||||
resize_eva_pos_embed(visual_state_dict, model)
|
||||
|
||||
visual_incompatible_keys = model.visual.load_state_dict(visual_state_dict, strict=strict)
|
||||
logging.info(f"num of loaded visual_state_dict keys: {len(visual_state_dict.keys())}")
|
||||
logging.info(f"visual_incompatible_keys.missing_keys: {visual_incompatible_keys.missing_keys}")
|
||||
|
||||
if text_checkpoint_path:
|
||||
if text_tag == "eva_clip" or text_tag == "open_clip":
|
||||
text_state_dict = load_clip_text_state_dict(text_checkpoint_path, is_openai=False, skip_list=skip_list)
|
||||
elif text_tag == "clip":
|
||||
text_state_dict = load_clip_text_state_dict(text_checkpoint_path, is_openai=True, skip_list=skip_list)
|
||||
else:
|
||||
text_state_dict = load_state_dict(visual_checkpoint_path, model_key=model_key, is_openai=False, skip_list=skip_list)
|
||||
|
||||
text_incompatible_keys = model.text.load_state_dict(text_state_dict, strict=strict)
|
||||
|
||||
logging.info(f"num of loaded text_state_dict keys: {len(text_state_dict.keys())}")
|
||||
logging.info(f"text_incompatible_keys.missing_keys: {text_incompatible_keys.missing_keys}")
|
||||
|
||||
return visual_incompatible_keys, text_incompatible_keys
|
||||
|
||||
def create_model(
|
||||
model_name: str,
|
||||
pretrained: Optional[str] = None,
|
||||
precision: str = 'fp32',
|
||||
device: Union[str, torch.device] = 'cpu',
|
||||
jit: bool = False,
|
||||
force_quick_gelu: bool = False,
|
||||
force_custom_clip: bool = False,
|
||||
force_patch_dropout: Optional[float] = None,
|
||||
pretrained_image: str = '',
|
||||
pretrained_text: str = '',
|
||||
pretrained_hf: bool = True,
|
||||
pretrained_visual_model: str = None,
|
||||
pretrained_text_model: str = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
skip_list: list = [],
|
||||
):
|
||||
model_name = model_name.replace('/', '-') # for callers using old naming with / in ViT names
|
||||
if isinstance(device, str):
|
||||
device = torch.device(device)
|
||||
|
||||
if pretrained and pretrained.lower() == 'openai':
|
||||
logging.info(f'Loading pretrained {model_name} from OpenAI.')
|
||||
model = load_openai_model(
|
||||
model_name,
|
||||
precision=precision,
|
||||
device=device,
|
||||
jit=jit,
|
||||
cache_dir=cache_dir,
|
||||
)
|
||||
else:
|
||||
model_cfg = get_model_config(model_name)
|
||||
if model_cfg is not None:
|
||||
logging.info(f'Loaded {model_name} model config.')
|
||||
else:
|
||||
logging.error(f'Model config for {model_name} not found; available models {list_models()}.')
|
||||
raise RuntimeError(f'Model config for {model_name} not found.')
|
||||
|
||||
if 'rope' in model_cfg.get('vision_cfg', {}):
|
||||
if model_cfg['vision_cfg']['rope']:
|
||||
os.environ['RoPE'] = "1"
|
||||
else:
|
||||
os.environ['RoPE'] = "0"
|
||||
|
||||
if force_quick_gelu:
|
||||
# override for use of QuickGELU on non-OpenAI transformer models
|
||||
model_cfg["quick_gelu"] = True
|
||||
|
||||
if force_patch_dropout is not None:
|
||||
# override the default patch dropout value
|
||||
model_cfg['vision_cfg']["patch_dropout"] = force_patch_dropout
|
||||
|
||||
cast_dtype = get_cast_dtype(precision)
|
||||
custom_clip = model_cfg.pop('custom_text', False) or force_custom_clip or ('hf_model_name' in model_cfg['text_cfg'])
|
||||
|
||||
|
||||
if custom_clip:
|
||||
if 'hf_model_name' in model_cfg.get('text_cfg', {}):
|
||||
model_cfg['text_cfg']['hf_model_pretrained'] = pretrained_hf
|
||||
model = CustomCLIP(**model_cfg, cast_dtype=cast_dtype)
|
||||
else:
|
||||
model = CLIP(**model_cfg, cast_dtype=cast_dtype)
|
||||
|
||||
pretrained_cfg = {}
|
||||
if pretrained:
|
||||
checkpoint_path = ''
|
||||
pretrained_cfg = get_pretrained_cfg(model_name, pretrained)
|
||||
if pretrained_cfg:
|
||||
checkpoint_path = download_pretrained(pretrained_cfg, cache_dir=cache_dir)
|
||||
elif os.path.exists(pretrained):
|
||||
checkpoint_path = pretrained
|
||||
|
||||
if checkpoint_path:
|
||||
logging.info(f'Loading pretrained {model_name} weights ({pretrained}).')
|
||||
load_checkpoint(model,
|
||||
checkpoint_path,
|
||||
model_key="model|module|state_dict",
|
||||
strict=False
|
||||
)
|
||||
else:
|
||||
error_str = (
|
||||
f'Pretrained weights ({pretrained}) not found for model {model_name}.'
|
||||
f'Available pretrained tags ({list_pretrained_tags_by_model(model_name)}.')
|
||||
logging.warning(error_str)
|
||||
raise RuntimeError(error_str)
|
||||
else:
|
||||
visual_checkpoint_path = ''
|
||||
text_checkpoint_path = ''
|
||||
|
||||
if pretrained_image:
|
||||
pretrained_visual_model = pretrained_visual_model.replace('/', '-') # for callers using old naming with / in ViT names
|
||||
pretrained_image_cfg = get_pretrained_cfg(pretrained_visual_model, pretrained_image)
|
||||
if 'timm_model_name' in model_cfg.get('vision_cfg', {}):
|
||||
# pretrained weight loading for timm models set via vision_cfg
|
||||
model_cfg['vision_cfg']['timm_model_pretrained'] = True
|
||||
elif pretrained_image_cfg:
|
||||
visual_checkpoint_path = download_pretrained(pretrained_image_cfg, cache_dir=cache_dir)
|
||||
elif os.path.exists(pretrained_image):
|
||||
visual_checkpoint_path = pretrained_image
|
||||
else:
|
||||
logging.warning(f'Pretrained weights ({visual_checkpoint_path}) not found for model {model_name}.visual.')
|
||||
raise RuntimeError(f'Pretrained weights ({visual_checkpoint_path}) not found for model {model_name}.visual.')
|
||||
|
||||
if pretrained_text:
|
||||
pretrained_text_model = pretrained_text_model.replace('/', '-') # for callers using old naming with / in ViT names
|
||||
pretrained_text_cfg = get_pretrained_cfg(pretrained_text_model, pretrained_text)
|
||||
if pretrained_image_cfg:
|
||||
text_checkpoint_path = download_pretrained(pretrained_text_cfg, cache_dir=cache_dir)
|
||||
elif os.path.exists(pretrained_text):
|
||||
text_checkpoint_path = pretrained_text
|
||||
else:
|
||||
logging.warning(f'Pretrained weights ({text_checkpoint_path}) not found for model {model_name}.text.')
|
||||
raise RuntimeError(f'Pretrained weights ({text_checkpoint_path}) not found for model {model_name}.text.')
|
||||
|
||||
if visual_checkpoint_path:
|
||||
logging.info(f'Loading pretrained {model_name}.visual weights ({visual_checkpoint_path}).')
|
||||
if text_checkpoint_path:
|
||||
logging.info(f'Loading pretrained {model_name}.text weights ({text_checkpoint_path}).')
|
||||
|
||||
if visual_checkpoint_path or text_checkpoint_path:
|
||||
load_pretrained_checkpoint(
|
||||
model,
|
||||
visual_checkpoint_path,
|
||||
text_checkpoint_path,
|
||||
strict=False,
|
||||
visual_model=pretrained_visual_model,
|
||||
text_model=pretrained_text_model,
|
||||
model_key="model|module|state_dict",
|
||||
skip_list=skip_list
|
||||
)
|
||||
|
||||
if "fp16" in precision or "bf16" in precision:
|
||||
logging.info(f'convert precision to {precision}')
|
||||
model = model.to(torch.bfloat16) if 'bf16' in precision else model.to(torch.float16)
|
||||
|
||||
model.to(device=device)
|
||||
|
||||
# set image / mean metadata from pretrained_cfg if available, or use default
|
||||
model.visual.image_mean = pretrained_cfg.get('mean', None) or OPENAI_DATASET_MEAN
|
||||
model.visual.image_std = pretrained_cfg.get('std', None) or OPENAI_DATASET_STD
|
||||
|
||||
if jit:
|
||||
model = torch.jit.script(model)
|
||||
|
||||
return model
|
||||
|
||||
|
||||
def create_model_and_transforms(
|
||||
model_name: str,
|
||||
pretrained: Optional[str] = None,
|
||||
precision: str = 'fp32',
|
||||
device: Union[str, torch.device] = 'cpu',
|
||||
jit: bool = False,
|
||||
force_quick_gelu: bool = False,
|
||||
force_custom_clip: bool = False,
|
||||
force_patch_dropout: Optional[float] = None,
|
||||
pretrained_image: str = '',
|
||||
pretrained_text: str = '',
|
||||
pretrained_hf: bool = True,
|
||||
pretrained_visual_model: str = None,
|
||||
pretrained_text_model: str = None,
|
||||
image_mean: Optional[Tuple[float, ...]] = None,
|
||||
image_std: Optional[Tuple[float, ...]] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
skip_list: list = [],
|
||||
):
|
||||
model = create_model(
|
||||
model_name,
|
||||
pretrained,
|
||||
precision=precision,
|
||||
device=device,
|
||||
jit=jit,
|
||||
force_quick_gelu=force_quick_gelu,
|
||||
force_custom_clip=force_custom_clip,
|
||||
force_patch_dropout=force_patch_dropout,
|
||||
pretrained_image=pretrained_image,
|
||||
pretrained_text=pretrained_text,
|
||||
pretrained_hf=pretrained_hf,
|
||||
pretrained_visual_model=pretrained_visual_model,
|
||||
pretrained_text_model=pretrained_text_model,
|
||||
cache_dir=cache_dir,
|
||||
skip_list=skip_list,
|
||||
)
|
||||
|
||||
image_mean = image_mean or getattr(model.visual, 'image_mean', None)
|
||||
image_std = image_std or getattr(model.visual, 'image_std', None)
|
||||
preprocess_train = image_transform(
|
||||
model.visual.image_size,
|
||||
is_train=True,
|
||||
mean=image_mean,
|
||||
std=image_std
|
||||
)
|
||||
preprocess_val = image_transform(
|
||||
model.visual.image_size,
|
||||
is_train=False,
|
||||
mean=image_mean,
|
||||
std=image_std
|
||||
)
|
||||
|
||||
return model, preprocess_train, preprocess_val
|
||||
|
||||
|
||||
def create_transforms(
|
||||
model_name: str,
|
||||
pretrained: Optional[str] = None,
|
||||
precision: str = 'fp32',
|
||||
device: Union[str, torch.device] = 'cpu',
|
||||
jit: bool = False,
|
||||
force_quick_gelu: bool = False,
|
||||
force_custom_clip: bool = False,
|
||||
force_patch_dropout: Optional[float] = None,
|
||||
pretrained_image: str = '',
|
||||
pretrained_text: str = '',
|
||||
pretrained_hf: bool = True,
|
||||
pretrained_visual_model: str = None,
|
||||
pretrained_text_model: str = None,
|
||||
image_mean: Optional[Tuple[float, ...]] = None,
|
||||
image_std: Optional[Tuple[float, ...]] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
skip_list: list = [],
|
||||
):
|
||||
model = create_model(
|
||||
model_name,
|
||||
pretrained,
|
||||
precision=precision,
|
||||
device=device,
|
||||
jit=jit,
|
||||
force_quick_gelu=force_quick_gelu,
|
||||
force_custom_clip=force_custom_clip,
|
||||
force_patch_dropout=force_patch_dropout,
|
||||
pretrained_image=pretrained_image,
|
||||
pretrained_text=pretrained_text,
|
||||
pretrained_hf=pretrained_hf,
|
||||
pretrained_visual_model=pretrained_visual_model,
|
||||
pretrained_text_model=pretrained_text_model,
|
||||
cache_dir=cache_dir,
|
||||
skip_list=skip_list,
|
||||
)
|
||||
|
||||
|
||||
image_mean = image_mean or getattr(model.visual, 'image_mean', None)
|
||||
image_std = image_std or getattr(model.visual, 'image_std', None)
|
||||
preprocess_train = image_transform(
|
||||
model.visual.image_size,
|
||||
is_train=True,
|
||||
mean=image_mean,
|
||||
std=image_std
|
||||
)
|
||||
preprocess_val = image_transform(
|
||||
model.visual.image_size,
|
||||
is_train=False,
|
||||
mean=image_mean,
|
||||
std=image_std
|
||||
)
|
||||
del model
|
||||
|
||||
return preprocess_train, preprocess_val
|
||||
|
||||
def create_model_from_pretrained(
|
||||
model_name: str,
|
||||
pretrained: str,
|
||||
precision: str = 'fp32',
|
||||
device: Union[str, torch.device] = 'cpu',
|
||||
jit: bool = False,
|
||||
force_quick_gelu: bool = False,
|
||||
force_custom_clip: bool = False,
|
||||
force_patch_dropout: Optional[float] = None,
|
||||
return_transform: bool = True,
|
||||
image_mean: Optional[Tuple[float, ...]] = None,
|
||||
image_std: Optional[Tuple[float, ...]] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
is_frozen: bool = False,
|
||||
):
|
||||
if not is_pretrained_cfg(model_name, pretrained) and not os.path.exists(pretrained):
|
||||
raise RuntimeError(
|
||||
f'{pretrained} is not a valid pretrained cfg or checkpoint for {model_name}.'
|
||||
f' Use open_clip.list_pretrained() to find one.')
|
||||
|
||||
model = create_model(
|
||||
model_name,
|
||||
pretrained,
|
||||
precision=precision,
|
||||
device=device,
|
||||
jit=jit,
|
||||
force_quick_gelu=force_quick_gelu,
|
||||
force_custom_clip=force_custom_clip,
|
||||
force_patch_dropout=force_patch_dropout,
|
||||
cache_dir=cache_dir,
|
||||
)
|
||||
|
||||
if is_frozen:
|
||||
for param in model.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
if not return_transform:
|
||||
return model
|
||||
|
||||
image_mean = image_mean or getattr(model.visual, 'image_mean', None)
|
||||
image_std = image_std or getattr(model.visual, 'image_std', None)
|
||||
preprocess = image_transform(
|
||||
model.visual.image_size,
|
||||
is_train=False,
|
||||
mean=image_mean,
|
||||
std=image_std
|
||||
)
|
||||
|
||||
return model, preprocess
|
||||
@@ -0,0 +1,57 @@
|
||||
# HF architecture dict:
|
||||
arch_dict = {
|
||||
# https://huggingface.co/docs/transformers/model_doc/roberta#roberta
|
||||
"roberta": {
|
||||
"config_names": {
|
||||
"context_length": "max_position_embeddings",
|
||||
"vocab_size": "vocab_size",
|
||||
"width": "hidden_size",
|
||||
"heads": "num_attention_heads",
|
||||
"layers": "num_hidden_layers",
|
||||
"layer_attr": "layer",
|
||||
"token_embeddings_attr": "embeddings"
|
||||
},
|
||||
"pooler": "mean_pooler",
|
||||
},
|
||||
# https://huggingface.co/docs/transformers/model_doc/xlm-roberta#transformers.XLMRobertaConfig
|
||||
"xlm-roberta": {
|
||||
"config_names": {
|
||||
"context_length": "max_position_embeddings",
|
||||
"vocab_size": "vocab_size",
|
||||
"width": "hidden_size",
|
||||
"heads": "num_attention_heads",
|
||||
"layers": "num_hidden_layers",
|
||||
"layer_attr": "layer",
|
||||
"token_embeddings_attr": "embeddings"
|
||||
},
|
||||
"pooler": "mean_pooler",
|
||||
},
|
||||
# https://huggingface.co/docs/transformers/model_doc/mt5#mt5
|
||||
"mt5": {
|
||||
"config_names": {
|
||||
# unlimited seqlen
|
||||
# https://github.com/google-research/text-to-text-transfer-transformer/issues/273
|
||||
# https://github.com/huggingface/transformers/blob/v4.24.0/src/transformers/models/t5/modeling_t5.py#L374
|
||||
"context_length": "",
|
||||
"vocab_size": "vocab_size",
|
||||
"width": "d_model",
|
||||
"heads": "num_heads",
|
||||
"layers": "num_layers",
|
||||
"layer_attr": "block",
|
||||
"token_embeddings_attr": "embed_tokens"
|
||||
},
|
||||
"pooler": "mean_pooler",
|
||||
},
|
||||
"bert": {
|
||||
"config_names": {
|
||||
"context_length": "max_position_embeddings",
|
||||
"vocab_size": "vocab_size",
|
||||
"width": "hidden_size",
|
||||
"heads": "num_attention_heads",
|
||||
"layers": "num_hidden_layers",
|
||||
"layer_attr": "layer",
|
||||
"token_embeddings_attr": "embeddings"
|
||||
},
|
||||
"pooler": "mean_pooler",
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
""" huggingface model adapter
|
||||
|
||||
Wraps HuggingFace transformers (https://github.com/huggingface/transformers) models for use as a text tower in CLIP model.
|
||||
"""
|
||||
|
||||
import re
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
from torch import TensorType
|
||||
try:
|
||||
import transformers
|
||||
from transformers import AutoModel, AutoModelForMaskedLM, AutoTokenizer, AutoConfig, PretrainedConfig
|
||||
from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling, \
|
||||
BaseModelOutputWithPoolingAndCrossAttentions
|
||||
except ImportError as e:
|
||||
transformers = None
|
||||
|
||||
|
||||
class BaseModelOutput:
|
||||
pass
|
||||
|
||||
|
||||
class PretrainedConfig:
|
||||
pass
|
||||
|
||||
from .hf_configs import arch_dict
|
||||
|
||||
# utils
|
||||
def _camel2snake(s):
|
||||
return re.sub(r'(?<!^)(?=[A-Z])', '_', s).lower()
|
||||
|
||||
# TODO: ?last - for gpt-like models
|
||||
_POOLERS = {}
|
||||
|
||||
def register_pooler(cls):
|
||||
"""Decorator registering pooler class"""
|
||||
_POOLERS[_camel2snake(cls.__name__)] = cls
|
||||
return cls
|
||||
|
||||
|
||||
@register_pooler
|
||||
class MeanPooler(nn.Module):
|
||||
"""Mean pooling"""
|
||||
def forward(self, x:BaseModelOutput, attention_mask:TensorType):
|
||||
masked_output = x.last_hidden_state * attention_mask.unsqueeze(-1)
|
||||
return masked_output.sum(dim=1) / attention_mask.sum(-1, keepdim=True)
|
||||
|
||||
@register_pooler
|
||||
class MaxPooler(nn.Module):
|
||||
"""Max pooling"""
|
||||
def forward(self, x:BaseModelOutput, attention_mask:TensorType):
|
||||
masked_output = x.last_hidden_state.masked_fill(attention_mask.unsqueeze(-1), -torch.inf)
|
||||
return masked_output.max(1).values
|
||||
|
||||
@register_pooler
|
||||
class ClsPooler(nn.Module):
|
||||
"""CLS token pooling"""
|
||||
def __init__(self, use_pooler_output=True):
|
||||
super().__init__()
|
||||
self.cls_token_position = 0
|
||||
self.use_pooler_output = use_pooler_output
|
||||
|
||||
def forward(self, x:BaseModelOutput, attention_mask:TensorType):
|
||||
|
||||
if (self.use_pooler_output and
|
||||
isinstance(x, (BaseModelOutputWithPooling, BaseModelOutputWithPoolingAndCrossAttentions)) and
|
||||
(x.pooler_output is not None)
|
||||
):
|
||||
return x.pooler_output
|
||||
|
||||
return x.last_hidden_state[:, self.cls_token_position, :]
|
||||
|
||||
class HFTextEncoder(nn.Module):
|
||||
"""HuggingFace model adapter"""
|
||||
def __init__(
|
||||
self,
|
||||
model_name_or_path: str,
|
||||
output_dim: int,
|
||||
tokenizer_name: str = None,
|
||||
config: PretrainedConfig = None,
|
||||
pooler_type: str = None,
|
||||
proj: str = None,
|
||||
pretrained: bool = True,
|
||||
masked_language_modeling: bool = False):
|
||||
super().__init__()
|
||||
|
||||
self.output_dim = output_dim
|
||||
|
||||
# TODO: find better way to get this information
|
||||
uses_transformer_pooler = (pooler_type == "cls_pooler")
|
||||
|
||||
if transformers is None:
|
||||
raise RuntimeError("Please `pip install transformers` to use pre-trained HuggingFace models")
|
||||
if config is None:
|
||||
self.config = AutoConfig.from_pretrained(model_name_or_path)
|
||||
if masked_language_modeling:
|
||||
create_func, model_args = (AutoModelForMaskedLM.from_pretrained, model_name_or_path) if pretrained else (
|
||||
AutoModelForMaskedLM.from_config, self.config)
|
||||
else:
|
||||
create_func, model_args = (AutoModel.from_pretrained, model_name_or_path) if pretrained else (
|
||||
AutoModel.from_config, self.config)
|
||||
# TODO: do all model configs have this attribute? PretrainedConfig does so yes??
|
||||
if hasattr(self.config, "is_encoder_decoder") and self.config.is_encoder_decoder:
|
||||
self.transformer = create_func(model_args)
|
||||
self.transformer = self.transformer.encoder
|
||||
else:
|
||||
self.transformer = create_func(model_args, add_pooling_layer=uses_transformer_pooler)
|
||||
else:
|
||||
self.config = config
|
||||
if masked_language_modeling:
|
||||
self.transformer = AutoModelForMaskedLM.from_config(config)
|
||||
else:
|
||||
self.transformer = AutoModel.from_config(config)
|
||||
|
||||
if pooler_type is None: # get default arch pooler
|
||||
self.pooler = _POOLERS[(arch_dict[self.config.model_type]["pooler"])]()
|
||||
else:
|
||||
self.pooler = _POOLERS[pooler_type]()
|
||||
|
||||
d_model = getattr(self.config, arch_dict[self.config.model_type]["config_names"]["width"])
|
||||
if (d_model == output_dim) and (proj is None): # do we always need a proj?
|
||||
self.proj = nn.Identity()
|
||||
elif proj == 'linear':
|
||||
self.proj = nn.Linear(d_model, output_dim, bias=False)
|
||||
elif proj == 'mlp':
|
||||
hidden_size = (d_model + output_dim) // 2
|
||||
self.proj = nn.Sequential(
|
||||
nn.Linear(d_model, hidden_size, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(hidden_size, output_dim, bias=False),
|
||||
)
|
||||
|
||||
# self.itm_proj = nn.Linear(d_model, 2, bias=False)
|
||||
# self.mlm_proj = nn.Linear(d_model, self.config.vocab_size), bias=False)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
|
||||
|
||||
# def forward_itm(self, x:TensorType, image_embeds:TensorType) -> TensorType:
|
||||
# image_atts = torch.ones(image_embeds.size()[:-1],dtype=torch.long).to(x.device)
|
||||
# attn_mask = (x != self.config.pad_token_id).long()
|
||||
# out = self.transformer(
|
||||
# input_ids=x,
|
||||
# attention_mask=attn_mask,
|
||||
# encoder_hidden_states = image_embeds,
|
||||
# encoder_attention_mask = image_atts,
|
||||
# )
|
||||
# pooled_out = self.pooler(out, attn_mask)
|
||||
|
||||
# return self.itm_proj(pooled_out)
|
||||
|
||||
def mask(self, input_ids, vocab_size, device, targets=None, masked_indices=None, probability_matrix=None):
|
||||
if masked_indices is None:
|
||||
masked_indices = torch.bernoulli(probability_matrix).bool()
|
||||
|
||||
masked_indices[input_ids == self.tokenizer.pad_token_id] = False
|
||||
masked_indices[input_ids == self.tokenizer.cls_token_id] = False
|
||||
|
||||
if targets is not None:
|
||||
targets[~masked_indices] = -100 # We only compute loss on masked tokens
|
||||
|
||||
# 80% of the time, we replace masked input tokens with tokenizer.mask_token ([MASK])
|
||||
indices_replaced = torch.bernoulli(torch.full(input_ids.shape, 0.8)).bool() & masked_indices
|
||||
input_ids[indices_replaced] = self.tokenizer.mask_token_id
|
||||
|
||||
# 10% of the time, we replace masked input tokens with random word
|
||||
indices_random = torch.bernoulli(torch.full(input_ids.shape, 0.5)).bool() & masked_indices & ~indices_replaced
|
||||
random_words = torch.randint(vocab_size, input_ids.shape, dtype=torch.long).to(device)
|
||||
input_ids[indices_random] = random_words[indices_random]
|
||||
# The rest of the time (10% of the time) we keep the masked input tokens unchanged
|
||||
|
||||
if targets is not None:
|
||||
return input_ids, targets
|
||||
else:
|
||||
return input_ids
|
||||
|
||||
def forward_mlm(self, input_ids, image_embeds, mlm_probability=0.25):
|
||||
labels = input_ids.clone()
|
||||
attn_mask = (input_ids != self.config.pad_token_id).long()
|
||||
image_atts = torch.ones(image_embeds.size()[:-1],dtype=torch.long).to(input_ids.device)
|
||||
vocab_size = getattr(self.config, arch_dict[self.config.model_type]["config_names"]["vocab_size"])
|
||||
probability_matrix = torch.full(labels.shape, mlm_probability)
|
||||
input_ids, labels = self.mask(input_ids, vocab_size, input_ids.device, targets=labels,
|
||||
probability_matrix = probability_matrix)
|
||||
mlm_output = self.transformer(input_ids,
|
||||
attention_mask = attn_mask,
|
||||
encoder_hidden_states = image_embeds,
|
||||
encoder_attention_mask = image_atts,
|
||||
return_dict = True,
|
||||
labels = labels,
|
||||
)
|
||||
return mlm_output.loss
|
||||
# mlm_output = self.transformer(input_ids,
|
||||
# attention_mask = attn_mask,
|
||||
# encoder_hidden_states = image_embeds,
|
||||
# encoder_attention_mask = image_atts,
|
||||
# return_dict = True,
|
||||
# ).last_hidden_state
|
||||
# logits = self.mlm_proj(mlm_output)
|
||||
|
||||
# # logits = logits[:, :-1, :].contiguous().view(-1, vocab_size)
|
||||
# logits = logits[:, 1:, :].contiguous().view(-1, vocab_size)
|
||||
# labels = labels[:, 1:].contiguous().view(-1)
|
||||
|
||||
# mlm_loss = F.cross_entropy(
|
||||
# logits,
|
||||
# labels,
|
||||
# # label_smoothing=0.1,
|
||||
# )
|
||||
# return mlm_loss
|
||||
|
||||
|
||||
def forward(self, x:TensorType) -> TensorType:
|
||||
attn_mask = (x != self.config.pad_token_id).long()
|
||||
out = self.transformer(input_ids=x, attention_mask=attn_mask)
|
||||
pooled_out = self.pooler(out, attn_mask)
|
||||
|
||||
return self.proj(pooled_out)
|
||||
|
||||
def lock(self, unlocked_layers:int=0, freeze_layer_norm:bool=True):
|
||||
if not unlocked_layers: # full freezing
|
||||
for n, p in self.transformer.named_parameters():
|
||||
p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False
|
||||
return
|
||||
|
||||
encoder = self.transformer.encoder if hasattr(self.transformer, 'encoder') else self.transformer
|
||||
layer_list = getattr(encoder, arch_dict[self.config.model_type]["config_names"]["layer_attr"])
|
||||
print(f"Unlocking {unlocked_layers}/{len(layer_list) + 1} layers of hf model")
|
||||
embeddings = getattr(
|
||||
self.transformer, arch_dict[self.config.model_type]["config_names"]["token_embeddings_attr"])
|
||||
modules = [embeddings, *layer_list][:-unlocked_layers]
|
||||
# freeze layers
|
||||
for module in modules:
|
||||
for n, p in module.named_parameters():
|
||||
p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False
|
||||
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
self.transformer.gradient_checkpointing_enable()
|
||||
|
||||
def get_num_layers(self):
|
||||
encoder = self.transformer.encoder if hasattr(self.transformer, 'encoder') else self.transformer
|
||||
layer_list = getattr(encoder, arch_dict[self.config.model_type]["config_names"]["layer_attr"])
|
||||
return len(layer_list)
|
||||
|
||||
def init_parameters(self):
|
||||
pass
|
||||
@@ -0,0 +1,138 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
try:
|
||||
import torch.distributed.nn
|
||||
from torch import distributed as dist
|
||||
has_distributed = True
|
||||
except ImportError:
|
||||
has_distributed = False
|
||||
|
||||
try:
|
||||
import horovod.torch as hvd
|
||||
except ImportError:
|
||||
hvd = None
|
||||
|
||||
from timm.loss import LabelSmoothingCrossEntropy
|
||||
|
||||
|
||||
def gather_features(
|
||||
image_features,
|
||||
text_features,
|
||||
local_loss=False,
|
||||
gather_with_grad=False,
|
||||
rank=0,
|
||||
world_size=1,
|
||||
use_horovod=False
|
||||
):
|
||||
assert has_distributed, 'torch.distributed did not import correctly, please use a PyTorch version with support.'
|
||||
if use_horovod:
|
||||
assert hvd is not None, 'Please install horovod'
|
||||
if gather_with_grad:
|
||||
all_image_features = hvd.allgather(image_features)
|
||||
all_text_features = hvd.allgather(text_features)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
all_image_features = hvd.allgather(image_features)
|
||||
all_text_features = hvd.allgather(text_features)
|
||||
if not local_loss:
|
||||
# ensure grads for local rank when all_* features don't have a gradient
|
||||
gathered_image_features = list(all_image_features.chunk(world_size, dim=0))
|
||||
gathered_text_features = list(all_text_features.chunk(world_size, dim=0))
|
||||
gathered_image_features[rank] = image_features
|
||||
gathered_text_features[rank] = text_features
|
||||
all_image_features = torch.cat(gathered_image_features, dim=0)
|
||||
all_text_features = torch.cat(gathered_text_features, dim=0)
|
||||
else:
|
||||
# We gather tensors from all gpus
|
||||
if gather_with_grad:
|
||||
all_image_features = torch.cat(torch.distributed.nn.all_gather(image_features), dim=0)
|
||||
all_text_features = torch.cat(torch.distributed.nn.all_gather(text_features), dim=0)
|
||||
# all_image_features = torch.cat(torch.distributed.nn.all_gather(image_features, async_op=True), dim=0)
|
||||
# all_text_features = torch.cat(torch.distributed.nn.all_gather(text_features, async_op=True), dim=0)
|
||||
else:
|
||||
gathered_image_features = [torch.zeros_like(image_features) for _ in range(world_size)]
|
||||
gathered_text_features = [torch.zeros_like(text_features) for _ in range(world_size)]
|
||||
dist.all_gather(gathered_image_features, image_features)
|
||||
dist.all_gather(gathered_text_features, text_features)
|
||||
if not local_loss:
|
||||
# ensure grads for local rank when all_* features don't have a gradient
|
||||
gathered_image_features[rank] = image_features
|
||||
gathered_text_features[rank] = text_features
|
||||
all_image_features = torch.cat(gathered_image_features, dim=0)
|
||||
all_text_features = torch.cat(gathered_text_features, dim=0)
|
||||
|
||||
return all_image_features, all_text_features
|
||||
|
||||
|
||||
class ClipLoss(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
local_loss=False,
|
||||
gather_with_grad=False,
|
||||
cache_labels=False,
|
||||
rank=0,
|
||||
world_size=1,
|
||||
use_horovod=False,
|
||||
smoothing=0.,
|
||||
):
|
||||
super().__init__()
|
||||
self.local_loss = local_loss
|
||||
self.gather_with_grad = gather_with_grad
|
||||
self.cache_labels = cache_labels
|
||||
self.rank = rank
|
||||
self.world_size = world_size
|
||||
self.use_horovod = use_horovod
|
||||
self.label_smoothing_cross_entropy = LabelSmoothingCrossEntropy(smoothing=smoothing) if smoothing > 0 else None
|
||||
|
||||
# cache state
|
||||
self.prev_num_logits = 0
|
||||
self.labels = {}
|
||||
|
||||
def forward(self, image_features, text_features, logit_scale=1.):
|
||||
device = image_features.device
|
||||
if self.world_size > 1:
|
||||
all_image_features, all_text_features = gather_features(
|
||||
image_features, text_features,
|
||||
self.local_loss, self.gather_with_grad, self.rank, self.world_size, self.use_horovod)
|
||||
|
||||
if self.local_loss:
|
||||
logits_per_image = logit_scale * image_features @ all_text_features.T
|
||||
logits_per_text = logit_scale * text_features @ all_image_features.T
|
||||
else:
|
||||
logits_per_image = logit_scale * all_image_features @ all_text_features.T
|
||||
logits_per_text = logits_per_image.T
|
||||
else:
|
||||
logits_per_image = logit_scale * image_features @ text_features.T
|
||||
logits_per_text = logit_scale * text_features @ image_features.T
|
||||
# calculated ground-truth and cache if enabled
|
||||
num_logits = logits_per_image.shape[0]
|
||||
if self.prev_num_logits != num_logits or device not in self.labels:
|
||||
labels = torch.arange(num_logits, device=device, dtype=torch.long)
|
||||
if self.world_size > 1 and self.local_loss:
|
||||
labels = labels + num_logits * self.rank
|
||||
if self.cache_labels:
|
||||
self.labels[device] = labels
|
||||
self.prev_num_logits = num_logits
|
||||
else:
|
||||
labels = self.labels[device]
|
||||
|
||||
if self.label_smoothing_cross_entropy:
|
||||
total_loss = (
|
||||
self.label_smoothing_cross_entropy(logits_per_image, labels) +
|
||||
self.label_smoothing_cross_entropy(logits_per_text, labels)
|
||||
) / 2
|
||||
else:
|
||||
total_loss = (
|
||||
F.cross_entropy(logits_per_image, labels) +
|
||||
F.cross_entropy(logits_per_text, labels)
|
||||
) / 2
|
||||
|
||||
acc = None
|
||||
i2t_acc = (logits_per_image.argmax(-1) == labels).sum() / len(logits_per_image)
|
||||
t2i_acc = (logits_per_text.argmax(-1) == labels).sum() / len(logits_per_text)
|
||||
acc = {"i2t": i2t_acc, "t2i": t2i_acc}
|
||||
return total_loss, acc
|
||||
@@ -0,0 +1,432 @@
|
||||
""" CLIP Model
|
||||
|
||||
Adapted from https://github.com/openai/CLIP. Originally MIT License, Copyright (c) 2021 OpenAI.
|
||||
"""
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple, Union
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
try:
|
||||
from .hf_model import HFTextEncoder
|
||||
except:
|
||||
HFTextEncoder = None
|
||||
from .modified_resnet import ModifiedResNet
|
||||
from .timm_model import TimmModel
|
||||
from .eva_vit_model import EVAVisionTransformer
|
||||
from .transformer import LayerNorm, QuickGELU, Attention, VisionTransformer, TextTransformer
|
||||
|
||||
try:
|
||||
from apex.normalization import FusedLayerNorm
|
||||
except:
|
||||
FusedLayerNorm = LayerNorm
|
||||
|
||||
@dataclass
|
||||
class CLIPVisionCfg:
|
||||
layers: Union[Tuple[int, int, int, int], int] = 12
|
||||
width: int = 768
|
||||
head_width: int = 64
|
||||
mlp_ratio: float = 4.0
|
||||
patch_size: int = 16
|
||||
image_size: Union[Tuple[int, int], int] = 224
|
||||
ls_init_value: Optional[float] = None # layer scale initial value
|
||||
patch_dropout: float = 0. # what fraction of patches to dropout during training (0 would mean disabled and no patches dropped) - 0.5 to 0.75 recommended in the paper for optimal results
|
||||
global_average_pool: bool = False # whether to global average pool the last embedding layer, instead of using CLS token (https://arxiv.org/abs/2205.01580)
|
||||
drop_path_rate: Optional[float] = None # drop path rate
|
||||
timm_model_name: str = None # a valid model name overrides layers, width, patch_size
|
||||
timm_model_pretrained: bool = False # use (imagenet) pretrained weights for named model
|
||||
timm_pool: str = 'avg' # feature pooling for timm model ('abs_attn', 'rot_attn', 'avg', '')
|
||||
timm_proj: str = 'linear' # linear projection for timm model output ('linear', 'mlp', '')
|
||||
timm_proj_bias: bool = False # enable bias final projection
|
||||
eva_model_name: str = None # a valid eva model name overrides layers, width, patch_size
|
||||
qkv_bias: bool = True
|
||||
fusedLN: bool = False
|
||||
xattn: bool = False
|
||||
postnorm: bool = False
|
||||
rope: bool = False
|
||||
pt_hw_seq_len: int = 16 # 224/14
|
||||
intp_freq: bool = False
|
||||
naiveswiglu: bool = False
|
||||
subln: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPTextCfg:
|
||||
context_length: int = 77
|
||||
vocab_size: int = 49408
|
||||
width: int = 512
|
||||
heads: int = 8
|
||||
layers: int = 12
|
||||
ls_init_value: Optional[float] = None # layer scale initial value
|
||||
hf_model_name: str = None
|
||||
hf_tokenizer_name: str = None
|
||||
hf_model_pretrained: bool = True
|
||||
proj: str = 'mlp'
|
||||
pooler_type: str = 'mean_pooler'
|
||||
masked_language_modeling: bool = False
|
||||
fusedLN: bool = False
|
||||
xattn: bool = False
|
||||
attn_mask: bool = True
|
||||
|
||||
def get_cast_dtype(precision: str):
|
||||
cast_dtype = None
|
||||
if precision == 'bf16':
|
||||
cast_dtype = torch.bfloat16
|
||||
elif precision == 'fp16':
|
||||
cast_dtype = torch.float16
|
||||
return cast_dtype
|
||||
|
||||
|
||||
def _build_vision_tower(
|
||||
embed_dim: int,
|
||||
vision_cfg: CLIPVisionCfg,
|
||||
quick_gelu: bool = False,
|
||||
cast_dtype: Optional[torch.dtype] = None
|
||||
):
|
||||
if isinstance(vision_cfg, dict):
|
||||
vision_cfg = CLIPVisionCfg(**vision_cfg)
|
||||
|
||||
# OpenAI models are pretrained w/ QuickGELU but native nn.GELU is both faster and more
|
||||
# memory efficient in recent PyTorch releases (>= 1.10).
|
||||
# NOTE: timm models always use native GELU regardless of quick_gelu flag.
|
||||
act_layer = QuickGELU if quick_gelu else nn.GELU
|
||||
|
||||
if vision_cfg.eva_model_name:
|
||||
vision_heads = vision_cfg.width // vision_cfg.head_width
|
||||
norm_layer = LayerNorm
|
||||
|
||||
visual = EVAVisionTransformer(
|
||||
img_size=vision_cfg.image_size,
|
||||
patch_size=vision_cfg.patch_size,
|
||||
num_classes=embed_dim,
|
||||
use_mean_pooling=vision_cfg.global_average_pool, #False
|
||||
init_values=vision_cfg.ls_init_value,
|
||||
patch_dropout=vision_cfg.patch_dropout,
|
||||
embed_dim=vision_cfg.width,
|
||||
depth=vision_cfg.layers,
|
||||
num_heads=vision_heads,
|
||||
mlp_ratio=vision_cfg.mlp_ratio,
|
||||
qkv_bias=vision_cfg.qkv_bias,
|
||||
drop_path_rate=vision_cfg.drop_path_rate,
|
||||
norm_layer= partial(FusedLayerNorm, eps=1e-6) if vision_cfg.fusedLN else partial(norm_layer, eps=1e-6),
|
||||
xattn=vision_cfg.xattn,
|
||||
rope=vision_cfg.rope,
|
||||
postnorm=vision_cfg.postnorm,
|
||||
pt_hw_seq_len= vision_cfg.pt_hw_seq_len, # 224/14
|
||||
intp_freq= vision_cfg.intp_freq,
|
||||
naiveswiglu= vision_cfg.naiveswiglu,
|
||||
subln= vision_cfg.subln
|
||||
)
|
||||
elif vision_cfg.timm_model_name:
|
||||
visual = TimmModel(
|
||||
vision_cfg.timm_model_name,
|
||||
pretrained=vision_cfg.timm_model_pretrained,
|
||||
pool=vision_cfg.timm_pool,
|
||||
proj=vision_cfg.timm_proj,
|
||||
proj_bias=vision_cfg.timm_proj_bias,
|
||||
embed_dim=embed_dim,
|
||||
image_size=vision_cfg.image_size
|
||||
)
|
||||
act_layer = nn.GELU # so that text transformer doesn't use QuickGELU w/ timm models
|
||||
elif isinstance(vision_cfg.layers, (tuple, list)):
|
||||
vision_heads = vision_cfg.width * 32 // vision_cfg.head_width
|
||||
visual = ModifiedResNet(
|
||||
layers=vision_cfg.layers,
|
||||
output_dim=embed_dim,
|
||||
heads=vision_heads,
|
||||
image_size=vision_cfg.image_size,
|
||||
width=vision_cfg.width
|
||||
)
|
||||
else:
|
||||
vision_heads = vision_cfg.width // vision_cfg.head_width
|
||||
norm_layer = LayerNormFp32 if cast_dtype in (torch.float16, torch.bfloat16) else LayerNorm
|
||||
visual = VisionTransformer(
|
||||
image_size=vision_cfg.image_size,
|
||||
patch_size=vision_cfg.patch_size,
|
||||
width=vision_cfg.width,
|
||||
layers=vision_cfg.layers,
|
||||
heads=vision_heads,
|
||||
mlp_ratio=vision_cfg.mlp_ratio,
|
||||
ls_init_value=vision_cfg.ls_init_value,
|
||||
patch_dropout=vision_cfg.patch_dropout,
|
||||
global_average_pool=vision_cfg.global_average_pool,
|
||||
output_dim=embed_dim,
|
||||
act_layer=act_layer,
|
||||
norm_layer=norm_layer,
|
||||
)
|
||||
|
||||
return visual
|
||||
|
||||
|
||||
def _build_text_tower(
|
||||
embed_dim: int,
|
||||
text_cfg: CLIPTextCfg,
|
||||
quick_gelu: bool = False,
|
||||
cast_dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
if isinstance(text_cfg, dict):
|
||||
text_cfg = CLIPTextCfg(**text_cfg)
|
||||
|
||||
if text_cfg.hf_model_name:
|
||||
text = HFTextEncoder(
|
||||
text_cfg.hf_model_name,
|
||||
output_dim=embed_dim,
|
||||
tokenizer_name=text_cfg.hf_tokenizer_name,
|
||||
proj=text_cfg.proj,
|
||||
pooler_type=text_cfg.pooler_type,
|
||||
masked_language_modeling=text_cfg.masked_language_modeling
|
||||
)
|
||||
else:
|
||||
act_layer = QuickGELU if quick_gelu else nn.GELU
|
||||
norm_layer = LayerNorm
|
||||
|
||||
text = TextTransformer(
|
||||
context_length=text_cfg.context_length,
|
||||
vocab_size=text_cfg.vocab_size,
|
||||
width=text_cfg.width,
|
||||
heads=text_cfg.heads,
|
||||
layers=text_cfg.layers,
|
||||
ls_init_value=text_cfg.ls_init_value,
|
||||
output_dim=embed_dim,
|
||||
act_layer=act_layer,
|
||||
norm_layer= FusedLayerNorm if text_cfg.fusedLN else norm_layer,
|
||||
xattn=text_cfg.xattn,
|
||||
attn_mask=text_cfg.attn_mask,
|
||||
)
|
||||
return text
|
||||
|
||||
class CLIP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int,
|
||||
vision_cfg: CLIPVisionCfg,
|
||||
text_cfg: CLIPTextCfg,
|
||||
quick_gelu: bool = False,
|
||||
cast_dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.visual = _build_vision_tower(embed_dim, vision_cfg, quick_gelu, cast_dtype)
|
||||
|
||||
text = _build_text_tower(embed_dim, text_cfg, quick_gelu, cast_dtype)
|
||||
self.transformer = text.transformer
|
||||
self.vocab_size = text.vocab_size
|
||||
self.token_embedding = text.token_embedding
|
||||
self.positional_embedding = text.positional_embedding
|
||||
self.ln_final = text.ln_final
|
||||
self.text_projection = text.text_projection
|
||||
self.register_buffer('attn_mask', text.attn_mask, persistent=False)
|
||||
|
||||
self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
|
||||
|
||||
def lock_image_tower(self, unlocked_groups=0, freeze_bn_stats=False):
|
||||
# lock image tower as per LiT - https://arxiv.org/abs/2111.07991
|
||||
self.visual.lock(unlocked_groups=unlocked_groups, freeze_bn_stats=freeze_bn_stats)
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
self.visual.set_grad_checkpointing(enable)
|
||||
self.transformer.grad_checkpointing = enable
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'logit_scale'}
|
||||
|
||||
def encode_image(self, image, normalize: bool = False):
|
||||
features = self.visual(image)
|
||||
return F.normalize(features, dim=-1) if normalize else features
|
||||
|
||||
def encode_text(self, text, normalize: bool = False):
|
||||
cast_dtype = self.transformer.get_cast_dtype()
|
||||
|
||||
x = self.token_embedding(text).to(cast_dtype) # [batch_size, n_ctx, d_model]
|
||||
|
||||
x = x + self.positional_embedding.to(cast_dtype)
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.transformer(x, attn_mask=self.attn_mask)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.ln_final(x) # [batch_size, n_ctx, transformer.width]
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection
|
||||
return F.normalize(x, dim=-1) if normalize else x
|
||||
|
||||
def forward(self, image, text):
|
||||
image_features = self.encode_image(image, normalize=True)
|
||||
text_features = self.encode_text(text, normalize=True)
|
||||
return image_features, text_features, self.logit_scale.exp()
|
||||
|
||||
|
||||
class CustomCLIP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int,
|
||||
vision_cfg: CLIPVisionCfg,
|
||||
text_cfg: CLIPTextCfg,
|
||||
quick_gelu: bool = False,
|
||||
cast_dtype: Optional[torch.dtype] = None,
|
||||
itm_task: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.visual = _build_vision_tower(embed_dim, vision_cfg, quick_gelu, cast_dtype)
|
||||
self.text = _build_text_tower(embed_dim, text_cfg, quick_gelu, cast_dtype)
|
||||
self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
|
||||
|
||||
def lock_image_tower(self, unlocked_groups=0, freeze_bn_stats=False):
|
||||
# lock image tower as per LiT - https://arxiv.org/abs/2111.07991
|
||||
self.visual.lock(unlocked_groups=unlocked_groups, freeze_bn_stats=freeze_bn_stats)
|
||||
|
||||
def lock_text_tower(self, unlocked_layers:int=0, freeze_layer_norm:bool=True):
|
||||
self.text.lock(unlocked_layers, freeze_layer_norm)
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
self.visual.set_grad_checkpointing(enable)
|
||||
self.text.set_grad_checkpointing(enable)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'logit_scale'}
|
||||
|
||||
def encode_image(self, image, normalize: bool = False):
|
||||
features = self.visual(image)
|
||||
return F.normalize(features, dim=-1) if normalize else features
|
||||
|
||||
def encode_text(self, text, normalize: bool = False):
|
||||
features = self.text(text)
|
||||
return F.normalize(features, dim=-1) if normalize else features
|
||||
|
||||
def forward(self, image, text):
|
||||
image_features = self.encode_image(image, normalize=True)
|
||||
text_features = self.encode_text(text, normalize=True)
|
||||
return image_features, text_features, self.logit_scale.exp()
|
||||
|
||||
|
||||
def convert_weights_to_lp(model: nn.Module, dtype=torch.float16):
|
||||
"""Convert applicable model parameters to low-precision (bf16 or fp16)"""
|
||||
|
||||
def _convert_weights(l):
|
||||
|
||||
if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Linear)):
|
||||
l.weight.data = l.weight.data.to(dtype)
|
||||
if l.bias is not None:
|
||||
l.bias.data = l.bias.data.to(dtype)
|
||||
|
||||
if isinstance(l, (nn.MultiheadAttention, Attention)):
|
||||
for attr in [*[f"{s}_proj_weight" for s in ["in", "q", "k", "v"]], "in_proj_bias", "bias_k", "bias_v"]:
|
||||
tensor = getattr(l, attr, None)
|
||||
if tensor is not None:
|
||||
tensor.data = tensor.data.to(dtype)
|
||||
|
||||
if isinstance(l, nn.Parameter):
|
||||
l.data = l.data.to(dtype)
|
||||
|
||||
for name in ["text_projection", "proj"]:
|
||||
if hasattr(l, name) and isinstance(l, nn.Parameter):
|
||||
attr = getattr(l, name, None)
|
||||
if attr is not None:
|
||||
attr.data = attr.data.to(dtype)
|
||||
|
||||
model.apply(_convert_weights)
|
||||
|
||||
|
||||
convert_weights_to_fp16 = convert_weights_to_lp # backwards compat
|
||||
|
||||
|
||||
# used to maintain checkpoint compatibility
|
||||
def convert_to_custom_text_state_dict(state_dict: dict):
|
||||
if 'text_projection' in state_dict:
|
||||
# old format state_dict, move text tower -> .text
|
||||
new_state_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if any(k.startswith(p) for p in (
|
||||
'text_projection',
|
||||
'positional_embedding',
|
||||
'token_embedding',
|
||||
'transformer',
|
||||
'ln_final',
|
||||
'logit_scale'
|
||||
)):
|
||||
k = 'text.' + k
|
||||
new_state_dict[k] = v
|
||||
return new_state_dict
|
||||
return state_dict
|
||||
|
||||
|
||||
def build_model_from_openai_state_dict(
|
||||
state_dict: dict,
|
||||
quick_gelu=True,
|
||||
cast_dtype=torch.float16,
|
||||
):
|
||||
vit = "visual.proj" in state_dict
|
||||
|
||||
if vit:
|
||||
vision_width = state_dict["visual.conv1.weight"].shape[0]
|
||||
vision_layers = len(
|
||||
[k for k in state_dict.keys() if k.startswith("visual.") and k.endswith(".attn.in_proj_weight")])
|
||||
vision_patch_size = state_dict["visual.conv1.weight"].shape[-1]
|
||||
grid_size = round((state_dict["visual.positional_embedding"].shape[0] - 1) ** 0.5)
|
||||
image_size = vision_patch_size * grid_size
|
||||
else:
|
||||
counts: list = [
|
||||
len(set(k.split(".")[2] for k in state_dict if k.startswith(f"visual.layer{b}"))) for b in [1, 2, 3, 4]]
|
||||
vision_layers = tuple(counts)
|
||||
vision_width = state_dict["visual.layer1.0.conv1.weight"].shape[0]
|
||||
output_width = round((state_dict["visual.attnpool.positional_embedding"].shape[0] - 1) ** 0.5)
|
||||
vision_patch_size = None
|
||||
assert output_width ** 2 + 1 == state_dict["visual.attnpool.positional_embedding"].shape[0]
|
||||
image_size = output_width * 32
|
||||
|
||||
embed_dim = state_dict["text_projection"].shape[1]
|
||||
context_length = state_dict["positional_embedding"].shape[0]
|
||||
vocab_size = state_dict["token_embedding.weight"].shape[0]
|
||||
transformer_width = state_dict["ln_final.weight"].shape[0]
|
||||
transformer_heads = transformer_width // 64
|
||||
transformer_layers = len(set(k.split(".")[2] for k in state_dict if k.startswith(f"transformer.resblocks")))
|
||||
|
||||
vision_cfg = CLIPVisionCfg(
|
||||
layers=vision_layers,
|
||||
width=vision_width,
|
||||
patch_size=vision_patch_size,
|
||||
image_size=image_size,
|
||||
)
|
||||
text_cfg = CLIPTextCfg(
|
||||
context_length=context_length,
|
||||
vocab_size=vocab_size,
|
||||
width=transformer_width,
|
||||
heads=transformer_heads,
|
||||
layers=transformer_layers
|
||||
)
|
||||
model = CLIP(
|
||||
embed_dim,
|
||||
vision_cfg=vision_cfg,
|
||||
text_cfg=text_cfg,
|
||||
quick_gelu=quick_gelu, # OpenAI models were trained with QuickGELU
|
||||
cast_dtype=cast_dtype,
|
||||
)
|
||||
|
||||
for key in ["input_resolution", "context_length", "vocab_size"]:
|
||||
state_dict.pop(key, None)
|
||||
|
||||
convert_weights_to_fp16(model) # OpenAI state dicts are partially converted to float16
|
||||
model.load_state_dict(state_dict)
|
||||
return model.eval()
|
||||
|
||||
|
||||
def trace_model(model, batch_size=256, device=torch.device('cpu')):
|
||||
model.eval()
|
||||
image_size = model.visual.image_size
|
||||
example_images = torch.ones((batch_size, 3, image_size, image_size), device=device)
|
||||
example_text = torch.zeros((batch_size, model.context_length), dtype=torch.int, device=device)
|
||||
model = torch.jit.trace_module(
|
||||
model,
|
||||
inputs=dict(
|
||||
forward=(example_images, example_text),
|
||||
encode_text=(example_text,),
|
||||
encode_image=(example_images,)
|
||||
))
|
||||
model.visual.image_size = image_size
|
||||
return model
|
||||
@@ -0,0 +1,19 @@
|
||||
{
|
||||
"embed_dim": 512,
|
||||
"vision_cfg": {
|
||||
"image_size": 224,
|
||||
"layers": 12,
|
||||
"width": 768,
|
||||
"patch_size": 16,
|
||||
"eva_model_name": "eva-clip-b-16",
|
||||
"ls_init_value": 0.1,
|
||||
"drop_path_rate": 0.0
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 512,
|
||||
"heads": 8,
|
||||
"layers": 12
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
{
|
||||
"embed_dim": 1024,
|
||||
"vision_cfg": {
|
||||
"image_size": 224,
|
||||
"layers": 40,
|
||||
"width": 1408,
|
||||
"head_width": 88,
|
||||
"mlp_ratio": 4.3637,
|
||||
"patch_size": 14,
|
||||
"eva_model_name": "eva-clip-g-14-x",
|
||||
"drop_path_rate": 0,
|
||||
"xattn": true,
|
||||
"fusedLN": true
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 1024,
|
||||
"heads": 16,
|
||||
"layers": 24,
|
||||
"xattn": false,
|
||||
"fusedLN": true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
{
|
||||
"embed_dim": 1024,
|
||||
"vision_cfg": {
|
||||
"image_size": 224,
|
||||
"layers": 40,
|
||||
"width": 1408,
|
||||
"head_width": 88,
|
||||
"mlp_ratio": 4.3637,
|
||||
"patch_size": 14,
|
||||
"eva_model_name": "eva-clip-g-14-x",
|
||||
"drop_path_rate": 0.4,
|
||||
"xattn": true,
|
||||
"fusedLN": true
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 768,
|
||||
"heads": 12,
|
||||
"layers": 12,
|
||||
"xattn": false,
|
||||
"fusedLN": true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
{
|
||||
"embed_dim": 512,
|
||||
"vision_cfg": {
|
||||
"image_size": 224,
|
||||
"layers": 12,
|
||||
"width": 768,
|
||||
"head_width": 64,
|
||||
"patch_size": 16,
|
||||
"mlp_ratio": 2.6667,
|
||||
"eva_model_name": "eva-clip-b-16-X",
|
||||
"drop_path_rate": 0.0,
|
||||
"xattn": true,
|
||||
"fusedLN": true,
|
||||
"rope": true,
|
||||
"pt_hw_seq_len": 16,
|
||||
"intp_freq": true,
|
||||
"naiveswiglu": true,
|
||||
"subln": true
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 512,
|
||||
"heads": 8,
|
||||
"layers": 12,
|
||||
"xattn": true,
|
||||
"fusedLN": true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
{
|
||||
"embed_dim": 768,
|
||||
"vision_cfg": {
|
||||
"image_size": 336,
|
||||
"layers": 24,
|
||||
"width": 1024,
|
||||
"drop_path_rate": 0,
|
||||
"head_width": 64,
|
||||
"mlp_ratio": 2.6667,
|
||||
"patch_size": 14,
|
||||
"eva_model_name": "eva-clip-l-14-336",
|
||||
"xattn": true,
|
||||
"fusedLN": true,
|
||||
"rope": true,
|
||||
"pt_hw_seq_len": 16,
|
||||
"intp_freq": true,
|
||||
"naiveswiglu": true,
|
||||
"subln": true
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 768,
|
||||
"heads": 12,
|
||||
"layers": 12,
|
||||
"xattn": false,
|
||||
"fusedLN": true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
{
|
||||
"embed_dim": 768,
|
||||
"vision_cfg": {
|
||||
"image_size": 224,
|
||||
"layers": 24,
|
||||
"width": 1024,
|
||||
"drop_path_rate": 0,
|
||||
"head_width": 64,
|
||||
"mlp_ratio": 2.6667,
|
||||
"patch_size": 14,
|
||||
"eva_model_name": "eva-clip-l-14",
|
||||
"xattn": true,
|
||||
"fusedLN": true,
|
||||
"rope": true,
|
||||
"pt_hw_seq_len": 16,
|
||||
"intp_freq": true,
|
||||
"naiveswiglu": true,
|
||||
"subln": true
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 768,
|
||||
"heads": 12,
|
||||
"layers": 12,
|
||||
"xattn": false,
|
||||
"fusedLN": true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
{
|
||||
"embed_dim": 1024,
|
||||
"vision_cfg": {
|
||||
"image_size": 224,
|
||||
"layers": 64,
|
||||
"width": 1792,
|
||||
"head_width": 112,
|
||||
"mlp_ratio": 8.571428571428571,
|
||||
"patch_size": 14,
|
||||
"eva_model_name": "eva-clip-4b-14-x",
|
||||
"drop_path_rate": 0,
|
||||
"xattn": true,
|
||||
"postnorm": true,
|
||||
"fusedLN": true
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 1280,
|
||||
"heads": 20,
|
||||
"layers": 32,
|
||||
"xattn": false,
|
||||
"fusedLN": true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
{
|
||||
"embed_dim": 1024,
|
||||
"vision_cfg": {
|
||||
"image_size": 224,
|
||||
"layers": 64,
|
||||
"width": 1792,
|
||||
"head_width": 112,
|
||||
"mlp_ratio": 8.571428571428571,
|
||||
"patch_size": 14,
|
||||
"eva_model_name": "eva-clip-4b-14-x",
|
||||
"drop_path_rate": 0,
|
||||
"xattn": true,
|
||||
"postnorm": true,
|
||||
"fusedLN": true
|
||||
},
|
||||
"text_cfg": {
|
||||
"context_length": 77,
|
||||
"vocab_size": 49408,
|
||||
"width": 1024,
|
||||
"heads": 16,
|
||||
"layers": 24,
|
||||
"xattn": false,
|
||||
"fusedLN": true
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from eva_clip.utils import freeze_batch_norm_2d
|
||||
|
||||
|
||||
class Bottleneck(nn.Module):
|
||||
expansion = 4
|
||||
|
||||
def __init__(self, inplanes, planes, stride=1):
|
||||
super().__init__()
|
||||
|
||||
# all conv layers have stride 1. an avgpool is performed after the second convolution when stride > 1
|
||||
self.conv1 = nn.Conv2d(inplanes, planes, 1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(planes)
|
||||
self.act1 = nn.ReLU(inplace=True)
|
||||
|
||||
self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
self.act2 = nn.ReLU(inplace=True)
|
||||
|
||||
self.avgpool = nn.AvgPool2d(stride) if stride > 1 else nn.Identity()
|
||||
|
||||
self.conv3 = nn.Conv2d(planes, planes * self.expansion, 1, bias=False)
|
||||
self.bn3 = nn.BatchNorm2d(planes * self.expansion)
|
||||
self.act3 = nn.ReLU(inplace=True)
|
||||
|
||||
self.downsample = None
|
||||
self.stride = stride
|
||||
|
||||
if stride > 1 or inplanes != planes * Bottleneck.expansion:
|
||||
# downsampling layer is prepended with an avgpool, and the subsequent convolution has stride 1
|
||||
self.downsample = nn.Sequential(OrderedDict([
|
||||
("-1", nn.AvgPool2d(stride)),
|
||||
("0", nn.Conv2d(inplanes, planes * self.expansion, 1, stride=1, bias=False)),
|
||||
("1", nn.BatchNorm2d(planes * self.expansion))
|
||||
]))
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
identity = x
|
||||
|
||||
out = self.act1(self.bn1(self.conv1(x)))
|
||||
out = self.act2(self.bn2(self.conv2(out)))
|
||||
out = self.avgpool(out)
|
||||
out = self.bn3(self.conv3(out))
|
||||
|
||||
if self.downsample is not None:
|
||||
identity = self.downsample(x)
|
||||
|
||||
out += identity
|
||||
out = self.act3(out)
|
||||
return out
|
||||
|
||||
|
||||
class AttentionPool2d(nn.Module):
|
||||
def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None):
|
||||
super().__init__()
|
||||
self.positional_embedding = nn.Parameter(torch.randn(spacial_dim ** 2 + 1, embed_dim) / embed_dim ** 0.5)
|
||||
self.k_proj = nn.Linear(embed_dim, embed_dim)
|
||||
self.q_proj = nn.Linear(embed_dim, embed_dim)
|
||||
self.v_proj = nn.Linear(embed_dim, embed_dim)
|
||||
self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim)
|
||||
self.num_heads = num_heads
|
||||
|
||||
def forward(self, x):
|
||||
x = x.reshape(x.shape[0], x.shape[1], x.shape[2] * x.shape[3]).permute(2, 0, 1) # NCHW -> (HW)NC
|
||||
x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (HW+1)NC
|
||||
x = x + self.positional_embedding[:, None, :].to(x.dtype) # (HW+1)NC
|
||||
x, _ = F.multi_head_attention_forward(
|
||||
query=x, key=x, value=x,
|
||||
embed_dim_to_check=x.shape[-1],
|
||||
num_heads=self.num_heads,
|
||||
q_proj_weight=self.q_proj.weight,
|
||||
k_proj_weight=self.k_proj.weight,
|
||||
v_proj_weight=self.v_proj.weight,
|
||||
in_proj_weight=None,
|
||||
in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]),
|
||||
bias_k=None,
|
||||
bias_v=None,
|
||||
add_zero_attn=False,
|
||||
dropout_p=0.,
|
||||
out_proj_weight=self.c_proj.weight,
|
||||
out_proj_bias=self.c_proj.bias,
|
||||
use_separate_proj_weight=True,
|
||||
training=self.training,
|
||||
need_weights=False
|
||||
)
|
||||
|
||||
return x[0]
|
||||
|
||||
|
||||
class ModifiedResNet(nn.Module):
|
||||
"""
|
||||
A ResNet class that is similar to torchvision's but contains the following changes:
|
||||
- There are now 3 "stem" convolutions as opposed to 1, with an average pool instead of a max pool.
|
||||
- Performs anti-aliasing strided convolutions, where an avgpool is prepended to convolutions with stride > 1
|
||||
- The final pooling layer is a QKV attention instead of an average pool
|
||||
"""
|
||||
|
||||
def __init__(self, layers, output_dim, heads, image_size=224, width=64):
|
||||
super().__init__()
|
||||
self.output_dim = output_dim
|
||||
self.image_size = image_size
|
||||
|
||||
# the 3-layer stem
|
||||
self.conv1 = nn.Conv2d(3, width // 2, kernel_size=3, stride=2, padding=1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(width // 2)
|
||||
self.act1 = nn.ReLU(inplace=True)
|
||||
self.conv2 = nn.Conv2d(width // 2, width // 2, kernel_size=3, padding=1, bias=False)
|
||||
self.bn2 = nn.BatchNorm2d(width // 2)
|
||||
self.act2 = nn.ReLU(inplace=True)
|
||||
self.conv3 = nn.Conv2d(width // 2, width, kernel_size=3, padding=1, bias=False)
|
||||
self.bn3 = nn.BatchNorm2d(width)
|
||||
self.act3 = nn.ReLU(inplace=True)
|
||||
self.avgpool = nn.AvgPool2d(2)
|
||||
|
||||
# residual layers
|
||||
self._inplanes = width # this is a *mutable* variable used during construction
|
||||
self.layer1 = self._make_layer(width, layers[0])
|
||||
self.layer2 = self._make_layer(width * 2, layers[1], stride=2)
|
||||
self.layer3 = self._make_layer(width * 4, layers[2], stride=2)
|
||||
self.layer4 = self._make_layer(width * 8, layers[3], stride=2)
|
||||
|
||||
embed_dim = width * 32 # the ResNet feature dimension
|
||||
self.attnpool = AttentionPool2d(image_size // 32, embed_dim, heads, output_dim)
|
||||
|
||||
self.init_parameters()
|
||||
|
||||
def _make_layer(self, planes, blocks, stride=1):
|
||||
layers = [Bottleneck(self._inplanes, planes, stride)]
|
||||
|
||||
self._inplanes = planes * Bottleneck.expansion
|
||||
for _ in range(1, blocks):
|
||||
layers.append(Bottleneck(self._inplanes, planes))
|
||||
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
def init_parameters(self):
|
||||
if self.attnpool is not None:
|
||||
std = self.attnpool.c_proj.in_features ** -0.5
|
||||
nn.init.normal_(self.attnpool.q_proj.weight, std=std)
|
||||
nn.init.normal_(self.attnpool.k_proj.weight, std=std)
|
||||
nn.init.normal_(self.attnpool.v_proj.weight, std=std)
|
||||
nn.init.normal_(self.attnpool.c_proj.weight, std=std)
|
||||
|
||||
for resnet_block in [self.layer1, self.layer2, self.layer3, self.layer4]:
|
||||
for name, param in resnet_block.named_parameters():
|
||||
if name.endswith("bn3.weight"):
|
||||
nn.init.zeros_(param)
|
||||
|
||||
def lock(self, unlocked_groups=0, freeze_bn_stats=False):
|
||||
assert unlocked_groups == 0, 'partial locking not currently supported for this model'
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
if freeze_bn_stats:
|
||||
freeze_batch_norm_2d(self)
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
# FIXME support for non-transformer
|
||||
pass
|
||||
|
||||
def stem(self, x):
|
||||
x = self.act1(self.bn1(self.conv1(x)))
|
||||
x = self.act2(self.bn2(self.conv2(x)))
|
||||
x = self.act3(self.bn3(self.conv3(x)))
|
||||
x = self.avgpool(x)
|
||||
return x
|
||||
|
||||
def forward(self, x):
|
||||
x = self.stem(x)
|
||||
x = self.layer1(x)
|
||||
x = self.layer2(x)
|
||||
x = self.layer3(x)
|
||||
x = self.layer4(x)
|
||||
x = self.attnpool(x)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,144 @@
|
||||
""" OpenAI pretrained model functions
|
||||
|
||||
Adapted from https://github.com/openai/CLIP. Originally MIT License, Copyright (c) 2021 OpenAI.
|
||||
"""
|
||||
|
||||
import os
|
||||
import warnings
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from .model import build_model_from_openai_state_dict, convert_weights_to_lp, get_cast_dtype
|
||||
from .pretrained import get_pretrained_url, list_pretrained_models_by_tag, download_pretrained_from_url
|
||||
|
||||
__all__ = ["list_openai_models", "load_openai_model"]
|
||||
|
||||
|
||||
def list_openai_models() -> List[str]:
|
||||
"""Returns the names of available CLIP models"""
|
||||
return list_pretrained_models_by_tag('openai')
|
||||
|
||||
|
||||
def load_openai_model(
|
||||
name: str,
|
||||
precision: Optional[str] = None,
|
||||
device: Optional[Union[str, torch.device]] = None,
|
||||
jit: bool = True,
|
||||
cache_dir: Optional[str] = None,
|
||||
):
|
||||
"""Load a CLIP model
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : str
|
||||
A model name listed by `clip.available_models()`, or the path to a model checkpoint containing the state_dict
|
||||
precision: str
|
||||
Model precision, if None defaults to 'fp32' if device == 'cpu' else 'fp16'.
|
||||
device : Union[str, torch.device]
|
||||
The device to put the loaded model
|
||||
jit : bool
|
||||
Whether to load the optimized JIT model (default) or more hackable non-JIT model.
|
||||
cache_dir : Optional[str]
|
||||
The directory to cache the downloaded model weights
|
||||
|
||||
Returns
|
||||
-------
|
||||
model : torch.nn.Module
|
||||
The CLIP model
|
||||
preprocess : Callable[[PIL.Image], torch.Tensor]
|
||||
A torchvision transform that converts a PIL image into a tensor that the returned model can take as its input
|
||||
"""
|
||||
if device is None:
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if precision is None:
|
||||
precision = 'fp32' if device == 'cpu' else 'fp16'
|
||||
|
||||
if get_pretrained_url(name, 'openai'):
|
||||
model_path = download_pretrained_from_url(get_pretrained_url(name, 'openai'), cache_dir=cache_dir)
|
||||
elif os.path.isfile(name):
|
||||
model_path = name
|
||||
else:
|
||||
raise RuntimeError(f"Model {name} not found; available models = {list_openai_models()}")
|
||||
|
||||
try:
|
||||
# loading JIT archive
|
||||
model = torch.jit.load(model_path, map_location=device if jit else "cpu").eval()
|
||||
state_dict = None
|
||||
except RuntimeError:
|
||||
# loading saved state dict
|
||||
if jit:
|
||||
warnings.warn(f"File {model_path} is not a JIT archive. Loading as a state dict instead")
|
||||
jit = False
|
||||
state_dict = torch.load(model_path, map_location="cpu")
|
||||
|
||||
if not jit:
|
||||
# Build a non-jit model from the OpenAI jitted model state dict
|
||||
cast_dtype = get_cast_dtype(precision)
|
||||
try:
|
||||
model = build_model_from_openai_state_dict(state_dict or model.state_dict(), cast_dtype=cast_dtype)
|
||||
except KeyError:
|
||||
sd = {k[7:]: v for k, v in state_dict["state_dict"].items()}
|
||||
model = build_model_from_openai_state_dict(sd, cast_dtype=cast_dtype)
|
||||
|
||||
# model from OpenAI state dict is in manually cast fp16 mode, must be converted for AMP/fp32/bf16 use
|
||||
model = model.to(device)
|
||||
if precision.startswith('amp') or precision == 'fp32':
|
||||
model.float()
|
||||
elif precision == 'bf16':
|
||||
convert_weights_to_lp(model, dtype=torch.bfloat16)
|
||||
|
||||
return model
|
||||
|
||||
# patch the device names
|
||||
device_holder = torch.jit.trace(lambda: torch.ones([]).to(torch.device(device)), example_inputs=[])
|
||||
device_node = [n for n in device_holder.graph.findAllNodes("prim::Constant") if "Device" in repr(n)][-1]
|
||||
|
||||
def patch_device(module):
|
||||
try:
|
||||
graphs = [module.graph] if hasattr(module, "graph") else []
|
||||
except RuntimeError:
|
||||
graphs = []
|
||||
|
||||
if hasattr(module, "forward1"):
|
||||
graphs.append(module.forward1.graph)
|
||||
|
||||
for graph in graphs:
|
||||
for node in graph.findAllNodes("prim::Constant"):
|
||||
if "value" in node.attributeNames() and str(node["value"]).startswith("cuda"):
|
||||
node.copyAttributes(device_node)
|
||||
|
||||
model.apply(patch_device)
|
||||
patch_device(model.encode_image)
|
||||
patch_device(model.encode_text)
|
||||
|
||||
# patch dtype to float32 (typically for CPU)
|
||||
if precision == 'fp32':
|
||||
float_holder = torch.jit.trace(lambda: torch.ones([]).float(), example_inputs=[])
|
||||
float_input = list(float_holder.graph.findNode("aten::to").inputs())[1]
|
||||
float_node = float_input.node()
|
||||
|
||||
def patch_float(module):
|
||||
try:
|
||||
graphs = [module.graph] if hasattr(module, "graph") else []
|
||||
except RuntimeError:
|
||||
graphs = []
|
||||
|
||||
if hasattr(module, "forward1"):
|
||||
graphs.append(module.forward1.graph)
|
||||
|
||||
for graph in graphs:
|
||||
for node in graph.findAllNodes("aten::to"):
|
||||
inputs = list(node.inputs())
|
||||
for i in [1, 2]: # dtype can be the second or third argument to aten::to()
|
||||
if inputs[i].node()["value"] == 5:
|
||||
inputs[i].node().copyAttributes(float_node)
|
||||
|
||||
model.apply(patch_float)
|
||||
patch_float(model.encode_image)
|
||||
patch_float(model.encode_text)
|
||||
model.float()
|
||||
|
||||
# ensure image_size attr available at consistent location for both jit and non-jit
|
||||
model.visual.image_size = model.input_resolution.item()
|
||||
return model
|
||||
@@ -0,0 +1,331 @@
|
||||
import hashlib
|
||||
import os
|
||||
import urllib
|
||||
import warnings
|
||||
from typing import Dict, Union
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
try:
|
||||
from huggingface_hub import hf_hub_download
|
||||
_has_hf_hub = True
|
||||
except ImportError:
|
||||
hf_hub_download = None
|
||||
_has_hf_hub = False
|
||||
|
||||
|
||||
def _pcfg(url='', hf_hub='', filename='', mean=None, std=None):
|
||||
return dict(
|
||||
url=url,
|
||||
hf_hub=hf_hub,
|
||||
mean=mean,
|
||||
std=std,
|
||||
)
|
||||
|
||||
_VITB32 = dict(
|
||||
openai=_pcfg(
|
||||
"https://openaipublic.azureedge.net/clip/models/40d365715913c9da98579312b702a82c18be219cc2a73407c4526f58eba950af/ViT-B-32.pt"),
|
||||
laion400m_e31=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e31-d867053b.pt"),
|
||||
laion400m_e32=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e32-46683a32.pt"),
|
||||
laion2b_e16=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-laion2b_e16-af8dbd0c.pth"),
|
||||
laion2b_s34b_b79k=_pcfg(hf_hub='laion/CLIP-ViT-B-32-laion2B-s34B-b79K/')
|
||||
)
|
||||
|
||||
_VITB32_quickgelu = dict(
|
||||
openai=_pcfg(
|
||||
"https://openaipublic.azureedge.net/clip/models/40d365715913c9da98579312b702a82c18be219cc2a73407c4526f58eba950af/ViT-B-32.pt"),
|
||||
laion400m_e31=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e31-d867053b.pt"),
|
||||
laion400m_e32=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e32-46683a32.pt"),
|
||||
)
|
||||
|
||||
_VITB16 = dict(
|
||||
openai=_pcfg(
|
||||
"https://openaipublic.azureedge.net/clip/models/5806e77cd80f8b59890b7e101eabd078d9fb84e6937f9e85e4ecb61988df416f/ViT-B-16.pt"),
|
||||
laion400m_e31=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16-laion400m_e31-00efa78f.pt"),
|
||||
laion400m_e32=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16-laion400m_e32-55e67d44.pt"),
|
||||
laion2b_s34b_b88k=_pcfg(hf_hub='laion/CLIP-ViT-B-16-laion2B-s34B-b88K/'),
|
||||
)
|
||||
|
||||
_EVAB16 = dict(
|
||||
eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_B_psz14to16.pt'),
|
||||
eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_B_psz14to16.pt'),
|
||||
eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_B_psz16_s8B.pt'),
|
||||
eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_B_psz16_s8B.pt'),
|
||||
)
|
||||
|
||||
_VITB16_PLUS_240 = dict(
|
||||
laion400m_e31=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16_plus_240-laion400m_e31-8fb26589.pt"),
|
||||
laion400m_e32=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16_plus_240-laion400m_e32-699c4b84.pt"),
|
||||
)
|
||||
|
||||
_VITL14 = dict(
|
||||
openai=_pcfg(
|
||||
"https://openaipublic.azureedge.net/clip/models/b8cca3fd41ae0c99ba7e8951adf17d267cdb84cd88be6f7c2e0eca1737a03836/ViT-L-14.pt"),
|
||||
laion400m_e31=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_l_14-laion400m_e31-69988bb6.pt"),
|
||||
laion400m_e32=_pcfg(
|
||||
"https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_l_14-laion400m_e32-3d133497.pt"),
|
||||
laion2b_s32b_b82k=_pcfg(
|
||||
hf_hub='laion/CLIP-ViT-L-14-laion2B-s32B-b82K/',
|
||||
mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),
|
||||
)
|
||||
|
||||
_EVAL14 = dict(
|
||||
eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_L_psz14.pt'),
|
||||
eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_L_psz14.pt'),
|
||||
eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_s4B.pt'),
|
||||
eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_s4B.pt'),
|
||||
)
|
||||
|
||||
_VITL14_336 = dict(
|
||||
openai=_pcfg(
|
||||
"https://openaipublic.azureedge.net/clip/models/3035c92b350959924f9f00213499208652fc7ea050643e8b385c2dac08641f02/ViT-L-14-336px.pt"),
|
||||
)
|
||||
|
||||
_EVAL14_336 = dict(
|
||||
eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_336_psz14_s6B.pt'),
|
||||
eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_336_psz14_s6B.pt'),
|
||||
eva_clip_224to336=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_224to336.pt'),
|
||||
eva02_clip_224to336=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_224to336.pt'),
|
||||
)
|
||||
|
||||
_VITH14 = dict(
|
||||
laion2b_s32b_b79k=_pcfg(hf_hub='laion/CLIP-ViT-H-14-laion2B-s32B-b79K/'),
|
||||
)
|
||||
|
||||
_VITg14 = dict(
|
||||
laion2b_s12b_b42k=_pcfg(hf_hub='laion/CLIP-ViT-g-14-laion2B-s12B-b42K/'),
|
||||
laion2b_s34b_b88k=_pcfg(hf_hub='laion/CLIP-ViT-g-14-laion2B-s34B-b88K/'),
|
||||
)
|
||||
|
||||
_EVAg14 = dict(
|
||||
eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/'),
|
||||
eva01=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_g_psz14.pt'),
|
||||
eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_psz14_s11B.pt'),
|
||||
eva01_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_psz14_s11B.pt'),
|
||||
)
|
||||
|
||||
_EVAg14_PLUS = dict(
|
||||
eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/'),
|
||||
eva01=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_g_psz14.pt'),
|
||||
eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_plus_psz14_s11B.pt'),
|
||||
eva01_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_plus_psz14_s11B.pt'),
|
||||
)
|
||||
|
||||
_VITbigG14 = dict(
|
||||
laion2b_s39b_b160k=_pcfg(hf_hub='laion/CLIP-ViT-bigG-14-laion2B-39B-b160k/'),
|
||||
)
|
||||
|
||||
_EVAbigE14 = dict(
|
||||
eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'),
|
||||
eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'),
|
||||
eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_s4B.pt'),
|
||||
eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_s4B.pt'),
|
||||
)
|
||||
|
||||
_EVAbigE14_PLUS = dict(
|
||||
eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'),
|
||||
eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'),
|
||||
eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_plus_s9B.pt'),
|
||||
eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_plus_s9B.pt'),
|
||||
)
|
||||
|
||||
|
||||
_PRETRAINED = {
|
||||
# "ViT-B-32": _VITB32,
|
||||
"OpenaiCLIP-B-32": _VITB32,
|
||||
"OpenCLIP-B-32": _VITB32,
|
||||
|
||||
# "ViT-B-32-quickgelu": _VITB32_quickgelu,
|
||||
"OpenaiCLIP-B-32-quickgelu": _VITB32_quickgelu,
|
||||
"OpenCLIP-B-32-quickgelu": _VITB32_quickgelu,
|
||||
|
||||
# "ViT-B-16": _VITB16,
|
||||
"OpenaiCLIP-B-16": _VITB16,
|
||||
"OpenCLIP-B-16": _VITB16,
|
||||
|
||||
"EVA02-B-16": _EVAB16,
|
||||
"EVA02-CLIP-B-16": _EVAB16,
|
||||
|
||||
# "ViT-B-16-plus-240": _VITB16_PLUS_240,
|
||||
"OpenCLIP-B-16-plus-240": _VITB16_PLUS_240,
|
||||
|
||||
# "ViT-L-14": _VITL14,
|
||||
"OpenaiCLIP-L-14": _VITL14,
|
||||
"OpenCLIP-L-14": _VITL14,
|
||||
|
||||
"EVA02-L-14": _EVAL14,
|
||||
"EVA02-CLIP-L-14": _EVAL14,
|
||||
|
||||
# "ViT-L-14-336": _VITL14_336,
|
||||
"OpenaiCLIP-L-14-336": _VITL14_336,
|
||||
|
||||
"EVA02-CLIP-L-14-336": _EVAL14_336,
|
||||
|
||||
# "ViT-H-14": _VITH14,
|
||||
# "ViT-g-14": _VITg14,
|
||||
"OpenCLIP-H-14": _VITH14,
|
||||
"OpenCLIP-g-14": _VITg14,
|
||||
|
||||
"EVA01-CLIP-g-14": _EVAg14,
|
||||
"EVA01-CLIP-g-14-plus": _EVAg14_PLUS,
|
||||
|
||||
# "ViT-bigG-14": _VITbigG14,
|
||||
"OpenCLIP-bigG-14": _VITbigG14,
|
||||
|
||||
"EVA02-CLIP-bigE-14": _EVAbigE14,
|
||||
"EVA02-CLIP-bigE-14-plus": _EVAbigE14_PLUS,
|
||||
}
|
||||
|
||||
|
||||
def _clean_tag(tag: str):
|
||||
# normalize pretrained tags
|
||||
return tag.lower().replace('-', '_')
|
||||
|
||||
|
||||
def list_pretrained(as_str: bool = False):
|
||||
""" returns list of pretrained models
|
||||
Returns a tuple (model_name, pretrain_tag) by default or 'name:tag' if as_str == True
|
||||
"""
|
||||
return [':'.join([k, t]) if as_str else (k, t) for k in _PRETRAINED.keys() for t in _PRETRAINED[k].keys()]
|
||||
|
||||
|
||||
def list_pretrained_models_by_tag(tag: str):
|
||||
""" return all models having the specified pretrain tag """
|
||||
models = []
|
||||
tag = _clean_tag(tag)
|
||||
for k in _PRETRAINED.keys():
|
||||
if tag in _PRETRAINED[k]:
|
||||
models.append(k)
|
||||
return models
|
||||
|
||||
|
||||
def list_pretrained_tags_by_model(model: str):
|
||||
""" return all pretrain tags for the specified model architecture """
|
||||
tags = []
|
||||
if model in _PRETRAINED:
|
||||
tags.extend(_PRETRAINED[model].keys())
|
||||
return tags
|
||||
|
||||
|
||||
def is_pretrained_cfg(model: str, tag: str):
|
||||
if model not in _PRETRAINED:
|
||||
return False
|
||||
return _clean_tag(tag) in _PRETRAINED[model]
|
||||
|
||||
|
||||
def get_pretrained_cfg(model: str, tag: str):
|
||||
if model not in _PRETRAINED:
|
||||
return {}
|
||||
model_pretrained = _PRETRAINED[model]
|
||||
return model_pretrained.get(_clean_tag(tag), {})
|
||||
|
||||
|
||||
def get_pretrained_url(model: str, tag: str):
|
||||
cfg = get_pretrained_cfg(model, _clean_tag(tag))
|
||||
return cfg.get('url', '')
|
||||
|
||||
|
||||
def download_pretrained_from_url(
|
||||
url: str,
|
||||
cache_dir: Union[str, None] = None,
|
||||
):
|
||||
if not cache_dir:
|
||||
cache_dir = os.path.expanduser("~/.cache/clip")
|
||||
os.makedirs(cache_dir, exist_ok=True)
|
||||
filename = os.path.basename(url)
|
||||
|
||||
if 'openaipublic' in url:
|
||||
expected_sha256 = url.split("/")[-2]
|
||||
elif 'mlfoundations' in url:
|
||||
expected_sha256 = os.path.splitext(filename)[0].split("-")[-1]
|
||||
else:
|
||||
expected_sha256 = ''
|
||||
|
||||
download_target = os.path.join(cache_dir, filename)
|
||||
|
||||
if os.path.exists(download_target) and not os.path.isfile(download_target):
|
||||
raise RuntimeError(f"{download_target} exists and is not a regular file")
|
||||
|
||||
if os.path.isfile(download_target):
|
||||
if expected_sha256:
|
||||
if hashlib.sha256(open(download_target, "rb").read()).hexdigest().startswith(expected_sha256):
|
||||
return download_target
|
||||
else:
|
||||
warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file")
|
||||
else:
|
||||
return download_target
|
||||
|
||||
with urllib.request.urlopen(url) as source, open(download_target, "wb") as output:
|
||||
with tqdm(total=int(source.headers.get("Content-Length")), ncols=80, unit='iB', unit_scale=True) as loop:
|
||||
while True:
|
||||
buffer = source.read(8192)
|
||||
if not buffer:
|
||||
break
|
||||
|
||||
output.write(buffer)
|
||||
loop.update(len(buffer))
|
||||
|
||||
if expected_sha256 and not hashlib.sha256(open(download_target, "rb").read()).hexdigest().startswith(expected_sha256):
|
||||
raise RuntimeError("Model has been downloaded but the SHA256 checksum does not not match")
|
||||
|
||||
return download_target
|
||||
|
||||
|
||||
def has_hf_hub(necessary=False):
|
||||
if not _has_hf_hub and necessary:
|
||||
# if no HF Hub module installed, and it is necessary to continue, raise error
|
||||
raise RuntimeError(
|
||||
'Hugging Face hub model specified but package not installed. Run `pip install huggingface_hub`.')
|
||||
return _has_hf_hub
|
||||
|
||||
|
||||
def download_pretrained_from_hf(
|
||||
model_id: str,
|
||||
filename: str = 'open_clip_pytorch_model.bin',
|
||||
revision=None,
|
||||
cache_dir: Union[str, None] = None,
|
||||
):
|
||||
has_hf_hub(True)
|
||||
cached_file = hf_hub_download(model_id, filename, revision=revision, cache_dir=cache_dir)
|
||||
return cached_file
|
||||
|
||||
|
||||
def download_pretrained(
|
||||
cfg: Dict,
|
||||
force_hf_hub: bool = False,
|
||||
cache_dir: Union[str, None] = None,
|
||||
):
|
||||
target = ''
|
||||
if not cfg:
|
||||
return target
|
||||
|
||||
download_url = cfg.get('url', '')
|
||||
download_hf_hub = cfg.get('hf_hub', '')
|
||||
if download_hf_hub and force_hf_hub:
|
||||
# use HF hub even if url exists
|
||||
download_url = ''
|
||||
|
||||
if download_url:
|
||||
target = download_pretrained_from_url(download_url, cache_dir=cache_dir)
|
||||
elif download_hf_hub:
|
||||
has_hf_hub(True)
|
||||
# we assume the hf_hub entries in pretrained config combine model_id + filename in
|
||||
# 'org/model_name/filename.pt' form. To specify just the model id w/o filename and
|
||||
# use 'open_clip_pytorch_model.bin' default, there must be a trailing slash 'org/model_name/'.
|
||||
model_id, filename = os.path.split(download_hf_hub)
|
||||
if filename:
|
||||
target = download_pretrained_from_hf(model_id, filename=filename, cache_dir=cache_dir)
|
||||
else:
|
||||
target = download_pretrained_from_hf(model_id, cache_dir=cache_dir)
|
||||
|
||||
return target
|
||||
@@ -0,0 +1,137 @@
|
||||
from math import pi
|
||||
import torch
|
||||
from torch import nn
|
||||
from einops import rearrange, repeat
|
||||
import logging
|
||||
|
||||
def broadcat(tensors, dim = -1):
|
||||
num_tensors = len(tensors)
|
||||
shape_lens = set(list(map(lambda t: len(t.shape), tensors)))
|
||||
assert len(shape_lens) == 1, 'tensors must all have the same number of dimensions'
|
||||
shape_len = list(shape_lens)[0]
|
||||
dim = (dim + shape_len) if dim < 0 else dim
|
||||
dims = list(zip(*map(lambda t: list(t.shape), tensors)))
|
||||
expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim]
|
||||
assert all([*map(lambda t: len(set(t[1])) <= 2, expandable_dims)]), 'invalid dimensions for broadcastable concatentation'
|
||||
max_dims = list(map(lambda t: (t[0], max(t[1])), expandable_dims))
|
||||
expanded_dims = list(map(lambda t: (t[0], (t[1],) * num_tensors), max_dims))
|
||||
expanded_dims.insert(dim, (dim, dims[dim]))
|
||||
expandable_shapes = list(zip(*map(lambda t: t[1], expanded_dims)))
|
||||
tensors = list(map(lambda t: t[0].expand(*t[1]), zip(tensors, expandable_shapes)))
|
||||
return torch.cat(tensors, dim = dim)
|
||||
|
||||
def rotate_half(x):
|
||||
x = rearrange(x, '... (d r) -> ... d r', r = 2)
|
||||
x1, x2 = x.unbind(dim = -1)
|
||||
x = torch.stack((-x2, x1), dim = -1)
|
||||
return rearrange(x, '... d r -> ... (d r)')
|
||||
|
||||
|
||||
class VisionRotaryEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
pt_seq_len,
|
||||
ft_seq_len=None,
|
||||
custom_freqs = None,
|
||||
freqs_for = 'lang',
|
||||
theta = 10000,
|
||||
max_freq = 10,
|
||||
num_freqs = 1,
|
||||
):
|
||||
super().__init__()
|
||||
if custom_freqs:
|
||||
freqs = custom_freqs
|
||||
elif freqs_for == 'lang':
|
||||
freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim))
|
||||
elif freqs_for == 'pixel':
|
||||
freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi
|
||||
elif freqs_for == 'constant':
|
||||
freqs = torch.ones(num_freqs).float()
|
||||
else:
|
||||
raise ValueError(f'unknown modality {freqs_for}')
|
||||
|
||||
if ft_seq_len is None: ft_seq_len = pt_seq_len
|
||||
t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len
|
||||
|
||||
freqs_h = torch.einsum('..., f -> ... f', t, freqs)
|
||||
freqs_h = repeat(freqs_h, '... n -> ... (n r)', r = 2)
|
||||
|
||||
freqs_w = torch.einsum('..., f -> ... f', t, freqs)
|
||||
freqs_w = repeat(freqs_w, '... n -> ... (n r)', r = 2)
|
||||
|
||||
freqs = broadcat((freqs_h[:, None, :], freqs_w[None, :, :]), dim = -1)
|
||||
|
||||
self.register_buffer("freqs_cos", freqs.cos())
|
||||
self.register_buffer("freqs_sin", freqs.sin())
|
||||
|
||||
logging.info(f'Shape of rope freq: {self.freqs_cos.shape}')
|
||||
|
||||
def forward(self, t, start_index = 0):
|
||||
rot_dim = self.freqs_cos.shape[-1]
|
||||
end_index = start_index + rot_dim
|
||||
assert rot_dim <= t.shape[-1], f'feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}'
|
||||
t_left, t, t_right = t[..., :start_index], t[..., start_index:end_index], t[..., end_index:]
|
||||
t = (t * self.freqs_cos) + (rotate_half(t) * self.freqs_sin)
|
||||
|
||||
return torch.cat((t_left, t, t_right), dim = -1)
|
||||
|
||||
class VisionRotaryEmbeddingFast(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
pt_seq_len,
|
||||
ft_seq_len=None,
|
||||
custom_freqs = None,
|
||||
freqs_for = 'lang',
|
||||
theta = 10000,
|
||||
max_freq = 10,
|
||||
num_freqs = 1,
|
||||
patch_dropout = 0.
|
||||
):
|
||||
super().__init__()
|
||||
if custom_freqs:
|
||||
freqs = custom_freqs
|
||||
elif freqs_for == 'lang':
|
||||
freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim))
|
||||
elif freqs_for == 'pixel':
|
||||
freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi
|
||||
elif freqs_for == 'constant':
|
||||
freqs = torch.ones(num_freqs).float()
|
||||
else:
|
||||
raise ValueError(f'unknown modality {freqs_for}')
|
||||
|
||||
if ft_seq_len is None: ft_seq_len = pt_seq_len
|
||||
t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len
|
||||
|
||||
freqs = torch.einsum('..., f -> ... f', t, freqs)
|
||||
freqs = repeat(freqs, '... n -> ... (n r)', r = 2)
|
||||
freqs = broadcat((freqs[:, None, :], freqs[None, :, :]), dim = -1)
|
||||
|
||||
freqs_cos = freqs.cos().view(-1, freqs.shape[-1])
|
||||
freqs_sin = freqs.sin().view(-1, freqs.shape[-1])
|
||||
|
||||
self.patch_dropout = patch_dropout
|
||||
|
||||
self.register_buffer("freqs_cos", freqs_cos)
|
||||
self.register_buffer("freqs_sin", freqs_sin)
|
||||
|
||||
logging.info(f'Shape of rope freq: {self.freqs_cos.shape}')
|
||||
|
||||
def forward(self, t, patch_indices_keep=None):
|
||||
if patch_indices_keep is not None:
|
||||
batch = t.size()[0]
|
||||
batch_indices = torch.arange(batch)
|
||||
batch_indices = batch_indices[..., None]
|
||||
|
||||
freqs_cos = repeat(self.freqs_cos, 'i j -> n i m j', n=t.shape[0], m=t.shape[1])
|
||||
freqs_sin = repeat(self.freqs_sin, 'i j -> n i m j', n=t.shape[0], m=t.shape[1])
|
||||
|
||||
freqs_cos = freqs_cos[batch_indices, patch_indices_keep]
|
||||
freqs_cos = rearrange(freqs_cos, 'n i m j -> n m i j')
|
||||
freqs_sin = freqs_sin[batch_indices, patch_indices_keep]
|
||||
freqs_sin = rearrange(freqs_sin, 'n i m j -> n m i j')
|
||||
|
||||
return t * freqs_cos + rotate_half(t) * freqs_sin
|
||||
|
||||
return t * self.freqs_cos + rotate_half(t) * self.freqs_sin
|
||||
@@ -0,0 +1,119 @@
|
||||
""" timm model adapter
|
||||
|
||||
Wraps timm (https://github.com/rwightman/pytorch-image-models) models for use as a vision tower in CLIP model.
|
||||
"""
|
||||
import logging
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
try:
|
||||
import timm
|
||||
from timm.models.layers import Mlp, to_2tuple
|
||||
try:
|
||||
# old timm imports < 0.8.1
|
||||
from timm.models.layers.attention_pool2d import RotAttentionPool2d
|
||||
from timm.models.layers.attention_pool2d import AttentionPool2d as AbsAttentionPool2d
|
||||
except ImportError:
|
||||
# new timm imports >= 0.8.1
|
||||
from timm.layers import RotAttentionPool2d
|
||||
from timm.layers import AttentionPool2d as AbsAttentionPool2d
|
||||
except ImportError:
|
||||
timm = None
|
||||
|
||||
from .utils import freeze_batch_norm_2d
|
||||
|
||||
|
||||
class TimmModel(nn.Module):
|
||||
""" timm model adapter
|
||||
# FIXME this adapter is a work in progress, may change in ways that break weight compat
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name,
|
||||
embed_dim,
|
||||
image_size=224,
|
||||
pool='avg',
|
||||
proj='linear',
|
||||
proj_bias=False,
|
||||
drop=0.,
|
||||
pretrained=False):
|
||||
super().__init__()
|
||||
|
||||
self.image_size = to_2tuple(image_size)
|
||||
self.trunk = timm.create_model(model_name, pretrained=pretrained)
|
||||
feat_size = self.trunk.default_cfg.get('pool_size', None)
|
||||
feature_ndim = 1 if not feat_size else 2
|
||||
if pool in ('abs_attn', 'rot_attn'):
|
||||
assert feature_ndim == 2
|
||||
# if attn pooling used, remove both classifier and default pool
|
||||
self.trunk.reset_classifier(0, global_pool='')
|
||||
else:
|
||||
# reset global pool if pool config set, otherwise leave as network default
|
||||
reset_kwargs = dict(global_pool=pool) if pool else {}
|
||||
self.trunk.reset_classifier(0, **reset_kwargs)
|
||||
prev_chs = self.trunk.num_features
|
||||
|
||||
head_layers = OrderedDict()
|
||||
if pool == 'abs_attn':
|
||||
head_layers['pool'] = AbsAttentionPool2d(prev_chs, feat_size=feat_size, out_features=embed_dim)
|
||||
prev_chs = embed_dim
|
||||
elif pool == 'rot_attn':
|
||||
head_layers['pool'] = RotAttentionPool2d(prev_chs, out_features=embed_dim)
|
||||
prev_chs = embed_dim
|
||||
else:
|
||||
assert proj, 'projection layer needed if non-attention pooling is used.'
|
||||
|
||||
# NOTE attention pool ends with a projection layer, so proj should usually be set to '' if such pooling is used
|
||||
if proj == 'linear':
|
||||
head_layers['drop'] = nn.Dropout(drop)
|
||||
head_layers['proj'] = nn.Linear(prev_chs, embed_dim, bias=proj_bias)
|
||||
elif proj == 'mlp':
|
||||
head_layers['mlp'] = Mlp(prev_chs, 2 * embed_dim, embed_dim, drop=drop, bias=(True, proj_bias))
|
||||
|
||||
self.head = nn.Sequential(head_layers)
|
||||
|
||||
def lock(self, unlocked_groups=0, freeze_bn_stats=False):
|
||||
""" lock modules
|
||||
Args:
|
||||
unlocked_groups (int): leave last n layer groups unlocked (default: 0)
|
||||
"""
|
||||
if not unlocked_groups:
|
||||
# lock full model
|
||||
for param in self.trunk.parameters():
|
||||
param.requires_grad = False
|
||||
if freeze_bn_stats:
|
||||
freeze_batch_norm_2d(self.trunk)
|
||||
else:
|
||||
# NOTE: partial freeze requires latest timm (master) branch and is subject to change
|
||||
try:
|
||||
# FIXME import here until API stable and in an official release
|
||||
from timm.models.helpers import group_parameters, group_modules
|
||||
except ImportError:
|
||||
raise RuntimeError('Please install latest timm `pip install git+https://github.com/rwightman/pytorch-image-models`')
|
||||
matcher = self.trunk.group_matcher()
|
||||
gparams = group_parameters(self.trunk, matcher)
|
||||
max_layer_id = max(gparams.keys())
|
||||
max_layer_id = max_layer_id - unlocked_groups
|
||||
for group_idx in range(max_layer_id + 1):
|
||||
group = gparams[group_idx]
|
||||
for param in group:
|
||||
self.trunk.get_parameter(param).requires_grad = False
|
||||
if freeze_bn_stats:
|
||||
gmodules = group_modules(self.trunk, matcher, reverse=True)
|
||||
gmodules = {k for k, v in gmodules.items() if v <= max_layer_id}
|
||||
freeze_batch_norm_2d(self.trunk, gmodules)
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
try:
|
||||
self.trunk.set_grad_checkpointing(enable)
|
||||
except Exception as e:
|
||||
logging.warning('grad checkpointing not supported for this timm image tower, continuing without...')
|
||||
|
||||
def forward(self, x):
|
||||
x = self.trunk(x)
|
||||
x = self.head(x)
|
||||
return x
|
||||
@@ -0,0 +1,201 @@
|
||||
""" CLIP tokenizer
|
||||
|
||||
Copied from https://github.com/openai/CLIP. Originally MIT License, Copyright (c) 2021 OpenAI.
|
||||
"""
|
||||
import gzip
|
||||
import html
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import Union, List
|
||||
|
||||
import ftfy
|
||||
import regex as re
|
||||
import torch
|
||||
|
||||
# https://stackoverflow.com/q/62691279
|
||||
import os
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def default_bpe():
|
||||
return os.path.join(os.path.dirname(os.path.abspath(__file__)), "bpe_simple_vocab_16e6.txt.gz")
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def bytes_to_unicode():
|
||||
"""
|
||||
Returns list of utf-8 byte and a corresponding list of unicode strings.
|
||||
The reversible bpe codes work on unicode strings.
|
||||
This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
|
||||
When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
|
||||
This is a signficant percentage of your normal, say, 32K bpe vocab.
|
||||
To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
|
||||
And avoids mapping to whitespace/control characters the bpe code barfs on.
|
||||
"""
|
||||
bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
|
||||
cs = bs[:]
|
||||
n = 0
|
||||
for b in range(2**8):
|
||||
if b not in bs:
|
||||
bs.append(b)
|
||||
cs.append(2**8+n)
|
||||
n += 1
|
||||
cs = [chr(n) for n in cs]
|
||||
return dict(zip(bs, cs))
|
||||
|
||||
|
||||
def get_pairs(word):
|
||||
"""Return set of symbol pairs in a word.
|
||||
Word is represented as tuple of symbols (symbols being variable-length strings).
|
||||
"""
|
||||
pairs = set()
|
||||
prev_char = word[0]
|
||||
for char in word[1:]:
|
||||
pairs.add((prev_char, char))
|
||||
prev_char = char
|
||||
return pairs
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
class SimpleTokenizer(object):
|
||||
def __init__(self, bpe_path: str = default_bpe(), special_tokens=None):
|
||||
self.byte_encoder = bytes_to_unicode()
|
||||
self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
|
||||
merges = gzip.open(bpe_path).read().decode("utf-8").split('\n')
|
||||
merges = merges[1:49152-256-2+1]
|
||||
merges = [tuple(merge.split()) for merge in merges]
|
||||
vocab = list(bytes_to_unicode().values())
|
||||
vocab = vocab + [v+'</w>' for v in vocab]
|
||||
for merge in merges:
|
||||
vocab.append(''.join(merge))
|
||||
if not special_tokens:
|
||||
special_tokens = ['<start_of_text>', '<end_of_text>']
|
||||
else:
|
||||
special_tokens = ['<start_of_text>', '<end_of_text>'] + special_tokens
|
||||
vocab.extend(special_tokens)
|
||||
self.encoder = dict(zip(vocab, range(len(vocab))))
|
||||
self.decoder = {v: k for k, v in self.encoder.items()}
|
||||
self.bpe_ranks = dict(zip(merges, range(len(merges))))
|
||||
self.cache = {t:t for t in special_tokens}
|
||||
special = "|".join(special_tokens)
|
||||
self.pat = re.compile(special + r"""|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", re.IGNORECASE)
|
||||
|
||||
self.vocab_size = len(self.encoder)
|
||||
self.all_special_ids = [self.encoder[t] for t in special_tokens]
|
||||
|
||||
def bpe(self, token):
|
||||
if token in self.cache:
|
||||
return self.cache[token]
|
||||
word = tuple(token[:-1]) + ( token[-1] + '</w>',)
|
||||
pairs = get_pairs(word)
|
||||
|
||||
if not pairs:
|
||||
return token+'</w>'
|
||||
|
||||
while True:
|
||||
bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
|
||||
if bigram not in self.bpe_ranks:
|
||||
break
|
||||
first, second = bigram
|
||||
new_word = []
|
||||
i = 0
|
||||
while i < len(word):
|
||||
try:
|
||||
j = word.index(first, i)
|
||||
new_word.extend(word[i:j])
|
||||
i = j
|
||||
except:
|
||||
new_word.extend(word[i:])
|
||||
break
|
||||
|
||||
if word[i] == first and i < len(word)-1 and word[i+1] == second:
|
||||
new_word.append(first+second)
|
||||
i += 2
|
||||
else:
|
||||
new_word.append(word[i])
|
||||
i += 1
|
||||
new_word = tuple(new_word)
|
||||
word = new_word
|
||||
if len(word) == 1:
|
||||
break
|
||||
else:
|
||||
pairs = get_pairs(word)
|
||||
word = ' '.join(word)
|
||||
self.cache[token] = word
|
||||
return word
|
||||
|
||||
def encode(self, text):
|
||||
bpe_tokens = []
|
||||
text = whitespace_clean(basic_clean(text)).lower()
|
||||
for token in re.findall(self.pat, text):
|
||||
token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
|
||||
bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
|
||||
return bpe_tokens
|
||||
|
||||
def decode(self, tokens):
|
||||
text = ''.join([self.decoder[token] for token in tokens])
|
||||
text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors="replace").replace('</w>', ' ')
|
||||
return text
|
||||
|
||||
|
||||
_tokenizer = SimpleTokenizer()
|
||||
|
||||
|
||||
def tokenize(texts: Union[str, List[str]], context_length: int = 77) -> torch.LongTensor:
|
||||
"""
|
||||
Returns the tokenized representation of given input string(s)
|
||||
|
||||
Parameters
|
||||
----------
|
||||
texts : Union[str, List[str]]
|
||||
An input string or a list of input strings to tokenize
|
||||
context_length : int
|
||||
The context length to use; all CLIP models use 77 as the context length
|
||||
|
||||
Returns
|
||||
-------
|
||||
A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length]
|
||||
"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
|
||||
sot_token = _tokenizer.encoder["<start_of_text>"]
|
||||
eot_token = _tokenizer.encoder["<end_of_text>"]
|
||||
all_tokens = [[sot_token] + _tokenizer.encode(text) + [eot_token] for text in texts]
|
||||
result = torch.zeros(len(all_tokens), context_length, dtype=torch.long)
|
||||
|
||||
for i, tokens in enumerate(all_tokens):
|
||||
if len(tokens) > context_length:
|
||||
tokens = tokens[:context_length] # Truncate
|
||||
tokens[-1] = eot_token
|
||||
result[i, :len(tokens)] = torch.tensor(tokens)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class HFTokenizer:
|
||||
"HuggingFace tokenizer wrapper"
|
||||
def __init__(self, tokenizer_name:str):
|
||||
from transformers import AutoTokenizer
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
|
||||
|
||||
def __call__(self, texts:Union[str, List[str]], context_length:int=77) -> torch.Tensor:
|
||||
# same cleaning as for default tokenizer, except lowercasing
|
||||
# adding lower (for case-sensitive tokenizers) will make it more robust but less sensitive to nuance
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
texts = [whitespace_clean(basic_clean(text)) for text in texts]
|
||||
input_ids = self.tokenizer(texts, return_tensors='pt', max_length=context_length, padding='max_length', truncation=True).input_ids
|
||||
return input_ids
|
||||
@@ -0,0 +1,103 @@
|
||||
from typing import Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchvision.transforms.functional as F
|
||||
|
||||
from torchvision.transforms import Normalize, Compose, RandomResizedCrop, InterpolationMode, ToTensor, Resize, \
|
||||
CenterCrop
|
||||
|
||||
from .constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD
|
||||
|
||||
|
||||
class ResizeMaxSize(nn.Module):
|
||||
|
||||
def __init__(self, max_size, interpolation=InterpolationMode.BICUBIC, fn='max', fill=0):
|
||||
super().__init__()
|
||||
if not isinstance(max_size, int):
|
||||
raise TypeError(f"Size should be int. Got {type(max_size)}")
|
||||
self.max_size = max_size
|
||||
self.interpolation = interpolation
|
||||
self.fn = min if fn == 'min' else min
|
||||
self.fill = fill
|
||||
|
||||
def forward(self, img):
|
||||
if isinstance(img, torch.Tensor):
|
||||
height, width = img.shape[:2]
|
||||
else:
|
||||
width, height = img.size
|
||||
scale = self.max_size / float(max(height, width))
|
||||
if scale != 1.0:
|
||||
new_size = tuple(round(dim * scale) for dim in (height, width))
|
||||
img = F.resize(img, new_size, self.interpolation)
|
||||
pad_h = self.max_size - new_size[0]
|
||||
pad_w = self.max_size - new_size[1]
|
||||
img = F.pad(img, padding=[pad_w//2, pad_h//2, pad_w - pad_w//2, pad_h - pad_h//2], fill=self.fill)
|
||||
return img
|
||||
|
||||
|
||||
def _convert_to_rgb(image):
|
||||
return image.convert('RGB')
|
||||
|
||||
|
||||
# class CatGen(nn.Module):
|
||||
# def __init__(self, num=4):
|
||||
# self.num = num
|
||||
# def mixgen_batch(image, text):
|
||||
# batch_size = image.shape[0]
|
||||
# index = np.random.permutation(batch_size)
|
||||
|
||||
# cat_images = []
|
||||
# for i in range(batch_size):
|
||||
# # image mixup
|
||||
# image[i,:] = lam * image[i,:] + (1 - lam) * image[index[i],:]
|
||||
# # text concat
|
||||
# text[i] = tokenizer((str(text[i]) + " " + str(text[index[i]])))[0]
|
||||
# text = torch.stack(text)
|
||||
# return image, text
|
||||
|
||||
|
||||
def image_transform(
|
||||
image_size: int,
|
||||
is_train: bool,
|
||||
mean: Optional[Tuple[float, ...]] = None,
|
||||
std: Optional[Tuple[float, ...]] = None,
|
||||
resize_longest_max: bool = False,
|
||||
fill_color: int = 0,
|
||||
):
|
||||
mean = mean or OPENAI_DATASET_MEAN
|
||||
if not isinstance(mean, (list, tuple)):
|
||||
mean = (mean,) * 3
|
||||
|
||||
std = std or OPENAI_DATASET_STD
|
||||
if not isinstance(std, (list, tuple)):
|
||||
std = (std,) * 3
|
||||
|
||||
if isinstance(image_size, (list, tuple)) and image_size[0] == image_size[1]:
|
||||
# for square size, pass size as int so that Resize() uses aspect preserving shortest edge
|
||||
image_size = image_size[0]
|
||||
|
||||
normalize = Normalize(mean=mean, std=std)
|
||||
if is_train:
|
||||
return Compose([
|
||||
RandomResizedCrop(image_size, scale=(0.9, 1.0), interpolation=InterpolationMode.BICUBIC),
|
||||
_convert_to_rgb,
|
||||
ToTensor(),
|
||||
normalize,
|
||||
])
|
||||
else:
|
||||
if resize_longest_max:
|
||||
transforms = [
|
||||
ResizeMaxSize(image_size, fill=fill_color)
|
||||
]
|
||||
else:
|
||||
transforms = [
|
||||
Resize(image_size, interpolation=InterpolationMode.BICUBIC),
|
||||
CenterCrop(image_size),
|
||||
]
|
||||
transforms.extend([
|
||||
_convert_to_rgb,
|
||||
ToTensor(),
|
||||
normalize,
|
||||
])
|
||||
return Compose(transforms)
|
||||
@@ -0,0 +1,721 @@
|
||||
import os
|
||||
import logging
|
||||
from collections import OrderedDict
|
||||
import math
|
||||
from typing import Callable, Optional, Sequence
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
try:
|
||||
from timm.models.layers import trunc_normal_
|
||||
except:
|
||||
from timm.layers import trunc_normal_
|
||||
|
||||
from .rope import VisionRotaryEmbedding, VisionRotaryEmbeddingFast
|
||||
from .utils import to_2tuple
|
||||
|
||||
|
||||
class LayerNormFp32(nn.LayerNorm):
|
||||
"""Subclass torch's LayerNorm to handle fp16 (by casting to float32 and back)."""
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
output = F.layer_norm(
|
||||
x.float(),
|
||||
self.normalized_shape,
|
||||
self.weight.float() if self.weight is not None else None,
|
||||
self.bias.float() if self.bias is not None else None,
|
||||
self.eps,
|
||||
)
|
||||
return output.type_as(x)
|
||||
|
||||
|
||||
class LayerNorm(nn.LayerNorm):
|
||||
"""Subclass torch's LayerNorm (with cast back to input dtype)."""
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
orig_type = x.dtype
|
||||
x = F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)
|
||||
return x.to(orig_type)
|
||||
|
||||
class QuickGELU(nn.Module):
|
||||
# NOTE This is slower than nn.GELU or nn.SiLU and uses more GPU memory
|
||||
def forward(self, x: torch.Tensor):
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
|
||||
class LayerScale(nn.Module):
|
||||
def __init__(self, dim, init_values=1e-5, inplace=False):
|
||||
super().__init__()
|
||||
self.inplace = inplace
|
||||
self.gamma = nn.Parameter(init_values * torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return x.mul_(self.gamma) if self.inplace else x * self.gamma
|
||||
|
||||
class PatchDropout(nn.Module):
|
||||
"""
|
||||
https://arxiv.org/abs/2212.00794
|
||||
"""
|
||||
|
||||
def __init__(self, prob, exclude_first_token=True):
|
||||
super().__init__()
|
||||
assert 0 <= prob < 1.
|
||||
self.prob = prob
|
||||
self.exclude_first_token = exclude_first_token # exclude CLS token
|
||||
logging.info(f"os.getenv('RoPE')={os.getenv('RoPE')}")
|
||||
|
||||
def forward(self, x):
|
||||
if not self.training or self.prob == 0.:
|
||||
return x
|
||||
|
||||
if self.exclude_first_token:
|
||||
cls_tokens, x = x[:, :1], x[:, 1:]
|
||||
else:
|
||||
cls_tokens = torch.jit.annotate(torch.Tensor, x[:, :1])
|
||||
|
||||
batch = x.size()[0]
|
||||
num_tokens = x.size()[1]
|
||||
|
||||
batch_indices = torch.arange(batch)
|
||||
batch_indices = batch_indices[..., None]
|
||||
|
||||
keep_prob = 1 - self.prob
|
||||
num_patches_keep = max(1, int(num_tokens * keep_prob))
|
||||
|
||||
rand = torch.randn(batch, num_tokens)
|
||||
patch_indices_keep = rand.topk(num_patches_keep, dim=-1).indices
|
||||
|
||||
x = x[batch_indices, patch_indices_keep]
|
||||
|
||||
if self.exclude_first_token:
|
||||
x = torch.cat((cls_tokens, x), dim=1)
|
||||
|
||||
if self.training and os.getenv('RoPE') == '1':
|
||||
return x, patch_indices_keep
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def _in_projection_packed(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
w: torch.Tensor,
|
||||
b: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""
|
||||
https://github.com/pytorch/pytorch/blob/db2a237763eb8693a20788be94f8c192e762baa8/torch/nn/functional.py#L4726
|
||||
"""
|
||||
E = q.size(-1)
|
||||
if k is v:
|
||||
if q is k:
|
||||
# self-attention
|
||||
return F.linear(q, w, b).chunk(3, dim=-1)
|
||||
else:
|
||||
# encoder-decoder attention
|
||||
w_q, w_kv = w.split([E, E * 2])
|
||||
if b is None:
|
||||
b_q = b_kv = None
|
||||
else:
|
||||
b_q, b_kv = b.split([E, E * 2])
|
||||
return (F.linear(q, w_q, b_q),) + F.linear(k, w_kv, b_kv).chunk(2, dim=-1)
|
||||
else:
|
||||
w_q, w_k, w_v = w.chunk(3)
|
||||
if b is None:
|
||||
b_q = b_k = b_v = None
|
||||
else:
|
||||
b_q, b_k, b_v = b.chunk(3)
|
||||
return F.linear(q, w_q, b_q), F.linear(k, w_k, b_k), F.linear(v, w_v, b_v)
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads=8,
|
||||
qkv_bias=True,
|
||||
scaled_cosine=False,
|
||||
scale_heads=False,
|
||||
logit_scale_max=math.log(1. / 0.01),
|
||||
attn_drop=0.,
|
||||
proj_drop=0.,
|
||||
xattn=False,
|
||||
rope=False
|
||||
):
|
||||
super().__init__()
|
||||
self.scaled_cosine = scaled_cosine
|
||||
self.scale_heads = scale_heads
|
||||
assert dim % num_heads == 0, 'dim should be divisible by num_heads'
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.scale = self.head_dim ** -0.5
|
||||
self.logit_scale_max = logit_scale_max
|
||||
|
||||
# keeping in_proj in this form (instead of nn.Linear) to match weight scheme of original
|
||||
self.in_proj_weight = nn.Parameter(torch.randn((dim * 3, dim)) * self.scale)
|
||||
if qkv_bias:
|
||||
self.in_proj_bias = nn.Parameter(torch.zeros(dim * 3))
|
||||
else:
|
||||
self.in_proj_bias = None
|
||||
|
||||
if self.scaled_cosine:
|
||||
self.logit_scale = nn.Parameter(torch.log(10 * torch.ones((num_heads, 1, 1))))
|
||||
else:
|
||||
self.logit_scale = None
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
if self.scale_heads:
|
||||
self.head_scale = nn.Parameter(torch.ones((num_heads, 1, 1)))
|
||||
else:
|
||||
self.head_scale = None
|
||||
self.out_proj = nn.Linear(dim, dim)
|
||||
self.out_drop = nn.Dropout(proj_drop)
|
||||
self.xattn = xattn
|
||||
self.xattn_drop = attn_drop
|
||||
self.rope = rope
|
||||
|
||||
def forward(self, x, attn_mask: Optional[torch.Tensor] = None):
|
||||
L, N, C = x.shape
|
||||
q, k, v = F.linear(x, self.in_proj_weight, self.in_proj_bias).chunk(3, dim=-1)
|
||||
if self.xattn:
|
||||
q = q.contiguous().view(L, N, self.num_heads, -1).transpose(0, 1)
|
||||
k = k.contiguous().view(L, N, self.num_heads, -1).transpose(0, 1)
|
||||
v = v.contiguous().view(L, N, self.num_heads, -1).transpose(0, 1)
|
||||
|
||||
x = xops.memory_efficient_attention(
|
||||
q, k, v,
|
||||
p=self.xattn_drop,
|
||||
scale=self.scale if self.logit_scale is None else None,
|
||||
attn_bias=xops.LowerTriangularMask() if attn_mask is not None else None,
|
||||
)
|
||||
else:
|
||||
q = q.contiguous().view(L, N * self.num_heads, -1).transpose(0, 1)
|
||||
k = k.contiguous().view(L, N * self.num_heads, -1).transpose(0, 1)
|
||||
v = v.contiguous().view(L, N * self.num_heads, -1).transpose(0, 1)
|
||||
|
||||
if self.logit_scale is not None:
|
||||
attn = torch.bmm(F.normalize(q, dim=-1), F.normalize(k, dim=-1).transpose(-1, -2))
|
||||
logit_scale = torch.clamp(self.logit_scale, max=self.logit_scale_max).exp()
|
||||
attn = attn.view(N, self.num_heads, L, L) * logit_scale
|
||||
attn = attn.view(-1, L, L)
|
||||
else:
|
||||
q = q * self.scale
|
||||
attn = torch.bmm(q, k.transpose(-1, -2))
|
||||
|
||||
if attn_mask is not None:
|
||||
if attn_mask.dtype == torch.bool:
|
||||
new_attn_mask = torch.zeros_like(attn_mask, dtype=q.dtype)
|
||||
new_attn_mask.masked_fill_(attn_mask, float("-inf"))
|
||||
attn_mask = new_attn_mask
|
||||
attn += attn_mask
|
||||
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = torch.bmm(attn, v)
|
||||
|
||||
if self.head_scale is not None:
|
||||
x = x.view(N, self.num_heads, L, C) * self.head_scale
|
||||
x = x.view(-1, L, C)
|
||||
x = x.transpose(0, 1).reshape(L, N, C)
|
||||
x = self.out_proj(x)
|
||||
x = self.out_drop(x)
|
||||
return x
|
||||
|
||||
class CustomAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads=8,
|
||||
qkv_bias=True,
|
||||
scaled_cosine=True,
|
||||
scale_heads=False,
|
||||
logit_scale_max=math.log(1. / 0.01),
|
||||
attn_drop=0.,
|
||||
proj_drop=0.,
|
||||
xattn=False
|
||||
):
|
||||
super().__init__()
|
||||
self.scaled_cosine = scaled_cosine
|
||||
self.scale_heads = scale_heads
|
||||
assert dim % num_heads == 0, 'dim should be divisible by num_heads'
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.scale = self.head_dim ** -0.5
|
||||
self.logit_scale_max = logit_scale_max
|
||||
|
||||
# keeping in_proj in this form (instead of nn.Linear) to match weight scheme of original
|
||||
self.in_proj_weight = nn.Parameter(torch.randn((dim * 3, dim)) * self.scale)
|
||||
if qkv_bias:
|
||||
self.in_proj_bias = nn.Parameter(torch.zeros(dim * 3))
|
||||
else:
|
||||
self.in_proj_bias = None
|
||||
|
||||
if self.scaled_cosine:
|
||||
self.logit_scale = nn.Parameter(torch.log(10 * torch.ones((num_heads, 1, 1))))
|
||||
else:
|
||||
self.logit_scale = None
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
if self.scale_heads:
|
||||
self.head_scale = nn.Parameter(torch.ones((num_heads, 1, 1)))
|
||||
else:
|
||||
self.head_scale = None
|
||||
self.out_proj = nn.Linear(dim, dim)
|
||||
self.out_drop = nn.Dropout(proj_drop)
|
||||
self.xattn = xattn
|
||||
self.xattn_drop = attn_drop
|
||||
|
||||
def forward(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attn_mask: Optional[torch.Tensor] = None):
|
||||
q, k, v = _in_projection_packed(query, key, value, self.in_proj_weight, self.in_proj_bias)
|
||||
N_q, B_q, C_q = q.shape
|
||||
N_k, B_k, C_k = k.shape
|
||||
N_v, B_v, C_v = v.shape
|
||||
if self.xattn:
|
||||
# B, N, C -> B, N, num_heads, C
|
||||
q = q.permute(1, 0, 2).reshape(B_q, N_q, self.num_heads, -1)
|
||||
k = k.permute(1, 0, 2).reshape(B_k, N_k, self.num_heads, -1)
|
||||
v = v.permute(1, 0, 2).reshape(B_v, N_v, self.num_heads, -1)
|
||||
|
||||
x = xops.memory_efficient_attention(
|
||||
q, k, v,
|
||||
p=self.xattn_drop,
|
||||
scale=self.scale if self.logit_scale is None else None,
|
||||
attn_bias=xops.LowerTriangularMask() if attn_mask is not None else None
|
||||
)
|
||||
else:
|
||||
# B*H, L, C
|
||||
q = q.contiguous().view(N_q, B_q * self.num_heads, -1).transpose(0, 1)
|
||||
k = k.contiguous().view(N_k, B_k * self.num_heads, -1).transpose(0, 1)
|
||||
v = v.contiguous().view(N_v, B_v * self.num_heads, -1).transpose(0, 1)
|
||||
|
||||
if self.logit_scale is not None:
|
||||
# B*H, N_q, N_k
|
||||
attn = torch.bmm(F.normalize(q, dim=-1), F.normalize(k, dim=-1).transpose(-1, -2))
|
||||
logit_scale = torch.clamp(self.logit_scale, max=self.logit_scale_max).exp()
|
||||
attn = attn.view(B_q, self.num_heads, N_q, N_k) * logit_scale
|
||||
attn = attn.view(-1, N_q, N_k)
|
||||
else:
|
||||
q = q * self.scale
|
||||
attn = torch.bmm(q, k.transpose(-1, -2))
|
||||
|
||||
if attn_mask is not None:
|
||||
if attn_mask.dtype == torch.bool:
|
||||
new_attn_mask = torch.zeros_like(attn_mask, dtype=q.dtype)
|
||||
new_attn_mask.masked_fill_(attn_mask, float("-inf"))
|
||||
attn_mask = new_attn_mask
|
||||
attn += attn_mask
|
||||
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = torch.bmm(attn, v)
|
||||
|
||||
if self.head_scale is not None:
|
||||
x = x.view(B_q, self.num_heads, N_q, C_q) * self.head_scale
|
||||
x = x.view(-1, N_q, C_q)
|
||||
x = x.transpose(0, 1).reshape(N_q, B_q, C_q)
|
||||
x = self.out_proj(x)
|
||||
x = self.out_drop(x)
|
||||
return x
|
||||
|
||||
class CustomResidualAttentionBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_head: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
ls_init_value: float = None,
|
||||
act_layer: Callable = nn.GELU,
|
||||
norm_layer: Callable = LayerNorm,
|
||||
scale_cosine_attn: bool = False,
|
||||
scale_heads: bool = False,
|
||||
scale_attn: bool = False,
|
||||
scale_fc: bool = False,
|
||||
cross_attn: bool = False,
|
||||
xattn: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.ln_1 = norm_layer(d_model)
|
||||
self.ln_1_k = norm_layer(d_model) if cross_attn else self.ln_1
|
||||
self.ln_1_v = norm_layer(d_model) if cross_attn else self.ln_1
|
||||
self.attn = CustomAttention(
|
||||
d_model, n_head,
|
||||
qkv_bias=True,
|
||||
attn_drop=0.,
|
||||
proj_drop=0.,
|
||||
scaled_cosine=scale_cosine_attn,
|
||||
scale_heads=scale_heads,
|
||||
xattn=xattn
|
||||
)
|
||||
|
||||
self.ln_attn = norm_layer(d_model) if scale_attn else nn.Identity()
|
||||
self.ls_1 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity()
|
||||
|
||||
self.ln_2 = norm_layer(d_model)
|
||||
mlp_width = int(d_model * mlp_ratio)
|
||||
self.mlp = nn.Sequential(OrderedDict([
|
||||
("c_fc", nn.Linear(d_model, mlp_width)),
|
||||
('ln', norm_layer(mlp_width) if scale_fc else nn.Identity()),
|
||||
("gelu", act_layer()),
|
||||
("c_proj", nn.Linear(mlp_width, d_model))
|
||||
]))
|
||||
|
||||
self.ls_2 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity()
|
||||
|
||||
def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, attn_mask: Optional[torch.Tensor] = None):
|
||||
q = q + self.ls_1(self.ln_attn(self.attn(self.ln_1(q), self.ln_1_k(k), self.ln_1_v(v), attn_mask=attn_mask)))
|
||||
q = q + self.ls_2(self.mlp(self.ln_2(q)))
|
||||
return q
|
||||
|
||||
class CustomTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
width: int,
|
||||
layers: int,
|
||||
heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
ls_init_value: float = None,
|
||||
act_layer: Callable = nn.GELU,
|
||||
norm_layer: Callable = LayerNorm,
|
||||
scale_cosine_attn: bool = True,
|
||||
scale_heads: bool = False,
|
||||
scale_attn: bool = False,
|
||||
scale_fc: bool = False,
|
||||
cross_attn: bool = False,
|
||||
xattn: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.width = width
|
||||
self.layers = layers
|
||||
self.grad_checkpointing = False
|
||||
self.xattn = xattn
|
||||
|
||||
self.resblocks = nn.ModuleList([
|
||||
CustomResidualAttentionBlock(
|
||||
width,
|
||||
heads,
|
||||
mlp_ratio,
|
||||
ls_init_value=ls_init_value,
|
||||
act_layer=act_layer,
|
||||
norm_layer=norm_layer,
|
||||
scale_cosine_attn=scale_cosine_attn,
|
||||
scale_heads=scale_heads,
|
||||
scale_attn=scale_attn,
|
||||
scale_fc=scale_fc,
|
||||
cross_attn=cross_attn,
|
||||
xattn=xattn)
|
||||
for _ in range(layers)
|
||||
])
|
||||
|
||||
def get_cast_dtype(self) -> torch.dtype:
|
||||
return self.resblocks[0].mlp.c_fc.weight.dtype
|
||||
|
||||
def forward(self, q: torch.Tensor, k: torch.Tensor = None, v: torch.Tensor = None, attn_mask: Optional[torch.Tensor] = None):
|
||||
if k is None and v is None:
|
||||
k = v = q
|
||||
for r in self.resblocks:
|
||||
if self.grad_checkpointing and not torch.jit.is_scripting():
|
||||
q = checkpoint(r, q, k, v, attn_mask)
|
||||
else:
|
||||
q = r(q, k, v, attn_mask=attn_mask)
|
||||
return q
|
||||
|
||||
|
||||
class ResidualAttentionBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_head: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
ls_init_value: float = None,
|
||||
act_layer: Callable = nn.GELU,
|
||||
norm_layer: Callable = LayerNorm,
|
||||
xattn: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.ln_1 = norm_layer(d_model)
|
||||
if xattn:
|
||||
self.attn = Attention(d_model, n_head, xattn=True)
|
||||
else:
|
||||
self.attn = nn.MultiheadAttention(d_model, n_head)
|
||||
self.ls_1 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity()
|
||||
|
||||
self.ln_2 = norm_layer(d_model)
|
||||
mlp_width = int(d_model * mlp_ratio)
|
||||
self.mlp = nn.Sequential(OrderedDict([
|
||||
("c_fc", nn.Linear(d_model, mlp_width)),
|
||||
("gelu", act_layer()),
|
||||
("c_proj", nn.Linear(mlp_width, d_model))
|
||||
]))
|
||||
|
||||
self.ls_2 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity()
|
||||
self.xattn = xattn
|
||||
|
||||
def attention(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None):
|
||||
attn_mask = attn_mask.to(x.dtype) if attn_mask is not None else None
|
||||
if self.xattn:
|
||||
return self.attn(x, attn_mask=attn_mask)
|
||||
return self.attn(x, x, x, need_weights=False, attn_mask=attn_mask)[0]
|
||||
|
||||
def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None):
|
||||
x = x + self.ls_1(self.attention(self.ln_1(x), attn_mask=attn_mask))
|
||||
x = x + self.ls_2(self.mlp(self.ln_2(x)))
|
||||
return x
|
||||
|
||||
class Transformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
width: int,
|
||||
layers: int,
|
||||
heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
ls_init_value: float = None,
|
||||
act_layer: Callable = nn.GELU,
|
||||
norm_layer: Callable = LayerNorm,
|
||||
xattn: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.width = width
|
||||
self.layers = layers
|
||||
self.grad_checkpointing = False
|
||||
|
||||
self.resblocks = nn.ModuleList([
|
||||
ResidualAttentionBlock(
|
||||
width, heads, mlp_ratio, ls_init_value=ls_init_value, act_layer=act_layer, norm_layer=norm_layer, xattn=xattn)
|
||||
for _ in range(layers)
|
||||
])
|
||||
|
||||
def get_cast_dtype(self) -> torch.dtype:
|
||||
return self.resblocks[0].mlp.c_fc.weight.dtype
|
||||
|
||||
def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None):
|
||||
for r in self.resblocks:
|
||||
if self.grad_checkpointing and not torch.jit.is_scripting():
|
||||
x = checkpoint(r, x, attn_mask)
|
||||
else:
|
||||
x = r(x, attn_mask=attn_mask)
|
||||
return x
|
||||
|
||||
|
||||
class VisionTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
image_size: int,
|
||||
patch_size: int,
|
||||
width: int,
|
||||
layers: int,
|
||||
heads: int,
|
||||
mlp_ratio: float,
|
||||
ls_init_value: float = None,
|
||||
patch_dropout: float = 0.,
|
||||
global_average_pool: bool = False,
|
||||
output_dim: int = 512,
|
||||
act_layer: Callable = nn.GELU,
|
||||
norm_layer: Callable = LayerNorm,
|
||||
xattn: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.image_size = to_2tuple(image_size)
|
||||
self.patch_size = to_2tuple(patch_size)
|
||||
self.grid_size = (self.image_size[0] // self.patch_size[0], self.image_size[1] // self.patch_size[1])
|
||||
self.output_dim = output_dim
|
||||
self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)
|
||||
|
||||
scale = width ** -0.5
|
||||
self.class_embedding = nn.Parameter(scale * torch.randn(width))
|
||||
self.positional_embedding = nn.Parameter(scale * torch.randn(self.grid_size[0] * self.grid_size[1] + 1, width))
|
||||
|
||||
# setting a patch_dropout of 0. would mean it is disabled and this function would be the identity fn
|
||||
self.patch_dropout = PatchDropout(patch_dropout) if patch_dropout > 0. else nn.Identity()
|
||||
self.ln_pre = norm_layer(width)
|
||||
|
||||
self.transformer = Transformer(
|
||||
width,
|
||||
layers,
|
||||
heads,
|
||||
mlp_ratio,
|
||||
ls_init_value=ls_init_value,
|
||||
act_layer=act_layer,
|
||||
norm_layer=norm_layer,
|
||||
xattn=xattn
|
||||
)
|
||||
|
||||
self.global_average_pool = global_average_pool
|
||||
self.ln_post = norm_layer(width)
|
||||
self.proj = nn.Parameter(scale * torch.randn(width, output_dim))
|
||||
|
||||
def lock(self, unlocked_groups=0, freeze_bn_stats=False):
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
if unlocked_groups != 0:
|
||||
groups = [
|
||||
[
|
||||
self.conv1,
|
||||
self.class_embedding,
|
||||
self.positional_embedding,
|
||||
self.ln_pre,
|
||||
],
|
||||
*self.transformer.resblocks[:-1],
|
||||
[
|
||||
self.transformer.resblocks[-1],
|
||||
self.ln_post,
|
||||
],
|
||||
self.proj,
|
||||
]
|
||||
|
||||
def _unlock(x):
|
||||
if isinstance(x, Sequence):
|
||||
for g in x:
|
||||
_unlock(g)
|
||||
else:
|
||||
if isinstance(x, torch.nn.Parameter):
|
||||
x.requires_grad = True
|
||||
else:
|
||||
for p in x.parameters():
|
||||
p.requires_grad = True
|
||||
|
||||
_unlock(groups[-unlocked_groups:])
|
||||
|
||||
def get_num_layers(self):
|
||||
return self.transformer.layers
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
self.transformer.grad_checkpointing = enable
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'positional_embedding', 'class_embedding'}
|
||||
|
||||
def forward(self, x: torch.Tensor, return_all_features: bool=False):
|
||||
x = self.conv1(x) # shape = [*, width, grid, grid]
|
||||
x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
|
||||
x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
|
||||
x = torch.cat(
|
||||
[self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device),
|
||||
x], dim=1) # shape = [*, grid ** 2 + 1, width]
|
||||
x = x + self.positional_embedding.to(x.dtype)
|
||||
|
||||
# a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in
|
||||
x = self.patch_dropout(x)
|
||||
x = self.ln_pre(x)
|
||||
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.transformer(x)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
|
||||
if not return_all_features:
|
||||
if self.global_average_pool:
|
||||
x = x.mean(dim=1) #x = x[:,1:,:].mean(dim=1)
|
||||
else:
|
||||
x = x[:, 0]
|
||||
|
||||
x = self.ln_post(x)
|
||||
|
||||
if self.proj is not None:
|
||||
x = x @ self.proj
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class TextTransformer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
context_length: int = 77,
|
||||
vocab_size: int = 49408,
|
||||
width: int = 512,
|
||||
heads: int = 8,
|
||||
layers: int = 12,
|
||||
ls_init_value: float = None,
|
||||
output_dim: int = 512,
|
||||
act_layer: Callable = nn.GELU,
|
||||
norm_layer: Callable = LayerNorm,
|
||||
xattn: bool= False,
|
||||
attn_mask: bool = True
|
||||
):
|
||||
super().__init__()
|
||||
self.context_length = context_length
|
||||
self.vocab_size = vocab_size
|
||||
self.width = width
|
||||
self.output_dim = output_dim
|
||||
|
||||
self.token_embedding = nn.Embedding(vocab_size, width)
|
||||
self.positional_embedding = nn.Parameter(torch.empty(self.context_length, width))
|
||||
self.transformer = Transformer(
|
||||
width=width,
|
||||
layers=layers,
|
||||
heads=heads,
|
||||
ls_init_value=ls_init_value,
|
||||
act_layer=act_layer,
|
||||
norm_layer=norm_layer,
|
||||
xattn=xattn
|
||||
)
|
||||
|
||||
self.xattn = xattn
|
||||
self.ln_final = norm_layer(width)
|
||||
self.text_projection = nn.Parameter(torch.empty(width, output_dim))
|
||||
|
||||
if attn_mask:
|
||||
self.register_buffer('attn_mask', self.build_attention_mask(), persistent=False)
|
||||
else:
|
||||
self.attn_mask = None
|
||||
|
||||
self.init_parameters()
|
||||
|
||||
def init_parameters(self):
|
||||
nn.init.normal_(self.token_embedding.weight, std=0.02)
|
||||
nn.init.normal_(self.positional_embedding, std=0.01)
|
||||
|
||||
proj_std = (self.transformer.width ** -0.5) * ((2 * self.transformer.layers) ** -0.5)
|
||||
attn_std = self.transformer.width ** -0.5
|
||||
fc_std = (2 * self.transformer.width) ** -0.5
|
||||
for block in self.transformer.resblocks:
|
||||
nn.init.normal_(block.attn.in_proj_weight, std=attn_std)
|
||||
nn.init.normal_(block.attn.out_proj.weight, std=proj_std)
|
||||
nn.init.normal_(block.mlp.c_fc.weight, std=fc_std)
|
||||
nn.init.normal_(block.mlp.c_proj.weight, std=proj_std)
|
||||
|
||||
if self.text_projection is not None:
|
||||
nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5)
|
||||
|
||||
@torch.jit.ignore
|
||||
def set_grad_checkpointing(self, enable=True):
|
||||
self.transformer.grad_checkpointing = enable
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
# return {'positional_embedding', 'token_embedding'}
|
||||
return {'positional_embedding'}
|
||||
|
||||
def get_num_layers(self):
|
||||
return self.transformer.layers
|
||||
|
||||
def build_attention_mask(self):
|
||||
# lazily create causal attention mask, with full attention between the vision tokens
|
||||
# pytorch uses additive attention mask; fill with -inf
|
||||
mask = torch.empty(self.context_length, self.context_length)
|
||||
mask.fill_(float("-inf"))
|
||||
mask.triu_(1) # zero out the lower diagonal
|
||||
return mask
|
||||
|
||||
def forward(self, text, return_all_features: bool=False):
|
||||
cast_dtype = self.transformer.get_cast_dtype()
|
||||
x = self.token_embedding(text).to(cast_dtype) # [batch_size, n_ctx, d_model]
|
||||
|
||||
x = x + self.positional_embedding.to(cast_dtype)
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.transformer(x, attn_mask=self.attn_mask)
|
||||
# x = self.transformer(x) # no attention mask is applied
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.ln_final(x)
|
||||
|
||||
if not return_all_features:
|
||||
# x.shape = [batch_size, n_ctx, transformer.width]
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection
|
||||
return x
|
||||
@@ -0,0 +1,326 @@
|
||||
from itertools import repeat
|
||||
import collections.abc
|
||||
import logging
|
||||
import math
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
from torch import nn as nn
|
||||
from torchvision.ops.misc import FrozenBatchNorm2d
|
||||
import torch.nn.functional as F
|
||||
|
||||
# open CLIP
|
||||
def resize_clip_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1):
|
||||
# Rescale the grid of position embeddings when loading from state_dict
|
||||
old_pos_embed = state_dict.get('visual.positional_embedding', None)
|
||||
if old_pos_embed is None or not hasattr(model.visual, 'grid_size'):
|
||||
return
|
||||
grid_size = to_2tuple(model.visual.grid_size)
|
||||
extra_tokens = 1 # FIXME detect different token configs (ie no class token, or more)
|
||||
new_seq_len = grid_size[0] * grid_size[1] + extra_tokens
|
||||
if new_seq_len == old_pos_embed.shape[0]:
|
||||
return
|
||||
|
||||
if extra_tokens:
|
||||
pos_emb_tok, pos_emb_img = old_pos_embed[:extra_tokens], old_pos_embed[extra_tokens:]
|
||||
else:
|
||||
pos_emb_tok, pos_emb_img = None, old_pos_embed
|
||||
old_grid_size = to_2tuple(int(math.sqrt(len(pos_emb_img))))
|
||||
|
||||
logging.info('Resizing position embedding grid-size from %s to %s', old_grid_size, grid_size)
|
||||
pos_emb_img = pos_emb_img.reshape(1, old_grid_size[0], old_grid_size[1], -1).permute(0, 3, 1, 2)
|
||||
pos_emb_img = F.interpolate(
|
||||
pos_emb_img,
|
||||
size=grid_size,
|
||||
mode=interpolation,
|
||||
align_corners=True,
|
||||
)
|
||||
pos_emb_img = pos_emb_img.permute(0, 2, 3, 1).reshape(1, grid_size[0] * grid_size[1], -1)[0]
|
||||
if pos_emb_tok is not None:
|
||||
new_pos_embed = torch.cat([pos_emb_tok, pos_emb_img], dim=0)
|
||||
else:
|
||||
new_pos_embed = pos_emb_img
|
||||
state_dict['visual.positional_embedding'] = new_pos_embed
|
||||
|
||||
|
||||
def resize_visual_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1):
|
||||
# Rescale the grid of position embeddings when loading from state_dict
|
||||
old_pos_embed = state_dict.get('positional_embedding', None)
|
||||
if old_pos_embed is None or not hasattr(model.visual, 'grid_size'):
|
||||
return
|
||||
grid_size = to_2tuple(model.visual.grid_size)
|
||||
extra_tokens = 1 # FIXME detect different token configs (ie no class token, or more)
|
||||
new_seq_len = grid_size[0] * grid_size[1] + extra_tokens
|
||||
if new_seq_len == old_pos_embed.shape[0]:
|
||||
return
|
||||
|
||||
if extra_tokens:
|
||||
pos_emb_tok, pos_emb_img = old_pos_embed[:extra_tokens], old_pos_embed[extra_tokens:]
|
||||
else:
|
||||
pos_emb_tok, pos_emb_img = None, old_pos_embed
|
||||
old_grid_size = to_2tuple(int(math.sqrt(len(pos_emb_img))))
|
||||
|
||||
logging.info('Resizing position embedding grid-size from %s to %s', old_grid_size, grid_size)
|
||||
pos_emb_img = pos_emb_img.reshape(1, old_grid_size[0], old_grid_size[1], -1).permute(0, 3, 1, 2)
|
||||
pos_emb_img = F.interpolate(
|
||||
pos_emb_img,
|
||||
size=grid_size,
|
||||
mode=interpolation,
|
||||
align_corners=True,
|
||||
)
|
||||
pos_emb_img = pos_emb_img.permute(0, 2, 3, 1).reshape(1, grid_size[0] * grid_size[1], -1)[0]
|
||||
if pos_emb_tok is not None:
|
||||
new_pos_embed = torch.cat([pos_emb_tok, pos_emb_img], dim=0)
|
||||
else:
|
||||
new_pos_embed = pos_emb_img
|
||||
state_dict['positional_embedding'] = new_pos_embed
|
||||
|
||||
def resize_evaclip_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1):
|
||||
all_keys = list(state_dict.keys())
|
||||
# interpolate position embedding
|
||||
if 'visual.pos_embed' in state_dict:
|
||||
pos_embed_checkpoint = state_dict['visual.pos_embed']
|
||||
embedding_size = pos_embed_checkpoint.shape[-1]
|
||||
num_patches = model.visual.patch_embed.num_patches
|
||||
num_extra_tokens = model.visual.pos_embed.shape[-2] - num_patches
|
||||
# height (== width) for the checkpoint position embedding
|
||||
orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5)
|
||||
# height (== width) for the new position embedding
|
||||
new_size = int(num_patches ** 0.5)
|
||||
# class_token and dist_token are kept unchanged
|
||||
if orig_size != new_size:
|
||||
print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size))
|
||||
extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens]
|
||||
# only the position tokens are interpolated
|
||||
pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:]
|
||||
pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2)
|
||||
pos_tokens = torch.nn.functional.interpolate(
|
||||
pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False)
|
||||
pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
|
||||
new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1)
|
||||
state_dict['visual.pos_embed'] = new_pos_embed
|
||||
|
||||
patch_embed_proj = state_dict['visual.patch_embed.proj.weight']
|
||||
patch_size = model.visual.patch_embed.patch_size
|
||||
state_dict['visual.patch_embed.proj.weight'] = torch.nn.functional.interpolate(
|
||||
patch_embed_proj.float(), size=patch_size, mode='bicubic', align_corners=False)
|
||||
|
||||
|
||||
def resize_eva_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1):
|
||||
all_keys = list(state_dict.keys())
|
||||
# interpolate position embedding
|
||||
if 'pos_embed' in state_dict:
|
||||
pos_embed_checkpoint = state_dict['pos_embed']
|
||||
embedding_size = pos_embed_checkpoint.shape[-1]
|
||||
num_patches = model.visual.patch_embed.num_patches
|
||||
num_extra_tokens = model.visual.pos_embed.shape[-2] - num_patches
|
||||
# height (== width) for the checkpoint position embedding
|
||||
orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5)
|
||||
# height (== width) for the new position embedding
|
||||
new_size = int(num_patches ** 0.5)
|
||||
# class_token and dist_token are kept unchanged
|
||||
if orig_size != new_size:
|
||||
print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size))
|
||||
extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens]
|
||||
# only the position tokens are interpolated
|
||||
pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:]
|
||||
pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2)
|
||||
pos_tokens = torch.nn.functional.interpolate(
|
||||
pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False)
|
||||
pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
|
||||
new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1)
|
||||
state_dict['pos_embed'] = new_pos_embed
|
||||
|
||||
patch_embed_proj = state_dict['patch_embed.proj.weight']
|
||||
patch_size = model.visual.patch_embed.patch_size
|
||||
state_dict['patch_embed.proj.weight'] = torch.nn.functional.interpolate(
|
||||
patch_embed_proj.float(), size=patch_size, mode='bicubic', align_corners=False)
|
||||
|
||||
|
||||
def resize_rel_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1):
|
||||
all_keys = list(state_dict.keys())
|
||||
for key in all_keys:
|
||||
if "relative_position_index" in key:
|
||||
state_dict.pop(key)
|
||||
|
||||
if "relative_position_bias_table" in key:
|
||||
rel_pos_bias = state_dict[key]
|
||||
src_num_pos, num_attn_heads = rel_pos_bias.size()
|
||||
dst_num_pos, _ = model.visual.state_dict()[key].size()
|
||||
dst_patch_shape = model.visual.patch_embed.patch_shape
|
||||
if dst_patch_shape[0] != dst_patch_shape[1]:
|
||||
raise NotImplementedError()
|
||||
num_extra_tokens = dst_num_pos - (dst_patch_shape[0] * 2 - 1) * (dst_patch_shape[1] * 2 - 1)
|
||||
src_size = int((src_num_pos - num_extra_tokens) ** 0.5)
|
||||
dst_size = int((dst_num_pos - num_extra_tokens) ** 0.5)
|
||||
if src_size != dst_size:
|
||||
print("Position interpolate for %s from %dx%d to %dx%d" % (
|
||||
key, src_size, src_size, dst_size, dst_size))
|
||||
extra_tokens = rel_pos_bias[-num_extra_tokens:, :]
|
||||
rel_pos_bias = rel_pos_bias[:-num_extra_tokens, :]
|
||||
|
||||
def geometric_progression(a, r, n):
|
||||
return a * (1.0 - r ** n) / (1.0 - r)
|
||||
|
||||
left, right = 1.01, 1.5
|
||||
while right - left > 1e-6:
|
||||
q = (left + right) / 2.0
|
||||
gp = geometric_progression(1, q, src_size // 2)
|
||||
if gp > dst_size // 2:
|
||||
right = q
|
||||
else:
|
||||
left = q
|
||||
|
||||
# if q > 1.090307:
|
||||
# q = 1.090307
|
||||
|
||||
dis = []
|
||||
cur = 1
|
||||
for i in range(src_size // 2):
|
||||
dis.append(cur)
|
||||
cur += q ** (i + 1)
|
||||
|
||||
r_ids = [-_ for _ in reversed(dis)]
|
||||
|
||||
x = r_ids + [0] + dis
|
||||
y = r_ids + [0] + dis
|
||||
|
||||
t = dst_size // 2.0
|
||||
dx = np.arange(-t, t + 0.1, 1.0)
|
||||
dy = np.arange(-t, t + 0.1, 1.0)
|
||||
|
||||
print("Original positions = %s" % str(x))
|
||||
print("Target positions = %s" % str(dx))
|
||||
|
||||
all_rel_pos_bias = []
|
||||
|
||||
for i in range(num_attn_heads):
|
||||
z = rel_pos_bias[:, i].view(src_size, src_size).float().numpy()
|
||||
f = F.interpolate.interp2d(x, y, z, kind='cubic')
|
||||
all_rel_pos_bias.append(
|
||||
torch.Tensor(f(dx, dy)).contiguous().view(-1, 1).to(rel_pos_bias.device))
|
||||
|
||||
rel_pos_bias = torch.cat(all_rel_pos_bias, dim=-1)
|
||||
|
||||
new_rel_pos_bias = torch.cat((rel_pos_bias, extra_tokens), dim=0)
|
||||
state_dict[key] = new_rel_pos_bias
|
||||
|
||||
# interpolate position embedding
|
||||
if 'pos_embed' in state_dict:
|
||||
pos_embed_checkpoint = state_dict['pos_embed']
|
||||
embedding_size = pos_embed_checkpoint.shape[-1]
|
||||
num_patches = model.visual.patch_embed.num_patches
|
||||
num_extra_tokens = model.visual.pos_embed.shape[-2] - num_patches
|
||||
# height (== width) for the checkpoint position embedding
|
||||
orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5)
|
||||
# height (== width) for the new position embedding
|
||||
new_size = int(num_patches ** 0.5)
|
||||
# class_token and dist_token are kept unchanged
|
||||
if orig_size != new_size:
|
||||
print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size))
|
||||
extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens]
|
||||
# only the position tokens are interpolated
|
||||
pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:]
|
||||
pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2)
|
||||
pos_tokens = torch.nn.functional.interpolate(
|
||||
pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False)
|
||||
pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
|
||||
new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1)
|
||||
state_dict['pos_embed'] = new_pos_embed
|
||||
|
||||
patch_embed_proj = state_dict['patch_embed.proj.weight']
|
||||
patch_size = model.visual.patch_embed.patch_size
|
||||
state_dict['patch_embed.proj.weight'] = torch.nn.functional.interpolate(
|
||||
patch_embed_proj.float(), size=patch_size, mode='bicubic', align_corners=False)
|
||||
|
||||
|
||||
def freeze_batch_norm_2d(module, module_match={}, name=''):
|
||||
"""
|
||||
Converts all `BatchNorm2d` and `SyncBatchNorm` layers of provided module into `FrozenBatchNorm2d`. If `module` is
|
||||
itself an instance of either `BatchNorm2d` or `SyncBatchNorm`, it is converted into `FrozenBatchNorm2d` and
|
||||
returned. Otherwise, the module is walked recursively and submodules are converted in place.
|
||||
|
||||
Args:
|
||||
module (torch.nn.Module): Any PyTorch module.
|
||||
module_match (dict): Dictionary of full module names to freeze (all if empty)
|
||||
name (str): Full module name (prefix)
|
||||
|
||||
Returns:
|
||||
torch.nn.Module: Resulting module
|
||||
|
||||
Inspired by https://github.com/pytorch/pytorch/blob/a5895f85be0f10212791145bfedc0261d364f103/torch/nn/modules/batchnorm.py#L762
|
||||
"""
|
||||
res = module
|
||||
is_match = True
|
||||
if module_match:
|
||||
is_match = name in module_match
|
||||
if is_match and isinstance(module, (nn.modules.batchnorm.BatchNorm2d, nn.modules.batchnorm.SyncBatchNorm)):
|
||||
res = FrozenBatchNorm2d(module.num_features)
|
||||
res.num_features = module.num_features
|
||||
res.affine = module.affine
|
||||
if module.affine:
|
||||
res.weight.data = module.weight.data.clone().detach()
|
||||
res.bias.data = module.bias.data.clone().detach()
|
||||
res.running_mean.data = module.running_mean.data
|
||||
res.running_var.data = module.running_var.data
|
||||
res.eps = module.eps
|
||||
else:
|
||||
for child_name, child in module.named_children():
|
||||
full_child_name = '.'.join([name, child_name]) if name else child_name
|
||||
new_child = freeze_batch_norm_2d(child, module_match, full_child_name)
|
||||
if new_child is not child:
|
||||
res.add_module(child_name, new_child)
|
||||
return res
|
||||
|
||||
|
||||
# From PyTorch internals
|
||||
def _ntuple(n):
|
||||
def parse(x):
|
||||
if isinstance(x, collections.abc.Iterable):
|
||||
return x
|
||||
return tuple(repeat(x, n))
|
||||
return parse
|
||||
|
||||
|
||||
to_1tuple = _ntuple(1)
|
||||
to_2tuple = _ntuple(2)
|
||||
to_3tuple = _ntuple(3)
|
||||
to_4tuple = _ntuple(4)
|
||||
to_ntuple = lambda n, x: _ntuple(n)(x)
|
||||
|
||||
|
||||
def is_logging(args):
|
||||
def is_global_master(args):
|
||||
return args.rank == 0
|
||||
|
||||
def is_local_master(args):
|
||||
return args.local_rank == 0
|
||||
|
||||
def is_master(args, local=False):
|
||||
return is_local_master(args) if local else is_global_master(args)
|
||||
return is_master
|
||||
|
||||
|
||||
class AllGather(torch.autograd.Function):
|
||||
"""An autograd function that performs allgather on a tensor.
|
||||
Performs all_gather operation on the provided tensors.
|
||||
*** Warning ***: torch.distributed.all_gather has no gradient.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, tensor, rank, world_size):
|
||||
tensors_gather = [torch.empty_like(tensor) for _ in range(world_size)]
|
||||
torch.distributed.all_gather(tensors_gather, tensor)
|
||||
ctx.rank = rank
|
||||
ctx.batch_size = tensor.shape[0]
|
||||
return torch.cat(tensors_gather, 0)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return (
|
||||
grad_output[ctx.batch_size * ctx.rank: ctx.batch_size * (ctx.rank + 1)],
|
||||
None,
|
||||
None
|
||||
)
|
||||
|
||||
allgather = AllGather.apply
|
||||
@@ -0,0 +1,594 @@
|
||||
import math
|
||||
from scipy import integrate
|
||||
import torch
|
||||
from torch import nn
|
||||
from torchdiffeq import odeint
|
||||
import torchsde
|
||||
from tqdm.auto import trange
|
||||
|
||||
|
||||
def append_zero(x):
|
||||
return torch.cat([x, x.new_zeros([1])])
|
||||
|
||||
|
||||
def get_sigmas_karras(n, sigma_min, sigma_max, rho=7., device='cpu'):
|
||||
"""Constructs the noise schedule of Karras et al. (2022)."""
|
||||
ramp = torch.linspace(0, 1, n)
|
||||
min_inv_rho = sigma_min ** (1 / rho)
|
||||
max_inv_rho = sigma_max ** (1 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
|
||||
return append_zero(sigmas).to(device)
|
||||
|
||||
|
||||
def get_sigmas_exponential(n, sigma_min, sigma_max, device='cpu'):
|
||||
"""Constructs an exponential noise schedule."""
|
||||
sigmas = torch.linspace(math.log(sigma_max), math.log(sigma_min), n, device=device).exp()
|
||||
return append_zero(sigmas)
|
||||
|
||||
|
||||
def get_sigmas_polyexponential(n, sigma_min, sigma_max, rho=1., device='cpu'):
|
||||
"""Constructs an polynomial in log sigma noise schedule."""
|
||||
ramp = torch.linspace(1, 0, n, device=device) ** rho
|
||||
sigmas = torch.exp(ramp * (math.log(sigma_max) - math.log(sigma_min)) + math.log(sigma_min))
|
||||
return append_zero(sigmas)
|
||||
|
||||
|
||||
def get_sigmas_vp(n, beta_d=19.9, beta_min=0.1, eps_s=1e-3, device='cpu'):
|
||||
"""Constructs a continuous VP noise schedule."""
|
||||
t = torch.linspace(1, eps_s, n, device=device)
|
||||
sigmas = torch.sqrt(torch.exp(beta_d * t ** 2 / 2 + beta_min * t) - 1)
|
||||
return append_zero(sigmas)
|
||||
|
||||
|
||||
def append_dims(x, target_dims):
|
||||
"""Appends dimensions to the end of a tensor until it has target_dims dimensions."""
|
||||
dims_to_append = target_dims - x.ndim
|
||||
if dims_to_append < 0:
|
||||
raise ValueError(f'input has {x.ndim} dims but target_dims is {target_dims}, which is less')
|
||||
return x[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
def to_d(x, sigma, denoised):
|
||||
"""Converts a denoiser output to a Karras ODE derivative."""
|
||||
return (x - denoised) / append_dims(sigma, x.ndim)
|
||||
|
||||
|
||||
def get_ancestral_step(sigma_from, sigma_to, eta=1.):
|
||||
"""Calculates the noise level (sigma_down) to step down to and the amount
|
||||
of noise to add (sigma_up) when doing an ancestral sampling step."""
|
||||
if not eta:
|
||||
return sigma_to, 0.
|
||||
sigma_up = min(sigma_to, eta * (sigma_to ** 2 * (sigma_from ** 2 - sigma_to ** 2) / sigma_from ** 2) ** 0.5)
|
||||
sigma_down = (sigma_to ** 2 - sigma_up ** 2) ** 0.5
|
||||
return sigma_down, sigma_up
|
||||
|
||||
|
||||
def default_noise_sampler(x):
|
||||
return lambda sigma, sigma_next: torch.randn_like(x)
|
||||
|
||||
|
||||
def inpaint_mask(x, i, steps, mask_args):
|
||||
noised_original = mask_args["latent"].clone().to(x)
|
||||
latent_mask = mask_args["latent_mask"].to(x)
|
||||
if i < steps:
|
||||
noised_original += mask_args["noise"].to(x) * mask_args["sigmas"][i+1].to(x)
|
||||
x = (latent_mask * x) + ((1 - latent_mask) * noised_original.to(x))
|
||||
return x
|
||||
|
||||
|
||||
class BatchedBrownianTree:
|
||||
"""A wrapper around torchsde.BrownianTree that enables batches of entropy."""
|
||||
|
||||
def __init__(self, x, t0, t1, seed=None, **kwargs):
|
||||
t0, t1, self.sign = self.sort(t0, t1)
|
||||
w0 = kwargs.get('w0', torch.zeros_like(x))
|
||||
if seed is None:
|
||||
seed = torch.randint(0, 2 ** 63 - 1, []).item()
|
||||
self.batched = True
|
||||
try:
|
||||
assert len(seed) == x.shape[0]
|
||||
w0 = w0[0]
|
||||
except TypeError:
|
||||
seed = [seed]
|
||||
self.batched = False
|
||||
self.trees = [torchsde.BrownianTree(t0, w0, t1, entropy=s, **kwargs) for s in seed]
|
||||
|
||||
@staticmethod
|
||||
def sort(a, b):
|
||||
return (a, b, 1) if a < b else (b, a, -1)
|
||||
|
||||
def __call__(self, t0, t1):
|
||||
t0, t1, sign = self.sort(t0, t1)
|
||||
w = torch.stack([tree(t0, t1) for tree in self.trees]) * (self.sign * sign)
|
||||
return w if self.batched else w[0]
|
||||
|
||||
|
||||
class BrownianTreeNoiseSampler:
|
||||
"""A noise sampler backed by a torchsde.BrownianTree.
|
||||
|
||||
Args:
|
||||
x (Tensor): The tensor whose shape, device and dtype to use to generate
|
||||
random samples.
|
||||
sigma_min (float): The low end of the valid interval.
|
||||
sigma_max (float): The high end of the valid interval.
|
||||
seed (int or List[int]): The random seed. If a list of seeds is
|
||||
supplied instead of a single integer, then the noise sampler will
|
||||
use one BrownianTree per batch item, each with its own seed.
|
||||
transform (callable): A function that maps sigma to the sampler's
|
||||
internal timestep.
|
||||
"""
|
||||
|
||||
def __init__(self, x, sigma_min, sigma_max, seed=None, transform=lambda x: x):
|
||||
self.transform = transform
|
||||
t0, t1 = self.transform(torch.as_tensor(sigma_min)), self.transform(torch.as_tensor(sigma_max))
|
||||
self.tree = BatchedBrownianTree(x, t0, t1, seed)
|
||||
|
||||
def __call__(self, sigma, sigma_next):
|
||||
t0, t1 = self.transform(torch.as_tensor(sigma)), self.transform(torch.as_tensor(sigma_next))
|
||||
return self.tree(t0, t1) / (t1 - t0).abs().sqrt()
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1., mask_args=None):
|
||||
"""Implements Algorithm 2 (Euler steps) from Karras et al. (2022)."""
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
gamma = min(s_churn / (len(sigmas) - 1), 2 ** 0.5 - 1) if s_tmin <= sigmas[i] <= s_tmax else 0.
|
||||
eps = torch.randn_like(x) * s_noise
|
||||
sigma_hat = sigmas[i] * (gamma + 1)
|
||||
if gamma > 0:
|
||||
x = x + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5
|
||||
denoised = model(x, sigma_hat * s_in, **extra_args)
|
||||
d = to_d(x, sigma_hat, denoised)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigma_hat, 'denoised': denoised})
|
||||
dt = sigmas[i + 1] - sigma_hat
|
||||
# Euler method
|
||||
x = x + (d * dt).to(x.dtype)
|
||||
if mask_args is not None:
|
||||
x = inpaint_mask(x, i, len(sigmas) - 2, mask_args)
|
||||
return x
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_euler_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None):
|
||||
"""Ancestral sampling with Euler method steps."""
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
d = to_d(x, sigmas[i], denoised)
|
||||
# Euler method
|
||||
dt = sigma_down - sigmas[i]
|
||||
x = x + (d * dt).to(x.dtype)
|
||||
if sigmas[i + 1] > 0:
|
||||
x = x + (noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up).to(x.dtype)
|
||||
if mask_args is not None:
|
||||
x = inpaint_mask(x, i, len(sigmas) - 2, mask_args)
|
||||
return x
|
||||
|
||||
|
||||
def linear_multistep_coeff(order, t, i, j):
|
||||
if order - 1 > i:
|
||||
raise ValueError(f'Order {order} too high for step {i}')
|
||||
def fn(tau):
|
||||
prod = 1.
|
||||
for k in range(order):
|
||||
if j == k:
|
||||
continue
|
||||
prod *= (tau - t[i - k]) / (t[i - j] - t[i - k])
|
||||
return prod
|
||||
return integrate.quad(fn, t[i], t[i + 1], epsrel=1e-4)[0]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def log_likelihood(model, x, sigma_min, sigma_max, extra_args=None, atol=1e-4, rtol=1e-4):
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
v = torch.randint_like(x, 2) * 2 - 1
|
||||
fevals = 0
|
||||
def ode_fn(sigma, x):
|
||||
nonlocal fevals
|
||||
with torch.enable_grad():
|
||||
x = x[0].detach().requires_grad_()
|
||||
denoised = model(x, sigma * s_in, **extra_args)
|
||||
d = to_d(x, sigma, denoised)
|
||||
fevals += 1
|
||||
grad = torch.autograd.grad((d * v).sum(), x)[0]
|
||||
d_ll = (v * grad).flatten(1).sum(1)
|
||||
return d.detach(), d_ll
|
||||
x_min = x, x.new_zeros([x.shape[0]])
|
||||
t = x.new_tensor([sigma_min, sigma_max])
|
||||
sol = odeint(ode_fn, x_min, t, atol=atol, rtol=rtol, method='dopri5')
|
||||
latent, delta_ll = sol[0][-1], sol[1][-1]
|
||||
ll_prior = torch.distributions.Normal(0, sigma_max).log_prob(latent).flatten(1).sum(1)
|
||||
return ll_prior + delta_ll, {'fevals': fevals}
|
||||
|
||||
|
||||
class PIDStepSizeController:
|
||||
"""A PID controller for ODE adaptive step size control."""
|
||||
def __init__(self, h, pcoeff, icoeff, dcoeff, order=1, accept_safety=0.81, eps=1e-8):
|
||||
self.h = h
|
||||
self.b1 = (pcoeff + icoeff + dcoeff) / order
|
||||
self.b2 = -(pcoeff + 2 * dcoeff) / order
|
||||
self.b3 = dcoeff / order
|
||||
self.accept_safety = accept_safety
|
||||
self.eps = eps
|
||||
self.errs = []
|
||||
|
||||
def limiter(self, x):
|
||||
return 1 + math.atan(x - 1)
|
||||
|
||||
def propose_step(self, error):
|
||||
inv_error = 1 / (float(error) + self.eps)
|
||||
if not self.errs:
|
||||
self.errs = [inv_error, inv_error, inv_error]
|
||||
self.errs[0] = inv_error
|
||||
factor = self.errs[0] ** self.b1 * self.errs[1] ** self.b2 * self.errs[2] ** self.b3
|
||||
factor = self.limiter(factor)
|
||||
accept = factor >= self.accept_safety
|
||||
if accept:
|
||||
self.errs[2] = self.errs[1]
|
||||
self.errs[1] = self.errs[0]
|
||||
self.h *= factor
|
||||
return accept
|
||||
|
||||
|
||||
class DPMSolver(nn.Module):
|
||||
"""DPM-Solver. See https://arxiv.org/abs/2206.00927."""
|
||||
|
||||
def __init__(self, model, extra_args=None, eps_callback=None, info_callback=None):
|
||||
super().__init__()
|
||||
self.model = model
|
||||
self.extra_args = {} if extra_args is None else extra_args
|
||||
self.eps_callback = eps_callback
|
||||
self.info_callback = info_callback
|
||||
|
||||
def t(self, sigma):
|
||||
return -sigma.log()
|
||||
|
||||
def sigma(self, t):
|
||||
return t.neg().exp()
|
||||
|
||||
def eps(self, eps_cache, key, x, t, *args, **kwargs):
|
||||
if key in eps_cache:
|
||||
return eps_cache[key], eps_cache
|
||||
sigma = self.sigma(t) * x.new_ones([x.shape[0]])
|
||||
eps = (x - self.model(x, sigma, *args, **self.extra_args, **kwargs)) / self.sigma(t)
|
||||
if self.eps_callback is not None:
|
||||
self.eps_callback()
|
||||
return eps, {key: eps, **eps_cache}
|
||||
|
||||
def dpm_solver_1_step(self, x, t, t_next, eps_cache=None):
|
||||
eps_cache = {} if eps_cache is None else eps_cache
|
||||
h = t_next - t
|
||||
eps, eps_cache = self.eps(eps_cache, 'eps', x, t)
|
||||
x_1 = x - self.sigma(t_next) * h.expm1() * eps
|
||||
return x_1, eps_cache
|
||||
|
||||
def dpm_solver_2_step(self, x, t, t_next, r1=1 / 2, eps_cache=None):
|
||||
eps_cache = {} if eps_cache is None else eps_cache
|
||||
h = t_next - t
|
||||
eps, eps_cache = self.eps(eps_cache, 'eps', x, t)
|
||||
s1 = t + r1 * h
|
||||
u1 = x - self.sigma(s1) * (r1 * h).expm1() * eps
|
||||
eps_r1, eps_cache = self.eps(eps_cache, 'eps_r1', u1, s1)
|
||||
x_2 = x - self.sigma(t_next) * h.expm1() * eps - self.sigma(t_next) / (2 * r1) * h.expm1() * (eps_r1 - eps)
|
||||
return x_2, eps_cache
|
||||
|
||||
def dpm_solver_3_step(self, x, t, t_next, r1=1 / 3, r2=2 / 3, eps_cache=None):
|
||||
eps_cache = {} if eps_cache is None else eps_cache
|
||||
h = t_next - t
|
||||
eps, eps_cache = self.eps(eps_cache, 'eps', x, t)
|
||||
s1 = t + r1 * h
|
||||
s2 = t + r2 * h
|
||||
u1 = x - self.sigma(s1) * (r1 * h).expm1() * eps
|
||||
eps_r1, eps_cache = self.eps(eps_cache, 'eps_r1', u1, s1)
|
||||
u2 = x - self.sigma(s2) * (r2 * h).expm1() * eps - self.sigma(s2) * (r2 / r1) * ((r2 * h).expm1() / (r2 * h) - 1) * (eps_r1 - eps)
|
||||
eps_r2, eps_cache = self.eps(eps_cache, 'eps_r2', u2, s2)
|
||||
x_3 = x - self.sigma(t_next) * h.expm1() * eps - self.sigma(t_next) / r2 * (h.expm1() / h - 1) * (eps_r2 - eps)
|
||||
return x_3, eps_cache
|
||||
|
||||
def dpm_solver_fast(self, x, t_start, t_end, nfe, eta=0., s_noise=1., noise_sampler=None):
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
if not t_end > t_start and eta:
|
||||
raise ValueError('eta must be 0 for reverse sampling')
|
||||
|
||||
m = math.floor(nfe / 3) + 1
|
||||
ts = torch.linspace(t_start, t_end, m + 1, device=x.device)
|
||||
|
||||
if nfe % 3 == 0:
|
||||
orders = [3] * (m - 2) + [2, 1]
|
||||
else:
|
||||
orders = [3] * (m - 1) + [nfe % 3]
|
||||
|
||||
for i in range(len(orders)):
|
||||
eps_cache = {}
|
||||
t, t_next = ts[i], ts[i + 1]
|
||||
if eta:
|
||||
sd, su = get_ancestral_step(self.sigma(t), self.sigma(t_next), eta)
|
||||
t_next_ = torch.minimum(t_end, self.t(sd))
|
||||
su = (self.sigma(t_next) ** 2 - self.sigma(t_next_) ** 2) ** 0.5
|
||||
else:
|
||||
t_next_, su = t_next, 0.
|
||||
|
||||
eps, eps_cache = self.eps(eps_cache, 'eps', x, t)
|
||||
denoised = x - self.sigma(t) * eps
|
||||
if self.info_callback is not None:
|
||||
self.info_callback({'x': x, 'i': i, 't': ts[i], 't_up': t, 'denoised': denoised})
|
||||
|
||||
if orders[i] == 1:
|
||||
x, eps_cache = self.dpm_solver_1_step(x, t, t_next_, eps_cache=eps_cache)
|
||||
elif orders[i] == 2:
|
||||
x, eps_cache = self.dpm_solver_2_step(x, t, t_next_, eps_cache=eps_cache)
|
||||
else:
|
||||
x, eps_cache = self.dpm_solver_3_step(x, t, t_next_, eps_cache=eps_cache)
|
||||
|
||||
x = x + su * s_noise * noise_sampler(self.sigma(t), self.sigma(t_next))
|
||||
|
||||
return x
|
||||
|
||||
def dpm_solver_adaptive(self, x, t_start, t_end, order=3, rtol=0.05, atol=0.0078, h_init=0.05, pcoeff=0., icoeff=1., dcoeff=0., accept_safety=0.81, eta=0., s_noise=1., noise_sampler=None):
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
if order not in {2, 3}:
|
||||
raise ValueError('order should be 2 or 3')
|
||||
forward = t_end > t_start
|
||||
if not forward and eta:
|
||||
raise ValueError('eta must be 0 for reverse sampling')
|
||||
h_init = abs(h_init) * (1 if forward else -1)
|
||||
atol = torch.tensor(atol)
|
||||
rtol = torch.tensor(rtol)
|
||||
s = t_start
|
||||
x_prev = x
|
||||
accept = True
|
||||
pid = PIDStepSizeController(h_init, pcoeff, icoeff, dcoeff, 1.5 if eta else order, accept_safety)
|
||||
info = {'steps': 0, 'nfe': 0, 'n_accept': 0, 'n_reject': 0}
|
||||
|
||||
while s < t_end - 1e-5 if forward else s > t_end + 1e-5:
|
||||
eps_cache = {}
|
||||
t = torch.minimum(t_end, s + pid.h) if forward else torch.maximum(t_end, s + pid.h)
|
||||
if eta:
|
||||
sd, su = get_ancestral_step(self.sigma(s), self.sigma(t), eta)
|
||||
t_ = torch.minimum(t_end, self.t(sd))
|
||||
su = (self.sigma(t) ** 2 - self.sigma(t_) ** 2) ** 0.5
|
||||
else:
|
||||
t_, su = t, 0.
|
||||
|
||||
eps, eps_cache = self.eps(eps_cache, 'eps', x, s)
|
||||
denoised = x - self.sigma(s) * eps
|
||||
|
||||
if order == 2:
|
||||
x_low, eps_cache = self.dpm_solver_1_step(x, s, t_, eps_cache=eps_cache)
|
||||
x_high, eps_cache = self.dpm_solver_2_step(x, s, t_, eps_cache=eps_cache)
|
||||
else:
|
||||
x_low, eps_cache = self.dpm_solver_2_step(x, s, t_, r1=1 / 3, eps_cache=eps_cache)
|
||||
x_high, eps_cache = self.dpm_solver_3_step(x, s, t_, eps_cache=eps_cache)
|
||||
delta = torch.maximum(atol, rtol * torch.maximum(x_low.abs(), x_prev.abs()))
|
||||
error = torch.linalg.norm((x_low - x_high) / delta) / x.numel() ** 0.5
|
||||
accept = pid.propose_step(error)
|
||||
if accept:
|
||||
x_prev = x_low
|
||||
x = x_high + su * s_noise * noise_sampler(self.sigma(s), self.sigma(t))
|
||||
s = t
|
||||
info['n_accept'] += 1
|
||||
else:
|
||||
info['n_reject'] += 1
|
||||
info['nfe'] += order
|
||||
info['steps'] += 1
|
||||
|
||||
if self.info_callback is not None:
|
||||
self.info_callback({'x': x, 'i': info['steps'] - 1, 't': s, 't_up': s, 'denoised': denoised, 'error': error, 'h': pid.h, **info})
|
||||
|
||||
return x, info
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None):
|
||||
"""Ancestral sampling with DPM-Solver++(2S) second-order steps."""
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sigma_fn = lambda t: t.neg().exp() # pylint: disable=C3001
|
||||
t_fn = lambda sigma: sigma.log().neg() # pylint: disable=C3001
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
if sigma_down == 0:
|
||||
# Euler method
|
||||
d = to_d(x, sigmas[i], denoised)
|
||||
dt = sigma_down - sigmas[i]
|
||||
x = x + d * dt
|
||||
else:
|
||||
# DPM-Solver++(2S)
|
||||
t, t_next = t_fn(sigmas[i]), t_fn(sigma_down)
|
||||
r = 1 / 2
|
||||
h = t_next - t
|
||||
s = t + r * h
|
||||
x_2 = (sigma_fn(s) / sigma_fn(t)) * x - (-h * r).expm1() * denoised
|
||||
denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args)
|
||||
x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_2
|
||||
# Noise addition
|
||||
if sigmas[i + 1] > 0:
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
|
||||
if mask_args is not None:
|
||||
x = inpaint_mask(x, i, len(sigmas) - 2, mask_args)
|
||||
return x
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_dpmpp_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1 / 2, mask_args=None):
|
||||
"""DPM-Solver++ (stochastic)."""
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max) if noise_sampler is None else noise_sampler
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sigma_fn = lambda t: t.neg().exp() # pylint: disable=C3001
|
||||
t_fn = lambda sigma: sigma.log().neg() # pylint: disable=C3001
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
if sigmas[i + 1] == 0:
|
||||
# Euler method
|
||||
d = to_d(x, sigmas[i], denoised)
|
||||
dt = sigmas[i + 1] - sigmas[i]
|
||||
x = x + d * dt
|
||||
else:
|
||||
# DPM-Solver++
|
||||
t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1])
|
||||
h = t_next - t
|
||||
s = t + h * r
|
||||
fac = 1 / (2 * r)
|
||||
|
||||
# Step 1
|
||||
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta)
|
||||
s_ = t_fn(sd)
|
||||
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - (t - s_).expm1() * denoised
|
||||
x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su
|
||||
denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args)
|
||||
|
||||
# Step 2
|
||||
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta)
|
||||
t_next_ = t_fn(sd)
|
||||
denoised_d = (1 - fac) * denoised + fac * denoised_2
|
||||
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - (t - t_next_).expm1() * denoised_d
|
||||
x = x + noise_sampler(sigma_fn(t), sigma_fn(t_next)) * s_noise * su
|
||||
if mask_args is not None:
|
||||
x = inpaint_mask(x, i, len(sigmas) - 2, mask_args)
|
||||
return x
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=None, mask_args=None):
|
||||
"""DPM-Solver++(2M)."""
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
sigma_fn = lambda t: t.neg().exp() # pylint: disable=C3001
|
||||
t_fn = lambda sigma: sigma.log().neg() # pylint: disable=C3001
|
||||
old_denoised = None
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1])
|
||||
h = t_next - t
|
||||
if old_denoised is None or sigmas[i + 1] == 0:
|
||||
x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised
|
||||
else:
|
||||
h_last = t - t_fn(sigmas[i - 1])
|
||||
r = h_last / h
|
||||
denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised
|
||||
x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_d
|
||||
old_denoised = denoised
|
||||
if mask_args is not None:
|
||||
x = inpaint_mask(x, i, len(sigmas) - 2, mask_args)
|
||||
return x
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_dpmpp_2m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, solver_type='midpoint', mask_args=None):
|
||||
"""DPM-Solver++(2M) SDE."""
|
||||
|
||||
if solver_type not in {'heun', 'midpoint'}:
|
||||
raise ValueError('solver_type must be \'heun\' or \'midpoint\'')
|
||||
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max) if noise_sampler is None else noise_sampler
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
|
||||
old_denoised = None
|
||||
h_last = None
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
# DPM-Solver++(2M) SDE
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
eta_h = eta * h
|
||||
|
||||
x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * denoised
|
||||
|
||||
if old_denoised is not None:
|
||||
r = h_last / h
|
||||
if solver_type == 'heun':
|
||||
x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (denoised - old_denoised)
|
||||
elif solver_type == 'midpoint':
|
||||
x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * (denoised - old_denoised)
|
||||
|
||||
if eta:
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise
|
||||
|
||||
old_denoised = denoised
|
||||
h_last = h
|
||||
if mask_args is not None:
|
||||
x = inpaint_mask(x, i, len(sigmas) - 2, mask_args)
|
||||
return x
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_dpmpp_3m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None):
|
||||
"""DPM-Solver++(3M) SDE."""
|
||||
|
||||
sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()
|
||||
noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max) if noise_sampler is None else noise_sampler
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
|
||||
denoised_1, denoised_2 = None, None
|
||||
h_1, h_2 = None, None
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||
if callback is not None:
|
||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||
if sigmas[i + 1] == 0:
|
||||
# Denoising step
|
||||
x = denoised
|
||||
else:
|
||||
t, s = -sigmas[i].log(), -sigmas[i + 1].log()
|
||||
h = s - t
|
||||
h_eta = h * (eta + 1)
|
||||
|
||||
x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised
|
||||
|
||||
if h_2 is not None:
|
||||
r0 = h_1 / h
|
||||
r1 = h_2 / h
|
||||
d1_0 = (denoised - denoised_1) / r0
|
||||
d1_1 = (denoised_1 - denoised_2) / r1
|
||||
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
|
||||
d2 = (d1_0 - d1_1) / (r0 + r1)
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
phi_3 = phi_2 / h_eta - 0.5
|
||||
x = x + phi_2 * d1 - phi_3 * d2
|
||||
elif h_1 is not None:
|
||||
r = h_1 / h
|
||||
d = (denoised - denoised_1) / r
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
x = x + phi_2 * d
|
||||
|
||||
if eta:
|
||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * h * eta).expm1().neg().sqrt() * s_noise
|
||||
|
||||
denoised_1, denoised_2 = denoised, denoised_1
|
||||
h_1, h_2 = h, h_1
|
||||
if mask_args is not None:
|
||||
x = inpaint_mask(x, i, len(sigmas) - 2, mask_args)
|
||||
return x
|
||||
@@ -0,0 +1,450 @@
|
||||
from typing import Union
|
||||
import os
|
||||
import cv2
|
||||
import insightface
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from PIL import Image
|
||||
from diffusers import StableDiffusionXLPipeline, StableDiffusionXLImg2ImgPipeline, StableDiffusionXLInpaintPipeline
|
||||
from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput
|
||||
|
||||
from huggingface_hub import hf_hub_download, snapshot_download
|
||||
from safetensors.torch import load_file
|
||||
from torchvision.transforms import InterpolationMode
|
||||
from torchvision.transforms.functional import normalize, resize
|
||||
|
||||
from basicsr.utils import img2tensor, tensor2img
|
||||
from facexlib.parsing import init_parsing_model
|
||||
from facexlib.utils.face_restoration_helper import FaceRestoreHelper
|
||||
from insightface.app import FaceAnalysis
|
||||
|
||||
from eva_clip import create_model_and_transforms
|
||||
from eva_clip.constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD
|
||||
from encoders_transformer import IDFormer, IDEncoder
|
||||
from modules.errors import log
|
||||
|
||||
|
||||
debug = log.trace if os.environ.get('SD_PULID_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
|
||||
|
||||
class StableDiffusionXLPuLIDPipeline:
|
||||
def __init__(self,
|
||||
pipe: Union[StableDiffusionXLPipeline, StableDiffusionXLImg2ImgPipeline, StableDiffusionXLInpaintPipeline],
|
||||
device: torch.device,
|
||||
dtype: torch.dtype=None,
|
||||
providers: list=None,
|
||||
offload: bool=True,
|
||||
sampler=None,
|
||||
cache_dir=None,
|
||||
sdp: bool=True,
|
||||
version: str='v1.1',
|
||||
):
|
||||
super().__init__()
|
||||
self.device = device
|
||||
self.dtype = dtype or torch.float16
|
||||
self.pipe = pipe
|
||||
self.cache_dir = cache_dir
|
||||
self.offload = offload
|
||||
self.sdp = sdp
|
||||
self.version = version
|
||||
self.folder = 'models--ToTheBeginning--PuLID'
|
||||
debug(f'PulID init: device={self.device} dtype={self.dtype} dir={self.cache_dir} offload={self.offload} sdp={self.sdp} version={self.version}')
|
||||
|
||||
# self.pipe.scheduler = DPMSolverMultistepScheduler.from_config(self.pipe.scheduler.config)
|
||||
self.hack_unet_attn_layers(self.pipe.unet)
|
||||
if self.version == 'v1.1':
|
||||
self.id_adapter = IDFormer().to(self.device, self.dtype)
|
||||
else:
|
||||
self.id_adapter = IDEncoder().to(self.device, self.dtype)
|
||||
debug(f'PulID load: adapter={self.id_adapter.__class__.__name__}')
|
||||
self.providers = providers or ['CUDAExecutionProvider', 'CPUExecutionProvider']
|
||||
debug(f'PulID load: providers={self.providers}')
|
||||
|
||||
# preprocessors
|
||||
# face align and parsing
|
||||
self.face_helper = FaceRestoreHelper(
|
||||
upscale_factor=1,
|
||||
face_size=512,
|
||||
crop_ratio=(1, 1),
|
||||
det_model='retinaface_resnet50',
|
||||
save_ext='png',
|
||||
device=self.device,
|
||||
)
|
||||
self.face_helper.face_parse = init_parsing_model(model_name='bisenet', device=self.device)
|
||||
debug(f'PulID load: facehelper={self.face_helper.__class__.__name__}')
|
||||
|
||||
# clip-vit backbone
|
||||
eva_precision = 'fp16' if self.dtype == torch.float16 or self.dtype == torch.bfloat16 else 'fp32'
|
||||
eva_model, _, _ = create_model_and_transforms('EVA02-CLIP-L-14-336', 'eva_clip', force_custom_clip=True, precision=eva_precision, device=self.device)
|
||||
self.clip_vision_model = eva_model.visual.to(dtype=self.dtype)
|
||||
debug(f'PulID load: evaclip={self.clip_vision_model.__class__.__name__} precision={eva_precision}')
|
||||
eva_transform_mean = getattr(self.clip_vision_model, 'image_mean', OPENAI_DATASET_MEAN)
|
||||
eva_transform_std = getattr(self.clip_vision_model, 'image_std', OPENAI_DATASET_STD)
|
||||
if not isinstance(eva_transform_mean, (list, tuple)):
|
||||
eva_transform_mean = (eva_transform_mean,) * 3
|
||||
if not isinstance(eva_transform_std, (list, tuple)):
|
||||
eva_transform_std = (eva_transform_std,) * 3
|
||||
self.eva_transform_mean = eva_transform_mean
|
||||
self.eva_transform_std = eva_transform_std
|
||||
|
||||
# antelopev2
|
||||
local_dir = os.path.join(self.cache_dir, self.folder, 'models', 'antelopev2')
|
||||
_loc = snapshot_download('DIAMONIK7777/antelopev2', local_dir=local_dir)
|
||||
self.app = FaceAnalysis(
|
||||
name='antelopev2',
|
||||
root=os.path.join(self.cache_dir, self.folder),
|
||||
providers=self.providers,
|
||||
)
|
||||
debug(f'PulID load: faceanalysis={_loc}')
|
||||
self.app.prepare(ctx_id=0, det_size=(640, 640))
|
||||
self.handler_ante = insightface.model_zoo.get_model(os.path.join(local_dir, 'glintr100.onnx'))
|
||||
self.handler_ante.prepare(ctx_id=0)
|
||||
debug(f'PulID load: handler={self.handler_ante.__class__.__name__}')
|
||||
|
||||
self.load_pretrain()
|
||||
|
||||
# other configs
|
||||
self.debug_img_list = []
|
||||
|
||||
# karras schedule related code, borrow from lllyasviel/Omost
|
||||
linear_start = 0.00085
|
||||
linear_end = 0.012
|
||||
timesteps = 1000
|
||||
betas = torch.linspace(linear_start**0.5, linear_end**0.5, timesteps, dtype=torch.float64) ** 2
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = torch.tensor(np.cumprod(alphas, axis=0), dtype=torch.float32)
|
||||
|
||||
self.sigmas = ((1 - alphas_cumprod) / alphas_cumprod) ** 0.5
|
||||
self.log_sigmas = self.sigmas.log()
|
||||
self.sigma_data = 1.0
|
||||
|
||||
# default scheduler
|
||||
if sampler is not None:
|
||||
self.sampler = sampler
|
||||
else:
|
||||
from modules.pulid import sampling
|
||||
self.sampler = sampling.sample_dpmpp_sde
|
||||
|
||||
@property
|
||||
def sigma_min(self):
|
||||
return self.sigmas[0]
|
||||
|
||||
@property
|
||||
def sigma_max(self):
|
||||
return self.sigmas[-1]
|
||||
|
||||
def timestep(self, sigma):
|
||||
log_sigma = sigma.log()
|
||||
dists = log_sigma.to(self.log_sigmas.device) - self.log_sigmas[:, None]
|
||||
return dists.abs().argmin(dim=0).view(sigma.shape).to(sigma.device)
|
||||
|
||||
def get_sigmas_karras(self, n, rho=7.0):
|
||||
ramp = torch.linspace(0, 1, n)
|
||||
min_inv_rho = self.sigma_min ** (1 / rho)
|
||||
max_inv_rho = self.sigma_max ** (1 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
|
||||
return torch.cat([sigmas, sigmas.new_zeros([1])])
|
||||
|
||||
def hack_unet_attn_layers(self, unet):
|
||||
if self.sdp:
|
||||
from attention_processor import AttnProcessor2_0 as AttnProcessor
|
||||
from attention_processor import IDAttnProcessor2_0 as IDAttnProcessor
|
||||
else:
|
||||
from attention_processor import AttnProcessor
|
||||
from attention_processor import IDAttnProcessor
|
||||
id_adapter_attn_procs = {}
|
||||
for name, _ in unet.attn_processors.items():
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
else:
|
||||
hidden_size = None
|
||||
if cross_attention_dim is not None:
|
||||
id_adapter_attn_procs[name] = IDAttnProcessor(
|
||||
hidden_size=hidden_size,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
).to(unet.device, unet.dtype)
|
||||
else:
|
||||
id_adapter_attn_procs[name] = AttnProcessor()
|
||||
debug(f'PulID attention: cls={IDAttnProcessor} std={AttnProcessor} len={len(id_adapter_attn_procs.keys())}')
|
||||
unet.set_attn_processor(id_adapter_attn_procs)
|
||||
self.id_adapter_attn_layers = nn.ModuleList(unet.attn_processors.values())
|
||||
|
||||
def load_pretrain(self):
|
||||
if self.version == 'v1.1':
|
||||
ckpt_path = hf_hub_download('guozinan/PuLID', 'pulid_v1.1.safetensors', local_dir=os.path.join(self.cache_dir, self.folder))
|
||||
state_dict = load_file(ckpt_path)
|
||||
else:
|
||||
ckpt_path = hf_hub_download('guozinan/PuLID', 'pulid_v1.bin', local_dir=os.path.join(self.cache_dir, self.folder))
|
||||
state_dict = torch.load(ckpt_path, map_location="cpu")
|
||||
debug(f'PulID load: fn="{ckpt_path}"')
|
||||
state_dict_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
module = k.split('.')[0]
|
||||
state_dict_dict.setdefault(module, {})
|
||||
new_k = k[len(module) + 1 :]
|
||||
state_dict_dict[module][new_k] = v.to(self.dtype)
|
||||
|
||||
for module in state_dict_dict:
|
||||
getattr(self, module).load_state_dict(state_dict_dict[module], strict=True)
|
||||
|
||||
def to_gray(self, img):
|
||||
x = 0.299 * img[:, 0:1] + 0.587 * img[:, 1:2] + 0.114 * img[:, 2:3]
|
||||
x = x.repeat(1, 3, 1, 1)
|
||||
return x
|
||||
|
||||
def get_id_embedding(self, image_list):
|
||||
"""
|
||||
Args:
|
||||
image in image_list: numpy rgb image, range [0, 255]
|
||||
"""
|
||||
id_cond_list = []
|
||||
id_vit_hidden_list = []
|
||||
self.face_helper.face_det.to(self.device)
|
||||
self.clip_vision_model.to(self.device)
|
||||
for _ii, image in enumerate(image_list):
|
||||
self.face_helper.clean_all()
|
||||
image_bgr = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
|
||||
# get antelopev2 embedding
|
||||
face_info = self.app.get(image_bgr)
|
||||
if len(face_info) > 0:
|
||||
face_info = sorted(face_info, key=lambda x: (x['bbox'][2] - x['bbox'][0]) * (x['bbox'][3] - x['bbox'][1]))[-1] # only use the maximum face
|
||||
id_ante_embedding = face_info['embedding']
|
||||
self.debug_img_list.append(image[int(face_info['bbox'][1]) : int(face_info['bbox'][3]), int(face_info['bbox'][0]) : int(face_info['bbox'][2])])
|
||||
else:
|
||||
id_ante_embedding = None
|
||||
|
||||
# using facexlib to detect and align face
|
||||
self.face_helper.read_image(image_bgr)
|
||||
self.face_helper.get_face_landmarks_5(only_center_face=True)
|
||||
self.face_helper.align_warp_face()
|
||||
if len(self.face_helper.cropped_faces) == 0:
|
||||
raise RuntimeError('facexlib align face fail')
|
||||
align_face = self.face_helper.cropped_faces[0]
|
||||
# incase insightface didn't detect face
|
||||
if id_ante_embedding is None:
|
||||
id_ante_embedding = self.handler_ante.get_feat(align_face)
|
||||
|
||||
id_ante_embedding = torch.from_numpy(id_ante_embedding).to(self.device)
|
||||
if id_ante_embedding.ndim == 1:
|
||||
id_ante_embedding = id_ante_embedding.unsqueeze(0)
|
||||
|
||||
# parsing
|
||||
input = img2tensor(align_face, bgr2rgb=True).unsqueeze(0) / 255.0 # pylint: disable=redefined-builtin
|
||||
input = input.to(self.device)
|
||||
parsing_out = self.face_helper.face_parse(normalize(input, [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]))[0]
|
||||
parsing_out = parsing_out.argmax(dim=1, keepdim=True)
|
||||
bg_label = [0, 16, 18, 7, 8, 9, 14, 15]
|
||||
bg = sum(parsing_out == i for i in bg_label).bool()
|
||||
white_image = torch.ones_like(input)
|
||||
# only keep the face features
|
||||
face_features_image = torch.where(bg, white_image, self.to_gray(input))
|
||||
self.debug_img_list.append(tensor2img(face_features_image, rgb2bgr=False))
|
||||
|
||||
# transform img before sending to eva-clip-vit
|
||||
face_features_image = resize(face_features_image, self.clip_vision_model.image_size, InterpolationMode.BICUBIC)
|
||||
face_features_image = normalize(face_features_image, self.eva_transform_mean, self.eva_transform_std).to(self.dtype)
|
||||
id_cond_vit, id_vit_hidden = self.clip_vision_model(face_features_image, return_all_features=False, return_hidden=True, shuffle=False)
|
||||
id_cond_vit_norm = torch.norm(id_cond_vit, 2, 1, True)
|
||||
id_cond_vit = torch.div(id_cond_vit, id_cond_vit_norm)
|
||||
|
||||
id_cond = torch.cat([id_ante_embedding, id_cond_vit], dim=-1)
|
||||
|
||||
id_cond_list.append(id_cond)
|
||||
id_vit_hidden_list.append(id_vit_hidden)
|
||||
|
||||
self.id_adapter.to(self.device)
|
||||
id_uncond = torch.zeros_like(id_cond_list[0]).to(self.dtype)
|
||||
id_vit_hidden_uncond = []
|
||||
for layer_idx in range(0, len(id_vit_hidden_list[0])):
|
||||
id_vit_hidden_uncond.append(torch.zeros_like(id_vit_hidden_list[0][layer_idx]).to(self.dtype))
|
||||
|
||||
id_cond = torch.stack(id_cond_list, dim=1).to(self.dtype)
|
||||
id_vit_hidden = id_vit_hidden_list[0]
|
||||
for i in range(1, len(image_list)):
|
||||
for j, x in enumerate(id_vit_hidden_list[i]):
|
||||
id_vit_hidden[j] = torch.cat([id_vit_hidden[j], x], dim=1).to(self.dtype)
|
||||
id_embedding = self.id_adapter(id_cond, id_vit_hidden)
|
||||
uncond_id_embedding = self.id_adapter(id_uncond, id_vit_hidden_uncond)
|
||||
|
||||
if self.offload:
|
||||
self.face_helper.face_det.to('cpu')
|
||||
self.id_adapter.to('cpu')
|
||||
self.clip_vision_model.to('cpu')
|
||||
|
||||
# return id_embedding
|
||||
debug(f'PulID embedding: cond={id_embedding.shape} uncond={uncond_id_embedding.shape}')
|
||||
return uncond_id_embedding, id_embedding
|
||||
|
||||
def set_progress_bar_config(self, bar_format: str = None, ncols: int = 80, colour: str = None):
|
||||
import functools
|
||||
from tqdm.auto import trange as trange_orig
|
||||
import pulid_sampling
|
||||
pulid_sampling.trange = functools.partial(trange_orig, bar_format=bar_format, ncols=ncols, colour=colour)
|
||||
|
||||
def sample(self, x, sigma, **extra_args):
|
||||
t = self.timestep(sigma)
|
||||
x_ddim_space = x / (sigma[:, None, None, None] ** 2 + self.sigma_data**2) ** 0.5
|
||||
cfg_scale = extra_args['cfg_scale']
|
||||
# debug(f'PulID sample start: step={self.step+1} x={x.shape} dtype={x.dtype} timestep={t.item()} sigma={sigma.shape} cfg={cfg_scale} args={extra_args.keys()}')
|
||||
eps_positive = self.pipe.unet(x_ddim_space, t, return_dict=False, **extra_args['positive'])[0]
|
||||
eps_negative = self.pipe.unet(x_ddim_space, t, return_dict=False, **extra_args['negative'])[0]
|
||||
noise_pred = eps_negative + cfg_scale * (eps_positive - eps_negative)
|
||||
latent = x - noise_pred * sigma[:, None, None, None]
|
||||
if self.callback_on_step_end is not None:
|
||||
self.step += 1
|
||||
self.callback_on_step_end(self.pipe, step=self.step, timestep=t, kwargs={ 'latents': latent })
|
||||
# debug(f'PulID sample end: step={self.step} x={latent.shape} dtype={x.dtype} min={torch.amin(latent)} max={torch.amax(latent)}')
|
||||
return latent
|
||||
|
||||
def init_latent(self, seed, size, image, mask_image, strength, width, height): # pylint: disable=unused-argument
|
||||
# standard txt2img will full noise
|
||||
noise = torch.randn((size[0], 4, size[1] // 8, size[2] // 8), device="cpu", generator=torch.manual_seed(seed))
|
||||
noise = noise.to(dtype=self.pipe.unet.dtype, device=self.device)
|
||||
if strength > 0 and image is not None:
|
||||
image = self.pipe.image_processor.preprocess(image)
|
||||
if mask_image is not None: # Inpaint
|
||||
latents = self.pipe.prepare_latents(1, # batch_size,
|
||||
self.pipe.vae.config.latent_channels, # num_channels_latents
|
||||
height,
|
||||
width,
|
||||
noise.dtype,
|
||||
noise.device,
|
||||
None, # generator
|
||||
latents=None,
|
||||
image=image,
|
||||
timestep=1000,
|
||||
is_strength_max=False,
|
||||
add_noise=False,
|
||||
return_noise=False,
|
||||
return_image_latents=False,
|
||||
)
|
||||
latents = latents[0]
|
||||
debug(f'PulID noise: op=inpaint latent={latents.shape} image={image} mask={mask_image} dtype={latents.dtype}')
|
||||
else: # img2img
|
||||
latents = self.pipe.prepare_latents(image,
|
||||
None, # timestep (not needed)
|
||||
1, # batch_size
|
||||
1, # num_images_per_prompt
|
||||
noise.dtype,
|
||||
noise.device,
|
||||
None, # generator
|
||||
False, # add_noise
|
||||
)
|
||||
debug(f'PulID noise: op=img2img latent={latents.shape} image={image} dtype={latents.dtype}')
|
||||
else:
|
||||
latents = torch.zeros_like(noise)
|
||||
debug(f'PulID noise: op=txt2img latent={latents.shape} dtype={latents.dtype}')
|
||||
return latents, noise
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
prompt: str='',
|
||||
negative_prompt: str='',
|
||||
width: int=1024,
|
||||
height: int=1024,
|
||||
guidance_scale: float=7.0,
|
||||
num_inference_steps: int=50,
|
||||
seed: int=-1,
|
||||
image: np.ndarray=None,
|
||||
mask_image: np.ndarray=None,
|
||||
strength: float=0.3,
|
||||
id_embedding=None,
|
||||
uncond_id_embedding=None,
|
||||
id_scale: float=1.0,
|
||||
output_type: str='pil',
|
||||
callback_on_step_end=None,
|
||||
):
|
||||
debug(f'PulID call: width={width} height={height} cfg={guidance_scale} steps={num_inference_steps} seed={seed} strength={strength} id_scale={id_scale} output={output_type}')
|
||||
self.step = 0 # pylint: disable=attribute-defined-outside-init
|
||||
self.callback_on_step_end = callback_on_step_end # pylint: disable=attribute-defined-outside-init
|
||||
if isinstance(image, list) and len(image) > 0 and isinstance(image[0], Image.Image):
|
||||
if image[0].width != width or image[0].height != height: # override width/height if different
|
||||
width, height = image[0].width, image[0].height
|
||||
size = (1, height, width)
|
||||
# sigmas
|
||||
sigmas = self.get_sigmas_karras(num_inference_steps).to(self.device)
|
||||
if image is not None and strength > 0:
|
||||
_, num_inference_steps = self.pipe.get_timesteps(num_inference_steps, strength, self.device, None) # denoising_start disabled
|
||||
sigmas = sigmas[-(num_inference_steps + 1):].to(self.device) # shorten sigmas in i2i
|
||||
debug(f'PulID sigmas: sigmas={sigmas.shape} dtype={sigmas.dtype}')
|
||||
|
||||
# latents
|
||||
latent, noise = self.init_latent(seed, size, image, mask_image, strength, width, height)
|
||||
noisy_latent = latent + noise * sigmas[0].to(noise)
|
||||
debug(f'PulID noisy: latent={noisy_latent.shape} dtype={noisy_latent.dtype}')
|
||||
|
||||
(
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
pooled_prompt_embeds,
|
||||
negative_pooled_prompt_embeds,
|
||||
) = self.pipe.encode_prompt(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
)
|
||||
|
||||
add_time_ids = list((size[1], size[2]) + (0, 0) + (size[1], size[2]))
|
||||
add_time_ids = torch.tensor([add_time_ids], dtype=self.pipe.unet.dtype, device=self.device)
|
||||
add_neg_time_ids = add_time_ids.clone()
|
||||
|
||||
sampler_kwargs = dict(
|
||||
cfg_scale=guidance_scale,
|
||||
positive=dict(
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids},
|
||||
cross_attention_kwargs={'id_embedding': id_embedding, 'id_scale': id_scale},
|
||||
),
|
||||
negative=dict(
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
added_cond_kwargs={"text_embeds": negative_pooled_prompt_embeds, "time_ids": add_neg_time_ids},
|
||||
cross_attention_kwargs={'id_embedding': uncond_id_embedding, 'id_scale': id_scale},
|
||||
),
|
||||
)
|
||||
if mask_image is not None:
|
||||
latent_mask = torch.Tensor(np.asarray(mask_image.convert("L").resize((noisy_latent.shape[-1], noisy_latent.shape[-2])))).reshape((noisy_latent.shape[-2], noisy_latent.shape[-1]))
|
||||
latent_mask /= latent_mask.max()
|
||||
mask_args = dict(
|
||||
latent=latent,
|
||||
latent_mask=latent_mask,
|
||||
noise=noise,
|
||||
sigmas=sigmas,
|
||||
)
|
||||
else:
|
||||
mask_args = None
|
||||
|
||||
# actual sampling loop
|
||||
latents = self.sampler(self.sample, noisy_latent, sigmas, extra_args=sampler_kwargs, disable=False, mask_args=mask_args)
|
||||
|
||||
# process output
|
||||
latents = latents.to(dtype=self.pipe.vae.dtype, device=self.device)
|
||||
debug(f'PulID output: latent={latents.shape} dtype={latents.dtype}')
|
||||
if output_type == 'latent':
|
||||
images = self.pipe.image_processor.postprocess(latents, output_type='latent')
|
||||
elif output_type == 'np':
|
||||
images = self.pipe.image_processor.postprocess(latents, output_type='np')
|
||||
else:
|
||||
latents = latents / self.pipe.vae.config.scaling_factor
|
||||
images = self.pipe.vae.decode(latents).sample
|
||||
images = self.pipe.image_processor.postprocess(images, output_type='pil')
|
||||
debug(f'PulID output: type={type(images)} images={images.shape if hasattr(images, "shape") else images}')
|
||||
return StableDiffusionXLPipelineOutput(images)
|
||||
|
||||
|
||||
class StableDiffusionXLPuLIDPipelineImage(StableDiffusionXLPuLIDPipeline):
|
||||
def __init__(self, pipe: StableDiffusionXLPipeline, device: torch.device, sampler=None, cache_dir=None): # pylint: disable=useless-parent-delegation
|
||||
super().__init__(pipe, device, sampler, cache_dir)
|
||||
# we dont do anything special here, just having different class so task-type can be detected/assigned
|
||||
|
||||
|
||||
class StableDiffusionXLPuLIDPipelineInpaint(StableDiffusionXLPuLIDPipeline):
|
||||
def __init__(self, pipe: StableDiffusionXLPipeline, device: torch.device, sampler=None, cache_dir=None): # pylint: disable=useless-parent-delegation
|
||||
super().__init__(pipe, device, sampler, cache_dir)
|
||||
# we dont do anything special here, just having different class so task-type can be detected/assigned
|
||||
@@ -0,0 +1,161 @@
|
||||
import importlib
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from torchvision.utils import make_grid
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
|
||||
def seed_everything(seed):
|
||||
os.environ["PL_GLOBAL_SEED"] = str(seed)
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
|
||||
def instantiate_from_config(config):
|
||||
if "target" not in config:
|
||||
if config == '__is_first_stage__' or config == "__is_unconditional__":
|
||||
return None
|
||||
raise KeyError("Expected key `target` to instantiate.")
|
||||
return get_obj_from_str(config["target"])(**config.get("params", {}))
|
||||
|
||||
|
||||
def get_obj_from_str(string, reload=False):
|
||||
module, cls = string.rsplit(".", 1)
|
||||
if reload:
|
||||
module_imp = importlib.import_module(module)
|
||||
importlib.reload(module_imp)
|
||||
return getattr(importlib.import_module(module, package=None), cls)
|
||||
|
||||
|
||||
def drop_seq_token(seq, drop_rate=0.5):
|
||||
idx = torch.randperm(seq.size(1))
|
||||
num_keep_tokens = int(len(idx) * (1 - drop_rate))
|
||||
idx = idx[:num_keep_tokens]
|
||||
seq = seq[:, idx]
|
||||
return seq
|
||||
|
||||
|
||||
def import_model_class_from_model_name_or_path(
|
||||
pretrained_model_name_or_path: str, revision: str, subfolder: str = "text_encoder"
|
||||
):
|
||||
text_encoder_config = PretrainedConfig.from_pretrained(
|
||||
pretrained_model_name_or_path, subfolder=subfolder, revision=revision
|
||||
)
|
||||
model_class = text_encoder_config.architectures[0]
|
||||
|
||||
if model_class == "CLIPTextModel":
|
||||
from transformers import CLIPTextModel
|
||||
|
||||
return CLIPTextModel
|
||||
elif model_class == "CLIPTextModelWithProjection":
|
||||
from transformers import CLIPTextModelWithProjection
|
||||
|
||||
return CLIPTextModelWithProjection
|
||||
else:
|
||||
raise ValueError(f"{model_class} is not supported.")
|
||||
|
||||
|
||||
def resize_numpy_image_long(image, resize_long_edge=768):
|
||||
h, w = image.shape[:2]
|
||||
if max(h, w) <= resize_long_edge:
|
||||
return image
|
||||
k = resize_long_edge / max(h, w)
|
||||
h = int(h * k)
|
||||
w = int(w * k)
|
||||
image = cv2.resize(image, (w, h), interpolation=cv2.INTER_LANCZOS4)
|
||||
return image
|
||||
|
||||
|
||||
# from basicsr
|
||||
def img2tensor(imgs, bgr2rgb=True, float32=True):
|
||||
"""Numpy array to tensor.
|
||||
|
||||
Args:
|
||||
imgs (list[ndarray] | ndarray): Input images.
|
||||
bgr2rgb (bool): Whether to change bgr to rgb.
|
||||
float32 (bool): Whether to change to float32.
|
||||
|
||||
Returns:
|
||||
list[tensor] | tensor: Tensor images. If returned results only have
|
||||
one element, just return tensor.
|
||||
"""
|
||||
|
||||
def _totensor(img, bgr2rgb, float32):
|
||||
if img.shape[2] == 3 and bgr2rgb:
|
||||
if img.dtype == 'float64':
|
||||
img = img.astype('float32')
|
||||
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
img = torch.from_numpy(img.transpose(2, 0, 1))
|
||||
if float32:
|
||||
img = img.float()
|
||||
return img
|
||||
|
||||
if isinstance(imgs, list):
|
||||
return [_totensor(img, bgr2rgb, float32) for img in imgs]
|
||||
return _totensor(imgs, bgr2rgb, float32)
|
||||
|
||||
|
||||
def tensor2img(tensor, rgb2bgr=True, out_type=np.uint8, min_max=(0, 1)):
|
||||
"""Convert torch Tensors into image numpy arrays.
|
||||
|
||||
After clamping to [min, max], values will be normalized to [0, 1].
|
||||
|
||||
Args:
|
||||
tensor (Tensor or list[Tensor]): Accept shapes:
|
||||
1) 4D mini-batch Tensor of shape (B x 3/1 x H x W);
|
||||
2) 3D Tensor of shape (3/1 x H x W);
|
||||
3) 2D Tensor of shape (H x W).
|
||||
Tensor channel should be in RGB order.
|
||||
rgb2bgr (bool): Whether to change rgb to bgr.
|
||||
out_type (numpy type): output types. If ``np.uint8``, transform outputs
|
||||
to uint8 type with range [0, 255]; otherwise, float type with
|
||||
range [0, 1]. Default: ``np.uint8``.
|
||||
min_max (tuple[int]): min and max values for clamp.
|
||||
|
||||
Returns:
|
||||
(Tensor or list): 3D ndarray of shape (H x W x C) OR 2D ndarray of
|
||||
shape (H x W). The channel order is BGR.
|
||||
"""
|
||||
if not (torch.is_tensor(tensor) or (isinstance(tensor, list) and all(torch.is_tensor(t) for t in tensor))):
|
||||
raise TypeError(f'tensor or list of tensors expected, got {type(tensor)}')
|
||||
|
||||
if torch.is_tensor(tensor):
|
||||
tensor = [tensor]
|
||||
result = []
|
||||
for _tensor in tensor:
|
||||
_tensor = _tensor.squeeze(0).float().detach().cpu().clamp_(*min_max)
|
||||
_tensor = (_tensor - min_max[0]) / (min_max[1] - min_max[0])
|
||||
|
||||
n_dim = _tensor.dim()
|
||||
if n_dim == 4:
|
||||
img_np = make_grid(_tensor, nrow=int(math.sqrt(_tensor.size(0))), normalize=False).numpy()
|
||||
img_np = img_np.transpose(1, 2, 0)
|
||||
if rgb2bgr:
|
||||
img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR)
|
||||
elif n_dim == 3:
|
||||
img_np = _tensor.numpy()
|
||||
img_np = img_np.transpose(1, 2, 0)
|
||||
if img_np.shape[2] == 1: # gray image
|
||||
img_np = np.squeeze(img_np, axis=2)
|
||||
else:
|
||||
if rgb2bgr:
|
||||
img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR)
|
||||
elif n_dim == 2:
|
||||
img_np = _tensor.numpy()
|
||||
else:
|
||||
raise TypeError(f'Only support 4D, 3D or 2D tensor. But received with dimension: {n_dim}')
|
||||
if out_type == np.uint8:
|
||||
# Unlike MATLAB, numpy.unit8() WILL NOT round by default.
|
||||
img_np = (img_np * 255.0).round()
|
||||
img_np = img_np.astype(out_type)
|
||||
result.append(img_np)
|
||||
if len(result) == 1:
|
||||
result = result[0]
|
||||
return result
|
||||
@@ -0,0 +1,897 @@
|
||||
# Credits: @ukaprch <https://github.com/huggingface/diffusers/issues/9924>
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchsde
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
class BatchedBrownianTree:
|
||||
"""A wrapper around torchsde.BrownianTree that enables batches of entropy."""
|
||||
|
||||
def __init__(self, x, t0, t1, seed=None, **kwargs):
|
||||
t0, t1, self.sign = self.sort(t0, t1)
|
||||
w0 = kwargs.get("w0", torch.zeros_like(x))
|
||||
if seed is None:
|
||||
seed = torch.randint(0, 2**63 - 1, []).item()
|
||||
self.batched = True
|
||||
try:
|
||||
assert len(seed) == x.shape[0]
|
||||
w0 = w0[0]
|
||||
except TypeError:
|
||||
seed = [seed]
|
||||
self.batched = False
|
||||
self.trees = [
|
||||
torchsde.BrownianInterval(
|
||||
t0=t0,
|
||||
t1=t1,
|
||||
size=w0.shape,
|
||||
dtype=w0.dtype,
|
||||
device=w0.device,
|
||||
entropy=s,
|
||||
tol=1e-6,
|
||||
pool_size=24,
|
||||
halfway_tree=True,
|
||||
)
|
||||
for s in seed
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def sort(a, b):
|
||||
return (a, b, 1) if a < b else (b, a, -1)
|
||||
|
||||
def __call__(self, t0, t1):
|
||||
t0, t1, sign = self.sort(t0, t1)
|
||||
w = torch.stack([tree(t0, t1) for tree in self.trees]) * (self.sign * sign)
|
||||
return w if self.batched else w[0]
|
||||
|
||||
class BrownianTreeNoiseSampler:
|
||||
"""A noise sampler backed by a torchsde.BrownianTree.
|
||||
|
||||
Args:
|
||||
x (Tensor): The tensor whose shape, device and dtype to use to generate
|
||||
random samples.
|
||||
sigma_min (float): The low end of the valid interval.
|
||||
sigma_max (float): The high end of the valid interval.
|
||||
seed (int or List[int]): The random seed. If a list of seeds is
|
||||
supplied instead of a single integer, then the noise sampler will use one BrownianTree per batch item, each
|
||||
with its own seed.
|
||||
transform (callable): A function that maps sigma to the sampler's
|
||||
internal timestep.
|
||||
"""
|
||||
|
||||
def __init__(self, x, sigma_min, sigma_max, seed=None, transform=lambda x: x):
|
||||
self.transform = transform
|
||||
t0, t1 = self.transform(torch.as_tensor(sigma_min)), self.transform(torch.as_tensor(sigma_max))
|
||||
self.tree = BatchedBrownianTree(x, t0, t1, seed)
|
||||
|
||||
def __call__(self, sigma, sigma_next):
|
||||
t0, t1 = self.transform(torch.as_tensor(sigma)), self.transform(torch.as_tensor(sigma_next))
|
||||
return self.tree(t0, t1) / (t1 - t0).abs().sqrt()
|
||||
|
||||
@dataclass
|
||||
class FlowMatchDPMSolverMultistepSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||
denoising loop.
|
||||
"""
|
||||
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
`DPMSolverMultistepScheduler` is a fast dedicated high-order solver for diffusion ODEs.
|
||||
|
||||
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||
methods the library implements for all schedulers such as loading and saving.
|
||||
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
solver_order (`int`, defaults to 2):
|
||||
The DPMSolver order which can be `2` or `3`. It is recommended to use `solver_order=2` for guided
|
||||
sampling, and `solver_order=3` for unconditional sampling.
|
||||
thresholding (`bool`, defaults to `False`):
|
||||
Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such
|
||||
as Stable Diffusion.
|
||||
dynamic_thresholding_ratio (`float`, defaults to 0.995):
|
||||
The ratio for the dynamic thresholding method. Valid only when `thresholding=True`.
|
||||
sample_max_value (`float`, defaults to 1.0):
|
||||
The threshold value for dynamic thresholding. Valid only when `thresholding=True`.
|
||||
algorithm_type (`str`, defaults to `dpmsolver++2M`):
|
||||
Algorithm type for the solver; can be `dpmsolver2`, `dpmsolver2A`, `dpmsolver++2M`, `dpmsolver++2S`, `dpmsolver++sde`, `dpmsolver++2Msde`,
|
||||
or `dpmsolver++3Msde`.
|
||||
solver_type (`str`, defaults to `midpoint`):
|
||||
Solver type for the second-order solver; can be `midpoint` or `heun`. The solver type slightly affects the
|
||||
sample quality, especially for a small number of steps. It is recommended to use `midpoint` solvers.
|
||||
sigma_schedule (`str`, *optional*, defaults to None): Sigma schedule to compute the `sigmas`. Optionally, we use
|
||||
the schedule "karras" introduced in the EDM paper (https://arxiv.org/abs/2206.00364). Other acceptable values are
|
||||
"exponential". The exponential schedule was incorporated in this model: https://huggingface.co/stabilityai/cosxl.
|
||||
Other acceptable values are "lambdas". The uniform-logSNR for step sizes proposed by Lu's DPM-Solver in the
|
||||
noise schedule during the sampling process. The sigmas and time steps are determined according to a sequence of `lambda(t)`.
|
||||
use_noise_sampler for BrownianTreeNoiseSampler (only valid for `dpmsolver++2S`, `dpmsolver++sde`, `dpmsolver++2Msde`,
|
||||
or `dpmsolver++3Msde`): A noise sampler backed by a torchsde increasing the stability of convergence. Default strategy
|
||||
(random noise) has it jumping all over the place, but Brownian sampling is more stable. Utilizes the model generation seed provided.
|
||||
midpoint_ratio (`float`, *optional*, range: 0.4 to 0.6, default=0.5): Only valid for (`dpmsolver++sde`, `dpmsolver++2S`).
|
||||
Higher values may result in smoothing, more vivid colors and less noise at the expense of more detail and effect.
|
||||
s_noise (`float`, *optional*, defaults to 1.0): Sigma noise strength: range 0 - 1.1 (only valid for `dpmsolver++2S`, `dpmsolver++sde`,
|
||||
`dpmsolver++2Msde`, or `dpmsolver++3Msde`). The amount of additional noise to counteract loss of detail during sampling. A
|
||||
reasonable range is [1.000, 1.011]. Defaults to 1.0 from the original implementation.
|
||||
use_SD35_sigmas: (`bool` defaults to False for FLUX and True for SD3). Based on original interpretation of using beta values for determining sigmas.
|
||||
use_dynamic_shifting (`bool` defaults to False for SD3 and True for FLUX). When `True`, shift is ignored.
|
||||
shift (`float`, defaults to 3.0): The shift value for the timestep schedule for SD3 when not using dynamic shifting
|
||||
The remaining args are specific to Flux's dynamic shifting based on resolution
|
||||
"""
|
||||
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
solver_order: int = 2,
|
||||
thresholding: Optional[bool] = False,
|
||||
dynamic_thresholding_ratio: float = 0.995,
|
||||
sample_max_value: Optional[float] = 1.0,
|
||||
algorithm_type: str = "dpmsolver++2M",
|
||||
solver_type: str = "midpoint",
|
||||
sigma_schedule: Optional[str] = None,
|
||||
shift: float = 3.0,
|
||||
midpoint_ratio: Optional[float] = 0.5,
|
||||
s_noise: Optional[float] = 1.0,
|
||||
use_noise_sampler: Optional[bool] = True,
|
||||
use_SD35_sigmas: Optional[bool] = False,
|
||||
use_dynamic_shifting=False,
|
||||
base_shift: Optional[float] = 0.5,
|
||||
max_shift: Optional[float] = 1.15,
|
||||
base_image_seq_len: Optional[int] = 256,
|
||||
max_image_seq_len: Optional[int] = 4096,
|
||||
):
|
||||
# settings for DPM-Solver
|
||||
if algorithm_type not in ["dpmsolver2", "dpmsolver2A", "dpmsolver++2M", "dpmsolver++2S", "dpmsolver++sde", "dpmsolver++2Msde", "dpmsolver++3Msde"]:
|
||||
raise NotImplementedError(f"{algorithm_type} is not implemented for {self.__class__}")
|
||||
|
||||
if solver_type not in ["midpoint", "heun"]:
|
||||
raise NotImplementedError(f"{solver_type} is not implemented for {self.__class__}")
|
||||
|
||||
# setable values
|
||||
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
|
||||
sigmas = timesteps / num_train_timesteps
|
||||
if not use_dynamic_shifting:
|
||||
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
|
||||
self.timesteps = sigmas * num_train_timesteps
|
||||
self.h_last = None
|
||||
self.h_1 = None
|
||||
self.h_2 = None
|
||||
self.noise_sampler = None
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication
|
||||
self.model_outputs = [None] * solver_order
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||
"""
|
||||
return self._step_index
|
||||
|
||||
@property
|
||||
def begin_index(self):
|
||||
"""
|
||||
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||
"""
|
||||
return self._begin_index
|
||||
|
||||
def set_begin_index(self, begin_index: int = 0):
|
||||
"""
|
||||
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||
|
||||
Args:
|
||||
begin_index (`int`):
|
||||
The begin index for the scheduler.
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
||||
|
||||
def set_timesteps(self,
|
||||
num_inference_steps: int = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
mu: Optional[float] = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
"""
|
||||
if self.config.use_dynamic_shifting and mu is None:
|
||||
raise ValueError(" you have a pass a value for `mu` when `use_dynamic_shifting` is set to be `True`")
|
||||
|
||||
if sigmas is None:
|
||||
self.use_SD35_sigmas = True
|
||||
self.num_inference_steps = num_inference_steps
|
||||
sigmas1 = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps, dtype=np.float64)
|
||||
beta_start = 0.00085
|
||||
beta_end = 0.012
|
||||
betas = torch.linspace(beta_start**0.5, beta_end**0.5, self.config.num_train_timesteps, dtype=torch.float64) ** 2
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
sigmas = np.array(((1 - alphas_cumprod) / alphas_cumprod) ** 0.5)
|
||||
del alphas_cumprod
|
||||
del alphas
|
||||
del betas
|
||||
elif self.use_SD35_sigmas:
|
||||
num_inference_steps = len(sigmas)
|
||||
self.num_inference_steps = num_inference_steps
|
||||
sigmas1 = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps, dtype=np.float64)
|
||||
beta_start = 0.00085
|
||||
beta_end = 0.012
|
||||
betas = torch.linspace(beta_start**0.5, beta_end**0.5, self.config.num_train_timesteps, dtype=torch.float64) ** 2
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
sigmas = np.array(((1 - alphas_cumprod) / alphas_cumprod) ** 0.5)
|
||||
del alphas_cumprod
|
||||
del alphas
|
||||
del betas
|
||||
else:
|
||||
num_inference_steps = len(sigmas)
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
if self.config.sigma_schedule == "exponential":
|
||||
if self.use_SD35_sigmas:
|
||||
sigmas = np.flip(sigmas).copy()
|
||||
sigma_min = sigmas[-1]
|
||||
sigma_max = sigmas[0]
|
||||
sigmas = self._convert_to_exponential(sigma_min, sigma_max, num_inference_steps=num_inference_steps)
|
||||
OldRange = sigma_max - sigma_min
|
||||
NewRange = 1.0 - sigma_min
|
||||
sigmas = (((sigmas - sigma_min) * NewRange) / OldRange) + sigma_min
|
||||
del sigmas1
|
||||
else:
|
||||
sigma_min = sigmas[-1]
|
||||
sigma_max = sigmas[0]
|
||||
sigmas = self._convert_to_exponential(sigma_min, sigma_max, num_inference_steps=num_inference_steps)
|
||||
elif self.config.sigma_schedule == "karras":
|
||||
if self.use_SD35_sigmas:
|
||||
sigmas = np.flip(sigmas).copy()
|
||||
sigma_min = sigmas[-1]
|
||||
sigma_max = sigmas[0]
|
||||
sigmas = self._convert_to_karras(sigma_min, sigma_max, num_inference_steps=num_inference_steps)
|
||||
OldRange = sigma_max - sigma_min
|
||||
NewRange = 1.0 - sigma_min
|
||||
sigmas = (((sigmas - sigma_min) * NewRange) / OldRange) + sigma_min
|
||||
del sigmas1
|
||||
else:
|
||||
sigma_min = sigmas[-1]
|
||||
sigma_max = sigmas[0]
|
||||
sigmas = self._convert_to_karras(sigma_min, sigma_max, num_inference_steps=num_inference_steps)
|
||||
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float64, device=device)
|
||||
elif self.config.sigma_schedule == "lambdas":
|
||||
if self.use_SD35_sigmas:
|
||||
log_sigmas = np.log(sigmas)
|
||||
lambdas = np.flip(log_sigmas.copy())
|
||||
lambdas = self._convert_to_lu(in_lambdas=lambdas, num_inference_steps=num_inference_steps)
|
||||
sigmas = np.exp(lambdas)
|
||||
sigma_min = sigmas[-1]
|
||||
sigma_max = sigmas[0]
|
||||
OldRange = sigma_max - sigma_min
|
||||
NewRange = 1.0 - sigma_min
|
||||
sigmas = (((sigmas - sigma_min) * NewRange) / OldRange) + sigma_min
|
||||
del sigmas1
|
||||
del lambdas
|
||||
del log_sigmas
|
||||
else:
|
||||
log_sigmas = np.log(sigmas)
|
||||
lambdas = log_sigmas.copy()
|
||||
lambdas = self._convert_to_lu(in_lambdas=lambdas, num_inference_steps=num_inference_steps)
|
||||
sigmas = np.exp(lambdas)
|
||||
del lambdas
|
||||
del log_sigmas
|
||||
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float64, device=device)
|
||||
else:
|
||||
if self.use_SD35_sigmas:
|
||||
sigmas = np.flip(sigmas).copy()
|
||||
sigma_min = sigmas[-1]
|
||||
sigmas = np.linspace(1.0, sigma_min, num_inference_steps)
|
||||
del sigmas1
|
||||
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float64, device=device)
|
||||
|
||||
if self.config.use_dynamic_shifting:
|
||||
sigmas = self.time_shift(mu, 1.0, sigmas)
|
||||
else:
|
||||
sigmas = self.config.shift * sigmas / (1 + (self.config.shift - 1) * sigmas)
|
||||
|
||||
timesteps = sigmas * self.config.num_train_timesteps
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
|
||||
self.h_last = None
|
||||
self.h_1 = None
|
||||
self.h_2 = None
|
||||
self.noise_sampler = None
|
||||
self.model_outputs = [None] * self.config.solver_order
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample
|
||||
def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
"Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the
|
||||
prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by
|
||||
s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing
|
||||
pixels from saturation at each step. We find that dynamic thresholding results in significantly better
|
||||
photorealism as well as better image-text alignment, especially when using very large guidance weights."
|
||||
|
||||
https://arxiv.org/abs/2205.11487
|
||||
"""
|
||||
dtype = sample.dtype
|
||||
batch_size, channels, *remaining_dims = sample.shape
|
||||
|
||||
if dtype not in (torch.float32, torch.float64):
|
||||
sample = sample.float() # upcast for quantile calculation, and clamp not implemented for cpu half
|
||||
|
||||
# Flatten sample for doing quantile calculation along each image
|
||||
sample = sample.reshape(batch_size, channels * np.prod(remaining_dims))
|
||||
|
||||
abs_sample = sample.abs() # "a certain percentile absolute pixel value"
|
||||
|
||||
s = torch.quantile(abs_sample, self.config.dynamic_thresholding_ratio, dim=1)
|
||||
s = torch.clamp(
|
||||
s, min=1, max=self.config.sample_max_value
|
||||
) # When clamped to min=1, equivalent to standard clipping to [-1, 1]
|
||||
s = s.unsqueeze(1) # (batch_size, 1) because clamp will broadcast along dim=0
|
||||
sample = torch.clamp(sample, -s, s) / s # "we threshold xt0 to the range [-s, s] and then divide by s"
|
||||
|
||||
sample = sample.reshape(batch_size, channels, *remaining_dims)
|
||||
sample = sample.to(dtype)
|
||||
|
||||
return sample
|
||||
|
||||
def _convert_to_lu(self, in_lambdas: torch.Tensor, num_inference_steps) -> torch.Tensor:
|
||||
"""Constructs the noise schedule of Lu et al. (2022)."""
|
||||
|
||||
lambda_min: float = in_lambdas[-1].item()
|
||||
lambda_max: float = in_lambdas[0].item()
|
||||
|
||||
rho = 1.0 # 1.0 is the value used in the paper
|
||||
ramp = np.linspace(0, 1, num_inference_steps)
|
||||
min_inv_rho = lambda_min ** (1 / rho)
|
||||
max_inv_rho = lambda_max ** (1 / rho)
|
||||
lambdas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
|
||||
return lambdas
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
|
||||
def _convert_to_karras(self, sigma_min, sigma_max, num_inference_steps) -> torch.Tensor:
|
||||
rho = 7.0 # 7.0 is the value used in the paper
|
||||
ramp = np.linspace(0, 1, num_inference_steps)
|
||||
min_inv_rho = sigma_min ** (1 / rho)
|
||||
max_inv_rho = sigma_max ** (1 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
|
||||
return sigmas
|
||||
|
||||
def _convert_to_exponential(self, sigma_min, sigma_max, num_inference_steps) -> torch.Tensor:
|
||||
sigmas = torch.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps).exp()
|
||||
return sigmas
|
||||
|
||||
def convert_model_output(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
sample: torch.Tensor = None,
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Convert the model output to the corresponding type the DPMSolver/DPMSolver++ algorithm needs. DPM-Solver is
|
||||
designed to discretize an integral of the noise prediction model, and DPM-Solver++ is designed to discretize an
|
||||
integral of the data prediction model.
|
||||
|
||||
<Tip>
|
||||
|
||||
The algorithm and model type are decoupled. You can use either DPMSolver or DPMSolver++ for both noise
|
||||
prediction and data prediction models.
|
||||
|
||||
</Tip>
|
||||
|
||||
Args:
|
||||
model_output (`torch.Tensor`):
|
||||
The direct output from the learned diffusion model.
|
||||
sample (`torch.Tensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
The converted model output.
|
||||
"""
|
||||
timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None)
|
||||
if sample is None:
|
||||
if len(args) > 1:
|
||||
sample = args[1]
|
||||
else:
|
||||
raise ValueError("missing `sample` as a required keyward argument")
|
||||
|
||||
# Flow Match needs to solve an integral of the data prediction model.
|
||||
sigma = self.sigmas[self.step_index]
|
||||
x0_pred = sample - sigma * model_output
|
||||
|
||||
if self.config.thresholding:
|
||||
x0_pred = self._threshold_sample(x0_pred)
|
||||
|
||||
return x0_pred
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
indices = (schedule_timesteps == timestep).nonzero()
|
||||
|
||||
# The sigma index that is taken for the **very** first `step`
|
||||
# is always the second index (or the last index if there is only 1)
|
||||
# This way we can ensure we don't accidentally skip a sigma in
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
return indices[pos].item()
|
||||
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
self._step_index = self.index_for_timestep(timestep)
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
variance_noise: Optional[torch.FloatTensor] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[FlowMatchDPMSolverMultistepSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with
|
||||
the multistep DPMSolver.
|
||||
|
||||
Args:
|
||||
model_output (`torch.Tensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`int`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.Tensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
variance_noise (`torch.Tensor`):
|
||||
Alternative to generating noise with `generator` by directly providing the noise for the variance
|
||||
itself. Useful for methods such as [`LEdits++`].
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_utils.SchedulerOutput`] is returned, otherwise a
|
||||
tuple is returned where the first element is the sample tensor.
|
||||
|
||||
"""
|
||||
if self.num_inference_steps is None:
|
||||
raise ValueError(
|
||||
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
if self.config.algorithm_type in ["dpmsolver2", "dpmsolver2A"]:
|
||||
pass
|
||||
else:
|
||||
model_output = self.convert_model_output(model_output, sample=sample)
|
||||
for i in range(self.config.solver_order - 1):
|
||||
self.model_outputs[i] = self.model_outputs[i + 1]
|
||||
self.model_outputs[-1] = model_output
|
||||
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
if sample.dtype != model_output.dtype:
|
||||
sample = sample.to(model_output.dtype)
|
||||
|
||||
if self.config.algorithm_type in ["dpmsolver2A", "dpmsolver++2S", "dpmsolver++sde", "dpmsolver++2Msde", "dpmsolver++3Msde"] and variance_noise is None:
|
||||
# Create a noise sampler if it hasn't been created yet
|
||||
if self.config.use_noise_sampler:
|
||||
if self.noise_sampler is None:
|
||||
min_sigma, max_sigma = self.sigmas[self.sigmas > 0].min(), self.sigmas.max()
|
||||
self.noise_sampler = BrownianTreeNoiseSampler(sample, min_sigma, max_sigma, generator)
|
||||
else:
|
||||
noise = randn_tensor(model_output.shape, generator=generator, device=model_output.device, dtype=model_output.dtype)
|
||||
elif self.config.algorithm_type in ["dpmsolver2A", "dpmsolver++2S", "dpmsolver++sde", "dpmsolver++2Msde", "dpmsolver++3Msde"]:
|
||||
noise = variance_noise.to(device=model_output.device, dtype=model_output.dtype)
|
||||
else:
|
||||
noise = None
|
||||
|
||||
def sigma_fn(_t: torch.Tensor) -> torch.Tensor:
|
||||
return _t.neg().exp()
|
||||
def t_fn(_sigma: torch.Tensor) -> torch.Tensor:
|
||||
return _sigma.log().neg()
|
||||
sigma = self.sigmas[self.step_index]
|
||||
sigma_next = self.sigmas[self.step_index + 1]
|
||||
sigma_prev = self.sigmas[self.step_index - 1]
|
||||
if self.config.algorithm_type == "dpmsolver2":
|
||||
if self.config.solver_order == 2:
|
||||
if sigma_next == 0:
|
||||
# Euler method
|
||||
model_output = sample - sigma * model_output
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sigma_next - sigma
|
||||
sample = sample + d * dt
|
||||
else:
|
||||
# DPM-Solver2
|
||||
sigma_mid = sigma.log().lerp(sigma_next.log(), 0.5).exp()
|
||||
|
||||
#using epsilon for new model output:
|
||||
pred_original_sample = sample - sigma * model_output
|
||||
# 2. Convert to an ODE derivative for 1st order
|
||||
d = (sample - pred_original_sample) / sigma
|
||||
# 3. delta timestep
|
||||
dt = sigma_mid - sigma
|
||||
x_2 = sample + d * dt
|
||||
|
||||
#using epsilon for new model output:
|
||||
denoised_2 = x_2 - sigma_mid * model_output
|
||||
# 2. Convert to an ODE derivative for 2nd order
|
||||
d = (x_2 - denoised_2) / sigma_mid
|
||||
|
||||
# 3. delta timestep
|
||||
dt = sigma_next - sigma
|
||||
sample = sample + d * dt
|
||||
|
||||
del pred_original_sample
|
||||
del denoised_2
|
||||
del x_2
|
||||
del d
|
||||
elif self.config.algorithm_type == "dpmsolver2A":
|
||||
if self.config.solver_order == 2:
|
||||
# get ancestral step
|
||||
sigma_from = sigma
|
||||
sigma_to = sigma_next
|
||||
su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5)
|
||||
sd = (sigma_to**2 - su**2) ** 0.5
|
||||
if sd == 0:
|
||||
# Euler method
|
||||
model_output = sample - sigma * model_output
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sd - sigma
|
||||
sample = sample + d * dt
|
||||
else:
|
||||
# DPM-Solver2A
|
||||
sigma_mid = sigma.log().lerp(sd.log(), 0.5).exp()
|
||||
|
||||
#using epsilon for new model output:
|
||||
model_output = sample - sigma * model_output
|
||||
# 2. Convert to an ODE derivative for 1st order
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sd - sigma
|
||||
sample = sample + d * dt
|
||||
|
||||
#using epsilon for new model output:
|
||||
pred_original_sample = sample - sigma * model_output
|
||||
# 2. Convert to an ODE derivative for 1st order
|
||||
d = (sample - pred_original_sample) / sigma
|
||||
# 3. delta timestep
|
||||
dt_1 = sigma_mid - sigma
|
||||
x_2 = sample + d * dt_1
|
||||
|
||||
#using epsilon for new model output:
|
||||
denoised_2 = x_2 - sigma_mid * model_output
|
||||
# 2. Convert to an ODE derivative for 2nd order
|
||||
d_2 = (x_2 - denoised_2) / sigma_mid
|
||||
|
||||
# 3. delta timestep
|
||||
dt_2 = sd - sigma_mid
|
||||
sample = sample + d_2 * dt_2
|
||||
|
||||
if self.config.use_noise_sampler:
|
||||
sample = sample + self.noise_sampler(sigma, sigma_next) * self.config.s_noise * su
|
||||
else:
|
||||
sample = sample + noise * self.config.s_noise * su
|
||||
|
||||
del pred_original_sample
|
||||
del denoised_2
|
||||
del x_2
|
||||
del d
|
||||
elif self.config.algorithm_type == "dpmsolver++2M":
|
||||
if self.config.solver_order == 2:
|
||||
t, t_next = t_fn(sigma), t_fn(sigma_next)
|
||||
h = t_next - t
|
||||
if self.model_outputs[-2] is None or sigma_next == 0:
|
||||
sample = (sigma_fn(t_next) / sigma_fn(t)) * sample - (-h).expm1() * model_output
|
||||
else:
|
||||
# DPM-Solver++(2M)
|
||||
h_last = t - t_fn(sigma_prev)
|
||||
r = h_last / h
|
||||
denoised_d = (1 + 1 / (2 * r)) * model_output - (1 / (2 * r)) * self.model_outputs[-2]
|
||||
sample = (sigma_fn(t_next) / sigma_fn(t)) * sample - (-h).expm1() * denoised_d
|
||||
del denoised_d
|
||||
elif self.config.algorithm_type == "dpmsolver++2S":
|
||||
if self.config.solver_order == 2:
|
||||
# get ancestral step
|
||||
sigma_from = sigma
|
||||
sigma_to = sigma_next
|
||||
su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5)
|
||||
sd = (sigma_to**2 - su**2) ** 0.5
|
||||
if sd == 0:
|
||||
# Euler method
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sd - sigma
|
||||
sample = sample + d * dt
|
||||
else:
|
||||
# DPM-Solver++(2S)
|
||||
t, t_next = t_fn(sigma), t_fn(sd)
|
||||
r = self.config.midpoint_ratio
|
||||
h = t_next - t
|
||||
s = t + r * h
|
||||
|
||||
# Euler method
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sd - sigma
|
||||
sample = sample + d * dt
|
||||
|
||||
x_2 = (sigma_fn(s) / sigma_fn(t)) * sample - (-h * r).expm1() * model_output
|
||||
|
||||
#using epsilon for new model output:
|
||||
denoised_2 = x_2 - sigma_fn(s) * model_output
|
||||
# 2. Convert to an ODE derivative for 2nd order
|
||||
d = (x_2 - denoised_2) / sigma_fn(s)
|
||||
dt = sd - sigma_next
|
||||
sample = sample + d * dt
|
||||
|
||||
del x_2
|
||||
del denoised_2
|
||||
del d
|
||||
# Noise addition
|
||||
if sigma_next > 0:
|
||||
if self.config.use_noise_sampler:
|
||||
sample = sample + self.noise_sampler(sigma, sigma_next) * self.config.s_noise * su
|
||||
else:
|
||||
sample = sample + noise * self.config.s_noise * su
|
||||
elif self.config.algorithm_type == "dpmsolver++sde":
|
||||
if self.config.solver_order == 2:
|
||||
if sigma_next == 0:
|
||||
# Euler method
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sigma_next - sigma
|
||||
sample = sample + d * dt
|
||||
else:
|
||||
# DPM-Solver++(SDE)
|
||||
t, t_next = t_fn(sigma), t_fn(sigma_next)
|
||||
r = self.config.midpoint_ratio
|
||||
h = t_next - t
|
||||
s = t + r * h
|
||||
|
||||
# Euler method
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sigma_next - sigma
|
||||
sample = sample + d * dt
|
||||
|
||||
# Step 1
|
||||
# get ancestral step
|
||||
sigma_from = sigma_fn(t)
|
||||
sigma_to = sigma_fn(s)
|
||||
su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5)
|
||||
sd = (sigma_to**2 - su**2) ** 0.5
|
||||
|
||||
# Euler method
|
||||
d = (sample - model_output) / sigma
|
||||
dt = sd - sigma
|
||||
sample = sample + d * dt
|
||||
|
||||
s_ = t_fn(sd)
|
||||
x_2 = (sigma_fn(s_) / sigma_fn(t)) * sample - (t - s_).expm1() * model_output
|
||||
if self.config.use_noise_sampler:
|
||||
x_2 = x_2 + self.noise_sampler(sigma_fn(t), sigma_fn(s)) * self.config.s_noise * su
|
||||
else:
|
||||
x_2 = x_2 + noise * self.config.s_noise * su
|
||||
|
||||
# Step 2
|
||||
# get ancestral step
|
||||
sigma_from = sigma_fn(t)
|
||||
sigma_to = sigma_fn(t_next)
|
||||
su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5)
|
||||
sd = (sigma_to**2 - su**2) ** 0.5
|
||||
|
||||
#using epsilon for new model output:
|
||||
denoised_2 = x_2 - sigma_fn(s) * model_output
|
||||
# 2. Convert to an ODE derivative for 2nd order
|
||||
d = (x_2 - denoised_2) / sigma_fn(s)
|
||||
dt = sd - sigma_next
|
||||
sample = sample + d * dt
|
||||
|
||||
if self.config.use_noise_sampler:
|
||||
sample = sample + self.noise_sampler(sigma_fn(t), sigma_fn(t_next)) * self.config.s_noise * su
|
||||
else:
|
||||
sample = sample + noise * self.config.s_noise * su
|
||||
del x_2
|
||||
del denoised_2
|
||||
del d
|
||||
elif self.config.algorithm_type == "dpmsolver++2Msde":
|
||||
if self.config.solver_order == 2:
|
||||
if sigma_next == 0:
|
||||
sample = model_output
|
||||
else:
|
||||
# DPM-Solver++(2M) SDE
|
||||
t, s = -sigma.log(), -sigma_next.log()
|
||||
h = s - t
|
||||
eta_h = h * 1
|
||||
|
||||
# 3. Delta timestep
|
||||
dt = sigma_next - sigma
|
||||
sample = sample + model_output * dt
|
||||
|
||||
sample = sigma_next / sigma * (-eta_h).exp() * sample + (-h - eta_h).expm1().neg() * model_output
|
||||
|
||||
if self.model_outputs[-2] is not None:
|
||||
r = self.h_last / h
|
||||
if self.solver_type == 'heun':
|
||||
sample = sample + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (model_output - self.model_outputs[-2])
|
||||
elif self.solver_type == 'midpoint':
|
||||
sample = sample + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * (model_output - self.model_outputs[-2])
|
||||
|
||||
if self.config.use_noise_sampler:
|
||||
sample = sample + self.noise_sampler(sigma, sigma_next) * sigma_next * (-2 * eta_h).expm1().neg().sqrt() * self.config.s_noise
|
||||
else:
|
||||
sample = sample + noise * sigma_next * (-2 * eta_h).expm1().neg().sqrt() * self.config.s_noise
|
||||
|
||||
self.h_last = h
|
||||
elif self.config.algorithm_type == "dpmsolver++3Msde":
|
||||
if self.config.solver_order == 3:
|
||||
if sigma_next == 0:
|
||||
sample = model_output
|
||||
else:
|
||||
# DPM-Solver++(3M) SDE
|
||||
t, s = -sigma.log(), -sigma_next.log()
|
||||
h = s - t
|
||||
h_eta = h * 2
|
||||
|
||||
# 3. Delta timestep
|
||||
dt = sigma_next - sigma
|
||||
sample = sample + model_output * dt
|
||||
|
||||
sample = torch.exp(-h_eta) * sample + (-h_eta).expm1().neg() * model_output
|
||||
|
||||
if self.h_2 is not None:
|
||||
r0 = self.h_1 / h
|
||||
r1 = self.h_2 / h
|
||||
d1_0 = (model_output - self.model_outputs[-2]) / r0
|
||||
d1_1 = (self.model_outputs[-2] - self.model_outputs[-3]) / r1
|
||||
d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)
|
||||
d2 = (d1_0 - d1_1) / (r0 + r1)
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
phi_3 = phi_2 / h_eta - 0.5
|
||||
sample = sample + phi_2 * d1 - phi_3 * d2
|
||||
del d1_0
|
||||
del d1_1
|
||||
del d1
|
||||
del d2
|
||||
del phi_2
|
||||
del phi_3
|
||||
elif self.h_1 is not None:
|
||||
r = self.h_1 / h
|
||||
d = (model_output - self.model_outputs[-2]) / r
|
||||
phi_2 = h_eta.neg().expm1() / h_eta + 1
|
||||
sample = sample + phi_2 * d
|
||||
del d
|
||||
del phi_2
|
||||
|
||||
if self.config.use_noise_sampler:
|
||||
sample = sample + self.noise_sampler(sigma, sigma_next) * sigma_next * (-2 * h).expm1().neg().sqrt() * self.config.s_noise
|
||||
else:
|
||||
sample = sample + noise * sigma_next * (-2 * h).expm1().neg().sqrt() * self.config.s_noise
|
||||
|
||||
self.h_2 = self.h_1
|
||||
self.h_1 = h
|
||||
if not self.config.use_noise_sampler and noise is not None:
|
||||
del noise
|
||||
prev_sample = sample
|
||||
|
||||
# Cast sample back to expected dtype
|
||||
prev_sample = prev_sample.to(model_output.dtype)
|
||||
|
||||
# upon completion increase step index by one
|
||||
self._step_index += 1
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample,)
|
||||
|
||||
return FlowMatchDPMSolverMultistepSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
"""
|
||||
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
|
||||
current timestep.
|
||||
|
||||
Args:
|
||||
sample (`torch.Tensor`):
|
||||
The input sample.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
return sample
|
||||
|
||||
def scale_noise(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
noise: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
Forward process in flow-matching
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The input sample.
|
||||
timestep (`int`, *optional*):
|
||||
The current timestep in the diffusion chain.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
# Make sure sigmas and timesteps have the same device and dtype as original_samples
|
||||
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
|
||||
|
||||
if sample.device.type == "mps" and torch.is_floating_point(timestep):
|
||||
# mps does not support float64
|
||||
schedule_timesteps = self.timesteps.to(sample.device, dtype=torch.float32)
|
||||
timestep = timestep.to(sample.device, dtype=torch.float32)
|
||||
else:
|
||||
schedule_timesteps = self.timesteps.to(sample.device)
|
||||
timestep = timestep.to(sample.device)
|
||||
|
||||
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
|
||||
if self.begin_index is None:
|
||||
step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep]
|
||||
elif self.step_index is not None:
|
||||
# add_noise is called after first denoising step (for inpainting)
|
||||
step_indices = [self.step_index] * timestep.shape[0]
|
||||
else:
|
||||
# add noise is called before first denoising step to create initial latent(img2img)
|
||||
step_indices = [self.begin_index] * timestep.shape[0]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < len(sample.shape):
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
|
||||
sample = sigma * noise + (1.0 - sigma) * sample
|
||||
|
||||
return sample
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
+3
-1
@@ -352,7 +352,9 @@ class ScriptRunner:
|
||||
self.selectable_scripts.clear()
|
||||
auto_processing_scripts = scripts_auto_postprocessing.create_auto_preprocessing_script_data()
|
||||
|
||||
for script_class, path, _basedir, _script_module in auto_processing_scripts + scripts_data:
|
||||
all_scripts = auto_processing_scripts + scripts_data
|
||||
sorted_scripts = sorted(all_scripts, key=lambda x: x.script_class().title().lower())
|
||||
for script_class, path, _basedir, _script_module in sorted_scripts:
|
||||
try:
|
||||
script = script_class()
|
||||
script.filename = path
|
||||
|
||||
@@ -295,13 +295,7 @@ def read_metadata_from_safetensors(filename):
|
||||
if k == 'format' and v == 'pt':
|
||||
continue
|
||||
large = True if len(v) > 2048 else False
|
||||
if large and k == 'ss_datasets':
|
||||
continue
|
||||
if large and k == 'workflow':
|
||||
continue
|
||||
if large and k == 'prompt':
|
||||
continue
|
||||
if large and k == 'ss_bucket_info':
|
||||
if large and k in ['ss_datasets', 'workflow', 'prompt', 'ss_bucket_info', 'sd_metadata_file']:
|
||||
continue
|
||||
if v[0:1] == '{':
|
||||
try:
|
||||
|
||||
+12
-1
@@ -1,7 +1,8 @@
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
import diffusers
|
||||
from modules import shared, shared_items, devices, errors
|
||||
from modules import shared, shared_items, devices, errors, model_tools
|
||||
|
||||
|
||||
debug_load = os.environ.get('SD_LOAD_DEBUG', None)
|
||||
@@ -81,6 +82,9 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
if 'meissonic' in f.lower():
|
||||
guess = 'Meissonic'
|
||||
pipeline = 'custom'
|
||||
if 'monetico' in f.lower():
|
||||
guess = 'Monetico'
|
||||
pipeline = 'custom'
|
||||
if 'omnigen' in f.lower():
|
||||
guess = 'OmniGen'
|
||||
pipeline = 'custom'
|
||||
@@ -103,6 +107,13 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False):
|
||||
pipeline = shared_items.get_pipelines().get(guess, None) if pipeline is None else pipeline
|
||||
if not quiet:
|
||||
shared.log.info(f'Autodetect {op}: detect="{guess}" class={getattr(pipeline, "__name__", None)} file="{f}" size={size}MB')
|
||||
t0 = time.time()
|
||||
keys = model_tools.get_safetensor_keys(f)
|
||||
if keys is not None and len(keys) > 0:
|
||||
modules = model_tools.list_to_dict(keys)
|
||||
modules = model_tools.remove_entries_after_depth(modules, 3)
|
||||
t1 = time.time()
|
||||
shared.log.debug(f'Autodetect modules: {modules} time={t1-t0:.2f}')
|
||||
except Exception as e:
|
||||
shared.log.error(f'Autodetect {op}: file="{f}" {e}')
|
||||
if debug_load:
|
||||
|
||||
+84
-42
@@ -291,10 +291,13 @@ def set_diffuser_options(sd_model, vae = None, op: str = 'model', offload=True):
|
||||
|
||||
|
||||
def set_accelerate_to_module(model):
|
||||
for k in model._internal_dict.keys(): # pylint: disable=protected-access
|
||||
component = getattr(model, k, None)
|
||||
if isinstance(component, torch.nn.Module):
|
||||
component.has_accelerate = True
|
||||
if hasattr(model, "pipe"):
|
||||
set_accelerate_to_module(model.pipe)
|
||||
if hasattr(model, "_internal_dict"):
|
||||
for k in model._internal_dict.keys(): # pylint: disable=protected-access
|
||||
component = getattr(model, k, None)
|
||||
if isinstance(component, torch.nn.Module):
|
||||
component.has_accelerate = True
|
||||
|
||||
|
||||
def set_accelerate(sd_model):
|
||||
@@ -397,7 +400,13 @@ def apply_balanced_offload(sd_model):
|
||||
return module
|
||||
|
||||
def apply_balanced_offload_to_module(pipe):
|
||||
for module_name in pipe._internal_dict.keys(): # pylint: disable=protected-access
|
||||
if hasattr(pipe, "pipe"):
|
||||
apply_balanced_offload_to_module(pipe.pipe)
|
||||
if hasattr(pipe, "_internal_dict"):
|
||||
keys = pipe._internal_dict.keys() # pylint: disable=protected-access
|
||||
else:
|
||||
keys = get_signature(shared.sd_model).keys()
|
||||
for module_name in keys: # pylint: disable=protected-access
|
||||
module = getattr(pipe, module_name, None)
|
||||
if isinstance(module, torch.nn.Module):
|
||||
checkpoint_name = pipe.sd_checkpoint_info.name if getattr(pipe, "sd_checkpoint_info", None) is not None else None
|
||||
@@ -416,6 +425,8 @@ def apply_balanced_offload(sd_model):
|
||||
devices.torch_gc(fast=True)
|
||||
|
||||
apply_balanced_offload_to_module(sd_model)
|
||||
if hasattr(sd_model, "pipe"):
|
||||
apply_balanced_offload_to_module(sd_model.pipe)
|
||||
if hasattr(sd_model, "prior_pipe"):
|
||||
apply_balanced_offload_to_module(sd_model.prior_pipe)
|
||||
if hasattr(sd_model, "decoder_pipe"):
|
||||
@@ -440,6 +451,9 @@ def move_model(model, device=None, force=False):
|
||||
devices.torch_gc()
|
||||
return
|
||||
|
||||
if hasattr(model, 'pipe'):
|
||||
move_model(model.pipe, device, force)
|
||||
|
||||
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
|
||||
if getattr(model, 'vae', None) is not None and get_diffusers_task(model) != DiffusersTaskType.TEXT_2_IMAGE:
|
||||
if device == devices.device and model.vae.device.type != "meta": # force vae back to gpu if not in txt2img mode
|
||||
@@ -467,7 +481,8 @@ def move_model(model, device=None, force=False):
|
||||
try:
|
||||
t0 = time.time()
|
||||
try:
|
||||
model.to(device)
|
||||
if hasattr(model, 'to'):
|
||||
model.to(device)
|
||||
if hasattr(model, "prior_pipe"):
|
||||
model.prior_pipe.to(device)
|
||||
except Exception as e0:
|
||||
@@ -477,7 +492,8 @@ def move_model(model, device=None, force=False):
|
||||
if hasattr(component, 'modules'):
|
||||
for module in component.modules():
|
||||
try:
|
||||
module.to(device)
|
||||
if hasattr(module, 'to'):
|
||||
module.to(device)
|
||||
except Exception as e2:
|
||||
if 'Cannot copy out of meta tensor' in str(e2):
|
||||
if os.environ.get('SD_MOVE_DEBUG', None):
|
||||
@@ -771,11 +787,11 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
if shared.opts.data.get('sd_model_checkpoint', '') == 'model.safetensors' or shared.opts.data.get('sd_model_checkpoint', '') == '':
|
||||
shared.opts.data['sd_model_checkpoint'] = "stabilityai/stable-diffusion-xl-base-1.0"
|
||||
|
||||
if op == 'model' or op == 'dict':
|
||||
if (model_data.sd_model is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model
|
||||
if (op == 'model' or op == 'dict'):
|
||||
if (model_data.sd_model is not None) and (checkpoint_info is not None) and (getattr(model_data.sd_model, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
else:
|
||||
if (model_data.sd_refiner is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model
|
||||
if (model_data.sd_refiner is not None) and (checkpoint_info is not None) and (getattr(model_data.sd_refiner, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
|
||||
sd_model = None
|
||||
@@ -797,6 +813,9 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
# preload vae so it can be used as param
|
||||
vae = None
|
||||
sd_vae.loaded_vae_file = None
|
||||
if model_type is None:
|
||||
shared.log.error(f'Load {op}: pipeline={shared.opts.diffusers_pipeline} not detected')
|
||||
return
|
||||
if model_type.startswith('Stable Diffusion') and (op == 'model' or op == 'refiner'): # preload vae for sd models
|
||||
vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename)
|
||||
vae = sd_vae.load_vae_diffusers(checkpoint_info.path, vae_file, vae_source)
|
||||
@@ -862,8 +881,9 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True)
|
||||
timer.record("embeddings")
|
||||
|
||||
from modules.prompt_parser_diffusers import insert_parser_highjack
|
||||
insert_parser_highjack(sd_model.__class__.__name__)
|
||||
from modules import prompt_parser_diffusers
|
||||
prompt_parser_diffusers.insert_parser_highjack(sd_model.__class__.__name__)
|
||||
prompt_parser_diffusers.cache.clear()
|
||||
|
||||
set_diffuser_options(sd_model, vae, op, offload=False)
|
||||
if shared.opts.nncf_compress_weights and not ('Model' in shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"):
|
||||
@@ -874,7 +894,8 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
|
||||
set_diffuser_offload(sd_model, op)
|
||||
if op == 'model' and not (os.path.isdir(checkpoint_info.path) or checkpoint_info.type == 'huggingface'):
|
||||
sd_vae.apply_vae_config(shared.sd_model.sd_checkpoint_info.filename, vae_file, sd_model)
|
||||
if getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None:
|
||||
sd_vae.apply_vae_config(shared.sd_model.sd_checkpoint_info.filename, vae_file, sd_model)
|
||||
if op == 'refiner' and shared.opts.diffusers_move_refiner:
|
||||
shared.log.debug('Moving refiner model to CPU')
|
||||
move_model(sd_model, devices.cpu)
|
||||
@@ -1049,6 +1070,7 @@ def set_diffuser_pipe(pipe, new_pipe_type):
|
||||
'AnimateDiffSDXLPipeline',
|
||||
'OmniGenPipeline',
|
||||
'StableDiffusion3ControlNetPipeline',
|
||||
'InstantIRPipeline',
|
||||
]
|
||||
|
||||
n = getattr(pipe.__class__, '__name__', '')
|
||||
@@ -1059,12 +1081,15 @@ def set_diffuser_pipe(pipe, new_pipe_type):
|
||||
return pipe
|
||||
|
||||
# skip specific pipelines
|
||||
cls = pipe.__class__.__name__
|
||||
if n in exclude:
|
||||
return pipe
|
||||
if 'Onnx' in pipe.__class__.__name__:
|
||||
if 'Onnx' in cls:
|
||||
return pipe
|
||||
|
||||
if new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE or new_pipe_type == DiffusersTaskType.INPAINTING: # in some cases we want to reset the pipeline as they dont have their own variants
|
||||
new_pipe = None
|
||||
# in some cases we want to reset the pipeline to parent as they dont have their own variants
|
||||
if new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE or new_pipe_type == DiffusersTaskType.INPAINTING:
|
||||
if n == 'StableDiffusionPAGPipeline':
|
||||
pipe = switch_pipe(diffusers.StableDiffusionPipeline, pipe)
|
||||
if n == 'StableDiffusionXLPAGPipeline':
|
||||
@@ -1080,19 +1105,38 @@ def set_diffuser_pipe(pipe, new_pipe_type):
|
||||
image_encoder = getattr(pipe, "image_encoder", None)
|
||||
feature_extractor = getattr(pipe, "feature_extractor", None)
|
||||
|
||||
try:
|
||||
if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE:
|
||||
new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe)
|
||||
elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE:
|
||||
new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe)
|
||||
elif new_pipe_type == DiffusersTaskType.INPAINTING:
|
||||
new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe)
|
||||
if new_pipe is None:
|
||||
if hasattr(pipe, 'config'): # real pipeline which can be auto-switched
|
||||
try:
|
||||
if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE:
|
||||
new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe)
|
||||
elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE:
|
||||
new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe)
|
||||
elif new_pipe_type == DiffusersTaskType.INPAINTING:
|
||||
new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe)
|
||||
else:
|
||||
shared.log.error(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls}')
|
||||
return pipe
|
||||
except Exception as e: # pylint: disable=unused-variable
|
||||
shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls} {e}')
|
||||
return pipe
|
||||
else:
|
||||
shared.log.error(f'Pipeline class change failed: type={new_pipe_type} pipeline={pipe.__class__.__name__}')
|
||||
return pipe
|
||||
except Exception as e: # pylint: disable=unused-variable
|
||||
shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={pipe.__class__.__name__} {e}')
|
||||
return pipe
|
||||
try: # maybe a wrapper pipeline so just change the class
|
||||
if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE:
|
||||
pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access
|
||||
new_pipe = pipe
|
||||
elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE:
|
||||
pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access
|
||||
new_pipe = pipe
|
||||
elif new_pipe_type == DiffusersTaskType.INPAINTING:
|
||||
pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING, cls) # pylint: disable=protected-access
|
||||
new_pipe = pipe
|
||||
else:
|
||||
shared.log.error(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls}')
|
||||
return pipe
|
||||
except Exception as e: # pylint: disable=unused-variable
|
||||
shared.log.warning(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls} {e}')
|
||||
return pipe
|
||||
|
||||
# if pipe.__class__ == new_pipe.__class__:
|
||||
# return pipe
|
||||
@@ -1108,10 +1152,14 @@ def set_diffuser_pipe(pipe, new_pipe_type):
|
||||
new_pipe.is_sdxl = getattr(pipe, 'is_sdxl', False) # a1111 compatibility item
|
||||
new_pipe.is_sd2 = getattr(pipe, 'is_sd2', False)
|
||||
new_pipe.is_sd1 = getattr(pipe, 'is_sd1', True)
|
||||
if hasattr(new_pipe, "watermark"):
|
||||
if hasattr(new_pipe, 'watermark'):
|
||||
new_pipe.watermark = NoWatermark()
|
||||
|
||||
if hasattr(new_pipe, 'pipe'): # also handle nested pipelines
|
||||
new_pipe.pipe = set_diffuser_pipe(new_pipe.pipe, new_pipe_type)
|
||||
|
||||
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
|
||||
shared.log.debug(f"Pipeline class change: original={pipe.__class__.__name__} target={new_pipe.__class__.__name__} device={pipe.device} fn={fn}") # pylint: disable=protected-access
|
||||
shared.log.debug(f"Pipeline class change: original={cls} target={new_pipe.__class__.__name__} device={pipe.device} fn={fn}") # pylint: disable=protected-access
|
||||
pipe = new_pipe
|
||||
return pipe
|
||||
|
||||
@@ -1140,6 +1188,9 @@ def set_diffusers_attention(pipe):
|
||||
else:
|
||||
module.set_attn_processor(attention)
|
||||
|
||||
# if hasattr(pipe, 'pipe'):
|
||||
# set_diffusers_attention(pipe.pipe)
|
||||
|
||||
if 'ControlNet' in pipe.__class__.__name__: # do not replace attention in ControlNet pipelines
|
||||
return
|
||||
shared.log.debug(f'Setting model: attention="{shared.opts.cross_attention_optimization}"')
|
||||
@@ -1180,10 +1231,10 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None,
|
||||
if checkpoint_info is None:
|
||||
return
|
||||
if op == 'model' or op == 'dict':
|
||||
if model_data.sd_model is not None and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model
|
||||
if (model_data.sd_model is not None) and (getattr(model_data.sd_model, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
else:
|
||||
if model_data.sd_refiner is not None and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model
|
||||
if (model_data.sd_refiner is not None) and (getattr(model_data.sd_refiner, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model
|
||||
return
|
||||
shared.log.debug(f'Load {op}: name={checkpoint_info.filename} dict={already_loaded_state_dict is not None}')
|
||||
if timer is None:
|
||||
@@ -1192,12 +1243,12 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None,
|
||||
if op == 'model' or op == 'dict':
|
||||
if model_data.sd_model is not None:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_model)
|
||||
current_checkpoint_info = model_data.sd_model.sd_checkpoint_info
|
||||
current_checkpoint_info = getattr(model_data.sd_model, 'sd_checkpoint_info', None)
|
||||
unload_model_weights(op=op)
|
||||
else:
|
||||
if model_data.sd_refiner is not None:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner)
|
||||
current_checkpoint_info = model_data.sd_refiner.sd_checkpoint_info
|
||||
current_checkpoint_info = getattr(model_data.sd_refiner, 'sd_checkpoint_info', None)
|
||||
unload_model_weights(op=op)
|
||||
|
||||
if not shared.native:
|
||||
@@ -1226,15 +1277,6 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None,
|
||||
sd_model = instantiate_from_config(sd_config.model)
|
||||
else:
|
||||
with contextlib.redirect_stdout(stdout):
|
||||
"""
|
||||
try:
|
||||
clip_is_included_into_sd = sd1_clip_weight in state_dict or sd2_clip_weight in state_dict
|
||||
with sd_disable_initialization.DisableInitialization(disable_clip=clip_is_included_into_sd):
|
||||
sd_model = instantiate_from_config(sd_config.model)
|
||||
except Exception as e:
|
||||
shared.log.error(f'LDM: instantiate from config: {e}')
|
||||
sd_model = instantiate_from_config(sd_config.model)
|
||||
"""
|
||||
sd_model = instantiate_from_config(sd_config.model)
|
||||
for line in stdout.getvalue().splitlines():
|
||||
if len(line) > 0:
|
||||
|
||||
+16
-17
@@ -47,6 +47,10 @@ def visible_sampler_names():
|
||||
|
||||
|
||||
def create_sampler(name, model):
|
||||
try:
|
||||
current = model.scheduler.__class__.__name__
|
||||
except Exception:
|
||||
current = None
|
||||
if name == 'Default' and hasattr(model, 'scheduler'):
|
||||
if getattr(model, "default_scheduler", None) is not None:
|
||||
model.scheduler = copy.deepcopy(model.default_scheduler)
|
||||
@@ -54,12 +58,13 @@ def create_sampler(name, model):
|
||||
model.prior_pipe.scheduler = copy.deepcopy(model.default_scheduler)
|
||||
model.prior_pipe.scheduler.config.clip_sample = False
|
||||
config = {k: v for k, v in model.scheduler.config.items() if not k.startswith('_')}
|
||||
shared.log.debug(f'Sampler: sampler=default class={model.scheduler.__class__.__name__}: {config}')
|
||||
shared.log.debug(f'Sampler: sampler=default class={current}: {config}')
|
||||
return model.scheduler
|
||||
config = find_sampler_config(name)
|
||||
if config is None or config.constructor is None:
|
||||
# shared.log.warning(f'Sampler: sampler="{name}" not found')
|
||||
return None
|
||||
sampler = None
|
||||
if not shared.native:
|
||||
sampler = config.constructor(model)
|
||||
sampler.config = config
|
||||
@@ -68,24 +73,18 @@ def create_sampler(name, model):
|
||||
shared.log.debug(f'Sampler: sampler="{name}" config={config.options}')
|
||||
return sampler
|
||||
elif shared.native:
|
||||
sampler = config.constructor(model)
|
||||
if 'Flux' in model.__class__.__name__:
|
||||
if 'base_image_seq_len' not in sampler.sampler.config or 'max_image_seq_len' not in sampler.sampler.config or 'base_shift' not in sampler.sampler.config or 'max_shift' not in sampler.sampler.config:
|
||||
shared.log.warning(f'FLUX: sampler="{name}" unsupported')
|
||||
# sampler.sampler.register_to_config(base_image_seq_len=256, max_image_seq_len=4096, base_shift=0.5, max_shift=1.15)
|
||||
return None
|
||||
if 'Lumina' in model.__class__.__name__:
|
||||
shared.log.warning(f'AlphaVLLM-Lumina: sampler="{name}" unsupported')
|
||||
return None
|
||||
if 'StableDiffusion3' in model.__class__.__name__:
|
||||
if sampler.name != 'Heun FlowMatch':
|
||||
return None
|
||||
return None
|
||||
if 'AuraFlow' in model.__class__.__name__:
|
||||
shared.log.warning(f'AuraFlow: sampler="{name}" unsupported')
|
||||
return None
|
||||
FlowModels = ['Flux', 'StableDiffusion3', 'Lumina', 'AuraFlow']
|
||||
if 'KDiffusion' in model.__class__.__name__:
|
||||
return None
|
||||
if any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' not in name:
|
||||
shared.log.warning(f'Sampler: default={current} target="{name}" class={model.__class__.__name__} linear scheduler unsupported')
|
||||
return None
|
||||
if not any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' in name:
|
||||
shared.log.warning(f'Sampler: default={current} target="{name}" class={model.__class__.__name__} flow-match scheduler unsupported')
|
||||
return None
|
||||
sampler = config.constructor(model)
|
||||
if sampler is None:
|
||||
sampler = config.constructor(model)
|
||||
if not hasattr(model, 'scheduler_config'):
|
||||
model.scheduler_config = sampler.sampler.config.copy() if hasattr(sampler.sampler, 'config') else {}
|
||||
model.scheduler = sampler.sampler
|
||||
|
||||
@@ -1,16 +1,16 @@
|
||||
import os
|
||||
import copy
|
||||
import re
|
||||
import copy
|
||||
import inspect
|
||||
import diffusers
|
||||
from modules import shared, errors
|
||||
from modules import sd_samplers_common
|
||||
from modules.tcd import TCDScheduler
|
||||
from modules.dcsolver import DCSolverMultistepScheduler
|
||||
from modules.vdm import VDMScheduler
|
||||
|
||||
|
||||
debug = shared.log.trace if os.environ.get('SD_SAMPLER_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
debug('Trace: SAMPLER')
|
||||
|
||||
|
||||
try:
|
||||
from diffusers import (
|
||||
CMStochasticIterativeScheduler,
|
||||
@@ -44,7 +44,15 @@ try:
|
||||
KDPM2AncestralDiscreteScheduler,
|
||||
)
|
||||
except Exception as e:
|
||||
import diffusers
|
||||
shared.log.error(f'Diffusers import error: version={diffusers.__version__} error: {e}')
|
||||
if os.environ.get('SD_SAMPLER_DEBUG', None) is not None:
|
||||
errors.display(e, 'Samplers')
|
||||
try:
|
||||
from modules.schedulers.scheduler_tcd import TCDScheduler # pylint: disable=ungrouped-imports
|
||||
from modules.schedulers.scheduler_dc import DCSolverMultistepScheduler # pylint: disable=ungrouped-imports
|
||||
from modules.schedulers.scheduler_vdm import VDMScheduler # pylint: disable=ungrouped-imports
|
||||
from modules.schedulers.scheduler_dpm_flowmatch import FlowMatchDPMSolverMultistepScheduler # pylint: disable=ungrouped-imports
|
||||
except Exception as e:
|
||||
shared.log.error(f'Diffusers import error: version={diffusers.__version__} error: {e}')
|
||||
if os.environ.get('SD_SAMPLER_DEBUG', None) is not None:
|
||||
errors.display(e, 'Samplers')
|
||||
@@ -63,33 +71,40 @@ config = {
|
||||
'Euler EDM': { 'sigma_schedule': "karras" },
|
||||
'Euler FlowMatch': { 'timestep_spacing': "linspace", 'shift': 1, 'use_dynamic_shifting': False },
|
||||
|
||||
'DPM++': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'sigma_min' },
|
||||
'DPM++ 1S': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 1 },
|
||||
'DPM++ 2M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 },
|
||||
'DPM++ 3M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 3 },
|
||||
'DPM++ 2M SDE': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "sde-dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 },
|
||||
'DPM++': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'final_sigmas_type': 'sigma_min' },
|
||||
'DPM++ 1S': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 1 },
|
||||
'DPM++ 2M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 },
|
||||
'DPM++ 3M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 3 },
|
||||
'DPM++ 2M SDE': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "sde-dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 },
|
||||
'DPM++ 2M EDM': { 'solver_order': 2, 'solver_type': 'midpoint', 'final_sigmas_type': 'zero', 'algorithm_type': 'dpmsolver++' },
|
||||
'DPM++ Cosine': { 'solver_order': 2, 'sigma_schedule': "exponential", 'prediction_type': "v-prediction" },
|
||||
'DPM SDE': { 'use_karras_sigmas': False, 'noise_sampler_seed': None, 'timestep_spacing': 'linspace', 'steps_offset': 0 },
|
||||
'DPM SDE': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'noise_sampler_seed': None, 'timestep_spacing': 'linspace', 'steps_offset': 0, },
|
||||
|
||||
'Heun': { 'use_beta_sigmas': False, 'use_karras_sigmas': False, 'timestep_spacing': 'linspace' },
|
||||
'DPM2 FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver2', 'use_noise_sampler': True },
|
||||
'DPM2a FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver2A', 'use_noise_sampler': True },
|
||||
'DPM2++ 2M FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++2M', 'use_noise_sampler': True },
|
||||
'DPM2++ 2S FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++2S', 'use_noise_sampler': True },
|
||||
'DPM2++ SDE FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++sde', 'use_noise_sampler': True },
|
||||
'DPM2++ 2M SDE FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++2Msde', 'use_noise_sampler': True },
|
||||
'DPM2++ 3M SDE FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 3, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++3Msde', 'use_noise_sampler': True },
|
||||
|
||||
'Heun': { 'use_beta_sigmas': False, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'timestep_spacing': 'linspace' },
|
||||
'Heun FlowMatch': { 'timestep_spacing': "linspace", 'shift': 1 },
|
||||
|
||||
'DEIS': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "deis", 'solver_type': "logrho", 'lower_order_final': True, 'timestep_spacing': 'linspace' },
|
||||
'SA Solver': {'predictor_order': 2, 'corrector_order': 2, 'thresholding': False, 'lower_order_final': True, 'use_karras_sigmas': False, 'timestep_spacing': 'linspace'},
|
||||
'DEIS': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "deis", 'solver_type': "logrho", 'lower_order_final': True, 'timestep_spacing': 'linspace', 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False },
|
||||
'SA Solver': {'predictor_order': 2, 'corrector_order': 2, 'thresholding': False, 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'timestep_spacing': 'linspace'},
|
||||
'DC Solver': { 'beta_start': 0.0001, 'beta_end': 0.02, 'solver_order': 2, 'prediction_type': "epsilon", 'thresholding': False, 'solver_type': 'bh2', 'lower_order_final': True, 'dc_order': 2, 'disable_corrector': [0] },
|
||||
'VDM Solver': { 'clip_sample_range': 2.0, },
|
||||
|
||||
'LCM': { 'beta_start': 0.00085, 'beta_end': 0.012, 'beta_schedule': "scaled_linear", 'set_alpha_to_one': True, 'rescale_betas_zero_snr': False, 'thresholding': False, 'timestep_spacing': 'linspace' },
|
||||
'TCD': { 'set_alpha_to_one': True, 'rescale_betas_zero_snr': False, 'beta_schedule': 'scaled_linear' },
|
||||
|
||||
'PNDM': { 'skip_prk_steps': False, 'set_alpha_to_one': False, 'steps_offset': 0, 'timestep_spacing': 'linspace' },
|
||||
'IPNDM': { },
|
||||
'DDPM': { 'variance_type': "fixed_small", 'clip_sample': False, 'thresholding': False, 'clip_sample_range': 1.0, 'sample_max_value': 1.0, 'timestep_spacing': 'linspace', 'rescale_betas_zero_snr': False },
|
||||
'LMSD': { 'use_karras_sigmas': False, 'timestep_spacing': 'linspace', 'steps_offset': 0 },
|
||||
'KDPM2': { 'steps_offset': 0, 'timestep_spacing': 'linspace' },
|
||||
'KDPM2 a': { 'steps_offset': 0, 'timestep_spacing': 'linspace' },
|
||||
'CMSI': { }, #{ 'sigma_min': 0.002, 'sigma_max': 80.0, 'sigma_data': 0.5, 's_noise': 1.0, 'rho': 7.0, 'clip_denoised': True },
|
||||
'LMSD': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'timestep_spacing': 'linspace', 'steps_offset': 0 },
|
||||
'KDPM2': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'steps_offset': 0, 'timestep_spacing': 'linspace' },
|
||||
'KDPM2 a': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'steps_offset': 0, 'timestep_spacing': 'linspace' },
|
||||
'CMSI': { },
|
||||
}
|
||||
|
||||
samplers_data_diffusers = [
|
||||
@@ -112,6 +127,14 @@ samplers_data_diffusers = [
|
||||
sd_samplers_common.SamplerData('DPM++ Cosine', lambda model: DiffusionSampler('DPM++ 2M EDM', CosineDPMSolverMultistepScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('DPM SDE', lambda model: DiffusionSampler('DPM SDE', DPMSolverSDEScheduler, model), [], {}),
|
||||
|
||||
sd_samplers_common.SamplerData('DPM2 FlowMatch', lambda model: DiffusionSampler('DPM2 FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('DPM2a FlowMatch', lambda model: DiffusionSampler('DPM2a FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('DPM2++ 2M FlowMatch', lambda model: DiffusionSampler('DPM2++ 2M FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('DPM2++ 2S FlowMatch', lambda model: DiffusionSampler('DPM2++ 2S FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('DPM2++ SDE FlowMatch', lambda model: DiffusionSampler('DPM2++ SDE FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('DPM2++ 2M SDE FlowMatch', lambda model: DiffusionSampler('DPM2++ 2M SDE FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('DPM2++ 3M SDE FlowMatch', lambda model: DiffusionSampler('DPM2++ 3M SDE FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}),
|
||||
|
||||
sd_samplers_common.SamplerData('Heun', lambda model: DiffusionSampler('Heun', HeunDiscreteScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('Heun FlowMatch', lambda model: DiffusionSampler('Heun FlowMatch', FlowMatchHeunDiscreteScheduler, model), [], {}),
|
||||
|
||||
@@ -183,6 +206,10 @@ class DiffusionSampler:
|
||||
self.config['use_karras_sigmas'] = shared.opts.schedulers_sigma == 'karras'
|
||||
if 'use_exponential_sigmas' in self.config:
|
||||
self.config['use_exponential_sigmas'] = shared.opts.schedulers_sigma == 'exponential'
|
||||
if 'use_lu_lambdas' in self.config:
|
||||
self.config['use_lu_lambdas'] = shared.opts.schedulers_sigma == 'lambdas'
|
||||
if 'sigma_schedule' in self.config:
|
||||
self.config['sigma_schedule'] = shared.opts.schedulers_sigma if shared.opts.schedulers_sigma != 'default' else None
|
||||
else:
|
||||
pass # timesteps are set using set_timesteps in set_pipeline_args
|
||||
|
||||
@@ -199,9 +226,18 @@ class DiffusionSampler:
|
||||
if 'beta_end' in self.config and shared.opts.schedulers_beta_end > 0:
|
||||
self.config['beta_end'] = shared.opts.schedulers_beta_end
|
||||
if 'shift' in self.config:
|
||||
self.config['shift'] = shared.opts.schedulers_shift
|
||||
if shared.opts.schedulers_shift == 0:
|
||||
if 'StableDiffusion3' in model.__class__.__name__:
|
||||
self.config['shift'] = 3
|
||||
if 'Flux' in model.__class__.__name__:
|
||||
self.config['shift'] = 1
|
||||
else:
|
||||
self.config['shift'] = shared.opts.schedulers_shift
|
||||
if 'use_dynamic_shifting' in self.config:
|
||||
self.config['use_dynamic_shifting'] = shared.opts.schedulers_dynamic_shift
|
||||
if 'Flux' in model.__class__.__name__:
|
||||
self.config['use_dynamic_shifting'] = shared.opts.schedulers_dynamic_shift
|
||||
if 'use_SD35_sigmas' in self.config:
|
||||
self.config['use_SD35_sigmas'] = 'StableDiffusion3' in model.__class__.__name__
|
||||
if 'rescale_betas_zero_snr' in self.config:
|
||||
self.config['rescale_betas_zero_snr'] = shared.opts.schedulers_rescale_betas
|
||||
if 'timestep_spacing' in self.config and shared.opts.schedulers_timestep_spacing != 'default' and shared.opts.schedulers_timestep_spacing is not None:
|
||||
|
||||
+1
-1
@@ -289,7 +289,7 @@ def reload_vae_weights(sd_model=None, vae_file=unspecified):
|
||||
if vae_file is not None:
|
||||
shared.log.info(f"VAE weights loaded: {vae_file}")
|
||||
else:
|
||||
if hasattr(sd_model, "vae") and hasattr(sd_model, "sd_checkpoint_info"):
|
||||
if hasattr(sd_model, "vae") and getattr(sd_model, "sd_checkpoint_info", None) is not None:
|
||||
vae = load_vae_diffusers(sd_model.sd_checkpoint_info.filename, vae_file, vae_source)
|
||||
if vae is not None:
|
||||
if not hasattr(sd_model, 'original_vae'):
|
||||
|
||||
+110
-65
@@ -1,3 +1,4 @@
|
||||
from functools import lru_cache
|
||||
import io
|
||||
import os
|
||||
import sys
|
||||
@@ -83,6 +84,8 @@ console = Console(log_time=True, log_time_format='%H:%M:%S-%f')
|
||||
dir_timestamps = {}
|
||||
dir_cache = {}
|
||||
max_workers = 8
|
||||
if os.environ.get("SD_HFCACHEDIR", None) is not None:
|
||||
hfcache_dir = os.environ.get("SD_HFCACHEDIR")
|
||||
if os.environ.get("HF_HUB_CACHE", None) is not None:
|
||||
hfcache_dir = os.environ.get("HF_HUB_CACHE")
|
||||
elif os.environ.get("HF_HUB", None) is not None:
|
||||
@@ -249,6 +252,7 @@ class OptionInfo:
|
||||
self.comment_before = comment_before # HTML text that will be added after label in UI
|
||||
self.comment_after = comment_after # HTML text that will be added before label in UI
|
||||
self.submit = submit
|
||||
self.exclude = ['sd_model_checkpoint', 'sd_model_refiner', 'sd_vae', 'sd_unet', 'sd_text_encoder', 'sd_model_dict']
|
||||
|
||||
def needs_reload_ui(self):
|
||||
return self
|
||||
@@ -273,6 +277,42 @@ class OptionInfo:
|
||||
self.comment_after += " <span class='info'>(requires restart)</span>"
|
||||
return self
|
||||
|
||||
def validate(self, opt, value):
|
||||
if opt in self.exclude:
|
||||
return True
|
||||
args = self.component_args if self.component_args is not None else {}
|
||||
if callable(args):
|
||||
try:
|
||||
args = args()
|
||||
except Exception:
|
||||
args = {}
|
||||
choices = args.get("choices", [])
|
||||
if callable(choices):
|
||||
try:
|
||||
choices = choices()
|
||||
except Exception:
|
||||
choices = []
|
||||
if len(choices) > 0:
|
||||
if not isinstance(value, list):
|
||||
value = [value]
|
||||
for v in value:
|
||||
if v not in choices:
|
||||
log.warning(f'Setting validation: "{opt}"="{v}" default="{self.default}" choices={choices}')
|
||||
return False
|
||||
minimum = args.get("minimum", None)
|
||||
maximum = args.get("maximum", None)
|
||||
if (minimum is not None and value < minimum) or (maximum is not None and value > maximum):
|
||||
log.error(f'Setting validation: "{opt}"={value} default={self.default} minimum={minimum} maximum={maximum}')
|
||||
return False
|
||||
return True
|
||||
|
||||
def __str__(self) -> str:
|
||||
args = self.component_args if self.component_args is not None else {}
|
||||
if callable(args):
|
||||
args = args()
|
||||
choices = args.get("choices", [])
|
||||
return f'OptionInfo: label="{self.label}" section="{self.section}" component="{self.component}" default="{self.default}" refresh="{self.refresh is not None}" change="{self.onchange is not None}" args={args} choices={choices}'
|
||||
|
||||
|
||||
def options_section(section_identifier, options_dict):
|
||||
for v in options_dict.values():
|
||||
@@ -436,17 +476,14 @@ options_templates.update(options_section(('sd', "Execution & Models"), {
|
||||
"sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles()}, refresh=refresh_checkpoints),
|
||||
"sd_checkpoint_autoload": OptionInfo(True, "Model autoload on start"),
|
||||
"sd_checkpoint_autodownload": OptionInfo(True, "Model auto-download on demand"),
|
||||
"sd_textencoder_cache": OptionInfo(True, "Cache text encoder results"),
|
||||
"sd_textencoder_cache": OptionInfo(True, "Cache text encoder results", gr.Checkbox, {"visible": False}),
|
||||
"sd_textencoder_cache_size": OptionInfo(4, "Text encoder results LRU cache size", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1}),
|
||||
"stream_load": OptionInfo(False, "Load models using stream loading method", gr.Checkbox, {"visible": not native }),
|
||||
"model_reuse_dict": OptionInfo(False, "Reuse loaded model dictionary", gr.Checkbox, {"visible": False}),
|
||||
"prompt_mean_norm": OptionInfo(False, "Prompt attention normalization", gr.Checkbox),
|
||||
"comma_padding_backtrack": OptionInfo(20, "Prompt padding", gr.Slider, {"minimum": 0, "maximum": 74, "step": 1, "visible": not native }),
|
||||
"prompt_attention": OptionInfo("Full parser", "Prompt attention parser", gr.Radio, {"choices": ["Full parser", "Compel parser", "xhinker parser", "A1111 parser", "Fixed attention"] }),
|
||||
"prompt_attention": OptionInfo("native", "Prompt attention parser", gr.Radio, {"choices": ["native", "compel", "xhinker", "a1111", "fixed"] }),
|
||||
"latent_history": OptionInfo(16, "Latent history size", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1}),
|
||||
"sd_checkpoint_cache": OptionInfo(0, "Cached models", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": not native }),
|
||||
"sd_vae_checkpoint_cache": OptionInfo(0, "Cached VAEs", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": False}),
|
||||
"sd_disable_ckpt": OptionInfo(False, "Disallow models in ckpt format", gr.Checkbox, {"visible": False}),
|
||||
"diffusers_version": OptionInfo("", "Diffusers version", gr.Textbox, {"visible": False}),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section(('cuda', "Compute Settings"), {
|
||||
@@ -460,8 +497,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), {
|
||||
"upcast_sampling": OptionInfo(False if sys.platform != "darwin" else True, "Upcast sampling"),
|
||||
"upcast_attn": OptionInfo(False, "Upcast attention layer"),
|
||||
"cuda_cast_unet": OptionInfo(False, "Fixed UNet precision"),
|
||||
"disable_nan_check": OptionInfo(True, "Disable NaN check", gr.Checkbox, {"visible": False}),
|
||||
"nan_skip": OptionInfo(False, "Skip Generation if NaN found in latents", gr.Checkbox, {"visible": True}),
|
||||
"nan_skip": OptionInfo(False, "Skip Generation if NaN found in latents", gr.Checkbox),
|
||||
"rollback_vae": OptionInfo(False, "Attempt VAE roll back for NaN values"),
|
||||
|
||||
"cross_attention_sep": OptionInfo("<h2>Cross Attention</h2>", "", gr.HTML),
|
||||
@@ -534,7 +570,6 @@ options_templates.update(options_section(('diffusers', "Diffusers Settings"), {
|
||||
"diffusers_eval": OptionInfo(True, "Force model eval"),
|
||||
"diffusers_to_gpu": OptionInfo(False, "Load model directly to GPU"),
|
||||
"disable_accelerate": OptionInfo(False, "Disable accelerate"),
|
||||
"diffusers_force_zeros": OptionInfo(False, "Force zeros for prompts when empty", gr.Checkbox, {"visible": False}),
|
||||
"diffusers_pooled": OptionInfo("default", "Diffusers SDXL pooled embeds", gr.Radio, {"choices": ['default', 'weighted']}),
|
||||
"diffusers_zeros_prompt_pad": OptionInfo(False, "Use zeros for prompt padding", gr.Checkbox),
|
||||
"huggingface_token": OptionInfo('', 'HuggingFace token'),
|
||||
@@ -615,9 +650,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), {
|
||||
"vae_dir": OptionInfo(os.path.join(paths.models_path, 'VAE'), "Folder with VAE files", folder=True),
|
||||
"unet_dir": OptionInfo(os.path.join(paths.models_path, 'UNET'), "Folder with UNET files", folder=True),
|
||||
"te_dir": OptionInfo(os.path.join(paths.models_path, 'Text-encoder'), "Folder with Text encoder files", folder=True),
|
||||
"sd_lora": OptionInfo("", "Add LoRA to prompt", gr.Textbox, {"visible": False}),
|
||||
"lora_dir": OptionInfo(os.path.join(paths.models_path, 'Lora'), "Folder with LoRA network(s)", folder=True),
|
||||
"lyco_dir": OptionInfo(os.path.join(paths.models_path, 'LyCORIS'), "Folder with LyCORIS network(s)", gr.Text, {"visible": False}),
|
||||
"styles_dir": OptionInfo(os.path.join(paths.data_path, 'styles.csv'), "File or Folder with user-defined styles", folder=True),
|
||||
"wildcards_dir": OptionInfo(os.path.join(paths.models_path, 'wildcards'), "Folder with user-defined wildcards", folder=True),
|
||||
"embeddings_dir": OptionInfo(os.path.join(paths.models_path, 'embeddings'), "Folder with textual inversion embeddings", folder=True),
|
||||
@@ -693,12 +726,10 @@ options_templates.update(options_section(('saving-paths', "Image Naming & Paths"
|
||||
"saving_sep_images": OptionInfo("<h2>Save options</h2>", "", gr.HTML),
|
||||
"save_images_add_number": OptionInfo(True, "Numbered filenames", component_args=hide_dirs),
|
||||
"use_original_name_batch": OptionInfo(True, "Batch uses original name"),
|
||||
"use_upscaler_name_as_suffix": OptionInfo(True, "Use upscaler as suffix", gr.Checkbox, {"visible": False}),
|
||||
"save_to_dirs": OptionInfo(False, "Save images to a subdirectory"),
|
||||
"directories_filename_pattern": OptionInfo("[date]", "Directory name pattern", component_args=hide_dirs),
|
||||
"samples_filename_pattern": OptionInfo("[seq]-[model_name]-[prompt_words]", "Images filename pattern", component_args=hide_dirs),
|
||||
"directories_max_prompt_words": OptionInfo(8, "Max words per pattern", gr.Slider, {"minimum": 1, "maximum": 99, "step": 1, **hide_dirs}),
|
||||
"use_save_to_dirs_for_ui": OptionInfo(False, "Save images to a subdirectory when using Save button", gr.Checkbox, {"visible": False}),
|
||||
|
||||
"outdir_sep_dirs": OptionInfo("<h2>Folders</h2>", "", gr.HTML),
|
||||
"outdir_samples": OptionInfo("", "Images folder", component_args=hide_dirs, folder=True),
|
||||
@@ -711,8 +742,6 @@ options_templates.update(options_section(('saving-paths', "Image Naming & Paths"
|
||||
"outdir_init_images": OptionInfo("outputs/init-images", "Folder for init images", component_args=hide_dirs, folder=True),
|
||||
|
||||
"outdir_sep_grids": OptionInfo("<h2>Grids</h2>", "", gr.HTML),
|
||||
"grid_extended_filename": OptionInfo(True, "Add extended info to filename when saving grid", gr.Checkbox, {"visible": False}),
|
||||
"grid_save_to_dirs": OptionInfo(False, "Save grids to a subdirectory", gr.Checkbox, {"visible": False}),
|
||||
"outdir_grids": OptionInfo("", "Grids folder", component_args=hide_dirs, folder=True),
|
||||
"outdir_txt2img_grids": OptionInfo("outputs/grids", 'Folder for txt2img grids', component_args=hide_dirs, folder=True),
|
||||
"outdir_img2img_grids": OptionInfo("outputs/grids", 'Folder for img2img grids', component_args=hide_dirs, folder=True),
|
||||
@@ -725,7 +754,6 @@ options_templates.update(options_section(('ui', "User Interface Options"), {
|
||||
"gradio_theme": OptionInfo("black-teal", "UI theme", gr.Dropdown, lambda: {"choices": theme.list_themes()}, refresh=theme.refresh_themes),
|
||||
"autolaunch": OptionInfo(False, "Autolaunch browser upon startup"),
|
||||
"font_size": OptionInfo(14, "Font size", gr.Slider, {"minimum": 8, "maximum": 32, "step": 1, "visible": True}),
|
||||
"tooltips": OptionInfo("UI Tooltips", "UI tooltips", gr.Radio, {"choices": ["None", "Browser default", "UI tooltips"], "visible": False}),
|
||||
"aspect_ratios": OptionInfo("1:1, 4:3, 3:2, 16:9, 16:10, 21:9, 2:3, 3:4, 9:16, 10:16, 9:21", "Allowed aspect ratios"),
|
||||
"motd": OptionInfo(True, "Show MOTD"),
|
||||
"compact_view": OptionInfo(False, "Compact view"),
|
||||
@@ -735,22 +763,14 @@ options_templates.update(options_section(('ui', "User Interface Options"), {
|
||||
"disable_weights_auto_swap": OptionInfo(True, "Do not change selected model when reading generation parameters"),
|
||||
"send_seed": OptionInfo(True, "Send seed when sending prompt or image to other interface"),
|
||||
"send_size": OptionInfo(True, "Send size when sending prompt or image to another interface"),
|
||||
"keyedit_precision_attention": OptionInfo(0.1, "Ctrl+up/down precision when editing (attention:1.1)", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}),
|
||||
"keyedit_precision_extra": OptionInfo(0.05, "Ctrl+up/down precision when editing <extra networks:0.9>", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}),
|
||||
"keyedit_delimiters": OptionInfo(r".,\/!?%^*;:{}=`~()", "Ctrl+up/down word delimiters", gr.Textbox, { "visible": False }),
|
||||
"quicksettings_list": OptionInfo(["sd_model_checkpoint"], "Quicksettings list", gr.Dropdown, lambda: {"multiselect":True, "choices": list(opts.data_labels.keys())}),
|
||||
"ui_scripts_reorder": OptionInfo("", "UI scripts order", gr.Textbox, { "visible": False }),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section(('live-preview', "Live Previews"), {
|
||||
"show_progressbar": OptionInfo(True, "Show progressbar", gr.Checkbox, {"visible": False}),
|
||||
"live_previews_enable": OptionInfo(True, "Show live previews", gr.Checkbox, {"visible": False}),
|
||||
"show_progress_grid": OptionInfo(True, "Show previews as a grid", gr.Checkbox, {"visible": False}),
|
||||
"notification_audio_enable": OptionInfo(False, "Play a notification upon completion"),
|
||||
"notification_audio_path": OptionInfo("html/notification.mp3","Path to notification sound", component_args=hide_dirs, folder=True),
|
||||
"show_progress_every_n_steps": OptionInfo(1, "Live preview display period", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1}),
|
||||
"show_progress_type": OptionInfo("Approximate", "Live preview method", gr.Radio, {"choices": ["Simple", "Approximate", "TAESD", "Full VAE"]}),
|
||||
"live_preview_content": OptionInfo("Combined", "Live preview subject", gr.Radio, {"choices": ["Combined", "Prompt", "Negative prompt"], "visible": False}),
|
||||
"live_preview_refresh_period": OptionInfo(500, "Progress update period", gr.Slider, {"minimum": 0, "maximum": 5000, "step": 25}),
|
||||
"live_preview_taesd_layers": OptionInfo(3, "TAESD decode layers", gr.Slider, {"minimum": 1, "maximum": 3, "step": 1}),
|
||||
"logmonitor_show": OptionInfo(True, "Show log view"),
|
||||
@@ -781,7 +801,7 @@ options_templates.update(options_section(('sampler-params', "Sampler Settings"),
|
||||
'schedulers_beta_start': OptionInfo(0, "Beta start", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.00001, "visible": native}),
|
||||
'schedulers_beta_end': OptionInfo(0, "Beta end", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.00001, "visible": native}),
|
||||
'schedulers_timesteps_range': OptionInfo(1000, "Timesteps range", gr.Slider, {"minimum": 250, "maximum": 4000, "step": 1, "visible": native}),
|
||||
'schedulers_shift': OptionInfo(1, "Sampler shift", gr.Slider, {"minimum": 0.1, "maximum": 10, "step": 0.1, "visible": native}),
|
||||
'schedulers_shift': OptionInfo(0, "Sampler shift", gr.Slider, {"minimum": 0.1, "maximum": 10, "step": 0.1, "visible": native}),
|
||||
'schedulers_dynamic_shift': OptionInfo(True, "Sampler dynamic shift", gr.Checkbox, {"visible": native}),
|
||||
|
||||
# managed from ui.py for backend original k-diffusion
|
||||
@@ -797,8 +817,6 @@ options_templates.update(options_section(('sampler-params', "Sampler Settings"),
|
||||
'uni_pc_variant': OptionInfo("bh2", "UniPC variant", gr.Radio, {"choices": ["bh1", "bh2", "vary_coeff"], "visible": not native}),
|
||||
'uni_pc_skip_type': OptionInfo("time_uniform", "UniPC skip type", gr.Radio, {"choices": ["time_uniform", "time_quadratic", "logSNR"], "visible": not native}),
|
||||
"ddim_discretize": OptionInfo('uniform', "DDIM discretize img2img", gr.Radio, {"choices": ['uniform', 'quad'], "visible": not native}),
|
||||
"pad_cond_uncond": OptionInfo(True, "Pad prompt and negative prompt to be same length", gr.Checkbox, {"visible": False}),
|
||||
"batch_cond_uncond": OptionInfo(True, "Do conditional and unconditional denoising in one batch", gr.Checkbox, {"visible": False}),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section(('postprocessing', "Postprocessing"), {
|
||||
@@ -808,12 +826,10 @@ options_templates.update(options_section(('postprocessing', "Postprocessing"), {
|
||||
"postprocessing_sep_img2img": OptionInfo("<h2>Img2Img & Inpainting</h2>", "", gr.HTML),
|
||||
"img2img_color_correction": OptionInfo(False, "Apply color correction"),
|
||||
"mask_apply_overlay": OptionInfo(True, "Apply mask as overlay"),
|
||||
"img2img_fix_steps": OptionInfo(False, "For image processing do exact number of steps as specified", gr.Checkbox, { "visible": False }),
|
||||
"img2img_background_color": OptionInfo("#ffffff", "Image transparent color fill", gr.ColorPicker, {}),
|
||||
"inpainting_mask_weight": OptionInfo(1.0, "Inpainting conditioning mask strength", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01}),
|
||||
"initial_noise_multiplier": OptionInfo(1.0, "Noise multiplier for image processing", gr.Slider, {"minimum": 0.1, "maximum": 1.5, "step": 0.01}),
|
||||
"img2img_extra_noise": OptionInfo(0.0, "Extra noise multiplier for img2img", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01}),
|
||||
"CLIP_stop_at_last_layers": OptionInfo(1, "Clip skip", gr.Slider, {"minimum": 1, "maximum": 8, "step": 1, "visible": False}),
|
||||
"initial_noise_multiplier": OptionInfo(1.0, "Noise multiplier for image processing", gr.Slider, {"minimum": 0.1, "maximum": 1.5, "step": 0.01, "visible": not native}),
|
||||
"img2img_extra_noise": OptionInfo(0.0, "Extra noise multiplier for img2img", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01, "visible": not native}),
|
||||
|
||||
# "postprocessing_sep_detailer": OptionInfo("<h2>Detailer</h2>", "", gr.HTML),
|
||||
"detailer_model": OptionInfo("Detailer", "Detailer model", gr.Radio, lambda: {"choices": [x.name() for x in detailers], "visible": False}),
|
||||
@@ -835,7 +851,6 @@ options_templates.update(options_section(('postprocessing', "Postprocessing"), {
|
||||
|
||||
"postprocessing_sep_upscalers": OptionInfo("<h2>Upscaling</h2>", "", gr.HTML),
|
||||
"upscaler_unload": OptionInfo(False, "Unload upscaler after processing"),
|
||||
"upscaler_for_img2img": OptionInfo("None", "Default upscaler for image resize operations", gr.Dropdown, lambda: {"choices": [x.name for x in sd_upscalers], "visible": False}, refresh=refresh_upscalers),
|
||||
"upscaler_tile_size": OptionInfo(192, "Upscaler tile size", gr.Slider, {"minimum": 0, "maximum": 512, "step": 16}),
|
||||
"upscaler_tile_overlap": OptionInfo(8, "Upscaler tile overlap", gr.Slider, {"minimum": 0, "maximum": 64, "step": 1}),
|
||||
}))
|
||||
@@ -846,28 +861,12 @@ options_templates.update(options_section(('control', "Control Options"), {
|
||||
"control_unload_processor": OptionInfo(False, "Processor unload after use"),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section(('interrogate', "Interrogate"), { # "Training" section disabled so just a placeholder
|
||||
"unload_models_when_training": OptionInfo(False, "Move VAE and CLIP to RAM when training", gr.Checkbox, { "visible": False }),
|
||||
"pin_memory": OptionInfo(True, "Pin training dataset to memory", gr.Checkbox, { "visible": False }),
|
||||
"save_optimizer_state": OptionInfo(False, "Save resumable optimizer state when training", gr.Checkbox, { "visible": False }),
|
||||
"save_training_settings_to_txt": OptionInfo(True, "Save training settings to a text file", gr.Checkbox, { "visible": False }),
|
||||
"dataset_filename_word_regex": OptionInfo("", "Filename word regex", gr.Textbox, { "visible": False }),
|
||||
"dataset_filename_join_string": OptionInfo(" ", "Filename join string", gr.Textbox, { "visible": False }),
|
||||
"embeddings_templates_dir": OptionInfo("", "Embeddings train templates directory", gr.Textbox, { "visible": False }),
|
||||
"training_image_repeats_per_epoch": OptionInfo(1, "Image repeats per epoch", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1, "visible": False }),
|
||||
"training_write_csv_every": OptionInfo(0, "Save loss CSV file every n steps", gr.Number, { "visible": False }),
|
||||
"training_enable_tensorboard": OptionInfo(False, "Enable tensorboard logging", gr.Checkbox, { "visible": False }),
|
||||
"training_tensorboard_save_images": OptionInfo(False, "Save generated images within tensorboard", gr.Checkbox, { "visible": False }),
|
||||
"training_tensorboard_flush_every": OptionInfo(120, "Tensorboard flush period", gr.Number, { "visible": False }),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section(('interrogate', "Interrogate"), {
|
||||
"interrogate_keep_models_in_memory": OptionInfo(False, "Interrogate: keep models in VRAM"),
|
||||
"interrogate_return_ranks": OptionInfo(True, "Interrogate: include ranks of model tags matches in results"),
|
||||
"interrogate_clip_num_beams": OptionInfo(1, "Interrogate: num_beams for BLIP", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1}),
|
||||
"interrogate_clip_min_length": OptionInfo(32, "Interrogate: minimum description length", gr.Slider, {"minimum": 1, "maximum": 128, "step": 1}),
|
||||
"interrogate_clip_max_length": OptionInfo(192, "Interrogate: maximum description length", gr.Slider, {"minimum": 1, "maximum": 256, "step": 1}),
|
||||
"interrogate_clip_dict_limit": OptionInfo(2048, "CLIP: maximum number of lines in text file", gr.Slider, { "visible": False }),
|
||||
"interrogate_clip_skip_categories": OptionInfo(["artists", "movements", "flavors"], "Interrogate: skip categories", gr.CheckboxGroup, lambda: {"choices": modules.interrogate.category_types()}, refresh=modules.interrogate.category_types),
|
||||
"interrogate_deepbooru_score_threshold": OptionInfo(0.65, "Interrogate: deepbooru score threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01}),
|
||||
"deepbooru_sort_alpha": OptionInfo(False, "Interrogate: deepbooru sort alphabetically"),
|
||||
@@ -888,7 +887,6 @@ options_templates.update(options_section(('extra_networks', "Networks"), {
|
||||
"extra_networks_card_size": OptionInfo(160, "UI card size (px)", gr.Slider, {"minimum": 20, "maximum": 2000, "step": 1}),
|
||||
"extra_networks_card_square": OptionInfo(True, "UI disable variable aspect ratio"),
|
||||
"extra_networks_fetch": OptionInfo(True, "UI fetch network info on mouse-over"),
|
||||
"extra_networks_card_fit": OptionInfo("cover", "UI image contain method", gr.Radio, {"choices": ["contain", "cover", "fill"], "visible": False}),
|
||||
"extra_network_skip_indexing": OptionInfo(False, "Build info on first access", gr.Checkbox),
|
||||
|
||||
"extra_networks_model_sep": OptionInfo("<h2>Models</h2>", "", gr.HTML),
|
||||
@@ -909,17 +907,59 @@ options_templates.update(options_section(('extra_networks', "Networks"), {
|
||||
"lora_apply_tags": OptionInfo(0, "LoRA auto-apply tags", gr.Slider, {"minimum": -1, "maximum": 32, "step": 1}),
|
||||
"lora_in_memory_limit": OptionInfo(0, "LoRA memory cache", gr.Slider, {"minimum": 0, "maximum": 24, "step": 1}),
|
||||
"lora_quant": OptionInfo("NF4","LoRA precision in quantized models", gr.Radio, {"choices": ["NF4", "FP4"]}),
|
||||
"lora_functional": OptionInfo(False, "Use Kohya method for handling multiple LoRA", gr.Checkbox, { "visible": False }),
|
||||
"lora_load_gpu": OptionInfo(True if not cmd_opts.lowvram else False, "Load LoRA directly to GPU"),
|
||||
}))
|
||||
|
||||
"hypernetwork_enabled": OptionInfo(False, "Enable Hypernetwork support", gr.Checkbox, {"visible": False}),
|
||||
"sd_hypernetwork": OptionInfo("None", "Add hypernetwork to prompt", gr.Dropdown, { "choices": ["None"], "visible": False }),
|
||||
options_templates.update(options_section((None, "Internal options"), {
|
||||
"diffusers_version": OptionInfo("", "Diffusers version", gr.Textbox, {"visible": False}),
|
||||
"disabled_extensions": OptionInfo([], "Disable these extensions"),
|
||||
"sd_checkpoint_hash": OptionInfo("", "SHA256 hash of the current checkpoint"),
|
||||
"tooltips": OptionInfo("UI Tooltips", "UI tooltips", gr.Radio, {"choices": ["None", "Browser default", "UI tooltips"], "visible": False}),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section((None, "Hidden options"), {
|
||||
"disabled_extensions": OptionInfo([], "Disable these extensions"),
|
||||
"batch_cond_uncond": OptionInfo(True, "Do conditional and unconditional denoising in one batch", gr.Checkbox, {"visible": False}),
|
||||
"CLIP_stop_at_last_layers": OptionInfo(1, "Clip skip", gr.Slider, {"minimum": 1, "maximum": 8, "step": 1, "visible": False}),
|
||||
"dataset_filename_join_string": OptionInfo(" ", "Filename join string", gr.Textbox, { "visible": False }),
|
||||
"dataset_filename_word_regex": OptionInfo("", "Filename word regex", gr.Textbox, { "visible": False }),
|
||||
"diffusers_force_zeros": OptionInfo(False, "Force zeros for prompts when empty", gr.Checkbox, {"visible": False}),
|
||||
"disable_all_extensions": OptionInfo("none", "Disable all extensions (preserves the list of disabled extensions)", gr.Radio, {"choices": ["none", "user", "all"]}),
|
||||
"sd_checkpoint_hash": OptionInfo("", "SHA256 hash of the current checkpoint"),
|
||||
"disable_nan_check": OptionInfo(True, "Disable NaN check", gr.Checkbox, {"visible": False}),
|
||||
"embeddings_templates_dir": OptionInfo("", "Embeddings train templates directory", gr.Textbox, { "visible": False }),
|
||||
"extra_networks_card_fit": OptionInfo("cover", "UI image contain method", gr.Radio, {"choices": ["contain", "cover", "fill"], "visible": False}),
|
||||
"grid_extended_filename": OptionInfo(True, "Add extended info to filename when saving grid", gr.Checkbox, {"visible": False}),
|
||||
"grid_save_to_dirs": OptionInfo(False, "Save grids to a subdirectory", gr.Checkbox, {"visible": False}),
|
||||
"hypernetwork_enabled": OptionInfo(False, "Enable Hypernetwork support", gr.Checkbox, {"visible": False}),
|
||||
"img2img_fix_steps": OptionInfo(False, "For image processing do exact number of steps as specified", gr.Checkbox, { "visible": False }),
|
||||
"interrogate_clip_dict_limit": OptionInfo(2048, "CLIP: maximum number of lines in text file", gr.Slider, { "visible": False }),
|
||||
"keyedit_delimiters": OptionInfo(r".,\/!?%^*;:{}=`~()", "Ctrl+up/down word delimiters", gr.Textbox, { "visible": False }),
|
||||
"keyedit_precision_attention": OptionInfo(0.1, "Ctrl+up/down precision when editing (attention:1.1)", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}),
|
||||
"keyedit_precision_extra": OptionInfo(0.05, "Ctrl+up/down precision when editing <extra networks:0.9>", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}),
|
||||
"live_preview_content": OptionInfo("Combined", "Live preview subject", gr.Radio, {"choices": ["Combined", "Prompt", "Negative prompt"], "visible": False}),
|
||||
"live_previews_enable": OptionInfo(True, "Show live previews", gr.Checkbox, {"visible": False}),
|
||||
"lora_functional": OptionInfo(False, "Use Kohya method for handling multiple LoRA", gr.Checkbox, { "visible": False }),
|
||||
"lyco_dir": OptionInfo(os.path.join(paths.models_path, 'LyCORIS'), "Folder with LyCORIS network(s)", gr.Text, {"visible": False}),
|
||||
"model_reuse_dict": OptionInfo(False, "Reuse loaded model dictionary", gr.Checkbox, {"visible": False}),
|
||||
"pad_cond_uncond": OptionInfo(True, "Pad prompt and negative prompt to be same length", gr.Checkbox, {"visible": False}),
|
||||
"pin_memory": OptionInfo(True, "Pin training dataset to memory", gr.Checkbox, { "visible": False }),
|
||||
"save_optimizer_state": OptionInfo(False, "Save resumable optimizer state when training", gr.Checkbox, { "visible": False }),
|
||||
"save_training_settings_to_txt": OptionInfo(True, "Save training settings to a text file", gr.Checkbox, { "visible": False }),
|
||||
"sd_disable_ckpt": OptionInfo(False, "Disallow models in ckpt format", gr.Checkbox, {"visible": False}),
|
||||
"sd_hypernetwork": OptionInfo("None", "Add hypernetwork to prompt", gr.Dropdown, { "choices": ["None"], "visible": False }),
|
||||
"sd_lora": OptionInfo("", "Add LoRA to prompt", gr.Textbox, {"visible": False}),
|
||||
"sd_vae_checkpoint_cache": OptionInfo(0, "Cached VAEs", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": False}),
|
||||
"show_progress_grid": OptionInfo(True, "Show previews as a grid", gr.Checkbox, {"visible": False}),
|
||||
"show_progressbar": OptionInfo(True, "Show progressbar", gr.Checkbox, {"visible": False}),
|
||||
"training_enable_tensorboard": OptionInfo(False, "Enable tensorboard logging", gr.Checkbox, { "visible": False }),
|
||||
"training_image_repeats_per_epoch": OptionInfo(1, "Image repeats per epoch", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1, "visible": False }),
|
||||
"training_tensorboard_flush_every": OptionInfo(120, "Tensorboard flush period", gr.Number, { "visible": False }),
|
||||
"training_tensorboard_save_images": OptionInfo(False, "Save generated images within tensorboard", gr.Checkbox, { "visible": False }),
|
||||
"training_write_csv_every": OptionInfo(0, "Save loss CSV file every n steps", gr.Number, { "visible": False }),
|
||||
"ui_scripts_reorder": OptionInfo("", "UI scripts order", gr.Textbox, { "visible": False }),
|
||||
"unload_models_when_training": OptionInfo(False, "Move VAE and CLIP to RAM when training", gr.Checkbox, { "visible": False }),
|
||||
"upscaler_for_img2img": OptionInfo("None", "Default upscaler for image resize operations", gr.Dropdown, lambda: {"choices": [x.name for x in sd_upscalers], "visible": False}, refresh=refresh_upscalers),
|
||||
"use_save_to_dirs_for_ui": OptionInfo(False, "Save images to a subdirectory when using Save button", gr.Checkbox, {"visible": False}),
|
||||
"use_upscaler_name_as_suffix": OptionInfo(True, "Use upscaler as suffix", gr.Checkbox, {"visible": False}),
|
||||
}))
|
||||
|
||||
options_templates.update()
|
||||
@@ -993,7 +1033,7 @@ class Options:
|
||||
if filename is None:
|
||||
filename = self.filename
|
||||
if cmd_opts.freeze:
|
||||
log.warning(f'Settings saving is disabled: {filename}')
|
||||
log.warning(f'Setting: fn="{filename}" save disabled')
|
||||
return
|
||||
try:
|
||||
# output = json.dumps(self.data, indent=2)
|
||||
@@ -1001,12 +1041,12 @@ class Options:
|
||||
unused_settings = []
|
||||
|
||||
if os.environ.get('SD_CONFIG_DEBUG', None) is not None:
|
||||
log.debug('Config: user settings')
|
||||
log.debug('Settings: user')
|
||||
for k, v in self.data.items():
|
||||
log.trace(f' Config: item={k} value={v} default={self.data_labels[k].default if k in self.data_labels else None}')
|
||||
log.debug('Config: default settings')
|
||||
log.debug('Settings: defaults')
|
||||
for k in self.data_labels.keys():
|
||||
log.trace(f' Config: item={k} default={self.data_labels[k].default}')
|
||||
log.trace(f' Setting: item={k} default={self.data_labels[k].default}')
|
||||
|
||||
for k, v in self.data.items():
|
||||
if k in self.data_labels:
|
||||
@@ -1021,9 +1061,9 @@ class Options:
|
||||
unused_settings.append(k)
|
||||
writefile(diff, filename, silent=silent)
|
||||
if len(unused_settings) > 0:
|
||||
log.debug(f"Unused settings: {unused_settings}")
|
||||
log.debug(f"Settings: unused={unused_settings}")
|
||||
except Exception as err:
|
||||
log.error(f'Save settings failed: {filename} {err}')
|
||||
log.error(f'Settings: fn="{filename}" {err}')
|
||||
|
||||
def save(self, filename=None, silent=False):
|
||||
threading.Thread(target=self.save_atomic, args=(filename, silent)).start()
|
||||
@@ -1039,7 +1079,7 @@ class Options:
|
||||
if filename is None:
|
||||
filename = self.filename
|
||||
if not os.path.isfile(filename):
|
||||
log.debug(f'Created default config: {filename}')
|
||||
log.debug(f'Settings: fn="{filename}" created')
|
||||
self.save(filename)
|
||||
return
|
||||
self.data = readfile(filename, lock=True)
|
||||
@@ -1047,13 +1087,17 @@ class Options:
|
||||
self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings').split(',')]
|
||||
unknown_settings = []
|
||||
for k, v in self.data.items():
|
||||
info = self.data_labels.get(k, None)
|
||||
info: OptionInfo = self.data_labels.get(k, None)
|
||||
if info is not None:
|
||||
if not info.validate(k, v):
|
||||
self.data[k] = info.default
|
||||
if info is not None and not self.same_type(info.default, v):
|
||||
log.error(f"Error: bad setting value: {k}: {v} ({type(v).__name__}; expected {type(info.default).__name__})")
|
||||
log.warning(f"Setting validation: {k}={v} ({type(v).__name__} expected={type(info.default).__name__})")
|
||||
self.data[k] = info.default
|
||||
if info is None and k not in compatibility_opts and not k.startswith('uiux_'):
|
||||
unknown_settings.append(k)
|
||||
if len(unknown_settings) > 0:
|
||||
log.debug(f"Unknown settings: {unknown_settings}")
|
||||
log.warning(f"Setting validation: unknown={unknown_settings}")
|
||||
|
||||
def onchange(self, key, func, call=True):
|
||||
item = self.data_labels.get(key)
|
||||
@@ -1111,7 +1155,7 @@ cmd_opts = cmd_args.settings_args(opts, cmd_opts)
|
||||
if cmd_opts.use_xformers:
|
||||
opts.data['cross_attention_optimization'] = 'xFormers'
|
||||
opts.data['uni_pc_lower_order_final'] = opts.schedulers_use_loworder # compatibility
|
||||
opts.data['uni_pc_order'] = opts.schedulers_solver_order # compatibility
|
||||
opts.data['uni_pc_order'] = max(2, opts.schedulers_solver_order) # compatibility
|
||||
log.info(f'Engine: backend={backend} compute={devices.backend} device={devices.get_optimal_device_name()} attention="{opts.cross_attention_optimization}" mode={devices.inference_context.__name__}')
|
||||
if not native:
|
||||
log.warning('Backend=original is in maintainance-only mode')
|
||||
@@ -1228,6 +1272,7 @@ def html(filename):
|
||||
return ""
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_version():
|
||||
version = None
|
||||
if version is None:
|
||||
|
||||
@@ -62,6 +62,38 @@ class State:
|
||||
}
|
||||
return obj
|
||||
|
||||
def status(self):
|
||||
from modules import progress
|
||||
from modules.api import models
|
||||
res = models.ResStatus(
|
||||
task=self.job,
|
||||
id=progress.current_task or '',
|
||||
job=max(self.job_no, 0),
|
||||
jobs=max(self.frame_count, self.job_count, self.job_no),
|
||||
total=self.total_jobs,
|
||||
timestamp=self.job_timestamp if self.job != '' else None,
|
||||
step=self.sampling_step,
|
||||
steps=self.sampling_steps,
|
||||
queued=len(progress.pending_tasks),
|
||||
status='unknown',
|
||||
uptime = round(time.time() - self.server_start)
|
||||
)
|
||||
res.step = res.steps * res.job + res.step
|
||||
res.steps = res.steps * res.jobs
|
||||
res.progress = round(min(1, abs(res.step / res.steps) if res.steps > 0 else 0), 2)
|
||||
res.elapsed = round(time.time() - self.time_start, 2) if self.time_start is not None else None
|
||||
predicted = round(res.elapsed / res.progress, 2) if res.progress > 0 and res.elapsed is not None else None
|
||||
res.eta = round(predicted - res.elapsed, 2) if predicted is not None else None
|
||||
if self.paused:
|
||||
res.status = 'paused'
|
||||
elif self.interrupted:
|
||||
res.status = 'interrupted'
|
||||
elif self.skipped:
|
||||
res.status = 'skipped'
|
||||
else:
|
||||
res.status = 'running' if self.job != '' else 'idle'
|
||||
return res
|
||||
|
||||
def begin(self, title="", api=None):
|
||||
import modules.devices
|
||||
self.total_jobs += 1
|
||||
|
||||
+22
-8
@@ -112,6 +112,8 @@ def apply_wildcards_to_prompt(prompt, all_wildcards, seed=-1, silent=False):
|
||||
|
||||
|
||||
def get_reference_style():
|
||||
if getattr(shared.sd_model, 'sd_checkpoint_info', None) is None:
|
||||
return None
|
||||
name = shared.sd_model.sd_checkpoint_info.name
|
||||
name = name.replace('\\', '/').replace('Diffusers/', '')
|
||||
for k, v in shared.reference_models.items():
|
||||
@@ -251,27 +253,33 @@ class StyleDatabase:
|
||||
return found[0] if len(found) > 0 else self.no_style
|
||||
|
||||
def get_style_prompts(self, styles):
|
||||
if styles is None or not isinstance(styles, list):
|
||||
if styles is None:
|
||||
return []
|
||||
if not isinstance(styles, list):
|
||||
shared.log.error(f'Styles invalid: {styles}')
|
||||
return []
|
||||
return [self.find_style(x).prompt for x in styles]
|
||||
|
||||
def get_negative_style_prompts(self, styles):
|
||||
if styles is None or not isinstance(styles, list):
|
||||
if styles is None:
|
||||
return []
|
||||
if not isinstance(styles, list):
|
||||
shared.log.error(f'Styles invalid: {styles}')
|
||||
return []
|
||||
return [self.find_style(x).negative_prompt for x in styles]
|
||||
|
||||
def apply_styles_to_prompts(self, prompts, negatives, styles, seeds):
|
||||
if styles is None or not isinstance(styles, list):
|
||||
if styles is None:
|
||||
return prompts, negatives
|
||||
if not isinstance(styles, list):
|
||||
shared.log.error(f'Styles invalid styles: {styles}')
|
||||
return prompts
|
||||
return prompts, negatives
|
||||
if prompts is None or not isinstance(prompts, list):
|
||||
shared.log.error(f'Styles invalid prompts: {prompts}')
|
||||
return prompts
|
||||
return prompts, negatives
|
||||
if seeds is None or not isinstance(prompts, list):
|
||||
shared.log.error(f'Styles invalid seeds: {seeds}')
|
||||
return prompts
|
||||
return prompts, negatives
|
||||
parsed_positive = []
|
||||
parsed_negative = []
|
||||
for i in range(len(prompts)):
|
||||
@@ -286,7 +294,9 @@ class StyleDatabase:
|
||||
return parsed_positive, parsed_negative
|
||||
|
||||
def apply_styles_to_prompt(self, prompt, styles):
|
||||
if styles is None or not isinstance(styles, list):
|
||||
if styles is None:
|
||||
return prompt
|
||||
if not isinstance(styles, list):
|
||||
shared.log.error(f'Styles invalid: {styles}')
|
||||
return prompt
|
||||
prompt = apply_styles_to_prompt(prompt, [self.find_style(x).prompt for x in styles])
|
||||
@@ -294,7 +304,9 @@ class StyleDatabase:
|
||||
return prompt
|
||||
|
||||
def apply_negative_styles_to_prompt(self, prompt, styles):
|
||||
if styles is None or not isinstance(styles, list):
|
||||
if styles is None:
|
||||
return prompt
|
||||
if not isinstance(styles, list):
|
||||
shared.log.error(f'Styles invalid: {styles}')
|
||||
return prompt
|
||||
prompt = apply_styles_to_prompt(prompt, [self.find_style(x).negative_prompt for x in styles])
|
||||
@@ -302,6 +314,8 @@ class StyleDatabase:
|
||||
return prompt
|
||||
|
||||
def apply_styles_to_extra(self, p):
|
||||
if p.styles is None:
|
||||
return
|
||||
if p.styles is None or not isinstance(p.styles, list):
|
||||
shared.log.error(f'Styles invalid: {p.styles}')
|
||||
return
|
||||
|
||||
@@ -49,7 +49,6 @@ def txt2img(id_task, state,
|
||||
subseed_strength=subseed_strength,
|
||||
seed_resize_from_h=seed_resize_from_h,
|
||||
seed_resize_from_w=seed_resize_from_w,
|
||||
seed_enable_extras=True,
|
||||
sampler_name = processing.get_sampler_name(sampler_index),
|
||||
hr_sampler_name = processing.get_sampler_name(hr_sampler_index),
|
||||
batch_size=batch_size,
|
||||
|
||||
+11
-16
@@ -354,22 +354,6 @@ def create_ui(startup_timer = None):
|
||||
from modules.onnx_impl import ui as ui_onnx
|
||||
ui_onnx.create_ui()
|
||||
|
||||
with gr.TabItem("Change log", id="change_log", elem_id="system_tab_changelog"):
|
||||
def get_changelog():
|
||||
with open('CHANGELOG.md', 'r', encoding='utf-8') as f:
|
||||
content = f.read()
|
||||
content = content.replace('# Change Log for SD.Next', ' ')
|
||||
return content
|
||||
|
||||
with gr.Column():
|
||||
get_changelog_btn = gr.Button(value='Get changelog', elem_id="get_changelog")
|
||||
with gr.Column():
|
||||
_changelog_search = gr.Textbox(label="Search", elem_id="changelog_search")
|
||||
_changelog_result = gr.HTML(elem_id="changelog_result")
|
||||
|
||||
changelog_markdown = gr.Markdown('', elem_id="changelog_markdown")
|
||||
get_changelog_btn.click(fn=get_changelog, outputs=[changelog_markdown], show_progress=True)
|
||||
|
||||
def unload_sd_weights():
|
||||
modules.sd_models.unload_model_weights(op='model')
|
||||
modules.sd_models.unload_model_weights(op='refiner')
|
||||
@@ -390,6 +374,16 @@ def create_ui(startup_timer = None):
|
||||
|
||||
timer.startup.record("ui-settings")
|
||||
|
||||
with gr.Blocks(analytics_enabled=False) as info_interface:
|
||||
with gr.Tabs(elem_id="tabs_info"):
|
||||
with gr.TabItem("Change log", id="change_log", elem_id="system_tab_changelog"):
|
||||
from modules import ui_docs
|
||||
ui_docs.create_ui_logs()
|
||||
|
||||
with gr.TabItem("Wiki", id="wiki", elem_id="system_tab_wiki"):
|
||||
from modules import ui_docs
|
||||
ui_docs.create_ui_wiki()
|
||||
|
||||
interfaces = []
|
||||
interfaces += [(txt2img_interface, "Text", "txt2img")]
|
||||
interfaces += [(img2img_interface, "Image", "img2img")]
|
||||
@@ -399,6 +393,7 @@ def create_ui(startup_timer = None):
|
||||
interfaces += [(models_interface, "Models", "models")]
|
||||
interfaces += script_callbacks.ui_tabs_callback()
|
||||
interfaces += [(settings_interface, "System", "system")]
|
||||
interfaces += [(info_interface, "Info", "info")]
|
||||
|
||||
from modules import ui_extensions
|
||||
extensions_interface = ui_extensions.create_ui()
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user