Atri
Add Anima-2.9B (40-block expanded) checkpoint
760c9a7
Raw History Blame Contribute Delete
12.4 kB
"""Anima 2B anime T2I on Gradio over the ComfyUI backend (ZeroGPU Space).
Workflow: Comfy-Org/workflow_templates image_anima_base_v1.json
Models: oai/civitai-collections (checkpoints/lora) + circlestone-labs/Anima (TE/VAE)
"""
import os
import random
import shutil
import subprocess
import sys
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # noqa: E402
COMFYUI_PATH = os.environ.get("COMFYUI_PATH", os.path.join(os.getcwd(), "ComfyUI"))
MODELS_ROOT = os.path.join(COMFYUI_PATH, "models")
CACHE_DIR = os.path.join(os.getcwd(), "hf_cache")
def ensure_comfyui() -> None:
if os.path.isfile(os.path.join(COMFYUI_PATH, "nodes.py")):
return
subprocess.run(
["git", "clone", "--depth", "1",
"https://github.com/comfyanonymous/ComfyUI.git", COMFYUI_PATH],
check=True,
)
ensure_comfyui()
if COMFYUI_PATH not in sys.path:
sys.path.insert(0, COMFYUI_PATH)
import comfy.options # noqa: E402
comfy.options.enable_args_parsing()
import numpy as np # noqa: E402
import torch # noqa: E402
torch.set_grad_enabled(False)
from huggingface_hub import hf_hub_download # noqa: E402
from comfy import model_management # noqa: E402
from comfy import sample as comfy_sample # noqa: E402
from nodes import ( # noqa: E402
CLIPLoader,
CLIPTextEncode,
EmptyLatentImage,
LoraLoaderModelOnly,
UNETLoader,
VAEDecode,
VAELoader,
)
# --------------------------------------------------------------------------
# Models
# --------------------------------------------------------------------------
COLLECTION_REPO = "oai/civitai-collections"
COMPANION_REPO = "circlestone-labs/Anima"
MODEL_CHOICES = [
("anima-base-v1.0.safetensors", COMPANION_REPO, "split_files/diffusion_models", 4.18),
("Anima-2.9B-preview-v1.safetensors", "Gazingstars123/Anima-2.9B", "", 5.84),
("hs-anima-2.0.safetensors", COLLECTION_REPO, "checkpoints/anima", 4.18),
("MiaoMiao RealSkin _Anima_v1.1zs_net.safetensors", COLLECTION_REPO, "checkpoints/anima", 4.18),
("MiaoMiao_Anima_lh3d_1.0_n.safetensors", COLLECTION_REPO, "checkpoints/anima", 4.18),
("One obsession_ anima3D_v1.0.safetensors", COLLECTION_REPO, "checkpoints/anima", 4.18),
]
MODEL_FILES = [fname for fname, _, _, _ in MODEL_CHOICES]
LORA_FILES = [
("age_slider.safetensors", "lora/anima"),
("real_skin.safetensors", "lora/anima"),
("wlop-2_v1_epoch15.safetensors", "lora/anima"),
]
TEXT_ENCODER = ("qwen_3_06b_base.safetensors", "split_files/text_encoders")
VAE_FILE = ("qwen_image_vae.safetensors", "split_files/vae")
def ensure_model_file(repo_id: str, subfolder: str, filename: str, dest_dir: str) -> str:
"""Download a model file into a ComfyUI models/ subdir (idempotent)."""
dest = os.path.join(dest_dir, filename)
if os.path.isfile(dest) and os.path.getsize(dest) > 1e6:
return dest
os.makedirs(CACHE_DIR, exist_ok=True)
os.makedirs(dest_dir, exist_ok=True)
src = hf_hub_download(
repo_id=repo_id,
subfolder=subfolder or None,
filename=filename,
local_dir=CACHE_DIR,
token=os.environ.get("HF_TOKEN") or None,
)
if not os.path.isfile(dest):
try:
os.replace(src, dest)
except OSError:
shutil.copy2(src, dest)
return dest
def ensure_models() -> None:
te_name, te_sub = TEXT_ENCODER
vae_name, vae_sub = VAE_FILE
ensure_model_file(COMPANION_REPO, te_sub, te_name, os.path.join(MODELS_ROOT, "text_encoders"))
ensure_model_file(COMPANION_REPO, vae_sub, vae_name, os.path.join(MODELS_ROOT, "vae"))
for fname, repo, sub, _ in MODEL_CHOICES:
ensure_model_file(repo, sub, fname, os.path.join(MODELS_ROOT, "diffusion_models"))
for fname, sub in LORA_FILES:
ensure_model_file(COLLECTION_REPO, sub, fname, os.path.join(MODELS_ROOT, "loras"))
ensure_models()
def get_value_at_index(obj, index: int):
try:
return obj[index]
except KeyError:
return obj["result"][index]
def list_loras() -> list[str]:
lora_dir = os.path.join(MODELS_ROOT, "loras")
return sorted(f for f in os.listdir(lora_dir) if f.endswith(".safetensors")) if os.path.isdir(lora_dir) else []
# --------------------------------------------------------------------------
# Load models at module scope so ZeroGPU packs them at startup.
# --------------------------------------------------------------------------
unet_loader = UNETLoader()
UNETS = {
fname: unet_loader.load_unet(unet_name=fname, weight_dtype="default")
for fname, _, _, _ in MODEL_CHOICES
}
clip_loader = CLIPLoader()
CLIP = clip_loader.load_clip(clip_name=TEXT_ENCODER[0], type="stable_diffusion")
vae_loader = VAELoader()
VAE = vae_loader.load_vae(vae_name=VAE_FILE[0])
lora_loader = LoraLoaderModelOnly()
text_encode = CLIPTextEncode()
empty_latent = EmptyLatentImage()
vae_decode = VAEDecode()
model_management.load_models_gpu(
[
getattr(get_value_at_index(unet, 0), "patcher", get_value_at_index(unet, 0))
for unet in UNETS.values()
]
+ [
getattr(get_value_at_index(CLIP, 0), "patcher", get_value_at_index(CLIP, 0)),
getattr(get_value_at_index(VAE, 0), "patcher", get_value_at_index(VAE, 0)),
]
)
# --------------------------------------------------------------------------
# Inference
# --------------------------------------------------------------------------
import gradio as gr # noqa: E402 # needed for gr.Progress default arg below
STATIC_INFO = (
f"**Loaded models:** {len(MODEL_CHOICES)}x Anima checkpoints (2B/2.9B) + Qwen-3 0.6B text encoder + Qwen-Image VAE "
f"(preloaded, ~{sum(size for _, _, _, size in MODEL_CHOICES) + 1.19 + 0.25:.0f} GB)\n"
"**Prompt weighting:** supported, e.g. `(tag)` (1.1x), `(tag:1.2)`, `(tag:0.7)` (per the Anima README, use higher weights than SDXL)"
)
# GPU stats are only real inside the ZeroGPU worker, so they are fetched at
# the end of each generation rather than at startup.
def gpu_stats() -> str:
try:
name = torch.cuda.get_device_name(0)
free, total = torch.cuda.mem_get_info(0)
used = total - free
except Exception as exc:
return f"**GPU:** unavailable ({exc})"
return (
f"**GPU:** {name}\n"
f"**VRAM:** {used / 1e9:.1f} GB used / {total / 1e9:.1f} GB total ({free / 1e9:.1f} GB free)"
)
@spaces.GPU(duration=int(os.environ.get("GPU_DURATION", "30")))
def generate_image(
model_name: str,
prompt: str,
negative_prompt: str,
width: int,
height: int,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
enable_lora: bool,
lora_name: str,
lora_strength: float,
progress: gr.Progress = gr.Progress(track_tqdm=False),
) -> tuple[np.ndarray, str]:
width = max(512, min(1536, int(width) // 16 * 16))
height = max(512, min(1536, int(height) // 16 * 16))
seed = int(seed) if int(seed) >= 0 else random.randint(1, 2**63)
steps = max(10, min(60, int(steps)))
cfg = float(max(1.0, min(8.0, cfg)))
lora_strength = float(lora_strength)
model = get_value_at_index(UNETS[model_name], 0)
if enable_lora and lora_name:
model = get_value_at_index(
lora_loader.load_lora_model_only(model=model, lora_name=lora_name, strength_model=lora_strength),
0,
)
positive = text_encode.encode(text=prompt, clip=get_value_at_index(CLIP, 0))
negative = text_encode.encode(text=negative_prompt, clip=get_value_at_index(CLIP, 0))
latent = empty_latent.generate(width=width, height=height, batch_size=1)
lat = get_value_at_index(latent, 0)
latent_image = comfy_sample.fix_empty_latent_channels(
model, lat["samples"],
lat.get("downscale_ratio_spacial", None),
lat.get("downscale_ratio_temporal", None),
)
progress(0, desc="Generating")
def step_cb(step: int, _denoised, _x, total_steps: int) -> None:
progress((step, total_steps), desc=f"Sampling {step + 1}/{total_steps}")
samples = comfy_sample.sample(
model=model,
noise=comfy_sample.prepare_noise(latent_image, seed, None),
steps=steps,
cfg=cfg,
sampler_name=sampler_name,
scheduler=scheduler,
positive=get_value_at_index(positive, 0),
negative=get_value_at_index(negative, 0),
latent_image=latent_image,
denoise=1.0,
callback=step_cb,
disable_pbar=True,
seed=seed,
)
sampled = {"samples": samples}
progress(1, desc="Decoding")
decoded = vae_decode.decode(samples=sampled, vae=get_value_at_index(VAE, 0))
image = get_value_at_index(decoded, 0)[0]
img_np = (
torch.nan_to_num(image, nan=0.0, posinf=1.0, neginf=0.0)
.mul(255)
.clamp_(0, 255)
.byte()
.cpu()
.numpy()
)
info = (
STATIC_INFO
+ f"\n{gpu_stats()}"
+ f"\n**Last run:** {model_name} | {width}x{height} | {steps} steps | CFG {cfg} | seed {seed}"
)
return img_np, info
# --------------------------------------------------------------------------
# UI
# --------------------------------------------------------------------------
RESOLUTIONS = {
"1:1 (1024x1024)": (1024, 1024),
"2:3 (832x1216)": (832, 1216),
"3:2 (1216x832)": (1216, 832),
"3:4 (896x1152)": (896, 1152),
"4:3 (1152x896)": (1152, 896),
"9:16 (768x1344)": (768, 1344),
"16:9 (1344x768)": (1344, 768),
}
SAMPLERS = ["er_sde", "euler", "euler_ancestral", "dpmpp_2m", "dpmpp_2m_sde", "dpmpp_2m_sde_gpu", "dpmpp_sde", "heun", "ddim", "uni_pc"]
SCHEDULERS = ["simple", "normal", "karras", "exponential", "sgm_uniform"]
LORA_CHOICES = list_loras() or ["real_skin.safetensors"]
output_image = gr.Image(label="Generated Image", format="png")
info_output = gr.Markdown(STATIC_INFO + "\n_Run a generation to see live GPU/VRAM stats._")
with gr.Blocks(title="Anima Collection") as app:
with gr.Row():
with gr.Column(scale=1):
model_input = gr.Dropdown(label="Checkpoint", choices=MODEL_FILES, value=MODEL_FILES[0])
prompt_input = gr.Textbox(label="Prompt", lines=3)
negative_input = gr.Textbox(
label="Negative prompt", lines=2,
value="worst quality, low quality, score_1, score_2, score_3, artist name, blurry, jpeg artifacts, chromatic aberration",
)
resolution = gr.Dropdown(label="Resolution", choices=list(RESOLUTIONS), value="1:1 (1024x1024)")
with gr.Row():
width_input = gr.Slider(512, 1536, value=1024, step=16, label="Width")
height_input = gr.Slider(512, 1536, value=1024, step=16, label="Height")
with gr.Row():
seed_input = gr.Number(label="Seed (-1 = random)", value=-1, precision=0)
steps_input = gr.Slider(label="Steps", minimum=10, maximum=60, value=30, step=1)
cfg_input = gr.Slider(label="CFG", minimum=1.0, maximum=8.0, value=4.0, step=0.1)
with gr.Accordion("Sampler", open=False):
sampler_input = gr.Dropdown(label="Sampler", choices=SAMPLERS, value="er_sde")
scheduler_input = gr.Dropdown(label="Scheduler", choices=SCHEDULERS, value="simple")
with gr.Accordion("LoRA", open=False):
lora_enable = gr.Checkbox(label="Enable", value=False)
lora_name_input = gr.Dropdown(label="File", choices=LORA_CHOICES, value=LORA_CHOICES[0])
lora_strength_input = gr.Slider(label="Strength", minimum=0.0, maximum=2.0, value=0.8, step=0.05)
generate_btn = gr.Button("Generate", variant="primary")
with gr.Column(scale=1):
output_image.render()
info_output.render()
resolution.change(
lambda name: list(RESOLUTIONS[name]),
inputs=[resolution],
outputs=[width_input, height_input],
)
generate_btn.click(
fn=generate_image,
inputs=[
model_input, prompt_input, negative_input, width_input, height_input, seed_input,
steps_input, cfg_input, sampler_input, scheduler_input, lora_enable, lora_name_input,
lora_strength_input,
],
outputs=[output_image, info_output],
show_progress=True,
)
if __name__ == "__main__":
app.launch()