Spaces:
Sleeping
Sleeping
| """AffectFlow-DINO: Uncertainty-Aware Multi-Task Affect Estimation. | |
| A Gradio demo that runs the AffectFlow-DINO model on a single face image and | |
| returns valence-arousal predictions, facial expression classification, and | |
| Action Unit detections — including uncertainty estimates from Monte Carlo | |
| sampling of the conditional rectified-flow head. | |
| """ | |
| import os | |
| os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") | |
| import spaces # MUST come before torch / any CUDA-touching import | |
| import json | |
| import math | |
| import time | |
| from pathlib import Path | |
| from typing import Any | |
| import torch | |
| import torch.nn.functional as F | |
| from torch import nn | |
| from torchvision import transforms | |
| from PIL import Image | |
| import gradio as gr | |
| from huggingface_hub import hf_hub_download | |
| # --------------------------------------------------------------------------- | |
| # Constants | |
| # --------------------------------------------------------------------------- | |
| HF_REPO_ID = "Bekhouche/AffectFlow-DINO" | |
| BEST_MODEL = "finetune-flow-retune-b10" | |
| IMAGE_SIZE = 224 | |
| IMAGENET_MEAN = (0.485, 0.456, 0.406) | |
| IMAGENET_STD = (0.229, 0.224, 0.225) | |
| EXPRESSION_NAMES = [ | |
| "Neutral", "Anger", "Disgust", "Fear", | |
| "Happiness", "Sadness", "Surprise", "Other", | |
| ] | |
| AU_NAMES = ["AU1", "AU2", "AU4", "AU6", "AU7", "AU10", | |
| "AU12", "AU15", "AU23", "AU24", "AU25", "AU26"] | |
| TARGET_DIM = 2 + len(EXPRESSION_NAMES) + len(AU_NAMES) # 22 | |
| # --------------------------------------------------------------------------- | |
| # Model architecture — self-contained, no import from the affectflow package | |
| # needed. We reconstruct the DINOv3 ViT-S/16 backbone from a scratch config so | |
| # the gated facebook/dinov3-vits16-pretrain-lvd1689m download is never | |
| # triggered; all weights come from the AffectFlow checkpoint. | |
| # --------------------------------------------------------------------------- | |
| class SinusoidalTimeEmbedding(nn.Module): | |
| def __init__(self, dim: int) -> None: | |
| super().__init__() | |
| self.dim = dim | |
| def forward(self, t: torch.Tensor) -> torch.Tensor: | |
| half = self.dim // 2 | |
| freqs = torch.exp( | |
| -math.log(10000.0) | |
| * torch.arange(half, device=t.device, dtype=t.dtype) / max(half - 1, 1) | |
| ) | |
| args = t[:, None] * freqs[None, :] | |
| emb = torch.cat([torch.sin(args), torch.cos(args)], dim=-1) | |
| if self.dim % 2 == 1: | |
| emb = F.pad(emb, (0, 1)) | |
| return emb | |
| class MLP(nn.Module): | |
| def __init__(self, in_dim: int, hidden_dim: int, out_dim: int, dropout: float = 0.1) -> None: | |
| super().__init__() | |
| self.net = nn.Sequential( | |
| nn.Linear(in_dim, hidden_dim), | |
| nn.LayerNorm(hidden_dim), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(hidden_dim, hidden_dim), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(hidden_dim, out_dim), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.net(x) | |
| def _build_dinov3_vits16() -> nn.Module: | |
| """Build a DINOv3 ViT-S/16 encoder from a scratch config (no weight download).""" | |
| from transformers import DINOv3ViTConfig, DINOv3ViTModel | |
| config = DINOv3ViTConfig( | |
| hidden_size=384, | |
| num_hidden_layers=12, | |
| num_attention_heads=6, | |
| intermediate_size=1536, | |
| patch_size=16, | |
| image_size=IMAGE_SIZE, | |
| num_channels=3, | |
| layer_scale_init_value=1e-5, | |
| use_mask_token=True, | |
| num_register_tokens=4, | |
| ) | |
| return DINOv3ViTModel(config) | |
| def normalize_backbone_output(output: Any) -> torch.Tensor: | |
| if hasattr(output, "pooler_output") and output.pooler_output is not None: | |
| output = output.pooler_output | |
| elif hasattr(output, "last_hidden_state"): | |
| output = output.last_hidden_state[:, 0] | |
| if isinstance(output, dict): | |
| if "x_norm_clstoken" in output: | |
| output = output["x_norm_clstoken"] | |
| elif "pooler_output" in output: | |
| output = output["pooler_output"] | |
| elif "last_hidden_state" in output: | |
| output = output["last_hidden_state"][:, 0] | |
| else: | |
| output = next(iter(output.values())) | |
| elif isinstance(output, (tuple, list)): | |
| output = output[0] | |
| if output.ndim == 4: | |
| output = output.mean(dim=(2, 3)) | |
| elif output.ndim == 3: | |
| output = output[:, 0] | |
| return output | |
| class DinoBackbone(nn.Module): | |
| """DINOv3 ViT-S/16 backbone wrapper.""" | |
| def __init__(self) -> None: | |
| super().__init__() | |
| self.encoder = _build_dinov3_vits16() | |
| self.feature_dim = 384 | |
| self.source = "scratch:dinov3_vits16" | |
| def forward(self, image: torch.Tensor) -> torch.Tensor: | |
| return normalize_backbone_output(self.encoder(image)) | |
| def forward_with_patches(self, image: torch.Tensor): | |
| output = self.encoder(image) | |
| if hasattr(output, "last_hidden_state"): | |
| hidden = output.last_hidden_state | |
| return hidden[:, 0], hidden[:, 1:] | |
| return normalize_backbone_output(output), None | |
| class AffectFlowModel(nn.Module): | |
| """AffectFlow-DINO model: DINOv3 backbone + deterministic heads + flow head.""" | |
| def __init__(self, hidden_dim: int = 768, time_dim: int = 128) -> None: | |
| super().__init__() | |
| self.backbone = DinoBackbone() | |
| feature_dim = self.backbone.feature_dim | |
| def _proj() -> nn.Sequential: | |
| return nn.Sequential( | |
| nn.Linear(feature_dim, hidden_dim), | |
| nn.LayerNorm(hidden_dim), | |
| nn.GELU(), | |
| ) | |
| self.proj = _proj() | |
| self.va_head = MLP(hidden_dim, hidden_dim, 2) | |
| self.expr_head = MLP(hidden_dim, hidden_dim, 8) | |
| self.au_head = MLP(hidden_dim, hidden_dim, 12) | |
| self.time_embedding = SinusoidalTimeEmbedding(time_dim) | |
| self.flow = MLP(hidden_dim + TARGET_DIM + time_dim, hidden_dim, TARGET_DIM) | |
| def _backbone_feat(self, image: torch.Tensor) -> torch.Tensor: | |
| return self.backbone(image) | |
| def encode(self, image: torch.Tensor) -> torch.Tensor: | |
| return self.proj(self._backbone_feat(image)) | |
| def deterministic(self, feature: torch.Tensor) -> dict[str, torch.Tensor]: | |
| return { | |
| "va": torch.tanh(self.va_head(feature)), | |
| "expr_logits": self.expr_head(feature), | |
| "au_logits": self.au_head(feature), | |
| } | |
| def velocity(self, feature: torch.Tensor, z_t: torch.Tensor, t: torch.Tensor) -> torch.Tensor: | |
| t_emb = self.time_embedding(t) | |
| return self.flow(torch.cat([feature, z_t, t_emb], dim=-1)) | |
| def sample_joint(self, image: torch.Tensor, steps: int = 30, samples: int = 16) -> torch.Tensor: | |
| """Draw `samples` flow trajectories per image; returns [B, samples, 22].""" | |
| self.eval() | |
| feature = self.encode(image) | |
| batch = image.shape[0] | |
| feature = ( | |
| feature[:, None, :] | |
| .expand(batch, samples, feature.shape[-1]) | |
| .reshape(batch * samples, -1) | |
| ) | |
| z = torch.randn(batch * samples, TARGET_DIM, device=image.device) | |
| dt = 1.0 / steps | |
| for idx in range(steps): | |
| t = torch.full((batch * samples,), idx / steps, device=image.device) | |
| z = z + dt * self.velocity(feature, z, t) | |
| return z.reshape(batch, samples, TARGET_DIM) | |
| # --------------------------------------------------------------------------- | |
| # Weight loading — handles the key remapping between checkpoint and model | |
| # --------------------------------------------------------------------------- | |
| def _remap_state_dict(model: nn.Module, state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: | |
| """Aligns checkpoint keys with the current module tree.""" | |
| model_state = model.state_dict() | |
| remapped: dict[str, torch.Tensor] = {} | |
| prefix = "backbone.encoder.model." | |
| for key, value in state_dict.items(): | |
| candidates = [key] | |
| # Handle transformers version drift: .encoder.model.layer. vs .encoder.model.model.layer. | |
| if ".encoder.model.layer." in key and ".encoder.model.model.layer." not in key: | |
| candidates.append(key.replace(".encoder.model.layer.", ".encoder.model.model.layer.", 1)) | |
| if ".encoder.model.model.layer." in key: | |
| candidates.append(key.replace(".encoder.model.model.layer.", ".encoder.model.layer.", 1)) | |
| for cand in candidates: | |
| if cand in model_state and model_state[cand].shape == value.shape: | |
| remapped[cand] = value | |
| break | |
| return remapped | |
| def load_model() -> AffectFlowModel: | |
| """Build the model and load weights from the HF checkpoint.""" | |
| ckpt_path = hf_hub_download(HF_REPO_ID, f"{BEST_MODEL}/model.pt") | |
| bundle = torch.load(ckpt_path, map_location="cpu", weights_only=False) | |
| model = AffectFlowModel() | |
| remapped = _remap_state_dict(model, bundle["model"]) | |
| missing, unexpected = model.load_state_dict(remapped, strict=False) | |
| critical_missing = [k for k in missing if not k.startswith("backbone.")] | |
| if critical_missing: | |
| raise RuntimeError(f"Missing critical keys: {critical_missing[:5]}") | |
| model.eval() | |
| return model | |
| def load_calibration() -> tuple[dict[str, float] | None, dict[str, float] | None]: | |
| """Download and return (au_thresholds, expr_weights) if available.""" | |
| au_thresholds = None | |
| expr_weights = None | |
| try: | |
| au_path = hf_hub_download(HF_REPO_ID, f"{BEST_MODEL}/au_thresholds.json") | |
| with open(au_path) as f: | |
| au_thresholds = json.loads(f.read())["thresholds"] | |
| except Exception: | |
| pass | |
| try: | |
| expr_path = hf_hub_download(HF_REPO_ID, f"{BEST_MODEL}/expr_weights.json") | |
| with open(expr_path) as f: | |
| expr_weights = json.loads(f.read())["weights"] | |
| except Exception: | |
| pass | |
| return au_thresholds, expr_weights | |
| # --------------------------------------------------------------------------- | |
| # Image preprocessing | |
| # --------------------------------------------------------------------------- | |
| _transform = transforms.Compose([ | |
| transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)), | |
| transforms.ToTensor(), | |
| transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), | |
| ]) | |
| def load_image_tensor(image: Any) -> torch.Tensor: | |
| """Load a PIL Image or file path into a normalized tensor.""" | |
| if isinstance(image, str): | |
| with Image.open(image) as img: | |
| return _transform(img.convert("RGB")) | |
| if isinstance(image, Image.Image): | |
| return _transform(image.convert("RGB")) | |
| raise TypeError(f"Unsupported image type: {type(image)}") | |
| # --------------------------------------------------------------------------- | |
| # Inference | |
| # --------------------------------------------------------------------------- | |
| def predict( | |
| image: Any, | |
| decode_mode: str = "flow", | |
| flow_steps: int = 30, | |
| flow_samples: int = 16, | |
| ) -> tuple[str, str, str, str, float]: | |
| """Run AffectFlow-DINO affect estimation on a face image. | |
| Args: | |
| image: A face image (PIL Image or file path). | |
| decode_mode: 'flow' for Monte Carlo sampling (uncertainty-aware), | |
| 'deterministic' for the direct head prediction. | |
| flow_steps: Number of rectified-flow integration steps. | |
| flow_samples: Number of flow trajectories to sample for uncertainty. | |
| Returns: | |
| A tuple of (va_plot_path, expr_plot_path, au_plot_path, summary_text). | |
| """ | |
| t0 = time.perf_counter() | |
| img_tensor = load_image_tensor(image).unsqueeze(0).to("cuda") | |
| model = _MODEL # noqa: F821 — module-scope global | |
| au_th = _AU_THRESHOLDS # noqa: F821 | |
| expr_w = _EXPR_WEIGHTS # noqa: F821 | |
| with torch.no_grad(): | |
| if decode_mode == "flow": | |
| samples = model.sample_joint(img_tensor, steps=flow_steps, samples=flow_samples) | |
| # [1, samples, 22] | |
| mean_joint = samples.mean(dim=1) # [1, 22] | |
| std_joint = samples.std(dim=1) # [1, 22] | |
| va = mean_joint[:, :2].clamp(-1.0, 1.0) | |
| va_std = std_joint[:, :2] | |
| expr_logits = mean_joint[:, 2:10] | |
| expr_probs = F.softmax(expr_logits, dim=-1) | |
| expr_probs_std = F.softmax(samples[:, :, 2:10], dim=-1).std(dim=1) | |
| au_logits = mean_joint[:, 10:] | |
| au_probs = torch.sigmoid(au_logits) | |
| au_probs_std = torch.sigmoid(samples[:, :, 10:]).std(dim=1) | |
| else: | |
| raw_feat = model._backbone_feat(img_tensor) | |
| feature = model.proj(raw_feat) | |
| out = model.deterministic(feature) | |
| va = out["va"].clamp(-1.0, 1.0) | |
| va_std = torch.zeros_like(va) | |
| expr_logits = out["expr_logits"] | |
| expr_probs = F.softmax(expr_logits, dim=-1) | |
| expr_probs_std = torch.zeros_like(expr_probs) | |
| au_logits = out["au_logits"] | |
| au_probs = torch.sigmoid(au_logits) | |
| au_probs_std = torch.zeros_like(au_probs) | |
| # Apply expression calibration | |
| if expr_w is not None: | |
| w = torch.tensor( | |
| [expr_w[n] for n in EXPRESSION_NAMES], device=expr_probs.device | |
| ) | |
| expr_idx = (expr_probs * w).argmax(dim=-1) | |
| else: | |
| expr_idx = expr_probs.argmax(dim=-1) | |
| # Apply AU calibration | |
| if au_th is not None: | |
| th = torch.tensor( | |
| [au_th[n] for n in AU_NAMES], device=au_probs.device | |
| ) | |
| aus = (au_probs >= th).long() | |
| else: | |
| aus = (au_probs >= 0.5).long() | |
| # Extract scalars | |
| valence = float(va[0, 0].cpu()) | |
| arousal = float(va[0, 1].cpu()) | |
| valence_std = float(va_std[0, 0].cpu()) | |
| arousal_std = float(va_std[0, 1].cpu()) | |
| expr_name = EXPRESSION_NAMES[int(expr_idx[0].cpu())] | |
| expr_probs_list = {EXPRESSION_NAMES[k]: float(expr_probs[0, k].cpu()) for k in range(8)} | |
| expr_probs_std_list = {EXPRESSION_NAMES[k]: float(expr_probs_std[0, k].cpu()) for k in range(8)} | |
| au_probs_list = {AU_NAMES[j]: float(au_probs[0, j].cpu()) for j in range(12)} | |
| au_probs_std_list = {AU_NAMES[j]: float(au_probs_std[0, j].cpu()) for j in range(12)} | |
| active_aus = [AU_NAMES[j] for j in range(12) if aus[0, j].item() == 1] | |
| elapsed = time.perf_counter() - t0 | |
| # Build the summary text | |
| summary = ( | |
| f"**Valence:** {valence:+.3f}" + (f" (±{valence_std:.3f})" if decode_mode == "flow" else "") + "\n\n" | |
| f"**Arousal:** {arousal:+.3f}" + (f" (±{arousal_std:.3f})" if decode_mode == "flow" else "") + "\n\n" | |
| f"**Expression:** {expr_name}\n\n" | |
| f"**Active AUs:** {', '.join(active_aus) if active_aus else 'None'}\n\n" | |
| f"**Decode mode:** {decode_mode} | **Inference time:** {elapsed:.2f}s" | |
| ) | |
| # Build plots | |
| va_plot = _make_va_plot(valence, arousal, valence_std, arousal_std) | |
| expr_plot = _make_expr_plot(expr_probs_list, expr_probs_std_list, decode_mode) | |
| au_plot = _make_au_plot(au_probs_list, au_probs_std_list, decode_mode, au_th) | |
| return va_plot, expr_plot, au_plot, summary | |
| def _make_va_plot(v: float, a: float, v_std: float, a_std: float) -> str: | |
| """Create a valence-arousal 2D plot and return the file path.""" | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import matplotlib.patches as mpatches | |
| import tempfile | |
| fig, ax = plt.subplots(1, 1, figsize=(4, 4)) | |
| # Draw the circumplex model axes | |
| ax.axhline(y=0, color="gray", linewidth=0.5, linestyle="--") | |
| ax.axvline(x=0, color="gray", linewidth=0.5, linestyle="--") | |
| # Quadrant labels | |
| quad_labels = { | |
| (1, 1): "Excited\nHappy", | |
| (-1, 1): "Tense\nAfraid", | |
| (-1, -1): "Sad\nBored", | |
| (1, -1): "Calm\nRelaxed", | |
| } | |
| for (qx, qy), label in quad_labels.items(): | |
| ax.text(qx * 0.7, qy * 0.7, label, ha="center", va="center", | |
| fontsize=8, color="lightgray", style="italic") | |
| # Plot the prediction point | |
| ax.plot(v, a, "ro", markersize=12, zorder=5, label="Mean prediction") | |
| # Uncertainty ellipse (if std > 0) | |
| if v_std > 0 or a_std > 0: | |
| ellipse = mpatches.Ellipse( | |
| (v, a), width=max(v_std * 2, 0.02), height=max(a_std * 2, 0.02), | |
| angle=0, fill=False, color="red", linewidth=1.5, linestyle="--", alpha=0.7, | |
| ) | |
| ax.add_patch(ellipse) | |
| ax.plot([], [], "r--", label="Uncertainty (±1σ)") | |
| ax.set_xlim(-1.1, 1.1) | |
| ax.set_ylim(-1.1, 1.1) | |
| ax.set_xlabel("Valence", fontsize=11) | |
| ax.set_ylabel("Arousal", fontsize=11) | |
| ax.set_title("Valence–Arousal Circumplex", fontsize=12, fontweight="bold") | |
| ax.set_aspect("equal") | |
| ax.legend(loc="upper right", fontsize=8, framealpha=0.8) | |
| ax.grid(True, alpha=0.2) | |
| plt.tight_layout() | |
| tmp = tempfile.NamedTemporaryFile(suffix=".png", delete=False, dir="/tmp") | |
| fig.savefig(tmp.name, dpi=150, bbox_inches="tight") | |
| plt.close(fig) | |
| return tmp.name | |
| def _make_expr_plot( | |
| probs: dict[str, float], stds: dict[str, float], decode_mode: str | |
| ) -> str: | |
| """Create a bar chart of expression probabilities.""" | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import tempfile | |
| names = list(probs.keys()) | |
| values = list(probs.values()) | |
| errors = list(stds.values()) if decode_mode == "flow" else None | |
| fig, ax = plt.subplots(1, 1, figsize=(6, 3.5)) | |
| colors = plt.cm.Set2(range(len(names))) | |
| bars = ax.barh(names, values, color=colors, xerr=errors, capsize=3, error_kw={"linewidth": 1, "alpha": 0.7}) | |
| # Highlight the predicted class | |
| max_idx = values.index(max(values)) | |
| bars[max_idx].set_edgecolor("red") | |
| bars[max_idx].set_linewidth(2) | |
| ax.set_xlim(0, 1) | |
| ax.set_xlabel("Probability", fontsize=11) | |
| ax.set_title("Expression Classification", fontsize=12, fontweight="bold") | |
| ax.invert_yaxis() | |
| plt.tight_layout() | |
| tmp = tempfile.NamedTemporaryFile(suffix=".png", delete=False, dir="/tmp") | |
| fig.savefig(tmp.name, dpi=150, bbox_inches="tight") | |
| plt.close(fig) | |
| return tmp.name | |
| def _make_au_plot( | |
| probs: dict[str, float], | |
| stds: dict[str, float], | |
| decode_mode: str, | |
| thresholds: dict[str, float] | None, | |
| ) -> str: | |
| """Create a bar chart of AU probabilities with threshold markers.""" | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import tempfile | |
| names = list(probs.keys()) | |
| values = list(probs.values()) | |
| errors = list(stds.values()) if decode_mode == "flow" else None | |
| ths = [thresholds.get(n, 0.5) if thresholds else 0.5 for n in names] | |
| fig, ax = plt.subplots(1, 1, figsize=(7, 3.5)) | |
| colors = ["#2ecc71" if v >= t else "#e74c3c" for v, t in zip(values, ths)] | |
| bars = ax.bar(names, values, color=colors, yerr=errors, capsize=2, error_kw={"linewidth": 0.8, "alpha": 0.7}) | |
| # Threshold markers | |
| ax.plot(names, ths, "k^", markersize=5, label="Calibrated threshold") | |
| ax.axhline(y=0.5, color="gray", linewidth=0.5, linestyle="--", alpha=0.5) | |
| ax.set_ylim(0, 1.05) | |
| ax.set_ylabel("Probability", fontsize=11) | |
| ax.set_title("Action Unit Detection", fontsize=12, fontweight="bold") | |
| ax.legend(loc="upper right", fontsize=8) | |
| plt.xticks(rotation=45, ha="right") | |
| plt.tight_layout() | |
| tmp = tempfile.NamedTemporaryFile(suffix=".png", delete=False, dir="/tmp") | |
| fig.savefig(tmp.name, dpi=150, bbox_inches="tight") | |
| plt.close(fig) | |
| return tmp.name | |
| # --------------------------------------------------------------------------- | |
| # Load model and calibration at module scope (ZeroGPU rule #2) | |
| # --------------------------------------------------------------------------- | |
| print("Loading AffectFlow-DINO model...") | |
| _MODEL = load_model() | |
| _MODEL = _MODEL.to("cuda") | |
| print(f"Model loaded on CUDA. Backbone: {_MODEL.backbone.source}") | |
| _AU_THRESHOLDS, _EXPR_WEIGHTS = load_calibration() | |
| if _AU_THRESHOLDS: | |
| print(f"Loaded AU thresholds: {len(_AU_THRESHOLDS)} AUs") | |
| if _EXPR_WEIGHTS: | |
| print(f"Loaded expression weights: {len(_EXPR_WEIGHTS)} classes") | |
| # --------------------------------------------------------------------------- | |
| # Gradio UI | |
| # --------------------------------------------------------------------------- | |
| CSS = """ | |
| #col-container { max-width: 1100px; margin: 0 auto; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| with gr.Blocks() as demo: | |
| gr.Markdown( | |
| "# AffectFlow-DINO: Uncertainty-Aware Multi-Task Affect Estimation\n\n" | |
| "Upload a cropped face image to predict **valence-arousal**, " | |
| "**facial expression** (8-way), and **12 Action Units** — " | |
| "with uncertainty estimates from conditional rectified flow.\n\n" | |
| "Model: [Bekhouche/AffectFlow-DINO](https://huggingface.co/Bekhouche/AffectFlow-DINO) · " | |
| "Paper: [arXiv:2607.13250](https://arxiv.org/abs/2607.13250)" | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| input_image = gr.Image( | |
| label="Face image", type="pil", | |
| ) | |
| run_btn = gr.Button("Predict Affect", variant="primary", scale=1) | |
| with gr.Accordion("Advanced settings", open=False): | |
| decode_mode = gr.Radio( | |
| choices=["flow", "deterministic"], | |
| value="flow", | |
| label="Decode mode", | |
| info="Flow: Monte Carlo sampling with uncertainty. " | |
| "Deterministic: direct head prediction (faster, no uncertainty).", | |
| ) | |
| flow_steps = gr.Slider( | |
| minimum=5, maximum=100, value=30, step=5, | |
| label="Flow steps", | |
| info="Rectified-flow integration steps (more = finer sampling).", | |
| ) | |
| flow_samples = gr.Slider( | |
| minimum=1, maximum=64, value=16, step=1, | |
| label="Flow samples", | |
| info="Number of MC trajectories for uncertainty estimation.", | |
| ) | |
| with gr.Column(scale=1): | |
| summary_out = gr.Markdown(label="Prediction Summary") | |
| va_plot = gr.Image(label="Valence–Arousal", show_label=True) | |
| with gr.Row(): | |
| expr_plot = gr.Image(label="Expression Probabilities", show_label=True) | |
| au_plot = gr.Image(label="Action Unit Probabilities", show_label=True) | |
| gr.Examples( | |
| examples=[ | |
| ["examples/astronaut.jpg"], | |
| ["examples/businessman_arms_crossed.jpg"], | |
| ["examples/elderly_man_beard.jpg"], | |
| ["examples/woman_sad.jpg"], | |
| ["examples/indian_woman_red.jpg"], | |
| ["examples/man_beach.jpg"], | |
| ], | |
| inputs=[input_image], | |
| outputs=[va_plot, expr_plot, au_plot, summary_out], | |
| fn=predict, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| ) | |
| run_btn.click( | |
| fn=predict, | |
| inputs=[input_image, decode_mode, flow_steps, flow_samples], | |
| outputs=[va_plot, expr_plot, au_plot, summary_out], | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS) |