如何清除MATLAB中Lambda函数内的Persistent变量以重置训练绘图
MATLAB神经网络训练中OutputFcn持久化变量重置问题解决方法
问题场景
在使用MATLAB的trainNetwork训练神经网络时,通过Lambda函数将plot_train_loss作为OutputFcn绘制训练进度,首次训练效果正常,但再次训练时,函数内的persistent变量(train_iteration、train_loss等)会保留上一次训练的数据,导致绘图初始显示历史训练记录。
训练配置代码
options = trainingOptions('sgdm', ... 'MaxEpochs',num_epochs,... 'InitialLearnRate',learning_rate, ... 'ValidationData',imdsTest, ... 'ValidationFrequency',validationFrequency, ... 'Verbose',false, ... 'Plots','none',... 'OutputFcn',@(info)plot_train_loss(info,3), ... 'MiniBatchSize', mini_batch_size); [net,info] = trainNetwork(imdsTrain,my_net,options);
plot_train_loss核心代码片段
function stop = plot_train_loss(info,N) % initialize persitent variable in order to save the past values persistent train_iteration persistent train_loss persistent train_accuracy % Keep track of the best validation accuracy and the number of validations for which % there has not been an improvement of the accuracy. persistent bestValAccuracy persistent valLag global training_figure; global num_iterations; stop = false; % check in which state the training is % if start: do nothing if info.State == "start" info.Iteration = {}; info.TrainingLoss = {}; info.TrainingAccuracy = {}; bestValAccuracy = 0; valLag = 0; return end % assign values of training to variable train_iteration(info.Iteration) = info.Iteration; train_loss(info.Iteration) = info.TrainingLoss; train_accuracy(info.Iteration) = info.TrainingAccuracy; % plot training progress t = tiledlayout(2,1); nexttile; plot(train_iteration,train_loss,'g'); xlim([0 num_iterations]); max_train_loss = max(train_loss(:)) + 0.2; ylim([0 max_train_loss]); grid on; title('Training loss'); nexttile; plot(train_iteration,train_accuracy,'b'); xlim([0 num_iterations]); ylim([0 100]); grid on; title('Training accuracy');
解决方法
核心思路是利用训练启动时的info.State == "start"状态,在此阶段重置所有persistent变量,保证每次新训练开始时变量都是空的,同一次训练周期内变量持续保留数据。
修改后的plot_train_loss关键代码
function stop = plot_train_loss(info,N) % 声明持久化变量 persistent train_iteration persistent train_loss persistent train_accuracy persistent bestValAccuracy persistent valLag global training_figure; global num_iterations; stop = false; % 训练开始时重置所有持久化变量 if info.State == "start" % 清空所有persistent变量,为新训练初始化 train_iteration = []; train_loss = []; train_accuracy = []; bestValAccuracy = 0; valLag = 0; return end % 赋值训练数据到持久化变量 train_iteration(info.Iteration) = info.Iteration; train_loss(info.Iteration) = info.TrainingLoss; train_accuracy(info.Iteration) = info.TrainingAccuracy; % 绘图逻辑保持不变 t = tiledlayout(2,1); nexttile; plot(train_iteration,train_loss,'g'); xlim([0 num_iterations]); max_train_loss = max(train_loss(:)) + 0.2; ylim([0 max_train_loss]); grid on; title('Training loss'); nexttile; plot(train_iteration,train_accuracy,'b'); xlim([0 num_iterations]); ylim([0 100]); grid on; title('Training accuracy');
额外说明
- 原代码中修改
info.Iteration等输入参数属性是无效的,MATLAB函数参数默认按值传递,修改不会影响训练过程的原始info对象,因此直接删除这部分冗余代码即可。 - 若需要在训练前强制清除函数的持久化变量,可在调用
trainNetwork前执行clear plot_train_loss,但这种方法灵活性不足,优先推荐在函数内部利用start状态处理。
内容的提问来源于stack exchange,提问作者beinando
相关产品推荐
相关产品推荐

