Team Ai
Apppublic

PCNUSMSE/transcript_service

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes
gradio_interface.py574 linesDownload Raw Back to api
1"""Gradio用户界面模块2 3提供基于Gradio的Web界面,支持文件上传、进度显示和结果展示。4"""5 6import asyncio7import json8from pathlib import Path9from typing import Dict, List, Optional, Tuple, Any10import gradio as gr11import pandas as pd12 13from ..core.config import get_config14from ..core.task_manager import get_task_manager, TaskStatus, TaskPriority15from ..utils.logger import get_task_logger16from ..services.file_validator import get_file_validator17 18 19class GradioInterface:20    """Gradio界面管理器"""21    22    def __init__(self):23        """初始化Gradio界面"""24        self.config = get_config()25        self.task_manager = get_task_manager()26        self.file_validator = get_file_validator()27        self.logger = get_task_logger(logger_name="transcript_service.gradio")28        29        # 当前任务ID30        self.current_task_id = None31        32        # 创建界面33        self.interface = self._create_interface()34        35        # 注册任务状态回调36        self.task_manager.add_status_callback(self._on_task_status_change)37    38    def _create_interface(self) -> gr.Blocks:39        """创建Gradio界面"""40        # 获取支持的格式信息41        supported_formats = self.file_validator.get_supported_formats()42        43        with gr.Blocks(44            title="音频转文字服务",45            theme=gr.themes.Soft(),46            css="""47            .main-container { max-width: 1000px; margin: 0 auto; }48            .upload-area { border: 2px dashed #ccc; border-radius: 10px; padding: 20px; text-align: center; }49            .result-area { margin-top: 20px; }50            .status-simple { font-size: 16px; font-weight: bold; }51            """52        ) as interface:53            # 简洁标题54            gr.Markdown("# 🎵 音频转文字服务")55            56            with gr.Row():57                with gr.Column(scale=3):58                    # 文件上传区59                    file_upload = gr.File(60                        label="📁 选择音频文件(支持多文件)",61                        file_count="multiple",62                        file_types=list(supported_formats['extensions']),63                        height=12064                    )65                    66                    # 简化的配置区67                    with gr.Row():68                        # 任务优先级69                        priority_select = gr.Radio(70                            label="优先级",71                            choices=[("普通", "NORMAL"), ("高优先级", "HIGH")],72                            value="NORMAL"73                        )74                    75                    # 参数设置区(默认隐藏)76                    with gr.Accordion("⚙️ 转录参数设置", open=False) as params_section:77                        # 语言选择78                        language_select = gr.CheckboxGroup(79                            label="识别语言",80                            choices=[81                                ("中文", "zh"), ("英文", "en"), ("日语", "ja"), 82                                ("粤语", "yue"), ("韩语", "ko"), ("德语", "de"),83                                ("法语", "fr"), ("俄语", "ru")84                            ],85                            value=["zh", "en"]86                        )87                        88                        with gr.Row():89                            # 基础选项90                            disfluency_removal = gr.Checkbox(91                                label="过滤语气词",92                                value=True93                            )94                            timestamp_alignment = gr.Checkbox(95                                label="时间戳校准",96                                value=True97                            )98                            diarization_enabled = gr.Checkbox(99                                label="说话人分离",100                                value=True101                            )102                        103                        with gr.Row():104                            speaker_count = gr.Number(105                                label="说话人数量(可选)",106                                value=None,107                                minimum=None,108                                maximum=100,109                                step=1,110                                info="留空则自动判断,如需指定请输入2-100之间的数值"111                            )112                            channel_select = gr.Textbox(113                                label="音轨索引",114                                value="0",115                                info="多音轨文件的音轨索引,用逗号分隔"116                            )117                        118                        # 高级选项(更深层折叠)119                        with gr.Accordion("高级选项", open=False):120                            vocabulary_id = gr.Textbox(121                                label="热词ID v2",122                                value="",123                                info="v2模型的热词ID"124                            )125                            phrase_id = gr.Textbox(126                                label="热词ID v1",127                                value="",128                                info="v1模型的热词ID"129                            )130                            special_word_filter = gr.Textbox(131                                label="敏感词过滤配置",132                                value="",133                                lines=2,134                                placeholder='JSON格式配置',135                                info="敏感词过滤的JSON配置"136                            )137                    138                    # 控制按钮139                    with gr.Row():140                        start_btn = gr.Button("🚀 开始转录", variant="primary", size="lg")141                        cancel_btn = gr.Button("❌ 取消", variant="secondary")142                        clear_btn = gr.Button("🗑️ 清空", variant="secondary")143                144                with gr.Column(scale=2):145                    # 简化的状态显示146                    status_text = gr.Textbox(147                        label="📊 当前状态",148                        value="等待上传文件...",149                        interactive=False,150                        elem_classes=["status-simple"]151                    )152                    153                    # 转录结果154                    result_text = gr.Textbox(155                        label="📝 转录结果",156                        placeholder="转录结果将在这里显示...",157                        lines=12,158                        max_lines=20,159                        show_copy_button=True,160                        elem_classes=["result-area"]161                    )162                    163                    # 文件统计表格164                    stats_df = gr.Dataframe(165                        headers=["文件名", "时长", "文本长度", "置信度"],166                        datatype=["str", "str", "number", "number"],167                        label="📈 处理统计",168                        visible=False169                    )170            171            # 折叠的详细信息区域172            with gr.Accordion("📋 详细信息", open=False) as detail_section:173                with gr.Tabs():174                    with gr.Tab("系统信息"):175                        system_info = gr.JSON(176                            label="服务状态",177                            value=self._get_system_info()178                        )179                        format_info = gr.JSON(180                            label="支持格式",181                            value=supported_formats182                        )183                    184                    with gr.Tab("任务信息"):185                        task_info = gr.JSON(186                            label="当前任务",187                            value={}188                        )189                    190                    with gr.Tab("完整结果"):191                        result_json = gr.JSON(192                            label="JSON结果",193                            value={}194                        )195                    196                    with gr.Tab("处理日志"):197                        log_text = gr.Textbox(198                            label="详细日志",199                            lines=8,200                            max_lines=12,201                            interactive=False,202                            show_copy_button=True203                        )204                        log_download = gr.File(205                            label="下载日志文件",206                            visible=False207                        )208            209 210            211            # 添加手动刷新按钮212            with gr.Row():213                refresh_btn = gr.Button("🔄 刷新状态", variant="secondary", size="sm")214                refresh_btn.click(215                    fn=self._update_interface,216                    outputs=[status_text, task_info, result_text, result_json, stats_df, system_info, log_text]217                )218            219            # 事件处理220            start_btn.click(221                fn=self._process_files,222                inputs=[223                    file_upload, priority_select, language_select, 224                    disfluency_removal, timestamp_alignment, diarization_enabled,225                    speaker_count, channel_select, vocabulary_id, 226                    phrase_id, special_word_filter227                ],228                outputs=[status_text, task_info, log_text]229            )230            231            cancel_btn.click(232                fn=self._cancel_current_task,233                outputs=[status_text, task_info]234            )235            236            clear_btn.click(237                fn=self._clear_interface,238                outputs=[file_upload, result_text, result_json, stats_df, log_text, status_text, task_info]239            )240            241            # 定时更新242            interface.load(243                fn=self._update_interface,244                outputs=[status_text, task_info, result_text, result_json, stats_df, system_info, log_text]245            )246        247        return interface248    249    def _get_custom_css(self) -> str:250        """获取自定义CSS样式"""251        return """252        .gradio-container {253            max-width: 1200px !important;254        }255        .gr-button-primary {256            background: linear-gradient(45deg, #FF6B6B, #4ECDC4) !important;257            border: none !important;258        }259        .gr-button-primary:hover {260            transform: translateY(-2px) !important;261            box-shadow: 0 4px 12px rgba(0,0,0,0.15) !important;262        }263        .progress-bar {264            background: linear-gradient(90deg, #FF6B6B, #4ECDC4) !important;265        }266        """267    268    def _get_system_info(self) -> Dict:269        """获取系统信息"""270        stats = self.task_manager.get_statistics()271        return {272            "服务状态": "运行中",273            "当前任务数": stats['total_tasks'],274            "待处理": stats['pending'],275            "处理中": stats['validating'] + stats['uploading'] + stats['transcribing'],276            "已完成": stats['completed'],277            "失败": stats['failed'],278            "队列大小": stats['queue_size']279        }280    281    def _get_timestamp(self) -> str:282        """获取当前时间戳"""283        from datetime import datetime284        return datetime.now().strftime("%Y-%m-%d %H:%M:%S")285    286    async def _process_files(287        self, 288        files: List, 289        priority: str,290        languages: List[str], 291        disfluency_removal: bool,292        timestamp_alignment: bool,293        diarization_enabled: bool,294        speaker_count: Optional[int] | None,295        channel_id: str,296        vocabulary_id: str,297        phrase_id: str,298        special_word_filter: str299    ) -> Tuple[str, Dict, str]:300        """处理上传的文件301        302        Args:303            files: 上传的文件列表304            languages: 选择的语言305            priority: 任务优先级306            channel_id: 音轨索引307            disfluency_removal: 是否过滤语气词308            timestamp_alignment: 是否启用时间戳校准309            diarization_enabled: 是否启用说话人分离310            speaker_count: 说话人数量参考值311            vocabulary_id: 热词ID v2312            phrase_id: 热词ID v1313            special_word_filter: 敏感词过滤配置314            315        Returns:316            (状态信息, 任务信息, 日志信息)317        """318        try:319            if not files:320                return "请先上传音频文件", {}, "错误: 未选择任何文件"321            322            # 记录详细日志323            log_messages = []324            log_messages.append(f"[{self._get_timestamp()}] 开始处理文件上传请求")325            log_messages.append(f"[{self._get_timestamp()}] 接收到 {len(files)} 个文件")326            327            # 转换文件路径328            file_paths = [Path(f.name) for f in files]329            log_messages.append(f"[{self._get_timestamp()}] 转换文件路径完成")330            331            # 显示文件信息332            for i, file_path in enumerate(file_paths):333                try:334                    file_size = file_path.stat().st_size335                    log_messages.append(f"[{self._get_timestamp()}] 文件 {i+1}: {file_path.name} (大小: {file_size} 字节)")336                except Exception as e:337                    log_messages.append(f"[{self._get_timestamp()}] 文件 {i+1}: {file_path.name} (无法获取文件信息: {str(e)})")338            339            # 解析音轨参数340            try:341                channel_list = [int(x.strip()) for x in channel_id.split(',') if x.strip()]342            except ValueError:343                channel_list = [0]  # 默认为第一条音轨344            345            # 验证说话人数量参数346            validated_speaker_count = None347            if speaker_count is not None:348                if isinstance(speaker_count, (int, float)) and speaker_count >= 2 and speaker_count <= 100:349                    validated_speaker_count = int(speaker_count)350                else:351                    log_messages.append(f"[{self._get_timestamp()}] 警告: 说话人数量无效({speaker_count}),将使用自动判断")352            353            # 解析敏感词过滤参数354            special_filter = None355            if special_word_filter.strip():356                try:357                    special_filter = json.loads(special_word_filter)358                except json.JSONDecodeError as e:359                    log_messages.append(f"[{self._get_timestamp()}] 警告: 敏感词过滤配置格式错误,将使用默认设置")360            361            # 创建任务362            task_priority = TaskPriority.HIGH if priority == "HIGH" else TaskPriority.NORMAL363            364            # 准备元数据,包含所有Paraformer参数365            metadata = {366                "languages": languages,367                "file_count": len(file_paths),368                "paraformer_params": {369                    "language_hints": languages,370                    "channel_id": channel_list,371                    "disfluency_removal_enabled": disfluency_removal,372                    "timestamp_alignment_enabled": timestamp_alignment,373                    "diarization_enabled": diarization_enabled,374                    "speaker_count": validated_speaker_count,375                    "vocabulary_id": vocabulary_id.strip() if vocabulary_id.strip() else None,376                    "phrase_id": phrase_id.strip() if phrase_id.strip() else None,377                    "special_word_filter": json.dumps(special_filter) if special_filter else None378                }379            }380            381            log_messages.append(f"[{self._get_timestamp()}] 创建任务,优先级: {task_priority.value}")382            log_messages.append(f"[{self._get_timestamp()}] 选择语言: {', '.join(languages) if languages else '自动识别'}")383            384            self.current_task_id = await self.task_manager.create_task(385                file_paths=file_paths,386                priority=task_priority,387                metadata=metadata388            )389            390            task = self.task_manager.get_task(self.current_task_id)391            392            log_messages.append(f"[{self._get_timestamp()}] 任务创建成功,任务ID: {self.current_task_id}")393            394            return (395                f"任务已创建: {self.current_task_id}",396                task.to_dict() if task else {},397                "\n".join(log_messages) + f"\n开始处理 {len(file_paths)} 个文件...\n"398            )399            400        except Exception as e:401            error_msg = f"创建任务失败: {str(e)}"402            self.logger.exception(error_msg)403            return error_msg, {}, f"错误: {error_msg}\n"404    405    def _cancel_current_task(self) -> Tuple[str, Dict]:406        """取消当前任务"""407        if not self.current_task_id:408            return "没有正在执行的任务", {}409        410        success = asyncio.create_task(411            self.task_manager.cancel_task(self.current_task_id)412        )413        414        if success:415            return f"任务 {self.current_task_id} 已取消", {}416        else:417            return "取消任务失败", {}418    419    def _clear_interface(self) -> Tuple[None, str, Dict, List, str, str, Dict]:420        """清空界面"""421        self.current_task_id = None422        return (423            None,  # file_upload424            "",    # result_text425            {},    # result_json426            [],    # stats_df427            "",    # log_text428            "界面已清空,等待上传文件...",  # status_text429            {}     # task_info430        )431    432    def _update_interface(self) -> Tuple[str, Dict, str, Dict, List, Dict, str]:433        """更新界面状态"""434        # 更新当前任务状态435        status_text = "等待上传文件..."436        task_info = {}437        result_text = ""438        result_json = {}439        stats_data = []440        log_text = ""441        442        if self.current_task_id:443            task = self.task_manager.get_task(self.current_task_id)444            if task:445                task_info = task.to_dict()446                status_text = f"[{task.status.value}] {task.progress.message}"447                448                # 收集详细日志449                log_text = self._collect_task_logs(task)450                451                # 如果任务完成,显示结果452                if task.status == TaskStatus.COMPLETED:453                    self.logger.debug(f"任务已完成,检查转录结果: {task.result.transcription_results}")454                    if task.result.transcription_results:455                        result_json = task.result.transcription_results456                        457                        # 提取转录文本458                        transcriptions = result_json.get('transcriptions', [])459                        self.logger.debug(f"转录结果: {transcriptions}")460                        result_text = "\n\n".join([461                            f"文件: {t.get('file_url', '').split('/')[-1]}\n{t.get('text', '')}"462                            for t in transcriptions if t.get('text')463                        ])464                        465                        # 生成统计表格466                        stats_data = []467                        for t in transcriptions:468                            if 'error' not in t:469                                stats_data.append([470                                    t.get('file_url', '').split('/')[-1],471                                    f"{t.get('duration', 0):.1f}s",472                                    len(t.get('text', '')),473                                    t.get('language', 'unknown'),474                                    round(t.get('confidence', 0), 3)475                                ])476                    else:477                        self.logger.debug("任务已完成但没有转录结果")478                elif task.status == TaskStatus.FAILED:479                    # 如果任务失败,显示错误信息480                    if task.result and task.result.error_message:481                        log_text += f"\n[{self._get_timestamp()}] 任务失败: {task.result.error_message}"482        483        # 更新系统信息484        system_info = self._get_system_info()485        486        return status_text, task_info, result_text, result_json, stats_data, system_info, log_text487    488    def _collect_task_logs(self, task) -> str:489        """收集任务的详细日志490        491        Args:492            task: 任务对象493            494        Returns:495            格式化的日志字符串496        """497        if not task:498            return "无任务信息"499        500        log_lines = []501        log_lines.append(f"[{self._get_timestamp()}] 任务ID: {task.id}")502        log_lines.append(f"[{self._get_timestamp()}] 任务状态: {task.status.value}")503        log_lines.append(f"[{self._get_timestamp()}] 任务创建时间: {task.created_at}")504        505        # 添加进度信息506        if task.progress:507            log_lines.append(f"[{self._get_timestamp()}] 进度信息: {task.progress.message}")508            # TaskProgress对象没有details属性,只使用message509        510        # 添加文件信息511        if hasattr(task, 'file_paths') and task.file_paths:512            log_lines.append(f"[{self._get_timestamp()}] 文件列表:")513            for i, file_path in enumerate(task.file_paths):514                try:515                    file_size = file_path.stat().st_size516                    log_lines.append(f"  {i+1}. {file_path.name} ({file_size} bytes)")517                except Exception as e:518                    log_lines.append(f"  {i+1}. {file_path.name} (无法获取文件信息: {str(e)})")519        520        # 添加结果信息(如果任务已完成)521        if task.status == TaskStatus.COMPLETED and task.result:522            log_lines.append(f"[{self._get_timestamp()}] 任务完成时间: {task.completed_at}")523            if hasattr(task.result, 'transcription_results') and task.result.transcription_results:524                transcriptions = task.result.transcription_results.get('transcriptions', [])525                log_lines.append(f"[{self._get_timestamp()}] 转录结果: {len(transcriptions)} 个文件")526        527        # 添加错误信息(如果有的话)528        # Task对象没有error属性,错误信息在result中529        530        return "\n".join(log_lines)531    532    def _on_task_status_change(self, task):533        """任务状态变化回调"""534        self.logger.debug(f"任务状态变化: {task.id} -> {task.status.value}")535        # 当任务状态变化时,不直接更新界面,而是依赖定时更新机制536        # Gradio的回调中不能直接更新界面组件537    538    def launch(self, **kwargs):539        """启动Gradio界面"""540        default_kwargs = {541            'server_name': '0.0.0.0',  # 改为0.0.0.0以允许外部访问542            'server_port': self.config.app.port,543            'share': True,  # 开启分享链接544            'debug': self.config.app.debug,545            'show_error': True,546            'quiet': not self.config.app.debug547        }548        default_kwargs.update(kwargs)549        550        self.logger.info(f"启动Gradio界面: http://{default_kwargs['server_name']}:{default_kwargs['server_port']}")551        552        return self.interface.launch(**default_kwargs)553 554 555# 全局界面实例556gradio_interface = GradioInterface()557 558 559def get_gradio_interface() -> GradioInterface:560    """获取Gradio界面实例561    562    Returns:563        Gradio界面实例564    """565    return gradio_interface566 567 568def create_demo_interface() -> gr.Blocks:569    """创建演示界面570    571    Returns:572        Gradio界面对象573    """574    return gradio_interface.interface