LSTM网络的建立、训练及应用实例

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)和三个门控机制:

2. 网络架构

LSTM网络通常由以下层组成:

  1. 序列输入层:接收输入序列
  2. LSTM层:处理序列数据,学习长期依赖
  3. 全连接层:将LSTM输出映射到目标空间
  4. 输出层:回归或分类

3. 训练过程

  1. 前向传播:计算预测值
  2. 损失计算:比较预测值与真实值
  3. 反向传播:计算梯度
  4. 参数更新:使用优化算法(如Adam)更新权重

参考代码 lstm网络的建立、训练及应用实例 www.youwenfan.com/contentcss/112771.html

应用场景

1. 时间序列预测

2. 自然语言处理

3. 异常检测

4. 语音识别

常见问题解决

  1. 梯度消失/爆炸

    • 使用梯度裁剪:'GradientThreshold', 1
    • 使用LSTM而不是普通RNN
    • 添加批量归一化层
  2. 过拟合

    • 增加dropout层
    • 使用正则化
    • 增加训练数据
    • 提前停止
  3. 训练缓慢

    • 使用GPU加速:trainingOptions('adam', 'ExecutionEnvironment', 'gpu')
    • 减小批大小
    • 简化网络结构
  4. 预测效果不佳

    • 增加网络深度/宽度
    • 调整学习率
    • 增加训练时间
    • 改进数据预处理

扩展功能

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网络的完整解决方案,具有以下特点:

  1. 模块化设计:清晰的类结构,易于扩展
  2. 多应用场景:支持时间序列预测、文本生成和异常检测
  3. 灵活配置:可自定义网络结构和训练参数
  4. 可视化支持:内置训练进度和结果可视化
  5. 实用演示:包含三种典型应用场景的完整示例

 

专注于matlab/simulink,电子电路,编程