big-sleep / app.py
multimodalart's picture
multimodalart HF Staff
The Big Sleep (BigGAN x CLIP) — advadnoun original notebook on ZeroGPU
dc04583 verified
Raw
History Blame Contribute Delete
8.84 kB
"""The Big Sleep — advadnoun's BigGAN x CLIP notebook (Jan 2021), on ZeroGPU.
Engine math is transcribed verbatim from the original Colab (see big_sleep.py):
the layer-wise BigGAN generator, `Pars`, 128 random crops, the latent/class
regularizers and Adam(lr 0.07). This file only adds the UI, live previews,
a timelapse and the notebook's diagnostics panel.
"""
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import time
import spaces # must precede torch
import gradio as gr
import imageio.v2 as imageio
import numpy as np
import torch
import big_sleep as bs
import clip
HERE = os.path.dirname(os.path.abspath(__file__))
GANS = {"512×512 (the original)": "biggan-deep-512",
"256×256": "biggan-deep-256",
"128×128 (fastest)": "biggan-deep-128"}
PERCEPTORS = {"ViT-B/32 (the original)": "ViT-B/32",
"ViT-B/16 (4× slower)": "ViT-B/16",
"ViT-L/14 (17× slower)": "ViT-L/14"}
PREVIEW_EVERY = 10
MAX_GPU_SECONDS = 600
# cutouts held in memory at once. B/32 at the notebook's 128 fits in one pass, so the
# default run takes the notebook's exact step; bigger perceptors need accumulation.
CHUNK = {"ViT-B/32": 128, "ViT-B/16": 32, "ViT-L/14": 8}
# cutouts a perceptor can afford before a run stops being watchable
DEFAULT_CUTN = {"ViT-B/32": 128, "ViT-B/16": 48, "ViT-L/14": 16}
# measured ms/iteration on an A100: generator forward+backward, and per cutout through CLIP
GAN_MS = {"biggan-deep-128": 40.0, "biggan-deep-256": 49.0, "biggan-deep-512": 57.0}
CUT_MS = {"ViT-B/32": 1.4, "ViT-B/16": 5.4, "ViT-L/14": 24.5}
print("Loading BigGAN generators...")
MODELS = {v: bs.load_biggan(v).to("cuda") for v in GANS.values()}
print("Loading CLIP perceptors...")
PERS = {v: clip.load(v, device="cpu", jit=False)[0].eval().requires_grad_(False).to("cuda")
for v in PERCEPTORS.values()}
DIAG_LABELS = {-1: "latent only (class vector zeroed)",
-2: "class only (fresh random latent)",
-3: "row 0 broadcast to every layer — plain BigGAN"}
def _step_ms(gan, perceptor, cutn):
return GAN_MS[GANS[gan]] + CUT_MS[PERCEPTORS[perceptor]] * int(cutn)
def _estimate(iterations, gan, perceptor, cutn):
return 25 + 1.3 * int(iterations) * _step_ms(gan, perceptor, cutn) / 1000
def _duration(prompt, iterations=500, gan=None, perceptor=None, seed=0, lr=bs.LR,
cutn=bs.CUTN, timelapse=True):
gan = gan or list(GANS)[0]
perceptor = perceptor or list(PERCEPTORS)[0]
return int(min(MAX_GPU_SECONDS, _estimate(iterations, gan, perceptor, cutn)))
# arg order starts with (prompt, iterations) so gr.Examples can pass just those two
@spaces.GPU(duration=_duration)
def dream(prompt, iterations=500, gan=None, perceptor=None, seed=0, lr=bs.LR,
cutn=bs.CUTN, timelapse=True):
gan = gan or list(GANS)[0]
perceptor = perceptor or list(PERCEPTORS)[0]
prompt = (prompt or "").strip() or "a cityscape in the style of Van Gogh"
iterations, cutn, seed = int(iterations), int(cutn), int(seed)
if seed < 0:
seed = int(torch.seed() % 2**31)
model, per = MODELS[GANS[gan]], PERS[PERCEPTORS[perceptor]]
chunk = CHUNK[PERCEPTORS[perceptor]]
note = f"{model.side}px · {PERCEPTORS[perceptor]} · {cutn} cutouts · lr {float(lr):g} · seed {seed}"
if chunk < cutn:
note += f" · cutouts accumulated {chunk} at a time"
frames, diagnostics, image = [], [], None
t0 = rate_t0 = time.time()
rate_i0 = 0
for i, img, losses in bs.dream(model, per, prompt, iterations=iterations, seed=seed,
lr=float(lr), cutn=cutn, chunk=chunk,
preview_every=PREVIEW_EVERY):
if i < 0:
diagnostics.append((img, DIAG_LABELS[i]))
continue
image = img
if timelapse:
frames.append(np.asarray(img))
if i == PREVIEW_EVERY: # start the rate clock after warmup
rate_t0, rate_i0 = time.time(), i
now = time.time()
rate = (f"{(i - rate_i0) / (now - rate_t0):.1f} it/s"
if i > rate_i0 else "warming up")
yield (image, gr.skip(), gr.skip(),
f"iteration {i}/{iterations} · {rate} · CLIP similarity {-losses[2] / 100:.3f}\n{note}")
dt = time.time() - rate_t0
n = max(iterations - rate_i0, 1)
video = None
if timelapse and len(frames) > 1:
video = os.path.join(HERE, "timelapse.mp4")
imageio.mimsave(video, frames, fps=12, quality=8, macro_block_size=None)
yield (image, video, diagnostics,
f"✨ done — {iterations} iterations in {time.time() - t0:.0f}s ({n / dt:.1f} it/s)\n{note}")
CSS = """
.gradio-container { max-width: 1180px !important; margin: 0 auto !important; }
#bs-title { text-align: center; }
"""
with gr.Blocks(title="The Big Sleep") as demo:
gr.Markdown(
"# 🛌 The Big Sleep\n"
"CLIP steers BigGAN toward your words — [advadnoun's original "
"*BigGANxCLIP* notebook](https://colab.research.google.com/drive/1NCceX2mbiKOSlAd_o7IU7nA9UskKN5WR) "
"(Ryan Murdock, January 2021), the notebook that started text-to-image as we know it, "
"running live on ZeroGPU. Nothing is being *sampled*: a batch of latent and class "
"vectors — **one pair per generator layer** — is optimized until CLIP agrees the picture "
"matches the prompt.",
elem_id="bs-title",
)
with gr.Row():
with gr.Column(scale=5):
prompt = gr.Textbox(label="Prompt", lines=2,
value="a cityscape in the style of Van Gogh")
with gr.Row():
gan = gr.Dropdown(list(GANS), value=list(GANS)[0], label="BigGAN")
perceptor = gr.Dropdown(list(PERCEPTORS), value=list(PERCEPTORS)[0],
label="Perceptor")
iterations = gr.Slider(50, 1500, 500, step=10, label="Iterations")
with gr.Accordion("Advanced — the notebook's magical numbers", open=False):
lr = gr.Slider(0.01, 0.2, bs.LR, step=0.005, label="Learning rate")
cutn = gr.Slider(16, 128, bs.CUTN, step=16, label="Cutouts per iteration")
seed = gr.Number(0, label="Seed (-1 = random)", precision=0)
timelapse = gr.Checkbox(True, label="Record timelapse")
estimate = gr.Markdown()
with gr.Row():
run_btn = gr.Button("🛌 Dream", variant="primary", size="lg")
stop_btn = gr.Button("Stop", variant="stop", size="lg")
with gr.Column(scale=6):
out_image = gr.Image(label="Dream", type="pil", format="jpeg", height=480)
status = gr.Textbox(label="Status", lines=2, interactive=False)
with gr.Accordion("Timelapse", open=False):
out_video = gr.Video(label="Timelapse", autoplay=True, loop=True)
with gr.Accordion("Diagnostics — what each half of the latent is doing", open=False):
out_diag = gr.Gallery(label="Diagnostics", columns=3, height=240,
object_fit="contain")
def _show_estimate(iterations, gan, perceptor, cutn):
secs = _estimate(iterations, gan, perceptor, cutn)
if secs > MAX_GPU_SECONDS:
return (f"⚠️ ~{secs / 60:.0f} min of GPU time — over the {MAX_GPU_SECONDS // 60}-minute "
"limit, so the run would be cut short. Lower the iterations or the cutouts.")
return f"≈ {secs:.0f} s of GPU time"
def _perceptor_changed(perceptor, iterations, gan):
n = DEFAULT_CUTN[PERCEPTORS[perceptor]]
return n, _show_estimate(iterations, gan, perceptor, n)
est_inputs = [iterations, gan, perceptor, cutn]
for c in (iterations, gan, cutn):
c.change(_show_estimate, est_inputs, estimate)
perceptor.change(_perceptor_changed, [perceptor, iterations, gan], [cutn, estimate])
demo.load(_show_estimate, est_inputs, estimate)
ev = run_btn.click(dream,
[prompt, iterations, gan, perceptor, seed, lr, cutn, timelapse],
[out_image, out_video, out_diag, status])
stop_btn.click(fn=None, cancels=[ev])
gr.Examples(
fn=dream,
examples=[
["a cityscape in the style of Van Gogh", 500],
["the shadow of a lighthouse falling on a wave", 700],
["an ancient library on fire, oil painting", 700],
["a demon made of stained glass", 700],
],
inputs=[prompt, iterations],
outputs=[out_image, out_video, out_diag, status],
cache_examples=True,
cache_mode="lazy",
)
if __name__ == "__main__":
demo.queue().launch(theme=gr.themes.Citrus(), css=CSS)