CID-Demo / app.py
fwerkor's picture
Replace third demo example
eec7dc5
Raw History Blame
22.6 kB
import spaces
import asyncio
import json
import time
from pathlib import Path
from typing import Any
import gradio as gr
import torch
from huggingface_hub import snapshot_download
from safetensors.torch import load_file as load_safetensors
from transformers import AutoTokenizer
from cid.accelerator import wrap_torch_autocast
from cid.data import TrajectoryExample
from cid.defaults import (
DEFAULT_ALLOCATION_THRESHOLD,
DEFAULT_BINDING_THRESHOLD,
DEFAULT_CONVERGENCE_THRESHOLD,
DEFAULT_NEED_THRESHOLD,
)
from cid.model import (
CIDMaterializerConfig,
ILLaDACIDConfig,
load_cid_adapter_from_pretrained,
load_cid_adapter_parameter_state,
)
from cid.model.benchmark import run_neural_benchmark_case
from cid.model.encoding import ILLaDATextEncoder
from cid.runtime.engine import RuntimeConfig
MODEL_ID = "fwerkor/CID-v1-0.4B"
MODEL_REVISION = "b291b799cf654f7074e6bebb48d49fee31c17824"
MODEL_CHOICES = [("CID-v1-0.4B · 398.8M", MODEL_ID)]
DEMO_CODE_REVISION = "dffbec3"
MAX_PROMPT_CHARS = 4_000
MAX_WORKSPACE_CHARS = 24_000
MAX_DOCUMENTS = 12
CUSTOM_CSS = r"""
.gradio-container {
max-width: 1280px !important;
margin: 0 auto !important;
padding: 28px 24px 36px !important;
}
#cid-hero {
padding: 10px 2px 24px;
margin-bottom: 20px;
border-bottom: 1px solid var(--border-color-primary);
}
#cid-hero .hero-kicker {
margin: 0 0 10px;
color: var(--primary-600);
font-size: 12px;
font-weight: 750;
letter-spacing: .11em;
text-transform: uppercase;
}
#cid-hero .hero-title {
margin: 0;
max-width: 980px;
font-size: clamp(32px, 5vw, 54px);
font-weight: 760;
line-height: 1.04;
letter-spacing: -.035em;
}
#cid-hero .hero-copy {
max-width: 900px;
margin: 14px 0 0;
color: var(--body-text-color-subdued);
font-size: 16px;
line-height: 1.65;
}
#cid-hero .hero-links {
display: flex;
flex-wrap: wrap;
gap: 8px;
margin-top: 18px;
}
#cid-hero .hero-links a {
display: inline-flex;
align-items: center;
min-height: 32px;
padding: 5px 10px;
border: 1px solid var(--border-color-primary);
border-radius: 999px;
color: var(--body-text-color);
background: var(--block-background-fill);
font-size: 13px;
font-weight: 600;
text-decoration: none;
}
#cid-hero .hero-links a:hover {
border-color: var(--primary-400);
color: var(--primary-600);
}
#prompt-panel, #result-panel {
padding: 18px;
border: 1px solid var(--border-color-primary);
border-radius: 18px;
background: var(--block-background-fill);
}
#prompt-panel h3, #result-panel h3 {
margin-top: 0;
}
#model-selector {
margin-bottom: 4px;
}
#model-selector label {
font-weight: 650;
}
#cid-prompt textarea {
min-height: 112px !important;
font-size: 15px !important;
line-height: 1.55 !important;
}
#cid-workspace textarea {
overflow-y: auto !important;
scrollbar-gutter: stable;
overscroll-behavior: contain;
}
#final-display textarea {
min-height: 176px !important;
font-size: 21px !important;
font-weight: 650 !important;
line-height: 1.5 !important;
}
#run-button {
min-height: 48px;
margin-top: 4px;
font-weight: 700;
}
#run-stats .stats {
display: grid;
grid-template-columns: repeat(4, minmax(0, 1fr));
gap: 8px;
margin-top: 10px;
}
#run-stats .stat {
padding: 10px 11px;
border: 1px solid var(--border-color-primary);
border-radius: 12px;
background: var(--background-fill-secondary);
}
#run-stats .stat-value {
display: block;
color: var(--body-text-color);
font-size: 16px;
font-weight: 720;
line-height: 1.2;
}
#run-stats .stat-label {
display: block;
margin-top: 3px;
color: var(--body-text-color-subdued);
font-size: 11px;
line-height: 1.25;
}
#examples {
margin-top: 14px;
}
#inspection {
margin-top: 22px;
}
#inspection .tab-nav {
border-bottom-color: var(--border-color-primary);
}
#runtime-trace table {
font-size: 12px !important;
}
#cid-footer {
margin-top: 28px;
padding-top: 16px;
border-top: 1px solid var(--border-color-primary);
color: var(--body-text-color-subdued);
font-size: 13px;
}
#cid-footer a {
color: var(--body-text-color);
font-weight: 600;
text-decoration: none;
}
#cid-footer a:hover {
color: var(--primary-600);
}
@media (max-width: 820px) {
.gradio-container {
padding: 18px 12px 28px !important;
}
#cid-hero {
padding-top: 2px;
}
#prompt-panel, #result-panel {
padding: 14px;
border-radius: 14px;
}
#run-stats .stats {
grid-template-columns: repeat(2, minmax(0, 1fr));
}
}
"""
WORKSPACE_DESCRIPTORS = (
{
"name": "workspace_search",
"description": (
"Search a task-local hidden document workspace and return candidate resource IDs "
"and titles. The workspace may be irrelevant to the task."
),
"arguments": ({"name": "query", "kind": "string", "required": True},),
"cacheable": True,
"dynamic": False,
"versioned": False,
},
{
"name": "workspace_read",
"description": "Read one task-local hidden resource by resource_id.",
"arguments": ({"name": "resource_id", "kind": "string", "required": True},),
"cacheable": True,
"dynamic": False,
"versioned": False,
},
)
def _load_model_bundle():
model_dir = Path(
snapshot_download(
MODEL_ID,
revision=MODEL_REVISION,
allow_patterns=[
"*.json",
"*.jinja",
"*.safetensors",
"*.pt",
"*.py",
],
)
)
release_config = json.loads((model_dir / "cid_config.json").read_text(encoding="utf-8"))
adapter_config = ILLaDACIDConfig(**release_config["adapter_config"])
tokenizer = AutoTokenizer.from_pretrained(model_dir, trust_remote_code=True)
device = torch.device("cuda")
adapter = load_cid_adapter_from_pretrained(
str(model_dir),
config=adapter_config,
freeze_backbone=True,
dtype=torch.float32,
low_cpu_mem_usage=True,
).to(device)
adapter.set_backbone_trainable(True)
load_cid_adapter_parameter_state(
adapter,
load_safetensors(str(model_dir / "cid_adapter.safetensors"), device="cpu"),
)
semantic_state = torch.load(
model_dir / "semantic-embedding.pt",
map_location="cpu",
weights_only=False,
)
text_encoder = ILLaDATextEncoder.from_frozen_snapshot_state(
adapter,
tokenizer,
semantic_state,
device=device,
embedding_device="cpu",
)
forward_model = wrap_torch_autocast(
torch,
adapter,
device_type="cuda",
dtype=torch.bfloat16,
)
adapter.eval()
forward_model.eval()
return adapter, tokenizer, text_encoder, forward_model
ADAPTER, TOKENIZER, TEXT_ENCODER, FORWARD_MODEL = _load_model_bundle()
def _parse_workspace(raw: str) -> tuple[dict[str, Any], ...]:
raw = (raw or "").strip()
if not raw:
return ()
if len(raw) > MAX_WORKSPACE_CHARS:
raise gr.Error(f"Workspace is limited to {MAX_WORKSPACE_CHARS:,} characters.")
blocks = [block.strip() for block in raw.replace("\r\n", "\n").split("\n---\n")]
documents: list[dict[str, Any]] = []
for index, block in enumerate(block for block in blocks if block):
if len(documents) >= MAX_DOCUMENTS:
break
lines = [line.strip() for line in block.splitlines() if line.strip()]
if not lines:
continue
first = lines[0]
if first.lower().startswith("title:"):
title = first.split(":", 1)[1].strip()
body_lines = lines[1:]
elif first.startswith("#"):
title = first.lstrip("#").strip()
body_lines = lines[1:]
else:
title = f"Document {index + 1}"
body_lines = lines
body = " ".join(body_lines).strip()
if not body:
body = title
sentences = [piece.strip() for piece in body.replace("!", ".").replace("?", ".").split(".")]
sentences = [piece for piece in sentences if piece]
documents.append(
{
"resource_id": f"doc-{index:02d}",
"title": title or f"Document {index + 1}",
"sentences": sentences or [body],
}
)
return tuple(documents)
def _build_example(prompt: str, workspace: str) -> tuple[TrajectoryExample, int]:
prompt = (prompt or "").strip()
if not prompt:
raise gr.Error("Enter a prompt first.")
if len(prompt) > MAX_PROMPT_CHARS:
raise gr.Error(f"Prompt is limited to {MAX_PROMPT_CHARS:,} characters.")
documents = _parse_workspace(workspace)
descriptors = WORKSPACE_DESCRIPTORS if documents else ()
metadata: dict[str, Any] = {}
if documents:
metadata = {
"benchmark_workspace_documents": list(documents),
"benchmark_workspace_search_top_k": min(5, len(documents)),
"benchmark_workspace_latency_steps": 2,
"interaction_pattern": "retrieval_qa",
}
example = TrajectoryExample(
example_id="space-demo",
prompt=prompt,
target_display="CID interactive demo",
source_descriptors=descriptors,
metadata=metadata,
)
return example, len(documents)
def _trace_rows(events: tuple[dict[str, Any], ...]) -> list[list[Any]]:
rows: list[list[Any]] = []
for event in events:
payload = dict(event.get("payload", {}))
display_text = payload.pop("display_materialized_text", None)
payload.pop("display_text", None)
payload.pop("display_token_ids", None)
payload.pop("display_visible_token_ids", None)
detail = json.dumps(payload, ensure_ascii=False, default=str)
if display_text:
detail = f"display={display_text!r} " + detail
rows.append(
[
int(event.get("step", 0)),
round(float(event.get("timestamp_s", 0.0)), 4),
str(event.get("kind", "")),
detail[:1200],
]
)
return rows
def _display_evolution(events: tuple[dict[str, Any], ...]) -> str:
changes: list[str] = []
previous: str | None = None
seen_step = False
for event in events:
if event.get("kind") != "model_step_finished":
continue
payload = event.get("payload", {})
current = str(payload.get("display_materialized_text", "")).strip()
if seen_step and current == previous:
continue
visible = current or "(empty Display)"
changes.append(f"**Step {event.get('step', 0)}** — {visible}")
previous = current
seen_step = True
if not changes:
return "No Display states were recorded before the run ended."
return "\n\n".join(changes)
def _summary_html(summary: dict[str, Any]) -> str:
status = "Converged" if summary["converged"] else "Not converged"
return f"""
<div class="stats">
<div class="stat"><span class="stat-value">{status}</span><span class="stat-label">run status</span></div>
<div class="stat"><span class="stat-value">{summary['runtime_steps']}</span><span class="stat-label">runtime steps</span></div>
<div class="stat"><span class="stat-value">{summary['wall_time_s']:.2f}s</span><span class="stat-label">wall time</span></div>
<div class="stat"><span class="stat-value">{summary['information_need_events']}</span><span class="stat-label">information needs</span></div>
</div>
"""
@spaces.GPU(duration=20)
def run_cid(
prompt: str,
workspace: str,
max_steps: int,
need_threshold: float,
convergence_threshold: float,
allocation_threshold: float,
binding_threshold: float,
model_id: str,
):
if model_id != MODEL_ID:
raise gr.Error(f"Model {model_id!r} is not available in this Space yet.")
example, document_count = _build_example(prompt, workspace)
materializer_config = CIDMaterializerConfig(
need_threshold=float(need_threshold),
convergence_threshold=float(convergence_threshold),
allocation_threshold=float(allocation_threshold),
)
runtime_config = RuntimeConfig(
max_steps=256,
max_total_steps=int(max_steps),
max_wall_time_s=55.0,
binding_threshold=float(binding_threshold),
trace_display=True,
)
started = time.perf_counter()
result = asyncio.run(
run_neural_benchmark_case(
ADAPTER,
TOKENIZER,
example,
text_encoder=TEXT_ENCODER,
forward_model=FORWARD_MODEL,
seed_teacher_state=False,
denoising_steps=8,
materializer_config=materializer_config,
runtime_config=runtime_config,
seed=0,
)
)
elapsed = time.perf_counter() - started
events = result.trace_events
need_events = [event for event in events if event.get("kind") == "information_need"]
tool_related = [
event
for event in events
if event.get("kind")
in {
"information_need",
"binding_active",
"job_started",
"job_finished",
"observation_available",
"cache_hit",
"quiescence_started",
"quiescence_resumed",
}
]
summary = {
"model": model_id,
"model_revision": MODEL_REVISION[:12],
"runtime_code_revision": DEMO_CODE_REVISION,
"runtime_steps": result.runtime_steps,
"converged": bool(result.evaluation.converged),
"wall_time_s": round(elapsed, 3),
"workspace_documents": document_count,
"information_need_events": len(need_events),
"tool_related_events": len(tool_related),
"thresholds": {
"need": float(need_threshold),
"convergence": float(convergence_threshold),
"allocation": float(allocation_threshold),
"binding": float(binding_threshold),
},
}
return (
result.final_text or "(empty Display)",
_summary_html(summary),
_display_evolution(events),
_trace_rows(events),
)
CYMBIDIUM_WORKSPACE = """Title: Patrinia
Patrinia is a genus of herbaceous plants in the honeysuckle family.
There are about 17 species native to grassy mountain habitats in China, Siberia and Japan.
These are unassuming clump-forming perennial plants having thin, erect stems with few leaves and bearing a terminal inflorescence with yellow or white flowers.
The use for this plant is to provide a flower through long hot summers.
---
Title: Cymbidium
Cymbidium , or boat orchid, is a genus of 52 evergreen species in the orchid family Orchidaceae.
The new Latin genus name is derived from the Latin "cymba" meaning boat.
Its first known use was in 1815."""
ADORABLE_WORKSPACE = """Title: Adorable (band)
Adorable was an alternative rock band, formed in Coventry in 1990.
The band consisted of band members Pete Fijalkowski (vocals, guitar), Robert Dillam (guitar), Stephen 'Wil' Williams (bass) and Kevin Gritton (drums).
---
Title: Lit (band)
Lit is an American rock band, formed in 1995 in Fullerton, California.
They are best known for their hit song "My Own Worst Enemy"."""
EXAMPLES = [
["How many days are in one week?", ""],
["Which genus has more species, Cymbidium or Patrinia?", CYMBIDIUM_WORKSPACE],
["Which band was formed first, Lit or Adorable?", ADORABLE_WORKSPACE],
]
WORKSPACE_PLACEHOLDER = """Title: Mission note
Project Orion launches on October 17, 2032.
---
Title: Logistics
The launch vehicle is Borealis."""
with gr.Blocks(
title="CID · Interactive Demo",
css=CUSTOM_CSS,
theme=gr.themes.Soft(
primary_hue=gr.themes.colors.indigo,
neutral_hue=gr.themes.colors.slate,
),
) as demo:
gr.HTML(
"""
<div class="hero-kicker">Diffusion-native tool-augmented reasoning</div>
<h1 class="hero-title">Continuous Interaction Diffusion</h1>
<p class="hero-copy">
Run the 398.8M-parameter CID reference checkpoint and inspect how its visible Display
changes across diffusion steps. CID maintains a latent Typed Cognitive Tensor (TCT),
can express information needs during generation, and revises output instead of committing
strictly left-to-right.
</p>
<div class="hero-links">
<a href="https://huggingface.co/collections/fwerkor/cid" target="_blank">CID Collection</a>
<a href="https://huggingface.co/fwerkor/CID-v1-0.4B" target="_blank">CID-v1-0.4B model</a>
<a href="https://huggingface.co/papers/2608.10438" target="_blank">Paper on Hugging Face</a>
<a href="https://github.com/fwerkor/continuous-interaction-diffusion" target="_blank">Source code</a>
</div>
""",
elem_id="cid-hero",
)
with gr.Row(equal_height=False):
with gr.Column(scale=5, min_width=360, elem_id="prompt-panel"):
gr.Markdown("### Try CID")
model_selector = gr.Dropdown(
choices=MODEL_CHOICES,
value=MODEL_ID,
label="Model",
interactive=True,
elem_id="model-selector",
)
prompt = gr.Textbox(
label="Prompt",
value="What is 7 + 5?",
lines=4,
max_lines=8,
placeholder="Ask a short question...",
elem_id="cid-prompt",
)
with gr.Accordion("Optional local workspace · experimental", open=False):
workspace = gr.Textbox(
label="Task-local documents",
value="",
lines=9,
max_lines=18,
placeholder=WORKSPACE_PLACEHOLDER,
info="Separate documents with a line containing ---. Start a block with 'Title:' or '#'. Retrieval is experimental on this compact checkpoint.",
elem_id="cid-workspace",
)
gr.Examples(
examples=EXAMPLES,
inputs=[prompt, workspace],
label="Examples",
elem_id="examples",
)
with gr.Accordion("Runtime controls", open=False):
max_steps = gr.Slider(1, 256, value=128, step=1, label="Max CID model steps")
with gr.Row():
need_threshold = gr.Slider(
0.05,
0.95,
value=DEFAULT_NEED_THRESHOLD,
step=0.05,
label="Need threshold",
)
convergence_threshold = gr.Slider(
0.05,
0.95,
value=DEFAULT_CONVERGENCE_THRESHOLD,
step=0.05,
label="Convergence threshold",
)
with gr.Row():
allocation_threshold = gr.Slider(
0.05,
0.95,
value=DEFAULT_ALLOCATION_THRESHOLD,
step=0.05,
label="Allocation threshold",
)
binding_threshold = gr.Slider(
0.05,
0.95,
value=DEFAULT_BINDING_THRESHOLD,
step=0.05,
label="Binding threshold",
)
run = gr.Button("Run CID", variant="primary", elem_id="run-button")
with gr.Column(scale=5, min_width=360, elem_id="result-panel"):
gr.Markdown("### Result")
final_display = gr.Textbox(
label="Final CID Display",
lines=7,
interactive=False,
elem_id="final-display",
)
gr.Markdown("#### Run at a glance")
summary = gr.HTML(
'<div class="stats"><div class="stat"><span class="stat-value">Ready</span><span class="stat-label">run status</span></div><div class="stat"><span class="stat-value">—</span><span class="stat-label">runtime steps</span></div><div class="stat"><span class="stat-value">—</span><span class="stat-label">wall time</span></div><div class="stat"><span class="stat-value">—</span><span class="stat-label">information needs</span></div></div>',
elem_id="run-stats",
)
with gr.Tabs(elem_id="inspection"):
with gr.Tab("Display evolution"):
gr.Markdown(
"The visible Display is materialized repeatedly; changes below expose revision across model steps."
)
display_evolution = gr.Markdown()
with gr.Tab("Runtime trace"):
gr.Markdown(
"Low-level runtime events, including TCT lifecycle transitions and any tool-related activity."
)
trace = gr.Dataframe(
headers=["step", "time_s", "event", "details"],
datatype=["number", "number", "str", "str"],
interactive=False,
wrap=True,
elem_id="runtime-trace",
)
run.click(
fn=run_cid,
inputs=[
prompt,
workspace,
max_steps,
need_threshold,
convergence_threshold,
allocation_threshold,
binding_threshold,
model_selector,
],
outputs=[final_display, summary, display_evolution, trace],
concurrency_limit=1,
api_name="run_cid",
)
gr.HTML(
"""
<div>
<strong>CID-v1-0.4B</strong> · 398.8M parameters · research prototype<br>
<a href="https://huggingface.co/collections/fwerkor/cid" target="_blank">Collection</a>
&nbsp;·&nbsp;
<a href="https://huggingface.co/fwerkor/CID-v1-0.4B" target="_blank">Model</a>
&nbsp;·&nbsp;
<a href="https://github.com/fwerkor/continuous-interaction-diffusion" target="_blank">Code</a>
&nbsp;·&nbsp;
<a href="https://huggingface.co/papers/2608.10438" target="_blank">Paper</a>
</div>
""",
elem_id="cid-footer",
)
if __name__ == "__main__":
demo.queue(default_concurrency_limit=1).launch()