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"""
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( '