bert+seq2seq 周公解梦,看AI如何解析你的梦境?【转】

2023-10-21 05:40

本文主要是介绍bert+seq2seq 周公解梦,看AI如何解析你的梦境?【转】,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

介绍

在参与的项目和产品中,涉及到模型和算法的需求,主要以自然语言处理(NLP)和知识图谱(KG)为主。NLP涉及面太广,而聚焦在具体场景下,想要生产落地的还需要花很多功夫。
作为NLP的主要方向,情感分析,文本多分类,实体识别等已经在项目中得到应用。例如
通过实体识别,抽取文本中提及到的公司、个人以及金融产品等。
通过情感分析,判别新闻资讯,对其提到的公司和个人是否利好?
通过文本多分类,判断资讯是否是高质量?判断资讯的行业和主题?
具体详情再找时间分享。而文本生成、序列到序列(Sequence to Sequence)在机器翻译、问答系统、聊天机器人中有较广的应用,在参与的项目中暂无涉及,本文主要通过tensorflow+bert+seq2seq实现一个简单的问答模型,旨在对seq2seq的了解和熟悉。

数据

关于seq2seq的demo数据有很多,例如小黄鸡聊天语料库,影视语料库,翻译语料库等等。由于最近总是做些奇怪的梦,便想着,做一个AI解梦的应用玩玩,just for fun。
通过采集从网上采集周公解梦数据,通过清洗,形成
dream:梦境;
decode:梦境解析结果。
这样的序列对,总计33000+ 条记录。数据集下载地址:后台回复“解梦”
{
"dream": "梦见商人或富翁",
"decode": "是个幸运的预兆,未来自己的事业很有机会成功,不过如果梦中的富翁是自己,则是一个凶兆。。"
}

模型准备

#下载 bert
$ git clone https://github.com/google-research/bert.git
#下载中文预训练模型
$ wget -c https://storage.googleapis.com/bert_models/2018_11_03/chinese_L-12_H-768_A-12.zip
$ unzip chinese_L-12_H-768_A-12.zip 

bert 的input:

self.input_ids = tf.placeholder(dtype=tf.int32,shape=[None, None],name="input_ids"
)
self.input_mask = tf.placeholder(dtype=tf.int32,shape=[None, None],name="input_mask"
)
self.segment_ids = tf.placeholder(dtype=tf.int32,shape=[None, None],name="segment_ids"
)
self.dropout = tf.placeholder(dtype=tf.float32,shape=None,name="dropout"
)

bert 的model :

self.bert_config = modeling.BertConfig.from_json_file(bert_config)model = modeling.BertModel(config=self.bert_config,is_training=self.is_training,input_ids=self.input_ids,input_mask=self.input_mask,token_type_ids=self.segment_ids,use_one_hot_embeddings=False)

seq2seq 的encoder_embedding 替换:

# 默认seq2seq model_inputs
# self.encoder_embedding = tf.Variable(tf.random_uniform([from_dict_size, embedded_size], -1, 1),name ="encoder_embedding")
# self.model_inputs = tf.nn.embedding_lookup(self.encoder_embedding, self.X),
#  替换成bert
self.embedded = model.get_sequence_output()
self.model_inputs = tf.nn.dropout(self.embedded, self.dropout)

seq2seq 的decoder_embedding 替换:

# 默认seq2seq decoder_embedding
# self.decoder_embedding = tf.Variable(tf.random_uniform([to_dict_size, embedded_size], -1, 1),name="decoder_embedding")
#  替换成bert
self.decoder_embedding = model.get_embedding_table()
self.decoder_input = tf.nn.embedding_lookup(self.decoder_embedding, decoder_input),

数据预处理

for i in range(len(inputs)):tokens = inputs[i]inputs_ids = model.tokenizer.convert_tokens_to_ids(inputs[i])segment_ids = [0] * len(inputs_ids)input_mask = [1] * len(inputs_ids)tag_ids = model.tokenizer.convert_tokens_to_ids(outputs[i])data.append([tokens, tag_ids, inputs_ids, segment_ids, input_mask])def pad_data(data):c_data = copy.deepcopy(data)max_x_length = max([len(i[0]) for i in c_data])max_y_length = max([len(i[1]) for i in c_data]) # 这里生成的序列的tag-id 和 input-id 长度要分开# print("max_x_length : {} ,max_y_length : {}".format( max_x_length,max_y_length))padded_data = []for i in c_data:tokens, tag_ids, inputs_ids, segment_ids, input_mask = itag_ids = tag_ids + (max_y_length - len(tag_ids)) * [0]# 注意tag-ids 的长度补充,和预测的序列长度一致。inputs_ids = inputs_ids + (max_x_length - len(inputs_ids)) * [0]segment_ids = segment_ids + (max_x_length - len(segment_ids)) * [0]input_mask = input_mask + (max_x_length - len(input_mask)) * [0]assert len(inputs_ids) == len(segment_ids) == len(input_mask)padded_data.append([tokens, tag_ids, inputs_ids, segment_ids, input_mask])return padded_data

训练

$ python3 model.py --task=train \--is_training=True \--epoch=100 \--size_layer=256 \--bert_config=chinese_L-12_H-768_A-12/bert_config.json \--vocab_file=chinese_L-12_H-768_A-12/vocab.txt \--num_layers=2 \--learning_rate=0.001 \--batch_size=16 \--checkpoint_dir=result

image

预测

$ python3 model.py --task=predict \--is_training=False \--epoch=100 \--size_layer=256 \--bert_config=chinese_L-12_H-768_A-12/bert_config.json \--vocab_file=chinese_L-12_H-768_A-12/vocab.txt \--num_layers=2 \--learning_rate=0.001 \--batch_size=16 \--checkpoint_dir=result

image

Just For Fun ^_^

本文代码: https://github.com/saiwaiyanyu/tensorflow-bert-seq2seq-dream-decoder

作者:saiwaiyanyu
链接:https://juejin.im/post/5dd9e07b51882572f00c4523
来源:掘金

8

本文由博客一文多发平台 OpenWrite 发布!

这篇关于bert+seq2seq 周公解梦,看AI如何解析你的梦境?【转】的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/252320

相关文章

线上Java OOM问题定位与解决方案超详细解析

《线上JavaOOM问题定位与解决方案超详细解析》OOM是JVM抛出的错误,表示内存分配失败,:本文主要介绍线上JavaOOM问题定位与解决方案的相关资料,文中通过代码介绍的非常详细,需要的朋... 目录一、OOM问题核心认知1.1 OOM定义与技术定位1.2 OOM常见类型及技术特征二、OOM问题定位工具

深度解析Python中递归下降解析器的原理与实现

《深度解析Python中递归下降解析器的原理与实现》在编译器设计、配置文件处理和数据转换领域,递归下降解析器是最常用且最直观的解析技术,本文将详细介绍递归下降解析器的原理与实现,感兴趣的小伙伴可以跟随... 目录引言:解析器的核心价值一、递归下降解析器基础1.1 核心概念解析1.2 基本架构二、简单算术表达

深度解析Java @Serial 注解及常见错误案例

《深度解析Java@Serial注解及常见错误案例》Java14引入@Serial注解,用于编译时校验序列化成员,替代传统方式解决运行时错误,适用于Serializable类的方法/字段,需注意签... 目录Java @Serial 注解深度解析1. 注解本质2. 核心作用(1) 主要用途(2) 适用位置3

Java MCP 的鉴权深度解析

《JavaMCP的鉴权深度解析》文章介绍JavaMCP鉴权的实现方式,指出客户端可通过queryString、header或env传递鉴权信息,服务器端支持工具单独鉴权、过滤器集中鉴权及启动时鉴权... 目录一、MCP Client 侧(负责传递,比较简单)(1)常见的 mcpServers json 配置

从原理到实战解析Java Stream 的并行流性能优化

《从原理到实战解析JavaStream的并行流性能优化》本文给大家介绍JavaStream的并行流性能优化:从原理到实战的全攻略,本文通过实例代码给大家介绍的非常详细,对大家的学习或工作具有一定的... 目录一、并行流的核心原理与适用场景二、性能优化的核心策略1. 合理设置并行度:打破默认阈值2. 避免装箱

Maven中生命周期深度解析与实战指南

《Maven中生命周期深度解析与实战指南》这篇文章主要为大家详细介绍了Maven生命周期实战指南,包含核心概念、阶段详解、SpringBoot特化场景及企业级实践建议,希望对大家有一定的帮助... 目录一、Maven 生命周期哲学二、default生命周期核心阶段详解(高频使用)三、clean生命周期核心阶

深入解析C++ 中std::map内存管理

《深入解析C++中std::map内存管理》文章详解C++std::map内存管理,指出clear()仅删除元素可能不释放底层内存,建议用swap()与空map交换以彻底释放,针对指针类型需手动de... 目录1️、基本清空std::map2️、使用 swap 彻底释放内存3️、map 中存储指针类型的对象

Java Scanner类解析与实战教程

《JavaScanner类解析与实战教程》JavaScanner类(java.util包)是文本输入解析工具,支持基本类型和字符串读取,基于Readable接口与正则分隔符实现,适用于控制台、文件输... 目录一、核心设计与工作原理1.底层依赖2.解析机制A.核心逻辑基于分隔符(delimiter)和模式匹

Java+AI驱动实现PDF文件数据提取与解析

《Java+AI驱动实现PDF文件数据提取与解析》本文将和大家分享一套基于AI的体检报告智能评估方案,详细介绍从PDF上传、内容提取到AI分析、数据存储的全流程自动化实现方法,感兴趣的可以了解下... 目录一、核心流程:从上传到评估的完整链路二、第一步:解析 PDF,提取体检报告内容1. 引入依赖2. 封装

深度解析Python yfinance的核心功能和高级用法

《深度解析Pythonyfinance的核心功能和高级用法》yfinance是一个功能强大且易于使用的Python库,用于从YahooFinance获取金融数据,本教程将深入探讨yfinance的核... 目录yfinance 深度解析教程 (python)1. 简介与安装1.1 什么是 yfinance?