import os import cv2 import numpy as np import torch import torch.nn.functional as F import matplotlib.pyplot as plt import matplotlib import tempfile import gradio as gr import spaces from easydict import EasyDict as edict from huggingface_hub import hf_hub_download from collections import OrderedDict # 导入你本地的模型构建和鱼眼投影工具 from utils.common_config import get_model from utils.panorama_utils import pano_to_fisheye_stereographic # ========================================== # 1. 完整的 ADE20K Palette # ========================================== ADE20K_PALETTE = [ [120, 120, 120], [180, 120, 120], [6, 230, 230], [80, 50, 50], [4, 200, 3], [120, 120, 80], [140, 140, 140], [204, 5, 255], [230, 230, 230], [4, 250, 7], [224, 5, 255], [235, 255, 7], [150, 5, 61], [120, 120, 70], [8, 255, 51], [255, 6, 82], [143, 255, 140], [204, 255, 4], [255, 51, 7], [204, 70, 3], [0, 102, 200], [61, 230, 250], [255, 6, 51], [11, 102, 255], [255, 7, 71], [255, 9, 224], [9, 7, 230], [220, 220, 220], [255, 9, 92], [112, 9, 255], [8, 255, 214], [7, 255, 224], [255, 184, 6], [10, 255, 71], [255, 41, 10], [7, 255, 255], [224, 255, 8], [102, 8, 255], [255, 61, 6], [255, 194, 7], [255, 122, 8], [0, 255, 20], [255, 8, 41], [255, 5, 153], [6, 51, 255], [235, 12, 255], [160, 150, 20], [0, 163, 255], [140, 140, 140], [250, 10, 15], [20, 255, 0], [31, 255, 0], [255, 31, 0], [255, 224, 0], [153, 255, 0], [0, 0, 255], [255, 71, 0], [0, 235, 255], [0, 174, 255], [0, 122, 255], [245, 0, 255], [255, 6, 122], [255, 245, 0], [10, 190, 212], [214, 255, 0], [0, 204, 255], [255, 0, 112], [0, 8, 255], [255, 0, 31], [255, 61, 0], [204, 0, 255], [255, 0, 204], [255, 255, 0], [0, 153, 255], [0, 102, 255], [0, 255, 245], [0, 255, 102], [255, 163, 0], [255, 153, 0], [0, 255, 10], [255, 112, 0], [143, 255, 0], [0, 41, 255], [0, 255, 174], [255, 0, 10], [174, 255, 0], [255, 245, 0], [255, 0, 20], [255, 0, 143], [255, 0, 82], [0, 245, 255], [0, 61, 255], [0, 255, 71], [0, 255, 153], [255, 0, 163], [255, 0, 174], [255, 0, 51], [255, 0, 71], [0, 204, 255], [255, 10, 0], [0, 255, 41], [0, 255, 51], [255, 204, 0], [255, 0, 194], [255, 102, 0], [0, 153, 255], [0, 102, 255], [0, 255, 204], [255, 0, 224], [255, 0, 92], [255, 0, 112], [255, 0, 122], [0, 255, 31], [0, 102, 255], [255, 0, 153], [255, 0, 143], [255, 0, 163], [255, 0, 6], [255, 0, 184], [0, 255, 214], [0, 255, 194], [255, 0, 71], [0, 255, 224], [255, 0, 143], [255, 0, 133], [122, 255, 0], [255, 0, 10], [255, 153, 0], [0, 112, 255], [255, 163, 0], [255, 204, 0], [255, 0, 41], [255, 0, 10], [255, 0, 20], [255, 0, 204], [255, 0, 194], [255, 0, 153], [255, 10, 0], [255, 0, 122], [255, 0, 71], [255, 0, 51], [255, 0, 31], [102, 255, 0], [0, 255, 10], [172, 255, 0], [255, 29, 0], [255, 0, 28], [255, 122, 0], [0, 255, 143], [255, 255, 184] ] # ========================================== # 2. 硬编码配置与模型初始化 # ========================================== def get_inference_config(): cfg = edict() cfg.model = 'TransformerBFE-DINO-DPT' cfg.backbone = 'dinov3L' cfg.head = 'dpt_head' cfg.embed_dim = 512 cfg.mtt_resolution_downsample_rate = 2 cfg.PRED_OUT_NUM_CONSTANT = 64 cfg.train_db_name = 'PanoMTDU' cfg.task_dictionary = edict({ 'include_semseg': True, 'include_depth': True, 'include_normals': True }) cfg.TASKS = edict() cfg.TASKS.NAMES = ['semseg', 'depth', 'normals'] cfg.TASKS.NUM_OUTPUT = { 'semseg': 150, 'depth': 1, 'normals': 3 } cfg.TEST = edict() cfg.TEST.SCALE = (512, 1024) cfg.TEST.PANO_SCALE = (512, 1024) cfg.TRAIN = edict() cfg.TRAIN.SCALE = (512, 1024) cfg.TRAIN.PANO_SCALE = (512, 1024) return cfg device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') cfg = get_inference_config() model = get_model(cfg) weight_filename = "mtpano_model_408k.pth.tar" print(f"Fetching {weight_filename} from Hugging Face (jdzhang0929/MTPano)...") try: checkpoint_path = hf_hub_download(repo_id="jdzhang0929/MTPano", filename=weight_filename) print(f"Successfully loaded checkpoint from: {checkpoint_path}") checkpoint = torch.load(checkpoint_path, map_location='cpu') state_dict = checkpoint.get('model', checkpoint) new_state_dict = OrderedDict() for k, v in state_dict.items(): name = k[7:] if k.startswith('module.') else k new_state_dict[name] = v try: model.load_state_dict(new_state_dict, strict=True) print("Model loaded (strict).") except RuntimeError as e: print(f"Strict loading failed, trying non-strict. Error: {e}") model.load_state_dict(new_state_dict, strict=False) print("Model loaded (non-strict).") except Exception as e: print(f"Error downloading or loading weights: {e}") print("Warning: Model failed to load. Demo will output garbage.") model.eval() # ========================================== # 3. 严格对齐的后处理核心函数 # ========================================== def smooth_step(t): return 3 * t**2 - 2 * t**3 def fix_panorama_seam(img_np, margin=2, task_type='depth'): H, W = img_np.shape[:2] rolled = np.roll(img_np, W // 2, axis=1) center = W // 2 left_idx = center - margin - 1 right_idx = center + margin left_anchors = rolled[:, left_idx] right_anchors = rolled[:, right_idx] if task_type == 'depth': valid_rows = (left_anchors > 1e-4) & (right_anchors > 1e-4) invalid_val = 0.0 elif task_type == 'normal': valid_left = np.any(left_anchors != 0.0, axis=-1) valid_right = np.any(right_anchors != 0.0, axis=-1) valid_rows = valid_left & valid_right invalid_val = 0.0 else: return img_np steps = 2 * margin + 1 for i, col in enumerate(range(center - margin, center + margin)): alpha = (i + 1) / steps if task_type == 'depth': interp = (1 - alpha) * left_anchors[valid_rows] + alpha * right_anchors[valid_rows] rolled[valid_rows, col] = interp elif task_type == 'normal': vec_left = left_anchors[valid_rows] vec_right = right_anchors[valid_rows] vec_interp = (1 - alpha) * vec_left + alpha * vec_right norms = np.linalg.norm(vec_interp, axis=-1, keepdims=True) norms[norms < 1e-6] = 1.0 vec_norm = vec_interp / norms rolled[valid_rows, col] = vec_norm rolled[~valid_rows, col] = invalid_val return np.roll(rolled, -(W // 2), axis=1) def colorize_semantic(semseg_tensor): semseg_np = semseg_tensor.squeeze().cpu().numpy().astype(np.uint8) h, w = semseg_np.shape palette = np.array(ADE20K_PALETTE, dtype=np.uint8) max_idx = len(palette) - 1 safe_semseg = np.clip(semseg_np, 0, max_idx) color_img = palette[safe_semseg] if semseg_np.max() > max_idx: mask = semseg_np > max_idx color_img[mask] = [0, 0, 0] return torch.from_numpy(color_img / 255.0).permute(2, 0, 1).unsqueeze(0).float() def colorize_depth_strict(depth_map_np, global_max, cmap_name="Spectral"): valid_mask = depth_map_np > 1e-4 if not valid_mask.any(): return np.zeros((depth_map_np.shape[0], depth_map_np.shape[1], 3), dtype=np.uint8) depth_norm = np.clip(depth_map_np / global_max, 0.0, 1.0) try: colormap_func = matplotlib.colormaps[cmap_name] except AttributeError: colormap_func = plt.get_cmap(cmap_name) colored = colormap_func(depth_norm)[..., :3] depth_rgb = (colored * 255).astype(np.uint8) depth_rgb[~valid_mask] = [0, 0, 0] return cv2.cvtColor(depth_rgb, cv2.COLOR_RGB2BGR) def blend_mask_rgb(mask_bgr, rgb_bgr, alpha=0.6): if mask_bgr.shape != rgb_bgr.shape: mask_bgr = cv2.resize(mask_bgr, (rgb_bgr.shape[1], rgb_bgr.shape[0]), interpolation=cv2.INTER_NEAREST) beta = 1.0 - alpha return cv2.addWeighted(mask_bgr, alpha, rgb_bgr, beta, 0.0) # ========================================== # 4. 推理及视频渲染引擎 # ========================================== def _run_inference_engine(model, img_tensor, device): """带 Mask 过滤和平滑接缝的核心推理""" with torch.no_grad(): inputs = img_tensor.to(device) outputs = model(inputs) results = {} if 'semseg' in outputs: sem_logits = outputs['semseg'] sem_pred = torch.argmax(sem_logits, dim=1, keepdim=True) results['semseg'] = sem_pred.float() if 'depth' in outputs: results['depth'] = outputs['depth'] results['raw_depth'] = outputs['depth'].clone() if 'normals' in outputs: norm_out = outputs['normals'] norm_out = F.normalize(norm_out, p=2, dim=1) results['normals'] = norm_out # Step 1: Filter out sky (ADE20K class 2) filter_class_ids_depth = [2] filter_class_ids_normals = [2, 8, 68] if 'semseg' in results and ('depth' in results or 'normals' in results): sem_map = results['semseg'] filter_mask_depth = torch.zeros_like(sem_map, dtype=torch.bool) filter_mask_normals = torch.zeros_like(sem_map, dtype=torch.bool) for cid in filter_class_ids_depth: filter_mask_depth = filter_mask_depth | (sem_map == cid) for cid in filter_class_ids_normals: filter_mask_normals = filter_mask_normals | (sem_map == cid) if 'depth' in results: results['depth'][filter_mask_depth] = 0.0 results['raw_depth'][filter_mask_depth] = 0.0 if 'normals' in results: filter_mask_3d = filter_mask_normals.repeat(1, 3, 1, 1) results['normals'][filter_mask_3d] = 0.0 # Step 2: Fix Seams if 'depth' in results: depth_np = results['depth'].squeeze().cpu().numpy() fixed_depth = fix_panorama_seam(depth_np, margin=2, task_type='depth') results['depth'] = torch.from_numpy(fixed_depth).unsqueeze(0).unsqueeze(0).to(device) raw_depth_np = results['raw_depth'].squeeze().cpu().numpy() fixed_raw_depth = fix_panorama_seam(raw_depth_np, margin=2, task_type='depth') results['raw_depth'] = torch.from_numpy(fixed_raw_depth).unsqueeze(0).unsqueeze(0).to(device) if 'normals' in results: norm_np = results['normals'].squeeze().permute(1, 2, 0).cpu().numpy() fixed_norm = fix_panorama_seam(norm_np, margin=2, task_type='normal') results['normals'] = torch.from_numpy(fixed_norm).permute(2, 0, 1).unsqueeze(0).to(device) return results def generate_video_engine(outputs_dict, raw_rgb, device, output_path): """电影级运镜渲染引擎""" pano_rgb = raw_rgb.to(device) pano_sem = outputs_dict.get('semseg').to(device) pano_depth_for_proj = outputs_dict.get('raw_depth', outputs_dict.get('depth')).to(device) pano_normal = outputs_dict.get('normals').to(device) PERSP_H, PERSP_W = 448, 448 FPS, DURATION = 30, 10 TOTAL_FRAMES = FPS * DURATION fourcc = cv2.VideoWriter_fourcc(*'mp4v') video_writer = cv2.VideoWriter(output_path, fourcc, FPS, (PERSP_W * 4, PERSP_H)) global_max = pano_depth_for_proj.max().item() if global_max == 0: global_max = 1.0 print(f"Rendering 10s Flythrough Video to {output_path}...") for i in range(TOTAL_FRAMES): progress = i / TOTAL_FRAMES # 运镜设计 if progress < 0.05: yaw, pitch, h_fov = -180.0, 180.0, 90.0 elif progress < 0.40: t = (progress - 0.05) / 0.35 yaw, pitch, h_fov = -180.0 - smooth_step(t) * 360.0, 180.0, 90.0 elif progress < 0.45: yaw, pitch, h_fov = -540.0, 180.0, 90.0 elif progress < 0.70: t = (progress - 0.45) / 0.25 ease = smooth_step(t) yaw = 180.0 + ease * 360.0 pitch = 180.0 - ease * 90.0 h_fov = 90.0 + ease * 130.0 elif progress < 0.75: yaw, pitch, h_fov = 540.0, 90.0, 220.0 else: t = (progress - 0.75) / 0.25 ease = smooth_step(t) yaw = 540.0 - ease * 360.0 pitch = 90.0 + ease * 90.0 h_fov = 220.0 - ease * 130.0 persp_rgb, persp_sem_label, persp_depth_z, persp_normal = pano_to_fisheye_stereographic( pano_rgb, pano_sem, pano_depth_for_proj, pano_normal, h_fov, yaw, pitch, PERSP_H, PERSP_W ) img_rgb_vis = (persp_rgb.squeeze().permute(1, 2, 0).cpu().numpy() * 255).clip(0, 255).astype(np.uint8) img_rgb_vis = cv2.cvtColor(img_rgb_vis, cv2.COLOR_RGB2BGR) persp_sem_color = colorize_semantic(persp_sem_label) img_sem_vis = (persp_sem_color.squeeze().permute(1, 2, 0).numpy() * 255).clip(0, 255).astype(np.uint8) img_sem_vis = cv2.cvtColor(img_sem_vis, cv2.COLOR_RGB2BGR) img_sem_blend = blend_mask_rgb(img_sem_vis, img_rgb_vis, alpha=0.6) # 处理天空 Mask persp_sem_np = persp_sem_label.squeeze().cpu().numpy() sky_mask = (persp_sem_np == 2) d_val = persp_depth_z.squeeze().cpu().numpy() img_depth = colorize_depth_strict(d_val, global_max, "Spectral") img_depth[sky_mask] = [0, 0, 0] n_np = persp_normal.squeeze().permute(1, 2, 0).cpu().numpy() img_norm = ((n_np * 0.5 + 0.5) * 255).clip(0, 255).astype(np.uint8) img_norm[sky_mask] = [127, 127, 127] img_norm = cv2.cvtColor(img_norm, cv2.COLOR_RGB2BGR) grid = np.hstack((img_rgb_vis, img_sem_blend, img_depth, img_norm)) cv2.putText(grid, "RGB", (10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 2) cv2.putText(grid, "Semantic Overlay", (PERSP_W + 10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 2) cv2.putText(grid, "Depth", (PERSP_W * 2 + 10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 2) cv2.putText(grid, "Normal", (PERSP_W * 3 + 10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 255, 255), 2) video_writer.write(grid) video_writer.release() print("Video Rendered!") # ========================================== # 5. 分离的 Gradio 推理接口 # ========================================== @spaces.GPU def predict_images(image_rgb): """只生成 2D 预测结果的接口""" if image_rgb is None: return None, None, None, None, None, gr.update(interactive=False, variant="secondary") # 处理带有 Alpha 通道的图片 if len(image_rgb.shape) == 3 and image_rgb.shape[2] == 4: image_rgb = cv2.cvtColor(image_rgb, cv2.COLOR_RGBA2RGB) model.to(device) target_size = cfg.TEST.PANO_SCALE # 预处理 img_resized = cv2.resize(image_rgb, (target_size[1], target_size[0]), interpolation=cv2.INTER_LINEAR) img_normalized = img_resized.astype(np.float32) / 255.0 mean = np.array([0.485, 0.456, 0.406], dtype=np.float32) std = np.array([0.229, 0.224, 0.225], dtype=np.float32) img_standardized = (img_normalized - mean) / std img_tensor = torch.from_numpy(img_standardized).permute(2, 0, 1).unsqueeze(0).float() # 核心推理 (含 Mask 及接缝修补) results = _run_inference_engine(model, img_tensor, device) # 1. Semantic Overlay 渲染 sem_color_tensor = colorize_semantic(results['semseg']) sem_vis_bgr = (sem_color_tensor.squeeze().permute(1,2,0).cpu().numpy() * 255).astype(np.uint8) sem_vis_bgr = cv2.cvtColor(sem_vis_bgr, cv2.COLOR_RGB2BGR) rgb_vis_bgr = cv2.cvtColor((img_resized).astype(np.uint8), cv2.COLOR_RGB2BGR) overlay_bgr = blend_mask_rgb(sem_vis_bgr, rgb_vis_bgr, alpha=0.6) overlay_rgb = cv2.cvtColor(overlay_bgr, cv2.COLOR_BGR2RGB) # 2. Depth 渲染 depth_np = results['depth'].squeeze().cpu().numpy() global_max = depth_np.max() if depth_np.max() > 0 else 1.0 depth_bgr = colorize_depth_strict(depth_np, global_max, cmap_name="Spectral") depth_color = cv2.cvtColor(depth_bgr, cv2.COLOR_BGR2RGB) depth_gray = (np.clip(depth_np / global_max, 0.0, 1.0) * 65535).astype(np.uint16) # 3. Normal 渲染 norm_np = results['normals'].squeeze().permute(1, 2, 0).cpu().numpy() norm_rgb = ((norm_np * 0.5 + 0.5) * 255).clip(0, 255).astype(np.uint8) return overlay_rgb, depth_color, depth_color, depth_gray, norm_rgb, gr.update(interactive=True, variant="primary") @spaces.GPU(duration=120) def predict_video(image_rgb): if image_rgb is None: return None # 处理带有 Alpha 通道的图片 if len(image_rgb.shape) == 3 and image_rgb.shape[2] == 4: image_rgb = cv2.cvtColor(image_rgb, cv2.COLOR_RGBA2RGB) model.to(device) target_size = cfg.TEST.PANO_SCALE # 预处理 img_resized = cv2.resize(image_rgb, (target_size[1], target_size[0]), interpolation=cv2.INTER_LINEAR) img_normalized = img_resized.astype(np.float32) / 255.0 mean = np.array([0.485, 0.456, 0.406], dtype=np.float32) std = np.array([0.229, 0.224, 0.225], dtype=np.float32) img_tensor = torch.from_numpy((img_normalized - mean) / std).permute(2, 0, 1).unsqueeze(0).float() raw_rgb_tensor = torch.from_numpy(img_normalized).permute(2, 0, 1).unsqueeze(0).float() # 重新推理 results = _run_inference_engine(model, img_tensor, device) # 渲染视频并保存到临时路径 temp_dir = tempfile.mkdtemp() video_path = os.path.join(temp_dir, "demo_flythrough.mp4") generate_video_engine(results, raw_rgb_tensor, device, video_path) return video_path # ========================================== # 6. 现代化 Gradio UI 界面 # ========================================== custom_css = """ /* 1. 强行放大外层按钮容器,并去除多余内边距 */ #pano-examples button { width: 180px !important; height: 90px !important; min-width: 180px !important; min-height: 90px !important; padding: 0 !important; border-radius: 6px !important; overflow: hidden !important; } /* 2. 强行击碎 Gradio 偷偷加在图片外层的 div/span 尺寸限制 */ #pano-examples button > div, #pano-examples button > span { width: 100% !important; height: 100% !important; max-width: none !important; max-height: none !important; padding: 0 !important; margin: 0 !important; display: block !important; } /* 3. 强行让图片本体撑满整个框,并完美保持比例 (cover) */ #pano-examples button img { width: 100% !important; height: 100% !important; max-width: none !important; max-height: none !important; object-fit: cover !important; margin: 0 !important; border-radius: 0 !important; } """ with gr.Blocks(title="MTPano Demo", theme=gr.themes.Soft(), css=custom_css) as demo: # 在后台创建两个不可见的变量,用于储存生成的彩色和灰度深度图 state_depth_color = gr.State() state_depth_gray = gr.State() gr.Markdown( """ # 🌐 MTPano: Multi-Task Panoramic Scene Understanding Upload a single panorama image to simultaneously generate **Semantic Segmentation**, **Depth Estimation**, and **Surface Normals Estimation**. You can also render a stunning **Cinematic Flythrough Video** based on the predictions! """ ) with gr.Row(): # ================= 左侧控制面板 (Input & Controls) ================= with gr.Column(scale=1): gr.Markdown("### 1. Upload & Settings") input_image = gr.Image(label="Input Panorama (RGB)", type="numpy") # 【英文说明】 gr.Markdown("💡 **Note**: Supports 360°×180° panoramas with Equirectangular Projection (ERP). The maximum supported resolution for inference is `1024×512` (larger images will be automatically resized).") # 按钮区 with gr.Row(): btn_infer = gr.Button("🚀 1. Run Inference", variant="primary") with gr.Row(): btn_video = gr.Button("🎥 2. Render Video", variant="secondary", interactive=False) gr.Markdown("*(Note: Video rendering takes ~10s on A100 GPU and will be enabled after Step 1.)*") with gr.Accordion("🖼️ Examples (Click an image to load)", open=True): gr.Examples( elem_id="pano-examples", examples=[ ["./examples/sample1.jpg"], ["./examples/sample2.png"], ["./examples/sample3.png"], ["./examples/sample4.jpg"], ["./examples/sample5.png"], ["./examples/sample6.jpg"], ["./examples/sample7.png"], ["./examples/sample8.jpg"], ["./examples/sample9.jpg"] ], inputs=input_image, label="Choose an example panorama" ) # ================= 右侧展示面板 (Outputs) ================= with gr.Column(scale=2): gr.Markdown("### 2. Predictions") # 👇 修改 1:给 Tabs 加上变量名 result_tabs with gr.Tabs() as result_tabs: # 👇 修改 2:给 Semantic 加上 id="tab_sem" with gr.Tab("🖼️ Semantic", id="tab_sem"): out_sem = gr.Image(label="Semantic Overlay (ADE20K)") with gr.Tab("🌌 Depth"): out_depth = gr.Image(label="Depth Map") depth_mode = gr.Radio( choices=["Colorized (Spectral)", "16-bit Grayscale"], value="Colorized (Spectral)", label="Visualization Mode" ) with gr.Tab("🧭 Normal"): out_norm = gr.Image(label="Surface Normal") # 👇 修改 3:给 Video 加上 id="tab_video" with gr.Tab("🎞️ Cinematic Video", id="tab_video"): out_video = gr.Video(label="10s Flythrough Video", autoplay=True) # ================= 交互事件绑定 ================= # 动作 1:点击 Run 2D Inference btn_infer.click( fn=predict_images, inputs=input_image, outputs=[out_sem, out_depth, state_depth_color, state_depth_gray, out_norm, btn_video] ) # 动作 2:切换深度图模式 def switch_depth_mode(mode, color_img, gray_img): if color_img is None and gray_img is None: return None if mode == "Colorized (Spectral)": return color_img else: return gray_img depth_mode.change( fn=switch_depth_mode, inputs=[depth_mode, state_depth_color, state_depth_gray], outputs=out_depth, queue=False ) # 动作 3:生成视频 btn_video.click( fn=lambda: gr.update(selected="tab_video"), # 瞬间触发 Tab 跳转 inputs=[], outputs=result_tabs, queue=False # 不排队,立即执行 ).then( fn=predict_video, # 跳转完后,紧接着开始生成视频 inputs=input_image, outputs=out_video ) # 动作 4:【终极修复】只监听 clear 事件,避开 Gradio 前端上传被打断的 Bug def reset_state(): return ( None, None, None, None, None, None, gr.update(interactive=False, variant="secondary"), gr.update(selected="tab_sem") # <--- 新增:清空图片时,自动跳回 Semantic 页面 ) # 只有当用户主动点击图片右上角的 ✖️ 清空时,才重置右侧的所有面板 input_image.change( fn=reset_state, inputs=[], outputs=[ out_sem, out_depth, state_depth_color, state_depth_gray, out_norm, out_video, btn_video, result_tabs # <--- 新增:把 result_tabs 加入输出列表,用于接收跳回指令 ], queue=False ) if __name__ == "__main__": demo.launch()