LSTM网络的建立、训练及应用实例
LSTM(长短期记忆)网络实现,包含网络构建、训练和多个应用实例(时间序列预测、文本生成、异常检测)。
classdef LSTMNetwork
% LSTM神经网络实现
% 支持时间序列预测、文本生成和异常检测
properties
layers; % 网络层结构
options; % 训练选项
net; % 训练好的网络
inputSize; % 输入维度
hiddenSize; % 隐藏单元数
outputSize; % 输出维度
sequenceLength; % 序列长度
vocabularySize; % 词汇表大小(文本应用时)
charToIndex; % 字符到索引映射(文本应用时)
indexToChar; % 索引到字符映射(文本应用时)
end
methods
function obj = LSTMNetwork(inputSize, hiddenSize, outputSize, sequenceLength)
% 构造函数
% 输入:
% inputSize - 输入特征维度
% hiddenSize - LSTM隐藏单元数量
% outputSize - 输出维度
% sequenceLength - 序列长度
obj.inputSize = inputSize;
obj.hiddenSize = hiddenSize;
obj.outputSize = outputSize;
obj.sequenceLength = sequenceLength;
obj.vocabularySize = 0;
end
function obj = buildNetwork(obj)
% 构建LSTM网络架构
layers = [
sequenceInputLayer(obj.inputSize) % 序列输入层
lstmLayer(obj.hiddenSize, 'OutputMode', 'sequence') % LSTM层
% 可以添加更多LSTM层
% lstmLayer(obj.hiddenSize, 'OutputMode', 'sequence')
fullyConnectedLayer(obj.outputSize) % 全连接层
regressionLayer % 回归层(用于预测)
];
obj.layers = layers;
end
function obj = buildTextGenerationNetwork(obj, vocabSize)
% 构建用于文本生成的LSTM网络
obj.vocabularySize = vocabSize;
layers = [
sequenceInputLayer(obj.inputSize)
lstmLayer(obj.hiddenSize, 'OutputMode', 'last') % 只输出最后一个时间步
fullyConnectedLayer(vocabSize)
softmaxLayer
classificationLayer
];
obj.layers = layers;
end
function obj = setTrainingOptions(obj, options)
% 设置训练选项
% 输入: options - 包含训练参数的结构体
defaultOptions = trainingOptions('adam', ...
'MaxEpochs', 100, ...
'GradientThreshold', 1, ...
'InitialLearnRate', 0.005, ...
'LearnRateSchedule', 'piecewise', ...
'LearnRateDropPeriod', 125, ...
'LearnRateDropFactor', 0.2, ...
'Verbose', 0, ...
'Plots', 'training-progress');
if nargin > 1
obj.options = options;
else
obj.options = defaultOptions;
end
end
function obj = trainNetwork(obj, XTrain, YTrain)
% 训练LSTM网络
% 输入:
% XTrain - 训练输入数据 (cell数组,每个元素是一个序列)
% YTrain - 训练目标数据 (cell数组)
if isempty(obj.layers)
obj.buildNetwork();
end
if isempty(obj.options)
obj.setTrainingOptions();
end
obj.net = trainNetwork(obj.layers, XTrain, YTrain, obj.options);
end
function predictions = predict(obj, XTest)
% 使用训练好的网络进行预测
% 输入: XTest - 测试输入数据
% 输出: predictions - 预测结果
if isempty(obj.net)
error('网络尚未训练');
end
predictions = classify(obj.net, XTest); % 分类任务
% 或使用 predict 用于回归任务
% predictions = predict(obj.net, XTest);
end
function generateText(obj, seedSequence, numChars)
% 生成文本
% 输入:
% seedSequence - 种子序列(字符索引)
% numChars - 生成字符数量
if isempty(obj.net)
error('网络尚未训练');
end
generated = seedSequence;
for i = 1:numChars
% 准备输入(最近sequenceLength个字符)
seqLen = min(obj.sequenceLength, length(generated));
inputSeq = generated(end-seqLen+1:end);
% 转换为one-hot编码
inputData = zeros(obj.vocabularySize, seqLen);
for j = 1:seqLen
idx = inputSeq(j);
if idx > 0 && idx <= obj.vocabularySize
inputData(idx, j) = 1;
end
end
% 预测下一个字符
[pred, scores] = classify(obj.net, inputData);
[~, nextIdx] = max(scores);
% 添加到生成序列
generated(end+1) = nextIdx;
end
% 将索引转换为字符
chars = '';
for i = 1:length(generated)
idx = generated(i);
if idx > 0 && idx <= length(obj.indexToChar)
chars = [chars, obj.indexToChar(idx)];
end
end
fprintf('生成文本:\n%s\n', chars);
end
function detectAnomalies(obj, testData, threshold)
% 异常检测
% 输入:
% testData - 测试数据
% threshold - 异常阈值(重构误差)
% 输出: anomalies - 异常检测结果
if isempty(obj.net)
error('网络尚未训练');
end
% 计算重构误差
reconstructed = predict(obj.net, testData);
errors = sqrt(sum((reconstructed - testData).^2, 2));
% 标记异常
anomalies = errors > threshold;
% 可视化结果
figure;
plot(errors, 'b-');
hold on;
plot(find(anomalies), errors(anomalies), 'ro');
hold off;
yline(threshold, 'r--', 'Threshold');
title('异常检测结果');
xlabel('样本索引');
ylabel('重构误差');
legend('重构误差', '异常点', '阈值');
fprintf('检测到 %d 个异常点\n', sum(anomalies));
end
function runTimeSeriesPredictionDemo()
% 时间序列预测演示
% 生成示例数据(正弦波加噪声)
t = 0:0.1:10;
data = sin(t) + 0.1*randn(size(t));
% 准备训练数据
sequenceLength = 10;
XTrain = {};
YTrain = {};
for i = 1:(length(data) - sequenceLength)
XTrain{end+1} = data(i:i+sequenceLength-1);
YTrain{end+1} = data(i+sequenceLength);
end
% 创建并训练LSTM网络
lstm = LSTMNetwork(1, 50, 1, sequenceLength);
lstm = lstm.buildNetwork();
lstm = lstm.setTrainingOptions();
lstm = lstm.trainNetwork(XTrain, YTrain);
% 预测
testInput = data(end-sequenceLength+1:end);
prediction = predict(lstm, {testInput});
% 可视化结果
figure;
plot(t, data, 'b-', 'DisplayName', '原始数据');
hold on;
plot(t(end)+0.1, prediction, 'ro', 'MarkerSize', 10, 'DisplayName', '预测值');
plot([t(end), t(end)+0.1], [data(end), prediction], 'r--', 'DisplayName', '预测连接');
title('时间序列预测');
xlabel('时间');
ylabel('值');
legend;
grid on;
end
function runTextGenerationDemo()
% 文本生成演示
% 加载示例文本(莎士比亚作品)
filename = 'shakespeare.txt';
if ~exist(filename, 'file')
websave(filename, 'https://www.gutenberg.org/files/100/100-0.txt');
end
textData = fileread(filename);
% 预处理文本
textData = lower(textData);
textData = regexprep(textData, '[^a-z .]', ''); % 移除非字母字符
uniqueChars = unique(textData);
vocabSize = length(uniqueChars);
% 创建字符映射
charToIndex = containers.Map(uniqueChars, 1:vocabSize);
indexToChar = containers.Map(1:vocabSize, uniqueChars);
% 准备训练数据
sequenceLength = 20;
XTrain = {};
YTrain = {};
for i = 1:(length(textData) - sequenceLength)
seq = textData(i:i+sequenceLength-1);
target = textData(i+sequenceLength);
% 转换为索引
inputIdx = zeros(1, sequenceLength);
for j = 1:sequenceLength
inputIdx(j) = charToIndex(seq(j));
end
% 转换为one-hot编码
inputData = zeros(vocabSize, sequenceLength);
for j = 1:sequenceLength
inputData(inputIdx(j), j) = 1;
end
XTrain{end+1} = inputData;
YTrain{end+1} = charToIndex(target);
end
% 创建并训练LSTM网络
lstm = LSTMNetwork(vocabSize, 128, vocabSize, sequenceLength);
lstm = lstm.buildTextGenerationNetwork(vocabSize);
lstm.charToIndex = charToIndex;
lstm.indexToChar = indexToChar;
% 设置训练选项
options = trainingOptions('adam', ...
'MaxEpochs', 20, ...
'MiniBatchSize', 64, ...
'GradientThreshold', 1, ...
'InitialLearnRate', 0.01, ...
'Verbose', 0, ...
'Plots', 'training-progress');
lstm = lstm.setTrainingOptions(options);
% 训练网络(可能需要较长时间)
fprintf('开始训练文本生成模型...\n');
lstm = lstm.trainNetwork(XTrain, YTrain);
% 生成文本
seedText = "to be or not to be";
seedSeq = zeros(1, length(seedText));
for i = 1:length(seedText)
if isKey(charToIndex, seedText(i))
seedSeq(i) = charToIndex(seedText(i));
else
seedSeq(i) = 1; % 未知字符用空格代替
end
end
fprintf('种子文本: %s\n', seedText);
lstm.generateText(seedSeq, 200);
end
function runAnomalyDetectionDemo()
% 异常检测演示
% 生成正常数据(二维高斯分布)
mu = [0, 0];
sigma = [1, 0.5; 0.5, 1];
normalData = mvnrnd(mu, sigma, 1000);
% 添加异常点
anomalyData = [5, 5; 6, -4; -5, 6; -7, -7];
allData = [normalData; anomalyData];
% 准备训练数据(仅使用正常数据)
sequenceLength = 1; % 使用单点作为输入
XTrain = {};
YTrain = {};
for i = 1:size(normalData, 1)
XTrain{end+1} = normalData(i, :)';
YTrain{end+1} = normalData(i, :)';
end
% 创建并训练LSTM网络(自编码器)
lstm = LSTMNetwork(2, 10, 2, sequenceLength);
lstm.layers = [
sequenceInputLayer(2)
lstmLayer(10, 'OutputMode', 'last')
fullyConnectedLayer(2)
regressionLayer
];
lstm = lstm.setTrainingOptions();
lstm = lstm.trainNetwork(XTrain, YTrain);
% 检测异常
lstm.detectAnomalies(allData', 1.5); % 阈值设为1.5
end
end
end
使用示例
1. 时间序列预测
% 创建LSTM网络
inputSize = 1; % 单变量时间序列
hiddenSize = 50; % 隐藏单元数量
outputSize = 1; % 预测单步
sequenceLength = 10; % 使用10个时间点预测下一个点
lstm = LSTMNetwork(inputSize, hiddenSize, outputSize, sequenceLength);
% 构建网络
lstm = lstm.buildNetwork();
% 设置训练选项
options = trainingOptions('adam', ...
'MaxEpochs', 100, ...
'GradientThreshold', 1, ...
'InitialLearnRate', 0.005, ...
'Verbose', 0, ...
'Plots', 'training-progress');
lstm = lstm.setTrainingOptions(options);
% 准备数据(示例:正弦波)
t = 0:0.1:20;
data = sin(t) + 0.2*randn(size(t));
% 创建输入输出序列
XTrain = {};
YTrain = {};
for i = 1:(length(data) - sequenceLength)
XTrain{end+1} = data(i:i+sequenceLength-1);
YTrain{end+1} = data(i+sequenceLength);
end
% 训练网络
lstm = lstm.trainNetwork(XTrain, YTrain);
% 预测未来值
testInput = data(end-sequenceLength+1:end);
prediction = predict(lstm, {testInput});
% 可视化结果
figure;
plot(t, data, 'b-', 'DisplayName', '原始数据');
hold on;
plot(t(end)+0.1, prediction, 'ro', 'MarkerSize', 10, 'DisplayName', '预测值');
plot([t(end), t(end)+0.1], [data(end), prediction], 'r--', 'DisplayName', '预测连接');
title('时间序列预测');
xlabel('时间');
ylabel('值');
legend;
grid on;
2. 文本生成
% 创建LSTM网络
lstm = LSTMNetwork(0, 128, 0, 20); % 参数将在构建时设置
% 加载文本数据
textData = fileread('shakespeare.txt');
textData = lower(textData);
textData = regexprep(textData, '[^a-z .]', '');
% 创建字符映射
uniqueChars = unique(textData);
vocabSize = length(uniqueChars);
charToIndex = containers.Map(uniqueChars, 1:vocabSize);
indexToChar = containers.Map(1:vocabSize, uniqueChars);
% 准备训练数据
sequenceLength = 20;
XTrain = {};
YTrain = {};
for i = 1:(length(textData) - sequenceLength)
seq = textData(i:i+sequenceLength-1);
target = textData(i+sequenceLength);
inputData = zeros(vocabSize, sequenceLength);
for j = 1:sequenceLength
char = seq(j);
if isKey(charToIndex, char)
inputData(charToIndex(char), j) = 1;
end
end
XTrain{end+1} = inputData;
YTrain{end+1} = charToIndex(target);
end
% 构建文本生成网络
lstm.inputSize = vocabSize;
lstm.outputSize = vocabSize;
lstm.sequenceLength = sequenceLength;
lstm = lstm.buildTextGenerationNetwork(vocabSize);
lstm.charToIndex = charToIndex;
lstm.indexToChar = indexToChar;
% 设置训练选项
options = trainingOptions('adam', ...
'MaxEpochs', 20, ...
'MiniBatchSize', 64, ...
'GradientThreshold', 1, ...
'InitialLearnRate', 0.01, ...
'Verbose', 0, ...
'Plots', 'training-progress');
lstm = lstm.setTrainingOptions(options);
% 训练网络
lstm = lstm.trainNetwork(XTrain, YTrain);
% 生成文本
seedText = "to be or not to be";
seedSeq = zeros(1, length(seedText));
for i = 1:length(seedText)
if isKey(charToIndex, seedText(i))
seedSeq(i) = charToIndex(seedText(i));
else
seedSeq(i) = 1; % 未知字符用空格代替
end
end
fprintf('种子文本: %s\n', seedText);
lstm.generateText(seedSeq, 200);
3. 异常检测
% 创建LSTM网络(自编码器)
lstm = LSTMNetwork(2, 10, 2, 1); % 输入和输出都是2维
lstm.layers = [
sequenceInputLayer(2)
lstmLayer(10, 'OutputMode', 'last')
fullyConnectedLayer(2)
regressionLayer
];
% 设置训练选项
lstm = lstm.setTrainingOptions();
% 生成数据
mu = [0, 0];
sigma = [1, 0.5; 0.5, 1];
normalData = mvnrnd(mu, sigma, 1000);
anomalyData = [5, 5; 6, -4; -5, 6; -7, -7];
allData = [normalData; anomalyData];
% 准备训练数据(仅正常数据)
XTrain = {};
YTrain = {};
for i = 1:size(normalData, 1)
XTrain{end+1} = normalData(i, :)';
YTrain{end+1} = normalData(i, :)';
end
% 训练网络
lstm = lstm.trainNetwork(XTrain, YTrain);
% 检测异常
lstm.detectAnomalies(allData', 1.5); % 阈值设为1.5
运行演示
% 运行时间序列预测演示
LSTMNetwork.runTimeSeriesPredictionDemo();
% 运行文本生成演示(需要下载文本文件)
LSTMNetwork.runTextGenerationDemo();
% 运行异常检测演示
LSTMNetwork.runAnomalyDetectionDemo();
LSTM网络原理详解
1. LSTM细胞结构
LSTM的核心是细胞状态(cell state)和三个门控机制:
-
遗忘门(Forget Gate):决定从细胞状态中丢弃哪些信息
-
输入门(Input Gate):决定哪些新信息存入细胞状态
-
细胞状态更新:
-
输出门(Output Gate):决定输出什么
2. 网络架构
LSTM网络通常由以下层组成:
- 序列输入层:接收输入序列
- LSTM层:处理序列数据,学习长期依赖
- 全连接层:将LSTM输出映射到目标空间
- 输出层:回归或分类
3. 训练过程
- 前向传播:计算预测值
- 损失计算:比较预测值与真实值
- 反向传播:计算梯度
- 参数更新:使用优化算法(如Adam)更新权重
参考代码 lstm网络的建立、训练及应用实例 www.youwenfan.com/contentcss/112771.html
应用场景
1. 时间序列预测
- 股票价格预测
- 天气预测
- 电力负荷预测
- 销售预测
2. 自然语言处理
- 机器翻译
- 文本摘要
- 情感分析
- 聊天机器人
3. 异常检测
- 网络入侵检测
- 设备故障预测
- 金融欺诈检测
- 医疗异常诊断
4. 语音识别
- 语音转文本
- 说话人识别
- 语音情感分析
常见问题解决
-
梯度消失/爆炸:
- 使用梯度裁剪:
'GradientThreshold', 1 - 使用LSTM而不是普通RNN
- 添加批量归一化层
- 使用梯度裁剪:
-
过拟合:
- 增加dropout层
- 使用正则化
- 增加训练数据
- 提前停止
-
训练缓慢:
- 使用GPU加速:
trainingOptions('adam', 'ExecutionEnvironment', 'gpu') - 减小批大小
- 简化网络结构
- 使用GPU加速:
-
预测效果不佳:
- 增加网络深度/宽度
- 调整学习率
- 增加训练时间
- 改进数据预处理
扩展功能
1. 多变量时间序列预测
% 创建多变量LSTM
inputSize = 3; % 3个输入特征
lstm = LSTMNetwork(inputSize, 64, 1, 20); % 预测1个输出
% 准备数据(3个相关时间序列)
% XTrain: 20×3矩阵(20个时间步,3个特征)
% YTrain: 标量值
2. 序列到序列学习
% 编码器-解码器结构
encoderLayers = [
sequenceInputLayer(inputSize)
lstmLayer(hiddenSize, 'OutputMode', 'last')
];
decoderLayers = [
sequenceInputLayer(outputSize)
lstmLayer(hiddenSize, 'OutputMode', 'sequence')
fullyConnectedLayer(outputSize)
softmaxLayer
classificationLayer
];
3. 迁移学习
% 加载预训练模型
net = load('pretrained_lstm.mat').net;
% 替换最后一层
layers = net.Layers;
layers(end) = fullyConnectedLayer(newOutputSize);
layers(end+1) = regressionLayer;
% 冻结部分层
trainableLayers = layers(2:end); % 冻结第一层
总结
本MATLAB实现提供了LSTM网络的完整解决方案,具有以下特点:
- 模块化设计:清晰的类结构,易于扩展
- 多应用场景:支持时间序列预测、文本生成和异常检测
- 灵活配置:可自定义网络结构和训练参数
- 可视化支持:内置训练进度和结果可视化
- 实用演示:包含三种典型应用场景的完整示例