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"""
{status}run status
{summary['runtime_steps']}runtime steps
{summary['wall_time_s']:.2f}swall time
{summary['information_need_events']}information needs
""" @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( """
Diffusion-native tool-augmented reasoning

Continuous Interaction Diffusion

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.

""", 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( '
Readyrun status
—runtime steps
—wall time
—information needs
', 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( """
CID-v1-0.4B · 398.8M parameters · research prototype
Collection  ·  Model  ·  Code  ·  Paper
""", elem_id="cid-footer", ) if __name__ == "__main__": demo.queue(default_concurrency_limit=1).launch()