diff --git a/Online/README.md b/Online/README.md index 2281591..7d78947 100644 --- a/Online/README.md +++ b/Online/README.md @@ -76,6 +76,7 @@ | [VideoClassification](https://github.com/mindspore-courses/orange-pi-mindspore/tree/master/Online/community/VideoClassification) | 8.0.0.beta1 | 2.6.0 |8T16G | | [MaskGeneration](https://github.com/mindspore-courses/orange-pi-mindspore/tree/master/Online/community/MaskGeneration) | 8.1.RC1 | 2.6.0 | 8T16G | | [DocumentQuestionAnswering](https://github.com/mindspore-courses/orange-pi-mindspore/tree/master/Online/community/DocumentQuestionAnswering) | 8.0.0.beta1 | 2.6.0 | 20T24G | +| [ThaiMultimodal](https://github.com/mindspore-courses/orange-pi-mindspore/tree/master/Online/community/ThaiMultimodal) | 8.1.RC1 | 2.5.0 | 20T12G | diff --git a/Online/community/README.md b/Online/community/README.md index 1230890..7d483b4 100644 Binary files a/Online/community/README.md and b/Online/community/README.md differ diff --git a/Online/community/ThaiMultimodal/README.md b/Online/community/ThaiMultimodal/README.md new file mode 100644 index 0000000..80d51ec --- /dev/null +++ b/Online/community/ThaiMultimodal/README.md @@ -0,0 +1,133 @@ +# 泰国文化多模态探索平台 + +基于昇思MindSpore框架和wongwian-micro-instruct模型实现的泰语多模态RAG问答系统 + +## 介绍 + +基于香橙派AIpro,利用「万卷·丝路」泰语多模态数据集,构建一个支持RAG问答、图文文化浏览和视频内容检索的多模态探索平台。 + +### 环境准备 + +开发者拿到香橙派开发板后,首先需要进行硬件资源确认,镜像烧录及CANN和MindSpore版本的升级,才可运行该案例,具体如下: + +开发板:香橙派AIpro或其他同硬件开发板 +开发板镜像:Ubuntu镜像 +`CANN Toolkit/Kernels:8.1RC1` +`MindSpore:2.5.0` +`MindNLP:0.4.1` +`Python:3.9` + +#### 镜像烧录 + +运行该案例需要烧录香橙派官网ubuntu镜像,烧录流程参考[昇思MindSpore官网--香橙派开发专区--环境搭建指南--镜像烧录](https://www.mindspore.cn/tutorials/zh-CN/r2.7.0rc1/orange_pi/environment_setup.html)章节。 + +#### CANN升级 + +CANN升级参考[昇思MindSpore官网--香橙派开发专区--环境搭建指南--CANN升级](https://www.mindspore.cn/tutorials/zh-CN/r2.7.0rc1/orange_pi/environment_setup.html)章节。 + +#### MindSpore升级 + +MindSpore升级参考[昇思MindSpore官网--香橙派开发专区--环境搭建指南--MindSpore升级](https://www.mindspore.cn/tutorials/zh-CN/r2.7.0rc1/orange_pi/environment_setup.html)章节。 + +### 核心库版本 + +``` +Python == 3.9 +MindSpore == 2.5.0 +mindnlp == 0.4.1 +gradio == 4.44.0 +scikit-learn == 1.4.0 +``` + +## 快速使用 + +建议下载 wongwian-micro-instruct 模型至本地路径(如 /home/HwHiAiUser): + +``` +git lfs install +git clone https://huggingface.co/wongwian/wongwian-micro-instruct +``` + +在 ipynb 文件中根据模型路径修改对应 step 中的模型路径,然后逐步运行即可。 + +## 1、使用模型 + +**wongwian-micro-instruct**:泰国的小模型,作为问题时输入需要泰语,输入中文时只会当作二进制进行处理无法正确回答 + +使用Qwen1.5模型运行时仅成功运行过一次(picture/0_1.jpg),可以中文输入,中文输出,其他大部分由于内存不足卡死(picture/0_2.jpg)。 + +![运行成功示例](picture/0_1.jpg) + +![内存不足示例](picture/0_2.jpg) + +## 2、界面效果 + +主要分为三个部分,如下图(picture/1.png)所示,分为多模态RAG回答、图文文化浏览、视频内容三个模块。 + +图文浏览主要就是筛选不同标签的图片进行展示,但由于部分图片链接是外网链接,香橙派无法访问到外网,所以图上有部分未显示出来;视频内容图像展示如下图(picture/2.png)。 + +![界面总览](picture/1.png) + +![视频内容展示](picture/2.png) + +由于香橙派内存一共仅12G,在进行数据集的加载及模型加载之后内存仅剩800多MB,所以后续没有记载相关视频。但第一版的代码可以加载视频,加载后效果是输入关键字,模型会筛选出相关视频的链接,在电脑本机上可以访问到外网YouTube视频。 + +## 3、模型问答 + +主要分为RAG问答和非RAG问答,RAG是检索增强生成,传统AI生成内容只依赖于其在大模型预训练阶段学习到的知识,而RAG在回答之前会先进行检索,对应于本项目: + +首先用户输入问题,然后系统会在后台匹配相关的泰文的图文并且结合上下文,再在对话框返回回答给用户。总的来说就是它会结合用户的历史问题进行回答。 + +以下就是进行了RAG回答和非RAG回答的结果对比图: + +### RAG回答 + +问:ต้มยำคืออะไร(冬阴功是什么?) + +答:ต้มยำคืออาหารไทยที่มีรสเปรี้ยวและเผ็ด มักใช้น้ำมะนาวหรือพริกเพื่อปรับรสให้กลมกล่อม(冬阴功是泰国的一道酸辣风味菜肴,通常会用青柠汁或辣椒来调和风味,让口感更平衡柔和。)如图(picture/3.jpg) + +![RAG回答示例1](picture/3.jpg) + +继续问:มันมีประโยชน์อะไรต่อร่างกายหรือเปล่า(它对身体有什么好处吗?) + +回答:ผัดมักกะโรนีกุ้งมักให้พลังงานสูงและสารอาหารครบ...(冬阴功通常能提供较高的能量和全面的营养,有助于增强饱腹感、减少零食摄入。)如图(picture/4.jpg) + +![RAG回答示例2](picture/4.jpg) + +继续问:วิธีทำอาหารจานนี้คืออะไร(这道菜的做法是什么?) + +答:ไอเดียคือทำต้มยำกุ้งแบบน้ำใส ใส่ปลาและผักเยอะๆ เสิร์ฟกับข้าวสวย.(建议的做法是做清汤冬阴功虾,放入鱼和大量蔬菜,搭配白米饭食用。)(picture/5.jpg) + +![RAG回答示例3](picture/5.jpg) + +**以上就是RAG回答:模型会根据用户所提过的问题进行检索并返回回答** + +### 非RAG回答 + +选择非RAG回答时,系统不用根据用户所提过的问题进行回答,因此回答时会牛头不对马嘴 + +例如(picture/6.jpg): + +问:อาหารไทย(泰国菜 / 泰国料理) + +答:รับทราบแล้ว จำได้ว่าเป็นเมนูส้มตำ(收到了,我记下来这道菜是青木瓜沙拉(宋丹)。) + +![非RAG回答示例1](picture/6.jpg) + +当继续问:แนะนำมันหน่อย(介绍一下它吧) + +答:โอเค งั้นขอลองเป็นแอปชื่อ "Find My Car" นะ...(好的,那我们来试试 "Find My Car" 这个应用吧。)—— 回答的根本不是所提问题,答非所问(picture/7.jpg) + +![非RAG回答示例2](picture/7.jpg) + +继续问:แนะนำอาหารที่คุณเพิ่งตอบกลับมาให้หน่อย(介绍一下你刚刚回答的那道菜吧。) + +答:ได้เลย ส่งข้อความที่คุยกันมาได้ไหม...(可以的,你把我们之前的聊天记录发给我一下可以吗?)—— 并未回答上文的问题(picture/8.jpg) + +![非RAG回答示例3](picture/8.jpg) + +## 预期输出 + +模型根据泰语多模态数据集进行RAG检索问答,RAG模式下结合历史上下文返回相关回答并附带图片链接;非RAG模式下直接生成回答,不参考历史记录,减少幻觉产生。 + +以上就是全部运行结果,对于图片显示,由于内存原因,所以在前期数据的加载对图片进行大部分消减,仅有3000条,因此只有问到特定的词才能找到相关图片,并显示图片链接,可在电脑本机访问外网。 diff --git a/Online/community/ThaiMultimodal/convert_to_bin.ipynb b/Online/community/ThaiMultimodal/convert_to_bin.ipynb new file mode 100644 index 0000000..2251fba --- /dev/null +++ b/Online/community/ThaiMultimodal/convert_to_bin.ipynb @@ -0,0 +1,75 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# safetensors 转 pytorch bin(float16)\n", + "\n", + "在 Windows 本地运行,把 safetensors 转成 pytorch bin 格式,同时转换为 float16。\n", + "\n", + "**前置依赖:**\n", + "```\n", + "pip install safetensors torch\n", + "```" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "!pip install safetensors torch" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "from safetensors.torch import load_file\n", + "import torch\n", + "import os\n", + "import json\n", + "\n", + "model_dir = './wongwian-micro-instruct'\n", + "safetensors_path = os.path.join(model_dir, 'model.safetensors')\n", + "bin_path = os.path.join(model_dir, 'pytorch_model.bin')\n", + "config_path = os.path.join(model_dir, 'config.json')\n", + "\n", + "print(f'读取: {safetensors_path}')\n", + "state_dict = load_file(safetensors_path)\n", + "\n", + "print('转换权重为 float16...')\n", + "state_dict_fp16 = {k: v.to(torch.float16) for k, v in state_dict.items()}\n", + "\n", + "print(f'写入: {bin_path}')\n", + "torch.save(state_dict_fp16, bin_path)\n", + "\n", + "with open(config_path, 'r') as f:\n", + " config = json.load(f)\n", + "config['torch_dtype'] = 'float16'\n", + "with open(config_path, 'w') as f:\n", + " json.dump(config, f, indent=2)\n", + "print('✓ config.json 已更新为 float16')\n", + "\n", + "print(f'文件大小: {os.path.getsize(bin_path) / 1024 / 1024:.1f} MB')" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python", + "version": "3.9.0" + } + }, + "nbformat": 4, + "nbformat_minor": 4 +} diff --git a/Online/community/ThaiMultimodal/picture/0_1.jpg b/Online/community/ThaiMultimodal/picture/0_1.jpg new file mode 100644 index 0000000..80c6da4 Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/0_1.jpg differ diff --git a/Online/community/ThaiMultimodal/picture/0_2.jpg b/Online/community/ThaiMultimodal/picture/0_2.jpg new file mode 100644 index 0000000..5f0920a Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/0_2.jpg differ diff --git a/Online/community/ThaiMultimodal/picture/1.png b/Online/community/ThaiMultimodal/picture/1.png new file mode 100644 index 0000000..3342e78 Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/1.png differ diff --git a/Online/community/ThaiMultimodal/picture/2.png b/Online/community/ThaiMultimodal/picture/2.png new file mode 100644 index 0000000..7514b12 Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/2.png differ diff --git a/Online/community/ThaiMultimodal/picture/3.jpg b/Online/community/ThaiMultimodal/picture/3.jpg new file mode 100644 index 0000000..552f268 Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/3.jpg differ diff --git a/Online/community/ThaiMultimodal/picture/4.jpg b/Online/community/ThaiMultimodal/picture/4.jpg new file mode 100644 index 0000000..8cb3fd5 Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/4.jpg differ diff --git a/Online/community/ThaiMultimodal/picture/5.jpg b/Online/community/ThaiMultimodal/picture/5.jpg new file mode 100644 index 0000000..c2f6381 Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/5.jpg differ diff --git a/Online/community/ThaiMultimodal/picture/6.jpg b/Online/community/ThaiMultimodal/picture/6.jpg new file mode 100644 index 0000000..011ca67 Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/6.jpg differ diff --git a/Online/community/ThaiMultimodal/picture/7.jpg b/Online/community/ThaiMultimodal/picture/7.jpg new file mode 100644 index 0000000..4e2127c Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/7.jpg differ diff --git a/Online/community/ThaiMultimodal/picture/8.jpg b/Online/community/ThaiMultimodal/picture/8.jpg new file mode 100644 index 0000000..99e29fe Binary files /dev/null and b/Online/community/ThaiMultimodal/picture/8.jpg differ diff --git a/Online/community/ThaiMultimodal/thai_multimodal_explorer.ipynb b/Online/community/ThaiMultimodal/thai_multimodal_explorer.ipynb new file mode 100644 index 0000000..3ee9d1c --- /dev/null +++ b/Online/community/ThaiMultimodal/thai_multimodal_explorer.ipynb @@ -0,0 +1,1292 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# 泰国文化多模态探索平台\n", + "\n", + "基于「万卷·丝路」泰语多模态数据集,在昇思MindSpore框架 + 香橙派AIpro开发板上运行。\n", + "\n", + "**数据来源:**\n", + "- SFT数据:泰语文化/生活/数学/代码指令问答\n", + "- 图文数据:泰国文化图片 + 泰语描述\n", + "- 视频数据:泰语视频元数据,含字幕/摘要/标签\n", + "\n", + "**功能模块:**\n", + "1. **多模态RAG问答** — 融合SFT文本知识库 + 图文caption,回答时附带相关图片\n", + "2. **图文文化浏览** — 按标签分类浏览泰国文化图片,展示泰语描述\n", + "3. **视频内容检索** — 关键词/标签检索泰语视频,展示摘要和跳转链接\n", + "\n", + "**运行环境:** 香橙派AIpro 12G | CANN 8.1RC1 | MindSpore 2.5.0 | MindNLP 0.4.1 | Python 3.9" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 环境准备" + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "metadata": { + "scrolled": true + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Defaulting to user installation because normal site-packages is not writeable\n", + "Requirement already satisfied: mindnlp==0.4.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (0.4.1)\n", + "Requirement already satisfied: safetensors in /usr/local/miniconda3/lib/python3.9/site-packages (from mindnlp==0.4.1) (0.4.5)\n", + "Requirement already satisfied: mindspore>=2.2.14 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (2.5.0)\n", + "Requirement already satisfied: requests in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (2.32.3)\n", + "Requirement already satisfied: evaluate in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (0.4.3)\n", + "Requirement already satisfied: regex in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (2024.9.11)\n", + "Requirement already satisfied: sentencepiece in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (0.2.0)\n", + "Requirement already satisfied: ml-dtypes in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (0.5.0)\n", + "Requirement already satisfied: datasets in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (3.0.1)\n", + "Requirement already satisfied: pytest==7.2.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (7.2.0)\n", + "Requirement already satisfied: tokenizers==0.19.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (0.19.1)\n", + "Requirement already satisfied: pillow>=10.0.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from mindnlp==0.4.1) (10.2.0)\n", + "Requirement already satisfied: addict in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (2.4.0)\n", + "Requirement already satisfied: tqdm in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (4.66.5)\n", + "Requirement already satisfied: pyctcdecode in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindnlp==0.4.1) (0.5.0)\n", + "Requirement already satisfied: exceptiongroup>=1.0.0rc8 in /usr/local/miniconda3/lib/python3.9/site-packages (from pytest==7.2.0->mindnlp==0.4.1) (1.2.0)\n", + "Requirement already satisfied: packaging in /usr/local/miniconda3/lib/python3.9/site-packages (from pytest==7.2.0->mindnlp==0.4.1) (23.1)\n", + "Requirement already satisfied: tomli>=1.0.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from pytest==7.2.0->mindnlp==0.4.1) (2.0.1)\n", + "Requirement already satisfied: pluggy<2.0,>=0.12 in /usr/local/miniconda3/lib/python3.9/site-packages (from pytest==7.2.0->mindnlp==0.4.1) (1.0.0)\n", + "Requirement already satisfied: attrs>=19.2.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from pytest==7.2.0->mindnlp==0.4.1) (23.2.0)\n", + "Requirement already satisfied: iniconfig in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from pytest==7.2.0->mindnlp==0.4.1) (2.0.0)\n", + "Requirement already satisfied: huggingface-hub<1.0,>=0.16.4 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from tokenizers==0.19.1->mindnlp==0.4.1) (0.25.2)\n", + "Requirement already satisfied: astunparse>=1.6.3 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindspore>=2.2.14->mindnlp==0.4.1) (1.6.3)\n", + "Requirement already satisfied: protobuf>=3.13.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from mindspore>=2.2.14->mindnlp==0.4.1) (3.20.0)\n", + "Requirement already satisfied: numpy<2.0.0,>=1.20.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from mindspore>=2.2.14->mindnlp==0.4.1) (1.22.4)\n", + "Requirement already satisfied: psutil>=5.6.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from mindspore>=2.2.14->mindnlp==0.4.1) (5.9.8)\n", + "Requirement already satisfied: asttokens>=2.0.4 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindspore>=2.2.14->mindnlp==0.4.1) (2.4.1)\n", + "Requirement already satisfied: scipy>=1.5.4 in /usr/local/miniconda3/lib/python3.9/site-packages (from mindspore>=2.2.14->mindnlp==0.4.1) (1.12.0)\n", + "Requirement already satisfied: dill>=0.3.7 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from mindspore>=2.2.14->mindnlp==0.4.1) (0.3.8)\n", + "Requirement already satisfied: pyyaml>=5.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (6.0.1)\n", + "Requirement already satisfied: filelock in /usr/local/miniconda3/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (3.13.1)\n", + "Requirement already satisfied: xxhash in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (3.5.0)\n", + "Requirement already satisfied: pyarrow>=15.0.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (17.0.0)\n", + "Requirement already satisfied: pandas in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (2.2.3)\n", + "Requirement already satisfied: fsspec[http]<=2024.6.1,>=2023.1.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (2023.12.2)\n", + "Requirement already satisfied: aiohttp in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (3.10.10)\n", + "Requirement already satisfied: multiprocess in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from datasets->mindnlp==0.4.1) (0.70.16)\n", + "Requirement already satisfied: charset-normalizer<4,>=2 in /usr/local/miniconda3/lib/python3.9/site-packages (from requests->mindnlp==0.4.1) (2.0.4)\n", + "Requirement already satisfied: certifi>=2017.4.17 in /usr/local/miniconda3/lib/python3.9/site-packages (from requests->mindnlp==0.4.1) (2023.11.17)\n", + "Requirement already satisfied: idna<4,>=2.5 in /usr/local/miniconda3/lib/python3.9/site-packages (from requests->mindnlp==0.4.1) (3.4)\n", + "Requirement already satisfied: urllib3<3,>=1.21.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from requests->mindnlp==0.4.1) (2.2.3)\n", + "Requirement already satisfied: pygtrie<3.0,>=2.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from pyctcdecode->mindnlp==0.4.1) (2.5.0)\n", + "Requirement already satisfied: hypothesis<7,>=6.14 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from pyctcdecode->mindnlp==0.4.1) (6.115.2)\n", + "Requirement already satisfied: six>=1.12.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from asttokens>=2.0.4->mindspore>=2.2.14->mindnlp==0.4.1) (1.16.0)\n", + "Requirement already satisfied: wheel<1.0,>=0.23.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from astunparse>=1.6.3->mindspore>=2.2.14->mindnlp==0.4.1) (0.37.1)\n", + "Requirement already satisfied: async-timeout<5.0,>=4.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from aiohttp->datasets->mindnlp==0.4.1) (4.0.3)\n", + "Requirement already satisfied: frozenlist>=1.1.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from aiohttp->datasets->mindnlp==0.4.1) (1.4.1)\n", + "Requirement already satisfied: multidict<7.0,>=4.5 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from aiohttp->datasets->mindnlp==0.4.1) (6.1.0)\n", + "Requirement already satisfied: aiosignal>=1.1.2 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from aiohttp->datasets->mindnlp==0.4.1) (1.3.1)\n", + "Requirement already satisfied: aiohappyeyeballs>=2.3.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from aiohttp->datasets->mindnlp==0.4.1) (2.4.3)\n", + "Requirement already satisfied: yarl<2.0,>=1.12.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from aiohttp->datasets->mindnlp==0.4.1) (1.15.3)\n", + "Requirement already satisfied: typing-extensions>=3.7.4.3 in /usr/local/miniconda3/lib/python3.9/site-packages (from huggingface-hub<1.0,>=0.16.4->tokenizers==0.19.1->mindnlp==0.4.1) (4.9.0)\n", + "Requirement already satisfied: sortedcontainers<3.0.0,>=2.1.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from hypothesis<7,>=6.14->pyctcdecode->mindnlp==0.4.1) (2.4.0)\n", + "Requirement already satisfied: tzdata>=2022.7 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from pandas->datasets->mindnlp==0.4.1) (2024.2)\n", + "Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/miniconda3/lib/python3.9/site-packages (from pandas->datasets->mindnlp==0.4.1) (2.8.2)\n", + "Requirement already satisfied: pytz>=2020.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from pandas->datasets->mindnlp==0.4.1) (2023.3.post1)\n", + "Requirement already satisfied: propcache>=0.2.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from yarl<2.0,>=1.12.0->aiohttp->datasets->mindnlp==0.4.1) (0.2.0)\n", + "Defaulting to user installation because normal site-packages is not writeable\n", + "Requirement already satisfied: gradio==4.44.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (4.44.0)\n", + "Requirement already satisfied: markupsafe~=2.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (2.1.4)\n", + "Requirement already satisfied: aiofiles<24.0,>=22.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (23.2.1)\n", + "Requirement already satisfied: pyyaml<7.0,>=5.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (6.0.1)\n", + "Requirement already satisfied: importlib-resources<7.0,>=1.3 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (6.1.1)\n", + "Requirement already satisfied: python-multipart>=0.0.9 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.0.12)\n", + "Requirement already satisfied: uvicorn>=0.14.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.32.0)\n", + "Requirement already satisfied: fastapi<1.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.115.2)\n", + "Requirement already satisfied: typing-extensions~=4.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (4.9.0)\n", + "Requirement already satisfied: anyio<5.0,>=3.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (4.2.0)\n", + "Requirement already satisfied: pandas<3.0,>=1.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (2.2.3)\n", + "Requirement already satisfied: matplotlib~=3.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (3.8.2)\n", + "Requirement already satisfied: numpy<3.0,>=1.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (1.22.4)\n", + "Requirement already satisfied: packaging in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (23.1)\n", + "Requirement already satisfied: pillow<11.0,>=8.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (10.2.0)\n", + "Requirement already satisfied: pydantic>=2.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (2.5.3)\n", + "Requirement already satisfied: semantic-version~=2.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (2.10.0)\n", + "Requirement already satisfied: httpx>=0.24.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.27.2)\n", + "Requirement already satisfied: jinja2<4.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio==4.44.0) (3.1.3)\n", + "Requirement already satisfied: tomlkit==0.12.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.12.0)\n", + "Requirement already satisfied: ffmpy in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.4.0)\n", + "Requirement already satisfied: gradio-client==1.3.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (1.3.0)\n", + "Requirement already satisfied: orjson~=3.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (3.10.7)\n", + "Requirement already satisfied: urllib3~=2.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (2.2.3)\n", + "Requirement already satisfied: ruff>=0.2.2 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.6.9)\n", + "Requirement already satisfied: typer<1.0,>=0.12 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.12.5)\n", + "Requirement already satisfied: huggingface-hub>=0.19.3 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.25.2)\n", + "Requirement already satisfied: pydub in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio==4.44.0) (0.25.1)\n", + "Requirement already satisfied: websockets<13.0,>=10.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from gradio-client==1.3.0->gradio==4.44.0) (12.0)\n", + "Requirement already satisfied: fsspec in /usr/local/miniconda3/lib/python3.9/site-packages (from gradio-client==1.3.0->gradio==4.44.0) (2023.12.2)\n", + "Requirement already satisfied: exceptiongroup>=1.0.2 in /usr/local/miniconda3/lib/python3.9/site-packages (from anyio<5.0,>=3.0->gradio==4.44.0) (1.2.0)\n", + "Requirement already satisfied: sniffio>=1.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from anyio<5.0,>=3.0->gradio==4.44.0) (1.3.0)\n", + "Requirement already satisfied: idna>=2.8 in /usr/local/miniconda3/lib/python3.9/site-packages (from anyio<5.0,>=3.0->gradio==4.44.0) (3.4)\n", + "Requirement already satisfied: starlette<0.41.0,>=0.37.2 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from fastapi<1.0->gradio==4.44.0) (0.40.0)\n", + "Requirement already satisfied: certifi in /usr/local/miniconda3/lib/python3.9/site-packages (from httpx>=0.24.1->gradio==4.44.0) (2023.11.17)\n", + "Requirement already satisfied: httpcore==1.* in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from httpx>=0.24.1->gradio==4.44.0) (1.0.6)\n", + "Requirement already satisfied: h11<0.15,>=0.13 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from httpcore==1.*->httpx>=0.24.1->gradio==4.44.0) (0.14.0)\n", + "Requirement already satisfied: filelock in /usr/local/miniconda3/lib/python3.9/site-packages (from huggingface-hub>=0.19.3->gradio==4.44.0) (3.13.1)\n", + "Requirement already satisfied: tqdm>=4.42.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from huggingface-hub>=0.19.3->gradio==4.44.0) (4.66.5)\n", + "Requirement already satisfied: requests in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from huggingface-hub>=0.19.3->gradio==4.44.0) (2.32.3)\n", + "Requirement already satisfied: zipp>=3.1.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from importlib-resources<7.0,>=1.3->gradio==4.44.0) (3.17.0)\n", + "Requirement already satisfied: cycler>=0.10 in /usr/local/miniconda3/lib/python3.9/site-packages (from matplotlib~=3.0->gradio==4.44.0) (0.12.1)\n", + "Requirement already satisfied: python-dateutil>=2.7 in /usr/local/miniconda3/lib/python3.9/site-packages (from matplotlib~=3.0->gradio==4.44.0) (2.8.2)\n", + "Requirement already satisfied: pyparsing>=2.3.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from matplotlib~=3.0->gradio==4.44.0) (3.1.1)\n", + "Requirement already satisfied: contourpy>=1.0.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from matplotlib~=3.0->gradio==4.44.0) (1.2.0)\n", + "Requirement already satisfied: kiwisolver>=1.3.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from matplotlib~=3.0->gradio==4.44.0) (1.4.5)\n", + "Requirement already satisfied: fonttools>=4.22.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from matplotlib~=3.0->gradio==4.44.0) (4.47.2)\n", + "Requirement already satisfied: pytz>=2020.1 in /usr/local/miniconda3/lib/python3.9/site-packages (from pandas<3.0,>=1.0->gradio==4.44.0) (2023.3.post1)\n", + "Requirement already satisfied: tzdata>=2022.7 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from pandas<3.0,>=1.0->gradio==4.44.0) (2024.2)\n", + "Requirement already satisfied: pydantic-core==2.14.6 in /usr/local/miniconda3/lib/python3.9/site-packages (from pydantic>=2.0->gradio==4.44.0) (2.14.6)\n", + "Requirement already satisfied: annotated-types>=0.4.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from pydantic>=2.0->gradio==4.44.0) (0.6.0)\n", + "Requirement already satisfied: click>=8.0.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from typer<1.0,>=0.12->gradio==4.44.0) (8.1.7)\n", + "Requirement already satisfied: shellingham>=1.3.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from typer<1.0,>=0.12->gradio==4.44.0) (1.5.4)\n", + "Requirement already satisfied: rich>=10.11.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from typer<1.0,>=0.12->gradio==4.44.0) (13.9.2)\n", + "Requirement already satisfied: six>=1.5 in /usr/local/miniconda3/lib/python3.9/site-packages (from python-dateutil>=2.7->matplotlib~=3.0->gradio==4.44.0) (1.16.0)\n", + "Requirement already satisfied: markdown-it-py>=2.2.0 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from rich>=10.11.0->typer<1.0,>=0.12->gradio==4.44.0) (3.0.0)\n", + "Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from rich>=10.11.0->typer<1.0,>=0.12->gradio==4.44.0) (2.17.2)\n", + "Requirement already satisfied: charset-normalizer<4,>=2 in /usr/local/miniconda3/lib/python3.9/site-packages (from requests->huggingface-hub>=0.19.3->gradio==4.44.0) (2.0.4)\n", + "Requirement already satisfied: mdurl~=0.1 in /home/HwHiAiUser/.local/lib/python3.9/site-packages (from markdown-it-py>=2.2.0->rich>=10.11.0->typer<1.0,>=0.12->gradio==4.44.0) (0.1.2)\n", + "Defaulting to user installation because normal site-packages is not writeable\n", + "Requirement already satisfied: scikit-learn in /usr/local/miniconda3/lib/python3.9/site-packages (1.4.0)\n", + "Requirement already satisfied: threadpoolctl>=2.0.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from scikit-learn) (3.2.0)\n", + "Requirement already satisfied: numpy<2.0,>=1.19.5 in /usr/local/miniconda3/lib/python3.9/site-packages (from scikit-learn) (1.22.4)\n", + "Requirement already satisfied: joblib>=1.2.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from scikit-learn) (1.3.2)\n", + "Requirement already satisfied: scipy>=1.6.0 in /usr/local/miniconda3/lib/python3.9/site-packages (from scikit-learn) (1.12.0)\n" + ] + } + ], + "source": [ + "!pip install mindnlp==0.4.1\n", + "!pip install gradio==4.44.0\n", + "!pip install scikit-learn" + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "metadata": {}, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "[WARNING] ME(9270:246290119323680,MainProcess):2026-05-06-16:12:20.605.30 [mindspore/run_check/_check_version.py:324] MindSpore version 2.5.0 and Ascend AI software package (Ascend Data Center Solution)version 7.7 does not match, the version of software package expect one of ['7.5', '7.6']. Please refer to the match info on: https://www.mindspore.cn/install\n", + "/usr/local/miniconda3/lib/python3.9/site-packages/numpy/core/getlimits.py:499: UserWarning: The value of the smallest subnormal for type is zero.\n", + " setattr(self, word, getattr(machar, word).flat[0])\n", + "/usr/local/miniconda3/lib/python3.9/site-packages/numpy/core/getlimits.py:89: UserWarning: The value of the smallest subnormal for type is zero.\n", + " return self._float_to_str(self.smallest_subnormal)\n", + "/usr/local/miniconda3/lib/python3.9/site-packages/numpy/core/getlimits.py:499: UserWarning: The value of the smallest subnormal for type is zero.\n", + " setattr(self, word, getattr(machar, word).flat[0])\n", + "/usr/local/miniconda3/lib/python3.9/site-packages/numpy/core/getlimits.py:89: UserWarning: The value of the smallest subnormal for type is zero.\n", + " return self._float_to_str(self.smallest_subnormal)\n", + "[WARNING] ME(9270:246290119323680,MainProcess):2026-05-06-16:12:27.565.446 [mindspore/run_check/_check_version.py:342] MindSpore version 2.5.0 and \"te\" wheel package version 7.7 does not match. For details, refer to the installation guidelines: https://www.mindspore.cn/install\n", + "[WARNING] ME(9270:246290119323680,MainProcess):2026-05-06-16:12:27.571.932 [mindspore/run_check/_check_version.py:349] MindSpore version 2.5.0 and \"hccl\" wheel package version 7.7 does not match. For details, refer to the installation guidelines: https://www.mindspore.cn/install\n", + "[WARNING] ME(9270:246290119323680,MainProcess):2026-05-06-16:12:27.574.440 [mindspore/run_check/_check_version.py:363] Please pay attention to the above warning, countdown: 3\n", + "[WARNING] ME(9270:246290119323680,MainProcess):2026-05-06-16:12:28.577.296 [mindspore/run_check/_check_version.py:363] Please pay attention to the above warning, countdown: 2\n", + "[WARNING] ME(9270:246290119323680,MainProcess):2026-05-06-16:12:29.579.736 [mindspore/run_check/_check_version.py:363] Please pay attention to the above warning, countdown: 1\n", + "[WARNING] ME(9270:246290119323680,MainProcess):2026-05-06-16:12:35.917.731 [mindspore/context.py:1335] For 'context.set_context', the parameter 'ascend_config' will be deprecated and removed in a future version. Please use the api mindspore.device_context.ascend.op_precision.precision_mode(),\n", + " mindspore.device_context.ascend.op_precision.op_precision_mode(),\n", + " mindspore.device_context.ascend.op_precision.matmul_allow_hf32(),\n", + " mindspore.device_context.ascend.op_precision.conv_allow_hf32(),\n", + " mindspore.device_context.ascend.op_tuning.op_compile() instead.\n", + "Building prefix dict from the default dictionary ...\n", + "Loading model from cache /tmp/jieba.cache\n", + "Loading model cost 2.289 seconds.\n", + "Prefix dict has been built successfully.\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "MindSpore: 2.5.0\n", + "MindNLP: 0.4.1\n", + "Gradio: 4.44.0\n", + "Sklearn: 1.4.0\n", + "NPU: 正常\n" + ] + } + ], + "source": [ + "import mindspore\n", + "import mindnlp\n", + "import sklearn\n", + "import gradio\n", + "import subprocess\n", + "import pkg_resources\n", + "\n", + "print(f'MindSpore: {mindspore.__version__}')\n", + "print(f'MindNLP: {pkg_resources.get_distribution(\"mindnlp\").version}')\n", + "print(f'Gradio: {gradio.__version__}')\n", + "print(f'Sklearn: {sklearn.__version__}')\n", + "\n", + "result = subprocess.run(['npu-smi', 'info'], capture_output=True, text=True)\n", + "print('NPU: 正常' if result.returncode == 0 else 'NPU: 未检测到,使用CPU')" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 数据加载与分析\n", + "\n", + "加载三种模态的泰语数据,分析数据分布。" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "SFT数据总条数: 2000 (限制 2000 条)\n", + "类型分布:\n", + " culture 2000条 ██████████\n", + "内存清理完成\n" + ] + } + ], + "source": [ + "import json\n", + "from collections import Counter\n", + "\n", + "# 只加载 2000 条 SFT 数据,减少内存占用\n", + "MAX_SFT = 2000\n", + "sft_data = []\n", + "with open('data/raw/sft/th/th.jsonl', 'r', encoding='utf-8') as f:\n", + " for i, line in enumerate(f):\n", + " if i >= MAX_SFT:\n", + " break\n", + " line = line.strip()\n", + " if line:\n", + " sft_data.append(json.loads(line))\n", + "\n", + "print(f'SFT数据总条数: {len(sft_data)} (限制 {MAX_SFT} 条)')\n", + "type_counter = Counter(d['type'] for d in sft_data)\n", + "print('类型分布:')\n", + "for t, c in sorted(type_counter.items(), key=lambda x: -x[1]):\n", + " print(f' {t:<12} {c:>5}条 {chr(9608) * (c // 200)}')\n", + "\n", + "import gc\n", + "gc.collect() \n", + "print('内存清理完成')" + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + " SFT数据各类型示例 \n", + "\n", + "[culture]\n", + " 问: ใครคือกษัตริย์ที่ยิ่งใหญ่ที่สุดในประวัติศาสตร์ไทย?\n", + " 答: กษัตริย์ที่ยิ่งใหญ่ที่สุดในประวัติศาสตร์ไทย เช่น รัชกาลที่ 5\n" + ] + } + ], + "source": [ + "# 展示各类型SFT数据示例\n", + "print(' SFT数据各类型示例 ')\n", + "shown = {}\n", + "for d in sft_data:\n", + " t = d['type']\n", + " if t not in shown:\n", + " shown[t] = d\n", + " print(f'\\n[{t}]')\n", + " print(f' 问: {d[\"prompt\"].strip()}')\n", + " print(f' 答: {d[\"completion\"].strip()}')\n", + " if len(shown) == 5:\n", + " break" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "图文数据总条数: 3000 (限制 3000 条)\n", + "图片标签分布(level2):\n", + " 食物 2090张\n", + " 乡村场景 946张\n", + "内存清理完成\n" + ] + } + ], + "source": [ + "# 只加载 3000 条图文数据,减少内存占用\n", + "MAX_IMAGE = 3000\n", + "image_data = []\n", + "with open('data/raw/image/th/th_image_text_pair.jsonl', 'r', encoding='utf-8') as f:\n", + " for i, line in enumerate(f):\n", + " if i >= MAX_IMAGE:\n", + " break\n", + " line = line.strip()\n", + " if line:\n", + " image_data.append(json.loads(line))\n", + "\n", + "print(f'图文数据总条数: {len(image_data)} (限制 {MAX_IMAGE} 条)')\n", + "label_counter = Counter()\n", + "for d in image_data:\n", + " for lv2 in d.get('labels', {}).get('pjwk_cates', {}).get('level2', []):\n", + " label_counter[lv2] += 1\n", + "print('图片标签分布(level2):')\n", + "for label, cnt in label_counter.most_common():\n", + " print(f' {label:<20} {cnt:>6}张')\n", + "\n", + "gc.collect() \n", + "print('内存清理完成')" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "=== 图文数据样本 ===\n", + "[1] 标签: ['食物'] 分辨率: [400, 600]\n", + " 描述: น้ำฝน\n", + " URL: https://www.khaosod.co.th/wpapp/uploads/2023/04/%E0%B8%A7%E0%B8%B1%E0%B8%99%E0%B\n", + "\n", + "[2] 标签: ['食物'] 分辨率: [267, 400]\n", + " 描述: กล้วยบวชชี\n", + " URL: https://food.mthai.com/app/uploads/2012/07/กล้วยบวชชี-3.jpg\n", + "\n", + "[3] 标签: ['食物'] 分辨率: [494, 800]\n", + " 描述: ตานุช ข้าวแกงปักษ์ใต้เสน่ห์ปลายจวักรสดั้งเดิม\n", + " URL: https://www.khaosod.co.th/wpapp/uploads/2020/03/d1-1.jpg\n", + "\n", + "[4] 标签: ['食物'] 分辨率: [332, 500]\n", + " 描述: รวมแบบกระทงต่างๆ กระทงกาบกล้วย กระทงดอกบัว กระทงกะลา\n", + " URL: http://scoop.mthai.com/app/uploads/2014/11/kratong-Lotus3.jpg\n", + "\n", + "[5] 标签: ['食物'] 分辨率: [392, 696]\n", + " 描述: เปิดใจ\n", + " URL: https://www.khaosod.co.th/wpapp/uploads/2023/06/20230620_163804-696x392.jpg\n", + "\n" + ] + } + ], + "source": [ + "# 展示前5条图文数据\n", + "print('=== 图文数据样本 ===')\n", + "for i, d in enumerate(image_data[:5]):\n", + " caption = d.get('captions', {}).get('content', '')\n", + " img_url = d.get('image', {}).get('path', '')\n", + " lv2 = d.get('labels', {}).get('pjwk_cates', {}).get('level2', [])\n", + " res = d.get('image', {}).get('resolution', [])\n", + " print(f'[{i+1}] 标签: {lv2} 分辨率: {res}')\n", + " print(f' 描述: {caption}')\n", + " print(f' URL: {img_url[:80]}')\n", + " print()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "image_caption 数据条数: 80035\n", + "\n", + "样本结构如下:\n", + "{\"img_id\": \"00177b4c605fa0f20b68b342cf492aef53b8f63005217dd7febf2588aebd2641\", \"image\": {\"path\": \"https://www.khaosod.co.th/wpapp/uploads/2021/05/388230-696x460.jpg\", \"resolution\": [696, 460], \"size\":\n", + "\n", + "{\"img_id\": \"001eba5de35c4d090d7dfb0a8adb2ab986a1d63db67babbe9e6b6a8a5640c4e5\", \"image\": {\"path\": \"https://www.m2fnews.com/media/content/2022/03/27/372654.jpg\", \"resolution\": [1500, 999], \"size\": 159.6\n", + "\n", + "{\"img_id\": \"002a95fd97d4ab30957d76a4df1593a5b85556be7f5d1a361650313a03a70d59\", \"image\": {\"path\": \"https://www.khaosod.co.th/wpapp/uploads/2023/01/%E0%B8%9B%E0%B8%A5%E0%B8%81%E0%B8%82%E0%B8%B2%E0%B8%97\n", + "\n" + ] + } + ], + "source": [ + "# 加载 th_image_caption.jsonl,图片+详细caption\n", + "caption_data = []\n", + "with open('data/raw/image/th/th_image_caption.jsonl', 'r', encoding='utf-8') as f:\n", + " for line in f:\n", + " line = line.strip()\n", + " if line:\n", + " try:\n", + " caption_data.append(json.loads(line))\n", + " except Exception:\n", + " continue\n", + "\n", + "print(f'image_caption 数据条数: {len(caption_data)}')\n", + "# 展示前3条结构\n", + "print('\\n样本结构如下:')\n", + "for d in caption_data[:3]:\n", + " print(json.dumps(d, ensure_ascii=False)[:200])\n", + " print()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "视频去重后总数: 1000(限制 1000 条)\n", + "视频标签分布(level2):\n", + " people 203个\n", + " unknown 173个\n", + " movie & animation 114个\n", + " howto 80个\n", + " interview 78个\n", + " game 67个\n", + " sports 57个\n", + " technology & military 55个\n", + " news 42个\n", + " music 39个\n", + "内存清理完成\n" + ] + } + ], + "source": [ + "MAX_VIDEOS = 1000 # 减少内存占用\n", + "\n", + "video_data = {}\n", + "with open('data/raw/video/th/multilingual_thai_meta_out.jsonl', 'r', encoding='utf-8') as f:\n", + " for line in f:\n", + " if len(video_data) >= MAX_VIDEOS:\n", + " break\n", + " line = line.strip()\n", + " if not line:\n", + " continue\n", + " try:\n", + " obj = json.loads(line)\n", + " except Exception:\n", + " continue\n", + " url = obj.get('url', '')\n", + " if not url or url in video_data:\n", + " continue\n", + " summary = ''\n", + " for cap in obj.get('caption', []):\n", + " if cap.get('type') == 'summary':\n", + " summary = cap.get('content', '')\n", + " break\n", + " labels = obj.get('labels', {}).get('pjwk_cates', {})\n", + " video_data[url] = {\n", + " 'url': url,\n", + " 'summary': summary,\n", + " 'level1': labels.get('level1', []),\n", + " 'level2': labels.get('level2', []),\n", + " }\n", + "\n", + "video_list = list(video_data.values())\n", + "print(f'视频去重后总数: {len(video_list)}(限制 {MAX_VIDEOS} 条)')\n", + "v_label_counter = Counter()\n", + "for v in video_list:\n", + " for lv2 in v['level2']:\n", + " v_label_counter[lv2] += 1\n", + "print('视频标签分布(level2):')\n", + "for label, cnt in v_label_counter.most_common(10):\n", + " print(f' {label:<20} {cnt:>5}个')\n", + "\n", + "del video_data \n", + "gc.collect() \n", + "print('内存清理完成')" + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "视频总数: 1000\n", + "\n", + "一级标签分布:\n", + " people 306个 ██████\n", + " scenary 285个 █████\n", + " general 236个 ████\n", + " unknown 173个 ███\n", + "\n", + "二级标签分布(全部):\n", + " people 203个 ████\n", + " unknown 173个 ███\n", + " movie & animation 114个 ██\n", + " howto 80个 █\n", + " interview 78个 █\n", + " game 67个 █\n", + " sports 57个 █\n", + " technology & military 55个 █\n", + " news 42个 \n", + " music 39个 \n", + " voyage 35个 \n", + " religion 32个 \n", + " animal 25个 \n", + "URL: https://www.youtube.com/watch?v=--9OlObKEo8\n", + "标签: ['people'] / ['animal']\n", + "摘要: ...\n", + "\n", + "URL: https://www.youtube.com/watch?v=--_xi0fogSk\n", + "标签: ['unknown'] / ['unknown']\n", + "摘要: ...\n", + "\n", + "URL: https://www.youtube.com/watch?v=--tAKJhnpfA\n", + "标签: ['scenary'] / ['game']\n", + "摘要: ...\n", + "\n" + ] + } + ], + "source": [ + "print(f'视频总数: {len(video_list)}')\n", + "\n", + "# level1 分布\n", + "lv1_counter = Counter()\n", + "for v in video_list:\n", + " for lv1 in v['level1']:\n", + " lv1_counter[lv1] += 1\n", + "print('\\n一级标签分布:')\n", + "for label, cnt in lv1_counter.most_common():\n", + " bar = chr(9608) * (cnt // 50)\n", + " print(f' {label:<20} {cnt:>5}个 {bar}')\n", + "\n", + "print('\\n二级标签分布(全部):')\n", + "for label, cnt in v_label_counter.most_common():\n", + " bar = chr(9608) * (cnt // 50)\n", + " print(f' {label:<20} {cnt:>5}个 {bar}')\n", + "\n", + "# 展示几条视频样本\n", + "for v in video_list[:3]:\n", + " print(f'URL: {v[\"url\"]}')\n", + " print(f'标签: {v[\"level1\"]} / {v[\"level2\"]}')\n", + " print(f'摘要: {v[\"summary\"][:100]}...')\n", + " print()" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 构建多模态知识库\n", + "\n", + "将SFT文本问答对 + 图文caption合并为统一知识库,使用TF-IDF字符级n-gram向量化,检索时同时返回文字答案和相关图片。" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "SFT 条目: 2000(将向量化)\n", + "图文条目: 3000(不向量化,实时匹配)\n" + ] + } + ], + "source": [ + "from sklearn.feature_extraction.text import TfidfVectorizer\n", + "from sklearn.metrics.pairwise import cosine_similarity\n", + "import numpy as np\n", + "\n", + "# 只对 SFT 数据做向量化\n", + "sft_entries = []\n", + "for d in sft_data:\n", + " sft_entries.append({\n", + " 'type': 'sft',\n", + " 'text': d['prompt'].strip() + ' ' + d['completion'].strip(),\n", + " 'prompt': d['prompt'].strip(),\n", + " 'answer': d['completion'].strip(),\n", + " 'category': d['type'],\n", + " })\n", + "\n", + "# 图文数据只存列表,不向量化,省内存\n", + "image_entries = []\n", + "for d in image_data:\n", + " caption = d.get('captions', {}).get('content', '').strip()\n", + " img_url = d.get('image', {}).get('path', '')\n", + " if caption and img_url:\n", + " image_entries.append({\n", + " 'text': caption,\n", + " 'image_url': img_url,\n", + " })\n", + "\n", + "print(f'SFT 条目: {len(sft_entries)}(将向量化)')\n", + "print(f'图文条目: {len(image_entries)}(不向量化,实时匹配)')" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "SFT 向量化完成: 2000 条,3000 维\n", + "内存清理完成\n" + ] + } + ], + "source": [ + "# 只向量化 SFT 数据\n", + "sft_texts = [e['text'] for e in sft_entries]\n", + "\n", + "sft_vectorizer = TfidfVectorizer(\n", + " analyzer='char',\n", + " ngram_range=(2, 3),\n", + " max_features=3000, # 减少内存占用\n", + " sublinear_tf=True,\n", + ")\n", + "sft_matrix = sft_vectorizer.fit_transform(sft_texts)\n", + "\n", + "print(f'SFT 向量化完成: {sft_matrix.shape[0]} 条,{sft_matrix.shape[1]} 维')\n", + "\n", + "del sft_texts \n", + "gc.collect() \n", + "print('内存清理完成')" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "SFT 检索: 2 条\n", + " score=0.5691 กีฬาประจำชาติของไทยคืออะไร?\n", + " score=0.5691 กีฬาประจำชาติของไทยคืออะไร?\n", + "图文检索: 0 条\n" + ] + } + ], + "source": [ + "def retrieve_sft(query, top_k=3):\n", + " \"\"\"SFT 用 TF-IDF 检索\"\"\"\n", + " qv = sft_vectorizer.transform([query])\n", + " scores = cosine_similarity(qv, sft_matrix).flatten()\n", + " idx_sorted = np.argsort(scores)[::-1]\n", + " results = []\n", + " for idx in idx_sorted[:top_k * 5]:\n", + " if scores[idx] <= 0:\n", + " break\n", + " e = sft_entries[idx]\n", + " results.append({**e, 'score': float(scores[idx])})\n", + " if len(results) >= top_k:\n", + " break\n", + " return results\n", + "\n", + "def retrieve_images(query, top_k=3):\n", + " # 图文用简单文本匹配,向量化,省内存\n", + " query_lower = query.lower()\n", + " results = []\n", + " for e in image_entries:\n", + " if query_lower in e['text'].lower():\n", + " results.append({**e, 'score': 1.0})\n", + " if len(results) >= top_k:\n", + " break\n", + " return results\n", + "\n", + "# 测试\n", + "test_q = 'มวยไทยคืออะไร'\n", + "r_sft = retrieve_sft(test_q, top_k=2)\n", + "r_img = retrieve_images(test_q, top_k=2)\n", + "print(f'SFT 检索: {len(r_sft)} 条')\n", + "for r in r_sft:\n", + " print(f' score={r[\"score\"]:.4f} {r[\"prompt\"][:50]}')\n", + "print(f'图文检索: {len(r_img)} 条')" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 加载wongwian-micro-instruct模型\n", + "\n", + "使用MindNLP加载wongwian-micro-instruct,支持泰语,适合在香橙派开发板上运行。" + ] + }, + { + "cell_type": "code", + "execution_count": 11, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "+--------------------------------------------------------------------------------------------------------+\n", + "| npu-smi 25.2.0 Version: 25.2.0 |\n", + "+-------------------------------+-----------------+------------------------------------------------------+\n", + "| NPU Name | Health | Power(W) Temp(C) Hugepages-Usage(page) |\n", + "| Chip Device | Bus-Id | AICore(%) Memory-Usage(MB) |\n", + "+===============================+=================+======================================================+\n", + "| 0 310B1 | Alarm | 0.0 66 15 / 15 |\n", + "| 0 0 | NA | 0 4572 / 11578 |\n", + "+===============================+=================+======================================================+\n", + "\n" + ] + } + ], + "source": [ + "import subprocess\n", + "result = subprocess.run(['npu-smi', 'info'], capture_output=True, text=True)\n", + "print(result.stdout)" + ] + }, + { + "cell_type": "code", + "execution_count": 12, + "metadata": { + "scrolled": true + }, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "[WARNING] ME(9206:246289758904352,MainProcess):2026-05-07-20:47:41.951.409 [mindspore/run_check/_check_version.py:324] MindSpore version 2.5.0 and Ascend AI software package (Ascend Data Center Solution)version 7.7 does not match, the version of software package expect one of ['7.5', '7.6']. Please refer to the match info on: https://www.mindspore.cn/install\n", + "[WARNING] ME(9206:246289758904352,MainProcess):2026-05-07-20:47:48.349.431 [mindspore/run_check/_check_version.py:342] MindSpore version 2.5.0 and \"te\" wheel package version 7.7 does not match. For details, refer to the installation guidelines: https://www.mindspore.cn/install\n", + "[WARNING] ME(9206:246289758904352,MainProcess):2026-05-07-20:47:48.359.043 [mindspore/run_check/_check_version.py:349] MindSpore version 2.5.0 and \"hccl\" wheel package version 7.7 does not match. For details, refer to the installation guidelines: https://www.mindspore.cn/install\n", + "[WARNING] ME(9206:246289758904352,MainProcess):2026-05-07-20:47:48.360.941 [mindspore/run_check/_check_version.py:363] Please pay attention to the above warning, countdown: 3\n", + "[WARNING] ME(9206:246289758904352,MainProcess):2026-05-07-20:47:49.369.405 [mindspore/run_check/_check_version.py:363] Please pay attention to the above warning, countdown: 2\n", + "[WARNING] ME(9206:246289758904352,MainProcess):2026-05-07-20:47:50.372.135 [mindspore/run_check/_check_version.py:363] Please pay attention to the above warning, countdown: 1\n", + "[WARNING] ME(9206:246289758904352,MainProcess):2026-05-07-20:47:57.752.7 [mindspore/context.py:1335] For 'context.set_context', the parameter 'ascend_config' will be deprecated and removed in a future version. Please use the api mindspore.device_context.ascend.op_precision.precision_mode(),\n", + " mindspore.device_context.ascend.op_precision.op_precision_mode(),\n", + " mindspore.device_context.ascend.op_precision.matmul_allow_hf32(),\n", + " mindspore.device_context.ascend.op_precision.conv_allow_hf32(),\n", + " mindspore.device_context.ascend.op_tuning.op_compile() instead.\n", + "Building prefix dict from the default dictionary ...\n", + "Loading model from cache /tmp/jieba.cache\n", + "Loading model cost 2.266 seconds.\n", + "Prefix dict has been built successfully.\n", + "The tokenizer class you load from this checkpoint is not the same type as the class this function is called from. It may result in unexpected tokenization. \n", + "The tokenizer class you load from this checkpoint is 'TokenizersBackend'. \n", + "The class this function is called from is 'LlamaTokenizer'.\n", + "You are using the default legacy behaviour of the . This is expected, and simply means that the `legacy` (previous) behavior will be used so nothing changes for you. If you want to use the new behaviour, set `legacy=False`. This should only be set if you understand what it means, and thoroughly read the reason why this was added as explained in https://github.com/huggingface/transformers/pull/24565 - if you loaded a llama tokenizer from a GGUF file you can ignore this message\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "加载 tokenizer...\n" + ] + }, + { + "name": "stderr", + "output_type": "stream", + "text": [ + "LlamaForCausalLM has generative capabilities, as `prepare_inputs_for_generation` is explicitly overwritten. However, it doesn't directly inherit from `GenerationMixin`.`PreTrainedModel` will NOT inherit from `GenerationMixin`, and this model will lose the ability to call `generate` and other related functions.\n", + " - If you are the owner of the model architecture code, please modify your model class such that it inherits from `GenerationMixin` (after `PreTrainedModel`, otherwise you'll get an exception).\n", + " - If you are not the owner of the model architecture class, please contact the model code owner to update it.\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "✓ tokenizer 加载完成\n", + "\n", + "加载模型...\n" + ] + }, + { + "name": "stderr", + "output_type": "stream", + "text": [ + "[WARNING] DEVICE(9206,dfffcd976020,python):2026-05-07-20:48:11.714.887 [mindspore/ccsrc/plugin/device/ascend/hal/device/ascend_memory_adapter.cc:118] Initialize] Free memory size is less than half of total memory size.Device 0 Device MOC total size:12140449792 Device MOC free size:5539405824 may be other processes occupying this card, check as: ps -ef|grep python\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "✓ 模型加载完成\n" + ] + } + ], + "source": [ + "import mindspore\n", + "from mindnlp.transformers import LlamaForCausalLM, LlamaTokenizer\n", + "import gc\n", + "from mindspore._c_expression import disable_multi_thread\n", + "\n", + "disable_multi_thread()\n", + "gc.collect()\n", + "\n", + "model_path = './models/wongwian-micro-instruct'\n", + "\n", + "print('加载 tokenizer...')\n", + "tokenizer = LlamaTokenizer.from_pretrained(model_path)\n", + "tokenizer.pad_token = tokenizer.unk_token\n", + "print('✓ tokenizer 加载完成')\n", + "\n", + "print('\\n加载模型...')\n", + "model = LlamaForCausalLM.from_pretrained(\n", + " model_path,\n", + " ms_dtype=mindspore.float16,\n", + " low_cpu_mem_usage=True,\n", + " use_safetensors=False\n", + ")\n", + "model.set_train(False)\n", + "print('✓ 模型加载完成')" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "输入形状: (1, 8)\n", + "开始生成...\n", + "..\n", + "✓ 推理完成!耗时: 67.56秒\n", + "问题: มวยไทยคืออะไร\n", + "回答: มวยไทยเป็นศิลปะการป้องกันตัวที่ใช้ร่างกายทั้งตัวและศีรษะปกป้องร่างกายจากอันตรายหรือการบาดเจ็บ โดยมักใช้ศอก เข่า เท้า และท่าศอก-เข่า\n" + ] + } + ], + "source": [ + "# 测试推理(泰语)\n", + "import time\n", + "import gc\n", + "\n", + "gc.collect()\n", + "question = 'มวยไทยคืออะไร' # 泰拳是什么\n", + "prompt = f'User: {question}\\nAssistant:'\n", + "\n", + "inputs = tokenizer(prompt, return_tensors='ms')\n", + "input_ids = inputs.input_ids\n", + "\n", + "print(f'输入形状: {input_ids.shape}')\n", + "print('开始生成...')\n", + "start = time.time()\n", + "\n", + "output = model.generate(\n", + " input_ids,\n", + " max_new_tokens=64,\n", + " do_sample=False,\n", + " num_beams=1,\n", + " eos_token_id=tokenizer.eos_token_id\n", + ")\n", + "\n", + "elapsed = time.time() - start\n", + "new_tokens = output[0][input_ids.shape[1]:].tolist()\n", + "response = tokenizer.decode(new_tokens, skip_special_tokens=True, clean_up_tokenization_spaces=True)\n", + "\n", + "print(f'\\n✓ 推理完成!耗时: {elapsed:.2f}秒')\n", + "print(f'问题: {question}')\n", + "print(f'回答: {response}')\n" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 多模态RAG推理\n", + "\n", + "RAG流程:用户提问 → 检索SFT知识库+ 图文知识库→ 注入Prompt → 模型生成回答 → 界面同时展示文字答案和相关图片。" + ] + }, + { + "cell_type": "code", + "execution_count": 13, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "RAG推理函数定义完成\n" + ] + } + ], + "source": [ + "from mindnlp.transformers import TextIteratorStreamer\n", + "from threading import Thread\n", + "import gc\n", + "\n", + "def rag_stream(query, history, use_rag=True):\n", + " retrieved_text = retrieve_sft(query, top_k=1) if use_rag else []\n", + " retrieved_img = retrieve_images(query, top_k=2) if use_rag else []\n", + "\n", + " # 构建 context\n", + " context = ''\n", + " if retrieved_text:\n", + " r = retrieved_text[0]\n", + " context = f\"ความรู้อ้างอิง:\\n{r['prompt'][:100]}\\n{r['answer'][:150]}\\n\\n\"\n", + "\n", + " # 按照 wongwian 的 chat template 格式手动构建 prompt\n", + " prompt = f\"{context}User: {query}\\nAssistant:\"\n", + "\n", + " gc.collect() # 推理前清理内存\n", + "\n", + " inputs = tokenizer(prompt, return_tensors='ms')\n", + " input_ids = inputs.input_ids\n", + " \n", + " streamer = TextIteratorStreamer(\n", + " tokenizer, timeout=120, skip_prompt=True, skip_special_tokens=True\n", + " )\n", + " t = Thread(target=model.generate, kwargs=dict(\n", + " input_ids=input_ids, streamer=streamer,\n", + " max_new_tokens=64,\n", + " do_sample=False,\n", + " num_beams=1,\n", + " eos_token_id=tokenizer.eos_token_id\n", + " ))\n", + " t.start()\n", + "\n", + " partial = ''\n", + " for token in streamer:\n", + " partial += token\n", + " yield partial.strip(), retrieved_img\n", + " \n", + " # 在回答后面添加图片链接或提示\n", + " if retrieved_img:\n", + " partial += '\\n\\n相关图片链接:'\n", + " for i, img in enumerate(retrieved_img[:3], 1):\n", + " url = img.get('url', '') if isinstance(img, dict) else img\n", + " partial += f'\\n图片{i}: {url}'\n", + " else:\n", + " partial += '\\n\\n未找到相关图片'\n", + " \n", + " yield partial.strip(), retrieved_img\n", + " \n", + " # 推理完成后清理\n", + " del inputs, input_ids\n", + " gc.collect()\n", + "\n", + "print('RAG推理函数定义完成')" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 视频检索模块\n", + "\n", + "对视频元数据构建TF-IDF索引,支持关键词检索和标签筛选,返回视频摘要和YouTube链接。" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "视频检索测试: 0 条\n" + ] + } + ], + "source": [ + "def search_video(query, label_filter='', top_k=5):\n", + " # 视频用关键词匹配,不向量化,省内存\n", + " if not query.strip():\n", + " results = [v for v in video_list if not label_filter or label_filter in v['level2']]\n", + " return results[:top_k]\n", + " query_lower = query.lower()\n", + " results = []\n", + " for v in video_list:\n", + " if label_filter and label_filter not in v['level2']:\n", + " continue\n", + " if query_lower in v['summary'].lower():\n", + " results.append({**v, 'score': 1.0})\n", + " if len(results) >= top_k:\n", + " break\n", + " return results\n", + "\n", + "# 测试\n", + "test_results = search_video('มวยไทย', top_k=3)\n", + "print(f'视频检索测试: {len(test_results)} 条')\n", + "for r in test_results:\n", + " print(f' {r[\"url\"]}')" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 启动Gradio多模态交互界面\n", + "\n", + "三个Tab:\n", + "- **多模态RAG问答**:文字回答 + 相关图片展示\n", + "- **图文文化浏览**:按标签分类浏览泰国文化图片\n", + "- **视频内容检索**:关键词/标签检索泰语视频\n", + "\n", + "在浏览器打开 [http://127.0.0.1:8090](http://127.0.0.1:8090) 开始使用" + ] + }, + { + "cell_type": "code", + "execution_count": 15, + "metadata": { + "scrolled": true + }, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "/home/HwHiAiUser/.local/lib/python3.9/site-packages/gradio/components/dropdown.py:188: UserWarning: The value passed into gr.Dropdown() is not in the list of choices. Please update the list of choices to include: culture or set allow_custom_value=True.\n", + " warnings.warn(\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Running on local URL: http://0.0.0.0:8090\n", + "\n", + "To create a public link, set `share=True` in `launch()`.\n" + ] + }, + { + "data": { + "text/html": [ + "
" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + }, + { + "data": { + "text/plain": [] + }, + "execution_count": 15, + "metadata": {}, + "output_type": "execute_result" + }, + { + "name": "stderr", + "output_type": "stream", + "text": [ + "/home/HwHiAiUser/.local/lib/python3.9/site-packages/gradio/analytics.py:106: UserWarning: IMPORTANT: You are using gradio version 4.44.0, however version 4.44.1 is available, please upgrade. \n", + "--------\n", + " warnings.warn(\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + ".." + ] + } + ], + "source": [ + "import gradio as gr\n", + "\n", + "all_img_labels = sorted(set(\n", + " lv2\n", + " for d in image_data\n", + " for lv2 in d.get('labels', {}).get('pjwk_cates', {}).get('level2', [])\n", + "))\n", + "all_video_labels = sorted(v_label_counter.keys())\n", + "\n", + "def browse_images(label, page):\n", + " filtered = [\n", + " d for d in image_data\n", + " if label in d.get('labels', {}).get('pjwk_cates', {}).get('level2', [])\n", + " ]\n", + " total = len(filtered)\n", + " start = int(page) * 9\n", + " end = min(start + 9, total)\n", + " gallery = [(d['image']['path'], d['captions']['content']) for d in filtered[start:end]]\n", + " info = f'标签「{label}」共 {total} 张,第 {int(page)+1} 页({start+1}-{end})'\n", + " return gallery, info\n", + "\n", + "def search_videos(query, label):\n", + " label_filter = '' if label == '全部' else label\n", + " results = search_video(query, label_filter=label_filter, top_k=8)\n", + " if not results:\n", + " return '

未找到相关视频

'\n", + " html = '
'\n", + " for r in results:\n", + " lv2 = ', '.join(r['level2']) if r['level2'] else '未知'\n", + " summary = r['summary'][:120] + '...' if len(r['summary']) > 120 else r['summary']\n", + " html += (\n", + " '
'\n", + " f'
🏷️ {lv2}
'\n", + " f'
{summary}
'\n", + " f'▶ 观看视频'\n", + " '
'\n", + " )\n", + " html += '
'\n", + " return html\n", + "\n", + "with gr.Blocks(title='泰国文化多模态探索平台', theme=gr.themes.Soft()) as demo:\n", + " gr.Markdown('# 🇹🇭 泰国文化多模态探索平台')\n", + " gr.Markdown('基于万卷·丝路泰语数据集| wongwian-micro-instruct + MindSpore')\n", + "\n", + " with gr.Tab('多模态RAG问答'):\n", + " gr.Markdown('输入泰语或中文问题,系统从知识库检索相关内容并附带相关图片。')\n", + " with gr.Row():\n", + " with gr.Column(scale=3):\n", + " chatbot = gr.Chatbot(height=400, label='对话')\n", + " msg_input = gr.Textbox(\n", + " placeholder='输入问题,如:มวยไทยคืออะไร',\n", + " label='问题'\n", + " )\n", + " with gr.Row():\n", + " submit_btn = gr.Button('发送(RAG)', variant='primary')\n", + " no_rag_btn = gr.Button('发送(无RAG对比)')\n", + " clear_btn = gr.Button('清空')\n", + " gr.Examples(\n", + " examples=[\n", + " 'กีฬาประจำชาติของไทยคืออะไร',\n", + " 'ต้มยำคืออะไร',\n", + " 'ประเพณีสงกรานต์คืออะไร',\n", + " 'ผัดไทยคืออะไร',\n", + " ],\n", + " inputs=msg_input\n", + " )\n", + " with gr.Column(scale=2):\n", + " ref_gallery = gr.Gallery(label='检索到的相关图片', columns=3, height=400)\n", + "\n", + " def user_submit(message, history):\n", + " return '', history + [[message, '']]\n", + "\n", + " def bot_rag(history):\n", + " message = history[-1][0]\n", + " prev = history[:-1]\n", + " img_urls = []\n", + " for partial, imgs in rag_stream(message, prev, use_rag=True):\n", + " history[-1][1] = partial\n", + " img_urls = imgs\n", + " yield history, [(url, '') for url in img_urls[:3]]\n", + "\n", + " def bot_no_rag(history):\n", + " message = history[-1][0]\n", + " prev = history[:-1]\n", + " for partial, _ in rag_stream(message, prev, use_rag=False):\n", + " history[-1][1] = partial\n", + " yield history, []\n", + "\n", + " submit_btn.click(user_submit, [msg_input, chatbot], [msg_input, chatbot]).then(\n", + " bot_rag, chatbot, [chatbot, ref_gallery]\n", + " )\n", + " no_rag_btn.click(user_submit, [msg_input, chatbot], [msg_input, chatbot]).then(\n", + " bot_no_rag, chatbot, [chatbot, ref_gallery]\n", + " )\n", + " clear_btn.click(lambda: ([], []), outputs=[chatbot, ref_gallery])\n", + "\n", + " with gr.Tab('图文文化浏览'):\n", + " gr.Markdown('按标签分类浏览泰国文化图片,展示泰语描述。')\n", + " with gr.Row():\n", + " img_label_dd = gr.Dropdown(\n", + " choices=all_img_labels,\n", + " value=all_img_labels[0] if all_img_labels else '',\n", + " label='选择标签'\n", + " )\n", + " img_page_slider = gr.Slider(0, 50, value=0, step=1, label='页码')\n", + " img_info = gr.Textbox(label='当前页信息', interactive=False)\n", + " img_gallery = gr.Gallery(label='图片', columns=3, height=500)\n", + "\n", + " img_label_dd.change(browse_images, [img_label_dd, img_page_slider], [img_gallery, img_info])\n", + " img_page_slider.change(browse_images, [img_label_dd, img_page_slider], [img_gallery, img_info])\n", + " demo.load(browse_images, [img_label_dd, img_page_slider], [img_gallery, img_info])\n", + "\n", + " with gr.Tab('视频内容'):\n", + " gr.Markdown('输入泰语关键词或选择标签,检索相关泰语视频,点击链接跳转YouTube。')\n", + " with gr.Row():\n", + " video_query = gr.Textbox(placeholder='输入泰语关键词,如:มวยไทย', label='关键词(可留空)')\n", + " video_label_dd = gr.Dropdown(choices=['全部'] + all_video_labels, value='全部', label='标签筛选')\n", + " video_btn = gr.Button('搜索', variant='primary')\n", + " video_results = gr.HTML(label='搜索结果')\n", + " video_btn.click(search_videos, [video_query, video_label_dd], video_results)\n", + " gr.Examples(\n", + " examples=[['มวยไทย', '全部'], ['อาหาร', '全部'], ['', 'culture']],\n", + " inputs=[video_query, video_label_dd]\n", + " )\n", + "\n", + "demo.launch(server_name='0.0.0.0', server_port=8090)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 总结\n", + "\n", + "本项目在香橙派AIpro开发板上实现了泰国文化多模态探索平台,完整利用了万卷·丝路泰语数据集的三种模态:\n", + "\n", + "- **SFT数据**:构建文字知识库,支持RAG检索增强问答\n", + "- **图文数据**:构建图片知识库,RAG回答时附带相关图片,真正体现多模态\n", + "- **视频数据**:构建视频检索索引,支持关键词+标签双重筛选\n", + "\n", + "硬件:昇腾AI开发板\n", + "开发板镜像: Ubuntu镜像\n", + "CANN 8.1RC1;MindSpore: 2.5.0或mindspore2.6;8.1RC1beta1\n", + "MindNLP: 0.4.1\n", + "Python: 3.9" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.9.2" + } + }, + "nbformat": 4, + "nbformat_minor": 4 +}