mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
civitai search redesign prototype
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Regular → Executable
+56
-30
@@ -1,84 +1,93 @@
|
||||
#!/usr/bin/env python
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
import sys
|
||||
import json
|
||||
import time
|
||||
import logging
|
||||
import bs4
|
||||
|
||||
|
||||
full_dct = False
|
||||
full_html = False
|
||||
debug = False
|
||||
logging.basicConfig(level = logging.INFO, format = '%(asctime)s %(levelname)s: %(message)s')
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ModelImage(object):
|
||||
@dataclass
|
||||
class ModelImage():
|
||||
def __init__(self, dct: dict):
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.dct: dict = dct
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.url: str = dct.get('url', '')
|
||||
self.width: int = dct.get('width', 0)
|
||||
self.height: int = dct.get('height', 0)
|
||||
self.type: str = dct.get('type', 'Unknown')
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'ModelImage(id={self.id} url="{self.url}" width={self.width} height={self.height} type="{self.type}")'
|
||||
|
||||
class ModelFile(object):
|
||||
@dataclass
|
||||
class ModelFile():
|
||||
def __init__(self, dct: dict):
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.dct: dict = dct
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.size: int = int(1024 * dct.get('sizeKB', 0))
|
||||
self.name: str = dct.get('name', 'Unknown')
|
||||
self.type: str = dct.get('type', 'Unknown')
|
||||
self.hashes: list[str] = dct.get('hashes', {}).values()
|
||||
self.url: str = dct.get('downloadUrl', '')
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'ModelFile(id={self.id} name="{self.name}" size={self.size} type="{self.type}" url="{self.url}")'
|
||||
|
||||
|
||||
class ModelVersion(object):
|
||||
@dataclass
|
||||
class ModelVersion():
|
||||
def __init__(self, dct: dict):
|
||||
import bs4
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.dct = dct
|
||||
self.id = dct.get('id', 0)
|
||||
self.name = dct.get('name', 'Unknown')
|
||||
self.base = dct.get('baseModel', 'Unknown')
|
||||
self.mtime = dct.get('publishedAt', '')
|
||||
self.downloads = dct.get('stats', {}).get('downloadCount', 0)
|
||||
self.availability = dct.get('availability', 'Unknown')
|
||||
self.html = dct.get('description', '') or ''
|
||||
self.desc = bs4.BeautifulSoup(self.html, features="html.parser").get_text()
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.name: str = dct.get('name', 'Unknown')
|
||||
self.base: str = dct.get('baseModel', 'Unknown')
|
||||
self.mtime: str = dct.get('publishedAt', '')
|
||||
self.downloads: int = dct.get('stats', {}).get('downloadCount', 0)
|
||||
self.availability: str = dct.get('availability', 'Unknown')
|
||||
self.html: str = dct.get('description', '') or '' if full_html else ''
|
||||
self.desc: str = bs4.BeautifulSoup(dct.get('description', '') or '', features="html.parser").get_text()
|
||||
self.files = [ModelFile(f) for f in dct.get('files', [])]
|
||||
self.images = [ModelImage(i) for i in dct.get('images', [])]
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'ModelVersion(id={self.id} name="{self.name}" base="{self.base}" mtime="{self.mtime}" downloads={self.downloads} availability={self.availability} desc="{self.desc[:30]}...")'
|
||||
|
||||
|
||||
class Model(object):
|
||||
@dataclass
|
||||
class Model():
|
||||
def __init__(self, dct: dict):
|
||||
import bs4
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.id = dct.get('id', 0)
|
||||
self.dct = dct
|
||||
self.url = f'https://civitai.com/models/{self.id}'
|
||||
self.type = dct.get('type', 'Unknown')
|
||||
self.name = dct.get('name', 'Unknown')
|
||||
self.html = dct.get('description', '')
|
||||
self.desc = bs4.BeautifulSoup(self.html, features="html.parser").get_text()
|
||||
self.tags = dct.get('tags', [])
|
||||
self.nsfw = dct.get('nsfw', False)
|
||||
self.level = dct.get('nsfwLevel', 0)
|
||||
self.availability = dct.get('availability', 'Unknown')
|
||||
self.downloads = dct.get('stats', {}).get('downloadCount', 0)
|
||||
self.creator = dct.get('creator', {}).get('username', 'Unknown')
|
||||
self.versions = [ModelVersion(v) for v in dct.get('modelVersions', [])]
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.url: str = f'https://civitai.com/models/{self.id}'
|
||||
self.type: str = dct.get('type', 'Unknown')
|
||||
self.name: str = dct.get('name', 'Unknown')
|
||||
self.html: str = dct.get('description', '') or '' if full_html else ''
|
||||
self.desc: str = bs4.BeautifulSoup(dct.get('description', '') or '', features="html.parser").get_text()
|
||||
self.tags: list[str] = dct.get('tags', [])
|
||||
self.nsfw: bool = dct.get('nsfw', False)
|
||||
self.level: str = dct.get('nsfwLevel', 0)
|
||||
self.availability: str = dct.get('availability', 'Unknown')
|
||||
self.downloads: int = dct.get('stats', {}).get('downloadCount', 0)
|
||||
self.creator: str = dct.get('creator', {}).get('username', 'Unknown')
|
||||
self.versions: list[ModelVersion] = [ModelVersion(v) for v in dct.get('modelVersions', [])]
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'Model(id={self.id} type={self.type} name="{self.name}" versions={len(self.versions)} nsfw={self.nsfw}/{self.level} downloads={self.downloads} author="{self.creator}" tags={self.tags} desc="{self.desc[:30]}...")'
|
||||
@@ -155,6 +164,23 @@ def search_civitai(
|
||||
return exact_models if len(exact_models) > 0 else models
|
||||
|
||||
|
||||
def models_to_dct(all_models:list, model_id:int=None):
|
||||
dct = []
|
||||
for model in all_models:
|
||||
if model_id is not None and model.id != model_id:
|
||||
continue
|
||||
model_dct = model.__dict__.copy()
|
||||
versions_dct = []
|
||||
for version in model.versions:
|
||||
version_dct = version.__dict__.copy()
|
||||
version_dct['files'] = [f.__dict__.copy() for f in version.files]
|
||||
version_dct['images'] = [i.__dict__.copy() for i in version.images]
|
||||
versions_dct.append(version_dct)
|
||||
model_dct['versions'] = versions_dct
|
||||
dct.append(model_dct)
|
||||
return dct
|
||||
|
||||
|
||||
def print_models(models: list[Model]):
|
||||
if debug:
|
||||
from rich import print as dbg
|
||||
|
||||
Submodule extensions-builtin/sdnext-modernui updated: 03e365e0b0...5403536134
@@ -0,0 +1,264 @@
|
||||
// hack to get pythons str.format in js
|
||||
String.prototype.format = function (arguments) { // eslint-disable-line no-extend-native, func-names
|
||||
let thisString = '';
|
||||
for (let charPos = 0; charPos < this.length; charPos++) thisString += this[charPos];
|
||||
for (const key in arguments) { // eslint-disable-line guard-for-in
|
||||
error(key, arguments[key]);
|
||||
const stringKey = `{${key}}`;
|
||||
thisString = thisString.replace(new RegExp(stringKey, 'g'), arguments[key]);
|
||||
}
|
||||
return thisString;
|
||||
};
|
||||
|
||||
const modelDetailsHTML = `
|
||||
<div id="model-details" class="model-details">
|
||||
<h3>{name}</h3>
|
||||
<p>Type: {type}</p>
|
||||
<p>Tags: {tags}</p>
|
||||
<p>NSFW: {nsfw}/{level}</p>
|
||||
<p>Availability: {availability}</p>
|
||||
<p>Downloads: {downloads}</p>
|
||||
<p>Author: {creator}</p>
|
||||
<div>{versions}</div>
|
||||
</div>
|
||||
`;
|
||||
|
||||
async function modelCardClick(id) {
|
||||
log('modelCardClick id', id);
|
||||
const el = gradioApp().getElementById('model-details');
|
||||
if (!el) return;
|
||||
const res = await fetch(`${window.api}/civitai?model_id=${encodeURI(id)}`);
|
||||
if (!res || res.status !== 200) {
|
||||
error(`modelCardClick: id=${id} status=${res ? res.status : 'unknown'}`);
|
||||
return;
|
||||
}
|
||||
let data = await res.json();
|
||||
log('modelCardClick data', data);
|
||||
if (!data || data.length === 0) return;
|
||||
data = data[0]; // assuming the first item is the one we want
|
||||
const obj = {
|
||||
name: data.name || 'unknown',
|
||||
type: data.type || 'unknown',
|
||||
tags: data.tags?.join(', ') || '',
|
||||
nsfw: data.nsfw ? 'yes' : 'no',
|
||||
level: data.level?.toString() || '',
|
||||
availability: data.availability || 'unknown',
|
||||
downloads: data.downloads?.toString() || '',
|
||||
creator: data.creator || 'unknown',
|
||||
versions: JSON.stringify(data.versions) || '[]',
|
||||
};
|
||||
log(obj);
|
||||
el.innerHTML = modelDetailsHTML.format({
|
||||
name: data.name || 'unknown',
|
||||
type: data.type || 'unknown',
|
||||
tags: data.tags?.join(', ') || '',
|
||||
nsfw: data.nsfw ? 'yes' : 'no',
|
||||
level: data.level?.toString() || '',
|
||||
availability: data.availability || 'unknown',
|
||||
downloads: data.downloads?.toString() || '',
|
||||
creator: data.creator || 'unknown',
|
||||
versions: JSON.stringify(data.versions) || '[]',
|
||||
});
|
||||
}
|
||||
|
||||
const example = {
|
||||
id: 1157409,
|
||||
url: 'https://civitai.com/models/1157409',
|
||||
type: 'Checkpoint',
|
||||
name: 'Tempest-by-Vlad',
|
||||
html: '',
|
||||
desc: 'Base versionFlexible SDXL model with custom encoder and finetuned for larger landscape resolutions with high details and high contrast.Recommended to use medium-low...',
|
||||
tags: [
|
||||
'base model',
|
||||
],
|
||||
nsfw: false,
|
||||
level: 15,
|
||||
availability: 'Public',
|
||||
downloads: 407,
|
||||
creator: 'vmandic',
|
||||
versions: [
|
||||
{
|
||||
id: 1301775,
|
||||
name: 'Base v0.1',
|
||||
base: 'SDXL 1.0',
|
||||
mtime: '2025-01-19T02:53:53.903Z',
|
||||
downloads: 346,
|
||||
availability: 'Public',
|
||||
html: '',
|
||||
desc: 'Initial release',
|
||||
files: [
|
||||
{
|
||||
id: 1206102,
|
||||
size: 6938089790,
|
||||
name: 'tempestByVlad_baseV01.safetensors',
|
||||
type: 'Model',
|
||||
hashes: [
|
||||
'79CB1E32',
|
||||
'8BFAD17222',
|
||||
'8BFAD1722243955B3F94103C69079C280D348B14729251E86824972C1063B616',
|
||||
'43E5E3BB',
|
||||
'DE83D56256411853AB6595CC3D8E865D5310D4A58D49A839DDC104C7F3429D4A',
|
||||
'4E933E1EBE61',
|
||||
],
|
||||
url: 'https://civitai.com/api/download/models/1301775',
|
||||
dct: {},
|
||||
},
|
||||
],
|
||||
images: [
|
||||
{
|
||||
id: 52503951,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/18c749f2-42ec-4024-9d20-0b1202b6bacc/width=1024/52503951.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 52508539,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/634d0ea8-ecdb-4ca6-a4ff-145319bc3fd3/width=1024/52508539.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 52508563,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/be529820-a89e-458f-8a3d-86cb43b154ac/width=1024/52508563.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 52508588,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/b2eb456a-a664-4de8-8c3e-6ecd1c4acb38/width=1024/52508588.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 52508654,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/f7b09e3f-4a48-459b-9b32-fa207904f74c/width=1024/52508654.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 52508659,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/fff89f18-628a-43b3-b9f5-44951dc078f7/width=1024/52508659.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 52508671,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/fdd3f774-ce2e-4ea7-82a0-bdb523fb86f6/width=1024/52508671.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 52512251,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/f26971c8-1123-45e3-a85b-2b97c6334b85/width=1024/52512251.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
],
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 1343512,
|
||||
name: 'Hyper v0.1',
|
||||
base: 'SDXL 1.0',
|
||||
mtime: '2025-01-28T22:54:12.734Z',
|
||||
downloads: 61,
|
||||
availability: 'Public',
|
||||
html: '',
|
||||
desc: 'Time-distilled version',
|
||||
files: [
|
||||
{
|
||||
id: 1246991,
|
||||
size: 6938085702,
|
||||
name: 'tempestByVlad_hyperV01.safetensors',
|
||||
type: 'Model',
|
||||
hashes: [
|
||||
'15943FD9',
|
||||
'4104FC6601',
|
||||
'4104FC6601F71C4C7A770AD422483FD700C8ECF72D06FCD8C4E8CD4B2D1C7DBB',
|
||||
'9F87BCEA',
|
||||
'CB52894625E9C13331285E4435799D707C4EAEF464974159C8B4B217EA32298E',
|
||||
'A0EE15E503DD',
|
||||
],
|
||||
url: 'https://civitai.com/api/download/models/1343512',
|
||||
dct: {},
|
||||
},
|
||||
],
|
||||
images: [
|
||||
{
|
||||
id: 54462987,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/1dd020c8-8a9a-4eb3-afec-fe83613217c5/width=1024/54462987.jpeg',
|
||||
width: 1024,
|
||||
height: 768,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 54462992,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/89edf233-939f-4b2e-97c8-175498704362/width=1536/54462992.jpeg',
|
||||
width: 1536,
|
||||
height: 640,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 54463002,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/87aea8fb-7687-48d0-98e6-27555b2ff87f/width=768/54463002.jpeg',
|
||||
width: 768,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 54463010,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/f1ee57cc-8920-4b8d-853a-ad1cfc7d9a5a/width=1024/54463010.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 54463011,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/dda6411b-fd36-484f-a9f0-0db847463128/width=1024/54463011.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 54463016,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/5f266a8c-f215-4e0f-8a64-aacc97f81d70/width=1024/54463016.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
{
|
||||
id: 54463019,
|
||||
url: 'https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/0ad04bdc-41bd-433c-9c25-f5365fbe082c/width=1024/54463019.jpeg',
|
||||
width: 1024,
|
||||
height: 1024,
|
||||
type: 'image',
|
||||
dct: {},
|
||||
},
|
||||
],
|
||||
dct: {},
|
||||
},
|
||||
],
|
||||
dct: {},
|
||||
};
|
||||
@@ -1515,12 +1515,6 @@ background: var(--background-color)
|
||||
min-height: 0;
|
||||
}
|
||||
|
||||
#models_error {
|
||||
font-family: monospace;
|
||||
|
||||
color: var(--body-text-color-subdued)
|
||||
}
|
||||
|
||||
#model_loader_df button {
|
||||
display: none !important;
|
||||
}
|
||||
|
||||
@@ -111,6 +111,10 @@ class Api:
|
||||
from modules.api import nudenet
|
||||
nudenet.register_api()
|
||||
|
||||
# civitai api
|
||||
from modules.civitai import api_civitai
|
||||
api_civitai.register_api()
|
||||
|
||||
|
||||
def add_api_route(self, path: str, endpoint, **kwargs):
|
||||
if (shared.cmd_opts.auth or shared.cmd_opts.auth_file) and shared.cmd_opts.api_only:
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
from starlette.responses import JSONResponse
|
||||
|
||||
|
||||
def models_to_json(all_models:list, model_id:int=None):
|
||||
dct = []
|
||||
for model in all_models:
|
||||
if model_id is not None and model.id != model_id:
|
||||
continue
|
||||
model_dct = model.__dict__.copy()
|
||||
versions_dct = []
|
||||
for version in model.versions:
|
||||
version_dct = version.__dict__.copy()
|
||||
version_dct['files'] = [f.__dict__.copy() for f in version.files]
|
||||
version_dct['images'] = [i.__dict__.copy() for i in version.images]
|
||||
versions_dct.append(version_dct)
|
||||
model_dct['versions'] = versions_dct
|
||||
dct.append(model_dct)
|
||||
# obj = json.dumps(dct, indent=2, ensure_ascii=False)
|
||||
return dct
|
||||
|
||||
|
||||
def get_civitai(
|
||||
model_id:int=None, # if model_id is provided assume fetch-from-cache
|
||||
query:str = '', # search query or tag is required
|
||||
tag:str = '', # search query or tag is required
|
||||
types:str = '', # Checkpoint, TextualInversion, Hypernetwork, AestheticGradient, LORA, Controlnet, Poses
|
||||
sort:str = '', # Highest Rated, Most Downloaded, Newest
|
||||
period:str = '', # AllTime, Year, Month, Week, Day
|
||||
nsfw:bool = None, # optional:bool
|
||||
limit:int = 0,
|
||||
base:list[str] = [], # list
|
||||
token:str = None,
|
||||
exact:bool = True,
|
||||
):
|
||||
from modules.civitai import search_civitai
|
||||
if model_id is not None:
|
||||
dct = models_to_json(search_civitai.models, model_id=model_id)
|
||||
return JSONResponse(content=dct, status_code=200)
|
||||
if len(query) > 0 or len(tag) > 0:
|
||||
models = search_civitai.search_civitai(
|
||||
query=query,
|
||||
tag=tag,
|
||||
types=types,
|
||||
sort=sort,
|
||||
period=period,
|
||||
nsfw=nsfw,
|
||||
limit=limit,
|
||||
base=base,
|
||||
token=token,
|
||||
exact=exact
|
||||
)
|
||||
dct = models_to_json(models)
|
||||
return JSONResponse(content=dct, status_code=200)
|
||||
return JSONResponse(content=[], status_code=200)
|
||||
|
||||
|
||||
def register_api():
|
||||
from modules.shared import api
|
||||
api.add_api_route("/sdapi/v1/civitai", get_civitai, methods=["GET"], response_model=list)
|
||||
@@ -0,0 +1,212 @@
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
import json
|
||||
import time
|
||||
from installer import install, log
|
||||
|
||||
|
||||
full_dct = False
|
||||
full_html = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelImage():
|
||||
def __init__(self, dct: dict):
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.url: str = dct.get('url', '')
|
||||
self.width: int = dct.get('width', 0)
|
||||
self.height: int = dct.get('height', 0)
|
||||
self.type: str = dct.get('type', 'Unknown')
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'ModelImage(id={self.id} url="{self.url}" width={self.width} height={self.height} type="{self.type}")'
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelFile():
|
||||
def __init__(self, dct: dict):
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.size: int = int(1024 * dct.get('sizeKB', 0))
|
||||
self.name: str = dct.get('name', 'Unknown')
|
||||
self.type: str = dct.get('type', 'Unknown')
|
||||
self.hashes: list[str] = [str(h) for h in dct.get('hashes', {}).values()]
|
||||
self.url: str = dct.get('downloadUrl', '')
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'ModelFile(id={self.id} name="{self.name}" size={self.size} type="{self.type}" url="{self.url}")'
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelVersion():
|
||||
def __init__(self, dct: dict):
|
||||
import bs4
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.name: str = dct.get('name', 'Unknown')
|
||||
self.base: str = dct.get('baseModel', 'Unknown')
|
||||
self.mtime: str = dct.get('publishedAt', '')
|
||||
self.downloads: int = dct.get('stats', {}).get('downloadCount', 0)
|
||||
self.availability: str = dct.get('availability', 'Unknown')
|
||||
self.html: str = dct.get('description', '') or '' if full_html else ''
|
||||
self.desc: str = bs4.BeautifulSoup(dct.get('description', '') or '', features="html.parser").get_text()
|
||||
self.files = [ModelFile(f) for f in dct.get('files', [])]
|
||||
self.images = [ModelImage(i) for i in dct.get('images', [])]
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'ModelVersion(id={self.id} name="{self.name}" base="{self.base}" mtime="{self.mtime}" downloads={self.downloads} availability={self.availability} desc="{self.desc[:30]}...")'
|
||||
|
||||
|
||||
@dataclass
|
||||
class Model():
|
||||
def __init__(self, dct: dict):
|
||||
import bs4
|
||||
if isinstance(dct, str):
|
||||
dct = json.loads(dct)
|
||||
self.id: int = dct.get('id', 0)
|
||||
self.url: str = f'https://civitai.com/models/{self.id}'
|
||||
self.type: str = dct.get('type', 'Unknown')
|
||||
self.name: str = dct.get('name', 'Unknown')
|
||||
self.html: str = dct.get('description', '') or '' if full_html else ''
|
||||
self.desc: str = bs4.BeautifulSoup(dct.get('description', '') or '', features="html.parser").get_text()
|
||||
self.tags: list[str] = dct.get('tags', [])
|
||||
self.nsfw: bool = dct.get('nsfw', False)
|
||||
self.level: str = dct.get('nsfwLevel', 0)
|
||||
self.availability: str = dct.get('availability', 'Unknown')
|
||||
self.downloads: int = dct.get('stats', {}).get('downloadCount', 0)
|
||||
self.creator: str = dct.get('creator', {}).get('username', 'Unknown')
|
||||
self.versions: list[ModelVersion] = [ModelVersion(v) for v in dct.get('modelVersions', [])]
|
||||
self.dct: dict = dct if full_dct else {}
|
||||
|
||||
def __str__(self):
|
||||
return f'Model(id={self.id} type={self.type} name="{self.name}" versions={len(self.versions)} nsfw={self.nsfw}/{self.level} downloads={self.downloads} author="{self.creator}" tags={self.tags} desc="{self.desc[:30]}...")'
|
||||
|
||||
|
||||
models: list[Model] = [] # global cache for civitai search results
|
||||
|
||||
|
||||
def search_civitai(
|
||||
query:str,
|
||||
tag:str = '', # optional:tag name
|
||||
types:str = '', # (Checkpoint, TextualInversion, Hypernetwork, AestheticGradient, LORA, Controlnet, Poses)
|
||||
sort:str = '', # (Highest Rated, Most Downloaded, Newest)
|
||||
period:str = '', # (AllTime, Year, Month, Week, Day)
|
||||
nsfw:bool = None, # optional:bool
|
||||
limit:int = 0,
|
||||
base:list[str] = [], # list
|
||||
token:str = None,
|
||||
exact:bool = True,
|
||||
):
|
||||
global models # pylint: disable=global-statement
|
||||
import requests
|
||||
from urllib.parse import urlencode
|
||||
install('bs4') # Ensure BeautifulSoup is installed
|
||||
|
||||
if len(query) == 0:
|
||||
log.error('CivitAI: empty query')
|
||||
return []
|
||||
|
||||
t0 = time.time()
|
||||
dct = { 'query': query }
|
||||
if len(tag) > 0:
|
||||
dct['tag'] = tag
|
||||
if nsfw is not None:
|
||||
dct['nsfw'] = 'true' if nsfw else 'false'
|
||||
if limit > 0:
|
||||
dct['limit'] = limit
|
||||
if len(types) > 0:
|
||||
dct['types'] = types
|
||||
if len(sort) > 0:
|
||||
dct['sort'] = sort
|
||||
if len(period) > 0:
|
||||
dct['period'] = period
|
||||
if len(base) > 0:
|
||||
dct['baseModels'] = ','.join(base)
|
||||
encoded = urlencode(dct)
|
||||
|
||||
headers = {}
|
||||
if token is None:
|
||||
token = os.environ.get('CIVITAI_TOKEN', None)
|
||||
if token is not None and len(token) > 0:
|
||||
headers['Authorization'] = f'Bearer {token}'
|
||||
|
||||
url = 'https://civitai.com/api/v1/models'
|
||||
uri = f'{url}?{encoded}'
|
||||
log.info(f'CivitAI request: uri="{uri}" dct={dct} token={token is not None}')
|
||||
result = requests.get(uri, headers=headers, timeout=60)
|
||||
|
||||
if result.status_code != 200:
|
||||
log.error(f'CivitAI: code={result.status_code} reason={result.reason} uri={result.url}')
|
||||
return []
|
||||
|
||||
all_models: list[Model] = []
|
||||
exact_models: list[Model] = []
|
||||
items = result.json().get('items', [])
|
||||
for item in items:
|
||||
all_models.append(Model(item))
|
||||
|
||||
if exact:
|
||||
for model in all_models:
|
||||
model_names = [model.name.lower()]
|
||||
version_names = [v.name.lower() for v in model.versions]
|
||||
file_names = [f.name.lower() for v in model.versions for f in v.files]
|
||||
if any([query.lower() in name for name in model_names + version_names + file_names]): # noqa: C419
|
||||
exact_models.append(model)
|
||||
|
||||
t1 = time.time()
|
||||
log.info(f'CivitAI result: code={result.status_code} exact={len(exact_models)} total={len(models)} time={t1-t0:.2f}')
|
||||
models = exact_models if len(exact_models) > 0 else all_models
|
||||
return models
|
||||
|
||||
|
||||
def create_model_cards(all_models: list[Model]) -> str:
|
||||
details = """
|
||||
<div id="model-details">
|
||||
</div>
|
||||
"""
|
||||
cards = """
|
||||
<div id="model-cards" class="extra-network-cards">
|
||||
{cards}
|
||||
</div>
|
||||
"""
|
||||
card = """
|
||||
<div class="card" data-id="{id}" onclick="modelCardClick({id})">
|
||||
<div class="overlay"><div class="name">{name}</div></div>
|
||||
<div class="version">{type}</div>
|
||||
<img class="preview" src="{preview}" alt="{name}" loading="lazy" />
|
||||
</div>
|
||||
"""
|
||||
all_cards = ''
|
||||
for model in all_models:
|
||||
previews = []
|
||||
for version in model.versions:
|
||||
for image in version.images:
|
||||
if image.url and len(image.url) > 0:
|
||||
previews.append(image.url)
|
||||
if len(previews) == 0:
|
||||
previews = ['./sd_extra_networks/thumb?filename=html/card-no-preview.png']
|
||||
all_cards += card.format(id=model.id, name=model.name, type=model.type, preview=previews[0])
|
||||
html = details + cards.format(cards=all_cards)
|
||||
return html
|
||||
|
||||
|
||||
def print_models(all_models: list[Model]):
|
||||
for model in all_models:
|
||||
log.info(f' {model}')
|
||||
log.trace('Model', model.dct)
|
||||
for version in model.versions:
|
||||
log.info(f' {version}')
|
||||
log.trace('ModelVersion', version.dct)
|
||||
for file in version.files:
|
||||
log.info(f' {file}')
|
||||
log.trace('ModelFile', file.dct)
|
||||
for image in version.images:
|
||||
log.info(f' {image}')
|
||||
log.trace('ModelImage', image.dct)
|
||||
+7
-4
@@ -87,10 +87,13 @@ timer.startup.record("transformers")
|
||||
import accelerate # pylint: disable=W0611,C0411
|
||||
timer.startup.record("accelerate")
|
||||
|
||||
import onnxruntime # pylint: disable=W0611,C0411
|
||||
onnxruntime.set_default_logger_severity(4)
|
||||
onnxruntime.set_default_logger_verbosity(1)
|
||||
onnxruntime.disable_telemetry_events()
|
||||
try:
|
||||
import onnxruntime # pylint: disable=W0611,C0411
|
||||
onnxruntime.set_default_logger_severity(4)
|
||||
onnxruntime.set_default_logger_verbosity(1)
|
||||
onnxruntime.disable_telemetry_events()
|
||||
except Exception as e:
|
||||
errors.log.warning(f'Torch onnxruntime: {e}')
|
||||
timer.startup.record("onnx")
|
||||
|
||||
from fastapi import FastAPI # pylint: disable=W0611,C0411
|
||||
|
||||
@@ -4,6 +4,7 @@ import torch
|
||||
import diffusers
|
||||
import onnxruntime as ort
|
||||
|
||||
|
||||
initialized = False
|
||||
|
||||
|
||||
|
||||
+40
-13
@@ -16,14 +16,12 @@ def create_ui():
|
||||
dummy_component = gr.Label(visible=False)
|
||||
with gr.Row(elem_id="models_tab"):
|
||||
with gr.Column(elem_id='models_output_container', scale=1):
|
||||
gr.HTML(elem_id="models_progress", value="")
|
||||
models_image = gr.Image(elem_id="models_image", show_label=False, interactive=False, type='pil')
|
||||
models_outcome = gr.HTML(elem_id="models_error", value="")
|
||||
models_outcome = gr.HTML(elem_id="models_outcome", value="")
|
||||
models_file = gr.File(label='', visible=False)
|
||||
|
||||
with gr.Column(elem_id='models_input_container', scale=3):
|
||||
|
||||
with gr.Tab(label="Current"):
|
||||
with gr.Tab(label="Current", elem_id="models_current_tab"):
|
||||
def create_modules_table(rows: list):
|
||||
html = """
|
||||
<table class="simple-table">
|
||||
@@ -78,7 +76,7 @@ def create_ui():
|
||||
|
||||
model_analyze.click(fn=analyze, inputs=[], outputs=[model_desc, model_meta])
|
||||
|
||||
with gr.Tab(label="List"):
|
||||
with gr.Tab(label="List", elem_id="models_list_tab"):
|
||||
def create_models_table(rows: list):
|
||||
from modules import sd_detect
|
||||
html = """
|
||||
@@ -137,8 +135,8 @@ def create_ui():
|
||||
model_checkhash_btn.click(fn=sd_models.update_model_hashes, inputs=[], outputs=[model_table])
|
||||
model_list_btn.click(fn=lambda: create_models_table(sd_models.checkpoints_list.values()), inputs=[], outputs=[model_table])
|
||||
|
||||
with gr.Tab(label="Metadata"):
|
||||
from modules.models_civitai import civit_search_metadata, civit_update_metadata
|
||||
with gr.Tab(label="Metadata", elem_id="models_metadata_tab"):
|
||||
from modules.civitai.metadata_civitai import civit_search_metadata, civit_update_metadata
|
||||
with gr.Row():
|
||||
gr.HTML('<h2>Fetch model preview metadata</h2><br>')
|
||||
with gr.Row():
|
||||
@@ -150,11 +148,11 @@ def create_ui():
|
||||
civit_update_btn.click(fn=civit_update_metadata, inputs=[], outputs=[civit_metadata])
|
||||
|
||||
|
||||
with gr.Tab(label="Loader"):
|
||||
with gr.Tab(label="Loader", elem_id="models_loader_tab"):
|
||||
from modules import ui_models_load
|
||||
ui_models_load.create_ui(models_outcome, models_file)
|
||||
|
||||
with gr.Tab(label="Merge"):
|
||||
with gr.Tab(label="Merge", elem_id="models_merge_tab"):
|
||||
from modules.merging import merge_methods
|
||||
from modules.merging.merge_utils import BETA_METHODS, TRIPLE_METHODS, interpolate
|
||||
from modules.merging.merge_presets import BLOCK_WEIGHTS_PRESETS, SDXL_BLOCK_WEIGHTS_PRESETS
|
||||
@@ -398,7 +396,7 @@ def create_ui():
|
||||
]
|
||||
)
|
||||
|
||||
with gr.Tab(label="Replace"):
|
||||
with gr.Tab(label="Replace", elem_id="models_replace_tab"):
|
||||
with gr.Row():
|
||||
gr.HTML('<h2> Replace model components<br></h2>')
|
||||
with gr.Row():
|
||||
@@ -470,8 +468,36 @@ def create_ui():
|
||||
outputs=[models_outcome]
|
||||
)
|
||||
|
||||
with gr.Tab(label="CivitAI"):
|
||||
from modules.models_civitai import civitai_update_token, civit_search_model, civit_search_metadata, civit_select1, civit_select2, civit_select3, civit_download_model
|
||||
with gr.Tab(label="CivitAI", elem_id="models_civitai_tab"):
|
||||
def civitai_search(civit_search_text, civit_search_tag, civit_nsfw, civit_type, civit_base, civit_token):
|
||||
from modules.civitai.search_civitai import search_civitai, create_model_cards
|
||||
results = search_civitai(query=civit_search_text, tag=civit_search_tag, nsfw=civit_nsfw, types=civit_type, base=civit_base, token=civit_token)
|
||||
html = create_model_cards(results)
|
||||
return html
|
||||
|
||||
with gr.Row():
|
||||
gr.HTML('<h2>Search & Download</h2>')
|
||||
with gr.Row(elem_id='civitai_search_row'):
|
||||
civit_search_text = gr.Textbox(label='', placeholder='keyword', elem_id="civit_search_text")
|
||||
civit_search_tag = gr.Textbox(label='', placeholder='tag', elem_id="civit_search_text")
|
||||
civit_search_text_btn = ToolButton(value=ui_symbols.search, interactive=True)
|
||||
with gr.Accordion(label='Search options', open=False, elem_id="civitai_search_options"):
|
||||
with gr.Row():
|
||||
civit_nsfw = gr.Checkbox(label='NSFW allowed', value=True)
|
||||
with gr.Row():
|
||||
civit_type = gr.Textbox(label='Model type', placeholder='Checkpoint, LORA, ...')
|
||||
with gr.Row():
|
||||
civit_base = gr.Textbox(label='Base model', placeholder='SDXL, ...')
|
||||
with gr.Row():
|
||||
civit_token = gr.Textbox(opts.civitai_token, label='CivitAI token', placeholder='optional access token for private or gated models')
|
||||
# sort, period, limit
|
||||
civit_inputs = [civit_search_text, civit_search_tag, civit_nsfw, civit_type, civit_base, civit_token]
|
||||
civit_search_text_btn.click(fn=civitai_search, inputs=civit_inputs, outputs=[models_outcome])
|
||||
civit_search_text.submit(fn=civitai_search, inputs=civit_inputs, outputs=[models_outcome])
|
||||
civit_search_tag.submit(fn=civitai_search, inputs=civit_inputs, outputs=[models_outcome])
|
||||
|
||||
"""
|
||||
from modules.civitai.legacy_civitai import civitai_update_token, civit_search_model, civit_search_metadata, civit_select1, civit_select2, civit_select3, civit_download_model
|
||||
|
||||
with gr.Row():
|
||||
gr.HTML('<h2>Search for models</h2>')
|
||||
@@ -527,8 +553,9 @@ def create_ui():
|
||||
civit_results2.change(fn=is_visible, inputs=[civit_results2], outputs=[civit_results2])
|
||||
civit_results3.change(fn=is_visible, inputs=[civit_results3], outputs=[civit_results3])
|
||||
civit_download_model_btn.click(fn=civit_download_model, inputs=[civit_selected, civit_name, civit_path, civit_model_type, civit_token], outputs=[models_outcome])
|
||||
"""
|
||||
|
||||
with gr.Tab(label="Huggingface"):
|
||||
with gr.Tab(label="Huggingface", elem_id="models_huggingface_tab"):
|
||||
from modules.models_hf import hf_search, hf_select, hf_download_model, hf_update_token
|
||||
with gr.Column(scale=6):
|
||||
with gr.Row():
|
||||
|
||||
+8
-4
@@ -2,7 +2,6 @@ import sys
|
||||
from typing import Union
|
||||
import torch
|
||||
from torch._prims_common import DeviceLikeType
|
||||
import onnxruntime as ort
|
||||
from modules import shared, devices, zluda_installer
|
||||
from modules.zluda_installer import core, default_agent # pylint: disable=unused-import
|
||||
from modules.onnx_impl.execution_providers import available_execution_providers, ExecutionProvider
|
||||
@@ -42,9 +41,14 @@ def initialize_zluda():
|
||||
torch.backends.cuda.enable_mem_efficient_sdp = do_nothing
|
||||
|
||||
# ONNX Runtime is not supported
|
||||
ort.capi._pybind_state.get_available_providers = lambda: [v for v in available_execution_providers if v != ExecutionProvider.CUDA] # pylint: disable=protected-access
|
||||
ort.get_available_providers = ort.capi._pybind_state.get_available_providers # pylint: disable=protected-access
|
||||
if shared.opts.onnx_execution_provider == ExecutionProvider.CUDA:
|
||||
try:
|
||||
import onnxruntime as ort
|
||||
ort.capi._pybind_state.get_available_providers = lambda: [v for v in available_execution_providers if v != ExecutionProvider.CUDA] # pylint: disable=protected-access
|
||||
ort.get_available_providers = ort.capi._pybind_state.get_available_providers # pylint: disable=protected-access
|
||||
if shared.opts.onnx_execution_provider == ExecutionProvider.CUDA:
|
||||
shared.opts.onnx_execution_provider = ExecutionProvider.CPU
|
||||
except Exception as e:
|
||||
shared.log.warning(f'ZLUDA ONNX runtime: {e}')
|
||||
shared.opts.onnx_execution_provider = ExecutionProvider.CPU
|
||||
|
||||
device = devices.get_optimal_device()
|
||||
|
||||
Reference in New Issue
Block a user