为什么你的 Agent 总是"听不懂人话"?意图解析的 5 个优化技巧


一、引言 (Introduction)

钩子 (The Hook)

你是否有过这样的经历:在拨打某公司客服热线时,你清晰地说"我想查一下我的订单",但电话那头的智能助手却回复"抱歉,我没有理解您的意思,请您换一种方式提问";或者你在使用某个智能家居 App 时,对语音助手说"把客厅的灯调暗一点",结果它却把卧室的灯打开了?

这些令人沮丧的交互体验,本质上都指向同一个问题:AI Agent 无法准确理解用户的真实意图。在人工智能技术飞速发展的今天,我们构建的 Agent 越来越多,但"听不懂人话"似乎仍然是一个普遍存在的痛点。

定义问题/阐述背景 (The “Why”)

在当今的数字化时代,对话式 AI 正在成为人机交互的主流方式。从智能客服到语音助手,从聊天机器人到智能家居控制,Agent 正在渗透到我们生活的方方面面。然而,这些 Agent 的表现却参差不齐,其中最关键的瓶颈之一就是**意图解析(Intent Recognition)**的准确性。

意图解析是自然语言理解(NLU)的核心任务之一,它的目标是从用户的输入文本或语音中,识别出用户的真实目的。这听起来简单,但实际上却充满挑战:人类语言的模糊性、歧义性、多样性,以及上下文的复杂性,都让意图解析成为一个极具挑战性的任务。

一个 Agent 如果无法准确理解用户的意图,那么无论它的后端逻辑多么强大,都无法提供令人满意的服务。它不仅会降低用户体验,还可能导致业务流失,甚至在某些关键场景下造成严重后果。

亮明观点/文章目标 (The “What” & How")

在本文中,我们将深入探讨意图解析的核心挑战,并分享 5 个实用的优化技巧,帮助你提升 Agent 的理解能力。我们将从基础概念讲起,然后逐步深入到算法实现和工程实践,最后还会探讨一些进阶话题和未来趋势。

读完本文,你将学到:

  1. 意图解析的核心概念和工作原理
  2. 导致 Agent"听不懂人话"的主要原因
  3. 5 个具体、可落地的意图解析优化技巧
  4. 如何在实际项目中应用这些技巧
  5. 意图解析领域的前沿趋势和最佳实践

无论你是刚刚接触 NLP 的初学者,还是有一定经验的算法工程师,相信这篇文章都能给你带来一些启发和帮助。


二、基础知识/背景铺垫 (Foundational Concepts)

核心概念定义

在深入探讨优化技巧之前,我们需要先明确一些核心概念,建立一个共同的理解基础。

1. 什么是意图(Intent)?

在对话系统中,意图是指用户通过输入想要达成的目标或执行的动作。它是对用户需求的一种抽象表示。

例如:

  • 用户说:“明天北京的天气怎么样?” - 意图可能是 查询天气
  • 用户说:“帮我订一张去上海的机票” - 意图可能是 预订机票
  • 用户说:“打开空调” - 意图可能是 控制设备

意图通常是预定义的,取决于 Agent 的应用场景。一个订票系统的意图集和一个智能家居控制系统的意图集会有很大的不同。

2. 什么是意图解析(Intent Recognition)?

意图解析,也称为意图识别,是指将用户的自然语言输入映射到预定义意图集合中的过程。它的输入是一段文本(或语音转文本后的结果),输出是一个或多个意图及其置信度。

形式化地说,给定一个用户输入 xxx,和一个预定义的意图集合 Y={y1,y2,...,yn}Y = \{y_1, y_2, ..., y_n\}Y={y1,y2,...,yn},意图解析的任务是找到一个函数 fff,使得:

f(x)=y^f(x) = \hat{y}f(x)=y^

其中 y^∈Y\hat{y} \in Yy^Y 是预测的意图。

在很多实际场景中,我们不仅需要知道最可能的意图,还需要知道每个意图的置信度,因此更常见的形式是:

f(x)=(p1,p2,...,pn)f(x) = (p_1, p_2, ..., p_n)f(x)=(p1,p2,...,pn)

其中 pip_ipi 是输入 xxx 属于意图 yiy_iyi 的概率,且 ∑i=1npi=1\sum_{i=1}^{n} p_i = 1i=1npi=1

3. 意图解析 vs. 其他 NLP 任务

意图解析不是一个孤立的任务,它通常与其他 NLP 任务配合使用,共同构成一个完整的自然语言理解系统。

实体识别(Entity Recognition):与意图解析密切相关的是实体识别。意图回答的是"用户想做什么",而实体回答的是"用户想对什么做这件事"。例如,在"订一张明天去北京的机票"中,"明天"是时间实体,"北京"是地点实体,"机票"是产品实体。

槽填充(Slot Filling):这通常是意图解析和实体识别的下游任务。一旦确定了用户的意图,系统需要收集完成该意图所需的所有信息(槽位)。例如,“预订机票"的意图可能需要"出发地”、“目的地”、“日期”、"舱位"等槽位。

情感分析(Sentiment Analysis):虽然不直接属于意图解析,但情感分析可以帮助系统更好地理解用户的状态,从而调整响应策略。

4. 意图解析的应用场景

意图解析技术在很多领域都有广泛的应用:

  1. 智能客服:理解用户的咨询意图,提供相应的解答或转接人工
  2. 语音助手:如 Siri、小爱同学等,理解用户的各种指令
  3. 聊天机器人:在社交、电商等场景中与用户进行交互
  4. 智能家居:理解用户对家庭设备的控制指令
  5. 汽车座舱:理解驾驶员的各种操作需求,如导航、播放音乐等
  6. 企业应用:如 ERP、CRM 系统中的自然语言查询接口

相关工具/技术概览

要实现意图解析,我们需要借助一些 NLP 技术和工具。让我们来概览一下当前主流的技术栈:

1. 传统机器学习方法

在深度学习兴起之前,意图解析主要依赖传统的机器学习方法:

  • 特征工程:人工设计各种文本特征,如词袋模型(Bag-of-Words)、TF-IDF、n-gram 等
  • 分类算法:使用 SVM、朴素贝叶斯、随机森林、逻辑回归等分类器进行意图分类

这些方法的优点是模型可解释性强,计算资源需求低;缺点是需要大量的特征工程,且难以捕获文本的语义信息。

2. 深度学习方法

随着深度学习的发展,特别是词向量(Word Embedding)技术的出现,意图解析的准确率得到了显著提升:

  • 词嵌入:如 Word2Vec、GloVe、FastText 等,将词映射到低维向量空间,捕获词的语义信息
  • 序列模型:如 RNN、LSTM、GRU 等,能够处理文本的序列特性,捕获上下文信息
  • 预训练语言模型:如 BERT、GPT、RoBERTa 等,通过在大规模文本上预训练,学习到丰富的语言表示,在各种 NLP 任务上都取得了 SOTA(State-of-the-Art)效果

深度学习方法的优点是自动学习特征,能够捕获复杂的语义信息,准确率高;缺点是模型复杂,计算资源需求大,可解释性相对较弱。

3. 主流工具和框架

有很多优秀的工具和框架可以帮助我们实现意图解析:

  • 通用 NLP 框架

    • spaCy:工业级的 NLP 库,提供了丰富的 NLP 功能
    • NLTK:Python 中最流行的 NLP 库之一,适合学习和研究
    • Transformers (Hugging Face):提供了大量预训练语言模型的接口,使用非常方便
  • 对话系统框架

    • Rasa:开源的对话机器学习框架,专门用于构建上下文 AI 助手
    • Dialogflow:Google 提供的对话式 AI 平台
    • LUIS:Microsoft 提供的语言理解智能服务
    • Alexa Skills Kit:Amazon 提供的用于构建 Alexa 技能的工具包
  • 深度学习框架

    • TensorFlow / Keras:Google 推出的深度学习框架
    • PyTorch:Facebook 推出的深度学习框架,在学术界和工业界都非常流行

在本文的后续部分,我们将使用 Python 和一些主流库来实现具体的优化技巧。


三、核心内容/实战演练 (The Core - “How-To”)

在了解了基础知识之后,让我们来探讨 5 个具体的意图解析优化技巧。每个技巧我们都会从问题背景、问题描述、解决方案、算法实现、代码示例等方面进行详细讲解。

技巧一:数据增强 - 用有限的数据创造无限的可能

核心概念

数据增强(Data Augmentation) 是指通过对现有数据进行各种变换,生成新的训练样本的技术。在计算机视觉领域,数据增强是一种非常成熟的技术,常见的操作有旋转、翻转、缩放、裁剪等。在自然语言处理领域,数据增强也越来越受到重视。

问题背景

意图解析模型的性能很大程度上依赖于训练数据的质量和数量。然而,在实际项目中,我们往往面临以下问题:

  1. 数据稀缺:对于很多垂直领域或新业务,很难收集到大量的标注数据
  2. 数据分布不均:某些意图的样本很多,而另一些意图的样本很少,导致模型偏向于样本多的意图
  3. 表达多样性不足:用户的表达方式千变万化,但我们收集到的数据往往只能覆盖其中的一小部分

这些问题都会导致模型的泛化能力不足,在实际应用中表现不佳。数据增强正是解决这些问题的有效手段之一。

问题描述

假设我们正在构建一个智能家居的意图解析系统,目前我们收集到了以下一些训练样本:

文本 意图
打开客厅的灯 控制设备
把空调打开 控制设备
关上卧室的窗户 控制设备
明天天气怎么样 查询天气
今天会下雨吗 查询天气
播放一首周杰伦的歌 播放音乐
来首歌听 播放音乐

每个意图只有 2-3 个样本,显然这样的数据量是远远不够的。如果我们直接用这些数据训练模型,模型的泛化能力会很差,对于"把客厅灯开一下"这样的简单变体,可能都无法正确识别。

我们的目标是通过数据增强技术,在不增加人工标注成本的情况下,生成更多高质量的训练样本,从而提升模型的性能。

问题解决

针对文本数据,有多种数据增强方法,我们可以根据具体场景选择合适的方法,或者组合使用多种方法。以下是一些常用的文本数据增强方法:

  1. 同义词替换(Synonym Replacement):随机选择句子中的几个非停用词,用它们的同义词替换
  2. 随机插入(Random Insertion):在句子中随机插入某个词的同义词
  3. 随机交换(Random Swap):随机交换句子中的两个词的位置
  4. 随机删除(Random Deletion):随机删除句子中的几个词
  5. 回译(Back Translation):将句子翻译成另一种语言,再翻译回原语言
  6. 上下文增强(Contextual Augmentation):使用预训练语言模型,根据上下文预测并替换句子中的某些词
  7. 模板生成(Template-based Generation):根据领域知识,定义一些模板,通过填充不同的槽值生成新样本

接下来,我们将详细介绍其中几种方法,并提供 Python 实现代码。

算法流程图

让我们先来看一个数据增强的整体流程图:

同义词替换

随机交换

回译

模板生成

通过

不通过

原始标注数据

选择增强方法

生成同义词替换样本

生成随机交换样本

生成回译样本

生成模板样本

数据清洗与过滤

质量评估

加入训练集

丢弃

算法实现与源代码

让我们来实现几种常用的数据增强方法。首先,我们需要安装一些必要的库:

pip install nltk transformers torch

然后,我们来编写代码:

import random
import nltk
from nltk.corpus import wordnet
from nltk.tokenize import word_tokenize
from transformers import pipeline, AutoTokenizer, AutoModelForMaskedLM
import torch

# 下载必要的 NLTK 数据
nltk.download('wordnet')
nltk.download('averaged_perceptron_tagger')
nltk.download('punkt')

class TextAugmenter:
    def __init__(self):
        # 初始化同义词替换所需的词性映射
        self.pos_map = {
            'NN': wordnet.NOUN,
            'VB': wordnet.VERB,
            'JJ': wordnet.ADJ,
            'RB': wordnet.ADV
        }
    
    def get_synonyms(self, word, pos):
        """获取一个词的同义词"""
        synonyms = set()
        for syn in wordnet.synsets(word, pos=pos):
            for lemma in syn.lemmas():
                synonym = lemma.name().replace('_', ' ')
                if synonym != word:
                    synonyms.add(synonym)
        return list(synonyms)
    
    def synonym_replacement(self, text, p=0.2):
        """
        同义词替换
        p: 替换词的比例
        """
        words = word_tokenize(text)
        pos_tags = nltk.pos_tag(words)
        
        # 选择可以替换的词(名词、动词、形容词、副词)
        replaceable_indices = []
        for i, (word, pos) in enumerate(pos_tags):
            for key in self.pos_map:
                if pos.startswith(key):
                    replaceable_indices.append(i)
                    break
        
        # 随机选择要替换的词
        num_replace = max(1, int(len(replaceable_indices) * p))
        replace_indices = random.sample(replaceable_indices, min(num_replace, len(replaceable_indices)))
        
        # 进行替换
        new_words = words.copy()
        for i in replace_indices:
            word = words[i]
            pos = pos_tags[i][1]
            # 映射到 WordNet 的词性
            wn_pos = None
            for key, value in self.pos_map.items():
                if pos.startswith(key):
                    wn_pos = value
                    break
            
            if wn_pos:
                synonyms = self.get_synonyms(word, wn_pos)
                if synonyms:
                    new_words[i] = random.choice(synonyms)
        
        return ' '.join(new_words)
    
    def random_swap(self, text, n=1):
        """
        随机交换
        n: 交换的次数
        """
        words = word_tokenize(text)
        if len(words) < 2:
            return text
        
        new_words = words.copy()
        for _ in range(n):
            idx1, idx2 = random.sample(range(len(new_words)), 2)
            new_words[idx1], new_words[idx2] = new_words[idx2], new_words[idx1]
        
        return ' '.join(new_words)
    
    def random_deletion(self, text, p=0.2):
        """
        随机删除
        p: 删除词的概率
        """
        words = word_tokenize(text)
        if len(words) == 1:
            return text
        
        new_words = []
        for word in words:
            if random.random() > p:
                new_words.append(word)
        
        # 确保不会删除所有词
        if not new_words:
            return random.choice(words)
        
        return ' '.join(new_words)
    
    def contextual_augmentation(self, text, p=0.2, model_name='bert-base-chinese'):
        """
        上下文增强
        使用预训练语言模型根据上下文预测并替换词
        """
        tokenizer = AutoTokenizer.from_pretrained(model_name)
        model = AutoModelForMaskedLM.from_pretrained(model_name)
        
        words = list(text)  # 对于中文,按字处理
        num_mask = max(1, int(len(words) * p))
        mask_indices = random.sample(range(len(words)), num_mask)
        
        # 创建掩码文本
        masked_words = words.copy()
        for idx in mask_indices:
            masked_words[idx] = tokenizer.mask_token
        
        masked_text = ''.join(masked_words)
        
        # 模型预测
        inputs = tokenizer(masked_text, return_tensors='pt')
        with torch.no_grad():
            outputs = model(**inputs)
            predictions = outputs.logits
        
        # 替换掩码位置
        for idx in mask_indices:
            mask_token_index = (inputs.input_ids[0] == tokenizer.mask_token_id).nonzero(as_tuple=True)[0]
            if idx < len(mask_token_index):
                mask_logits = predictions[0, mask_token_index[idx]]
                # 获取 top 5 预测,随机选择一个
                top_tokens = torch.topk(mask_logits, 5).indices
                # 排除原词
                original_token_id = tokenizer.convert_tokens_to_ids(words[idx])
                candidates = [t for t in top_tokens if t != original_token_id]
                if candidates:
                    new_token_id = random.choice(candidates)
                    words[idx] = tokenizer.decode(new_token_id, skip_special_tokens=True)
        
        return ''.join(words)


# 模板生成器
class TemplateGenerator:
    def __init__(self):
        # 定义模板和槽值
        self.templates = {
            '控制设备': [
                '{action}{location}的{device}',
                '把{location}的{device}{action}',
                '{action}{location}{device}',
                '{location}的{device}{action}'
            ],
            '查询天气': [
                '{time}{location}天气怎么样',
                '{time}{location}会{weather_condition}吗',
                '查一下{time}{location}的天气',
                '{time}{location}的天气如何'
            ],
            '播放音乐': [
                '播放一首{singer}的歌',
                '来首{singer}的歌听',
                '放一首{singer}的歌',
                '播放{singer}的音乐'
            ]
        }
        
        self.slot_values = {
            'action': ['打开', '关上', '开启', '关闭', '开一下', '关一下'],
            'location': ['客厅', '卧室', '厨房', '浴室', '书房', '阳台'],
            'device': ['灯', '空调', '电视', '风扇', '窗户', '窗帘'],
            'time': ['今天', '明天', '后天', '这周', '下周', '周末'],
            'location_weather': ['北京', '上海', '广州', '深圳', '杭州', '成都'],
            'weather_condition': ['下雨', '下雪', '晴天', '阴天', '刮风'],
            'singer': ['周杰伦', '林俊杰', '陈奕迅', '邓紫棋', '李荣浩', '薛之谦']
        }
    
    def generate(self, intent, num_samples=10):
        """生成指定意图的样本"""
        if intent not in self.templates:
            return []
        
        samples = []
        templates = self.templates[intent]
        
        for _ in range(num_samples):
            template = random.choice(templates)
            sample = template
            
            # 替换槽值
            for slot, values in self.slot_values.items():
                if '{' + slot + '}' in sample:
                    # 特殊处理天气查询中的地点,避免冲突
                    if intent == '查询天气' and slot == 'location':
                        value = random.choice(self.slot_values['location_weather'])
                    else:
                        value = random.choice(values)
                    sample = sample.replace('{' + slot + '}', value)
            
            samples.append(sample)
        
        return samples


# 使用示例
if __name__ == "__main__":
    # 原始数据
    original_data = [
        ("打开客厅的灯", "控制设备"),
        ("把空调打开", "控制设备"),
        ("关上卧室的窗户", "控制设备"),
        ("明天天气怎么样", "查询天气"),
        ("今天会下雨吗", "查询天气"),
        ("播放一首周杰伦的歌", "播放音乐"),
        ("来首歌听", "播放音乐")
    ]
    
    # 初始化增强器
    augmenter = TextAugmenter()
    template_generator = TemplateGenerator()
    
    # 增强后的数据
    augmented_data = []
    
    # 对每个原始样本进行增强
    for text, intent in original_data:
        augmented_data.append((text, intent))  # 保留原始样本
        
        # 同义词替换
        sr_text = augmenter.synonym_replacement(text)
        if sr_text != text:
            augmented_data.append((sr_text, intent))
        
        # 随机交换
        rs_text = augmenter.random_swap(text)
        if rs_text != text:
            augmented_data.append((rs_text, intent))
        
        # 随机删除
        rd_text = augmenter.random_deletion(text)
        if rd_text != text:
            augmented_data.append((rd_text, intent))
    
    # 使用模板生成更多样本
    intents = set([intent for _, intent in original_data])
    for intent in intents:
        template_samples = template_generator.generate(intent, num_samples=20)
        for sample in template_samples:
            augmented_data.append((sample, intent))
    
    # 输出增强后的数据
    print(f"原始样本数: {len(original_data)}")
    print(f"增强后样本数: {len(augmented_data)}")
    print("\n增强后的样本示例:")
    for i, (text, intent) in enumerate(augmented_data[:20], 1):
        print(f"{i}. [{intent}] {text}")
边界与外延

虽然数据增强是一种有效的技术,但它也有一些边界和注意事项:

  1. 保持语义一致性:在进行数据增强时,最重要的原则是确保生成的新样本的语义与原始样本保持一致,意图不变。如果生成的样本改变了原意,反而会损害模型的性能。

  2. 领域适用性:不同的增强方法适用于不同的领域和语言。例如,基于 WordNet 的同义词替换在通用领域可能效果不错,但在专业领域可能不太适用,因为 WordNet 可能不包含专业术语。

  3. 增强强度:过度增强可能会导致样本质量下降,甚至生成一些不自然或无意义的句子。需要根据具体任务调整增强的强度。

  4. 数据增强 vs. 主动学习:数据增强是利用现有数据生成新样本,而主动学习是选择最有价值的样本进行标注。这两种技术可以结合使用,达到更好的效果。

实际场景应用

在实际项目中,我们可以按照以下步骤应用数据增强技术:

  1. 数据分析:首先分析现有数据的特点,如样本数量、类别分布、表达方式多样性等,确定数据增强的重点。

  2. 选择增强方法:根据数据特点和任务需求,选择合适的增强方法。通常可以组合使用多种方法。

  3. 生成增强样本:使用选定的方法生成增强样本。

  4. 质量过滤:对生成的样本进行质量过滤,去除低质量或语义不一致的样本。可以使用规则过滤,也可以训练一个分类器来自动过滤。

  5. 实验验证:在增强后的数据集上训练模型,在验证集上评估性能,根据结果调整增强策略。

通过合理应用数据增强技术,我们可以在不增加太多人工成本的情况下,显著提升模型的性能。


技巧二:小样本学习 - 让 Agent 在少量数据下也能"举一反三"

核心概念

小样本学习(Few-Shot Learning) 是机器学习的一个子领域,它研究如何让模型在只有少量标注样本的情况下,也能快速学习并泛化到新的任务。在意图解析场景中,小样本学习可以帮助我们解决新意图冷启动的问题。

问题背景

在实际的对话系统开发中,我们经常会遇到以下场景:

  1. 新意图上线:产品经理提出了一个新的功能,需要支持一个新的意图,但我们目前只有几个标注样本,甚至没有标注样本。
  2. 长尾意图:大多数用户的查询集中在少数几个热门意图上,但还有大量的长尾意图,每个意图只有很少的样本。
  3. 快速迭代:业务需求变化很快,我们需要快速支持新的意图,没有时间收集大量标注数据。

在这些场景下,传统的监督学习方法往往表现不佳,因为它们需要大量的标注数据才能取得好的效果。小样本学习正是为了解决这些问题而提出的。

问题描述

假设我们的智能家居系统已经支持了"控制设备"、“查询天气”、"播放音乐"这三个意图,现在我们需要新增一个"设置闹钟"的意图,但我们只有以下几个样本:

文本 意图
明天早上7点叫我起床 设置闹钟
帮我定一个下午3点的闹钟 设置闹钟
定个闹钟,明天8点 设置闹钟

同时,我们还有之前三个意图的大量数据。我们的目标是利用这少量的"设置闹钟"样本,以及之前的大量数据,构建一个能够准确识别"设置闹钟"意图的模型。

问题解决

小样本学习有多种方法,我们将介绍几种在意图解析中常用的方法:

  1. 数据增强:我们在技巧一中已经介绍过,可以作为小样本学习的一个基础步骤。
  2. 迁移学习(Transfer Learning):在有大量数据的源任务上预训练模型,然后在只有少量数据的目标任务上微调。
  3. 元学习(Meta-Learning):也称为"学习如何学习",训练模型在多个任务上学习,使其能够快速适应新任务。
  4. 度量学习(Metric Learning):学习一个度量空间,使得相似样本的距离更近,不同样本的距离更远,然后通过比较距离来分类。
  5. 提示学习(Prompt Learning):通过精心设计的提示,激发预训练语言模型的知识,使其在少样本情况下也能表现良好。

接下来,我们将重点介绍迁移学习和度量学习在意图解析中的应用。

概念结构与核心要素组成

让我们首先了解一下小样本学习的基本概念结构:

consists_of

contains

contains

includes

includes

FEW-SHOT-LEARNING

string

name

string

description

EPISODE

string

name

string

description

SUPPORT-SET

string

name

int

K

string

description

QUERY-SET

string

name

int

Q

string

description

META-TRAINING

string

name

string

description

META-TESTING

string

name

string

description

在小样本学习中,我们通常会使用"情节(Episode)"的概念来组织训练和测试。每个情节包含:

  • 支持集(Support Set):包含 K 个类别,每个类别 N 个样本(因此也称为 N-way K-shot 学习)
  • 查询集(Query Set):包含从相同类别中采样的 Q 个样本,模型需要对这些样本进行分类
数学模型

让我们首先介绍一下基于度量学习的小样本学习方法,其中最经典的是原型网络(Prototypical Networks)。

原型网络(Prototypical Networks)

原型网络的核心思想是:对于每个类别,计算其在嵌入空间中的"原型"(即该类别所有样本嵌入的平均值),然后对于一个新的查询样本,将其分类到最近的原型所在的类别。

形式化定义如下:

  1. 嵌入函数:首先,我们有一个嵌入函数 fθ:X→RDf_\theta: \mathcal{X} \rightarrow \mathbb{R}^Dfθ:XRD,它将输入样本映射到 D 维嵌入空间,其中 θ\thetaθ 是函数的参数。

  2. 计算原型:对于支持集 S={(x1,y1),...,(xN×K,yN×K)}S = \{(x_1, y_1), ..., (x_{N \times K}, y_{N \times K})\}S={(x1,y1),...,(xN×K,yN×K)},其中包含 N 个类别,每个类别 K 个样本,我们计算每个类别的原型:

ck=1K∑(xi,yi)∈Skfθ(xi)c_k = \frac{1}{K} \sum_{(x_i, y_i) \in S_k} f_\theta(x_i)ck=K1(xi,yi)Skfθ(xi)

其中 SkS_kSk 是支持集中属于第 k 个类别的样本集合,ckc_kck 是第 k 个类别的原型。

  1. 分类:对于一个查询样本 xxx,我们计算其嵌入与每个原型之间的距离,然后通过 softmax 函数将距离转换为概率:

pθ(y=k∣x)=exp⁡(−d(fθ(x),ck))∑k′=1Nexp⁡(−d(fθ(x),ck′))p_\theta(y=k|x) = \frac{\exp(-d(f_\theta(x), c_k))}{\sum_{k'=1}^{N} \exp(-d(f_\theta(x), c_{k'}))}pθ(y=kx)=k=1Nexp(d(fθ(x),ck))exp(d(fθ(x),ck))

其中 d(⋅,⋅)d(\cdot, \cdot)d(,) 是距离函数,通常使用欧氏距离或余弦距离。

  1. 训练:在训练阶段,我们通过最小化查询样本的负对数似然来优化嵌入函数的参数 θ\thetaθ

J(θ)=−1Q∑(x,y)∈Qlog⁡pθ(y∣x)J(\theta) = -\frac{1}{Q} \sum_{(x, y) \in Q} \log p_\theta(y|x)J(θ)=Q1(x,y)Qlogpθ(yx)

其中 Q 是查询集。

算法流程图

下面是原型网络的训练和推理流程图:

推理阶段

加载训练好的模型

获取新类别的支持样本

计算新类别样本的嵌入

计算新类别的原型

获取查询样本

计算查询样本的嵌入

计算与各原型的距离

返回最近的类别

训练阶段

采样训练任务

构建支持集和查询集

计算支持集样本的嵌入

计算每个类别的原型

计算查询集样本的嵌入

计算查询样本与原型的距离

计算损失

更新模型参数

训练完成?

保存训练好的模型

算法源代码

让我们使用 PyTorch 来实现一个简单的原型网络,用于小样本意图分类:

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
from transformers import BertTokenizer, BertModel
import random
import numpy as np
from collections import defaultdict

# 设置随机种子
def set_seed(seed):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)

set_seed(42)

# 数据集类
class FewShotIntentDataset(Dataset):
    def __init__(self, data, tokenizer, max_len=128):
        self.data = data
        self.tokenizer = tokenizer
        self.max_len = max_len
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        text, label = self.data[idx]
        encoding = self.tokenizer(
            text,
            truncation=True,
            padding='max_length',
            max_length=self.max_len,
            return_tensors='pt'
        )
        return {
            'input_ids': encoding['input_ids'].flatten(),
            'attention_mask': encoding['attention_mask'].flatten(),
            'label': torch.tensor(label, dtype=torch.long)
        }

# 原型网络模型
class PrototypicalNetwork(nn.Module):
    def __init__(self, model_name='bert-base-chinese', hidden_size=768):
        super().__init__()
        self.bert = BertModel.from_pretrained(model_name)
        self.hidden_size = hidden_size
        # 可选:添加一个投影层,将 BERT 的输出投影到更低维度
        self.projection = nn.Linear(self.hidden_size, self.hidden_size)
    
    def forward(self, input_ids, attention_mask):
        outputs = self.bert(
            input_ids=input_ids,
            attention_mask=attention_mask
        )
        # 使用 <[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token 的表示作为句子嵌入
        pooled_output = outputs.last_hidden_state[:, 0, :]
        # 可选:投影
        # pooled_output = self.projection(pooled_output)
        return pooled_output
    
    def compute_prototypes(self, support_embeddings, support_labels):
        """计算每个类别的原型"""
        unique_labels = torch.unique(support_labels)
        prototypes = []
        
        for label in unique_labels:
            # 获取该类别的所有嵌入
            mask = (support_labels == label)
            class_embeddings = support_embeddings[mask]
            # 计算平均值作为原型
            prototype = torch.mean(class_embeddings, dim=0)
            prototypes.append(prototype)
        
        return torch.stack(prototypes), unique_labels
    
    def euclidean_distance(self, x, y):
        """计算欧氏距离"""
        # x: (num_query, hidden_size)
        # y: (num_classes, hidden_size)
        # 返回: (num_query, num_classes)
        n = x.size(0)
        m = y.size(0)
        d = x.size(1)
        
        x = x.unsqueeze(1).expand(n, m, d)
        y = y.unsqueeze(0).expand(n, m, d)
        
        return torch.pow(x - y, 2).sum(2)
    
    def cosine_distance(self, x, y):
        """计算余弦距离(1 - 余弦相似度)"""
        # x: (num_query, hidden_size)
        # y: (num_classes, hidden_size)
        # 返回: (num_query, num_classes)
        x_norm = F.normalize(x, p=2, dim=1)
        y_norm = F.normalize(y, p=2, dim=1)
        cosine_similarity = torch.mm(x_norm, y_norm.t())
        return 1 - cosine_similarity


# 小样本学习训练器
class FewShotTrainer:
    def __init__(self, model, tokenizer, device='cuda' if torch.cuda.is_available() else 'cpu'):
        self.model = model.to(device)
        self.tokenizer = tokenizer
        self.device = device
        self.optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5)
    
    def create_episode(self, data_by_label, N, K, Q):
        """
        创建一个训练/测试情节
        N: 类别数
        K: 每个类别的支持样本数
        Q: 每个类别的查询样本数
        """
        # 随机选择 N 个类别
        selected_labels = random.sample(list(data_by_label.keys()), N)
        
        support_data = []
        query_data = []
        
        for i, label in enumerate(selected_labels):
            # 从该类别中随机选择 K+Q 个样本
            samples = random.sample(data_by_label[label], K + Q)
            # 前 K 个作为支持集,后 Q 个作为查询集
            for text in samples[:K]:
                support_data.append((text, i))  # 使用相对标签
            for text in samples[K:]:
                query_data.append((text, i))
        
        # 创建数据集和数据加载器
        support_dataset = FewShotIntentDataset(support_data, self.tokenizer)
        query_dataset = FewShotIntentDataset(query_data, self.tokenizer)
        
        support_loader = DataLoader(support_dataset, batch_size=len(support_dataset))
        query_loader = DataLoader(query_dataset, batch_size=len(query_dataset))
        
        return support_loader, query_loader
    
    def train_episode(self, support_loader, query_loader):
        """在一个情节上训练"""
        self.model.train()
        
        # 获取支持集数据
        support_batch = next(iter(support_loader))
        support_input_ids = support_batch['input_ids'].to(self.device)
        support_attention_mask = support_batch['attention_mask'].to(self.device)
        support_labels = support_batch['label'].to(self.device)
        
        # 获取查询集数据
        query_batch = next(iter(query_loader))
        query_input_ids = query_batch['input_ids'].to(self.device)
        query_attention_mask = query_batch['attention_mask'].to(self.device)
        query_labels = query_batch['label'].to(self.device)
        
        # 前向传播
        self.optimizer.zero_grad()
        
        # 计算支持集和查询集的嵌入
        support_embeddings = self.model(support_input_ids, support_attention_mask)
        query_embeddings = self.model(query_input_ids, query_attention_mask)
        
        # 计算原型
        prototypes, _ = self.model.compute_prototypes(support_embeddings, support_labels)
        
        # 计算距离
        distances = self.model.euclidean_distance(query_embeddings, prototypes)
        
        # 计算损失(使用负对数似然)
        log_p_y = F.log_softmax(-distances, dim=1)
        loss = F.nll_loss(log_p_y, query_labels)
        
        # 计算准确率
        _, predictions = torch.max(-distances, 1)
        accuracy = torch.mean((predictions == query_labels).float())
        
        # 反向传播和优化
        loss.backward()
        self.optimizer.step()
        
        return loss.item(), accuracy.item()
    
    def train(self, data_by_label, N, K, Q, num_episodes=1000, eval_every=100):
        """训练模型"""
        train_losses = []
        train_accuracies = []
        
        for episode in range(num_episodes):
            # 创建一个训练情节
            support_loader, query_loader = self.create_episode(data_by_label, N, K, Q)
            
            # 训练
            loss, accuracy = self.train_episode(support_loader, query_loader)
            
            train_losses.append(loss)
            train_accuracies.append(accuracy)
            
            # 打印进度
            if (episode + 1) % eval_every == 0:
                avg_loss = np.mean(train_losses[-eval_every:])
                avg_accuracy = np.mean(train_accuracies[-eval_every:])
                print(f"Episode {episode+1}/{num_episodes}, Loss: {avg_loss:.4f}, Accuracy: {avg_accuracy:.4f}")
    
    def predict(self, support_texts, support_labels, query_texts, distance_metric='euclidean'):
        """
        预测查询文本的类别
        support_texts: 支持集文本列表
        support_labels: 支持集标签列表
        query_texts: 查询文本列表
        """
        self.model.eval()
        
        # 处理支持集
        support_encodings = self.tokenizer(
            support_texts,
            truncation=True,
            padding='max_length',
            max_length=128,
            return_tensors='pt'
        )
        support_input_ids = support_encodings['input_ids'].to(self.device)
        support_attention_mask = support_encodings['attention_mask'].to(self.device)
        support_labels = torch.tensor(support_labels).to(self.device)
        
        # 处理查询集
        query_encodings = self.tokenizer(
            query_texts,
            truncation=True,
            padding='max_length',
            max_length=128,
            return_tensors='pt'
        )
        query_input_ids = query_encodings['input_ids'].to(self.device)
        query_attention_mask = query_encodings['attention_mask'].to(self.device)
        
        with torch.no_grad():
            # 计算嵌入
            support_embeddings = self.model(support_input_ids, support_attention_mask)
            query_embeddings = self.model(query_input_ids, query_attention_mask)
            
            # 计算原型
            prototypes, unique_labels = self.model.compute_prototypes(support_embeddings, support_labels)
            
            # 计算距离
            if distance_metric == 'euclidean':
                distances = self.model.euclidean_distance(query_embeddings, prototypes)
            else:  # cosine
                distances = self.model.cosine_distance(query_embeddings, prototypes)
            
            # 预测
            _, predictions = torch.min(distances, 1)
            predicted_labels = [unique_labels[p].item() for p in predictions]
        
        return predicted_labels


# 示例使用
if __name__ == "__main__":
    # 模拟一些数据
    # 假设我们有 5 个已有意图,每个意图有 50 个样本
    existing_intents = {
        '控制设备': ['打开客厅的灯', '关上卧室的窗户', '把空调打开', '开一下书房的灯', '关闭浴室的排气扇'] * 10,
        '查询天气': ['今天北京的天气怎么样', '明天会下雨吗', '查一下上海的天气', '杭州今天的天气如何', '深圳周末会晴天吗'] * 10,
        '播放音乐': ['播放一首周杰伦的歌', '来首歌听', '放一首林俊杰的歌', '播放陈奕迅的音乐', '放一首邓紫棋的歌'] * 10,
        '设置提醒': ['提醒我明天开会', '下午3点提醒我吃药', '明天早上8点叫我起床', '提醒我给妈妈打电话', '周五提醒我交报告'] * 10,
        '查询时间': ['现在几点了', '今天是星期几', '现在是什么时间', '今天几号', '明天星期几'] * 10
    }
    
    # 将数据转换为 label -> texts 的格式
    label_to_texts = {}
    label_list = list(existing_intents.keys())
    for i, (intent, texts) in enumerate(existing_intents.items()):
        label_to_texts[i] = texts
    
    # 初始化模型和训练器
    tokenizer = BertTokenizer.from_pretrained('bert-base-chinese')
    model = PrototypicalNetwork('bert-base-chinese')
    trainer = FewShotTrainer(model, tokenizer)
    
    # 训练模型(元训练)
    print("开始训练...")
    trainer.train(
        data_by_label=label_to_texts,
        N=3,  # 每个情节 3 个类别
        K=5,  # 每个类别 5 个支持样本
        Q=5,  # 每个类别 5 个查询样本
        num_episodes=500,  # 训练 500 个情节
        eval_every=50
    )
    
    # 现在,假设我们有一个新的意图"设置闹钟",只有 3 个样本
    new_intent_samples = [
        "明天早上7点叫我起床",
        "帮我定一个下午3点的闹钟",
        "定个闹钟,明天8点"
    ]
    new_intent_label
Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐