在深度学习训练里,我吃过最大的亏就是“闷头跑实验”。模型扔上去,loss 一开始降得挺漂亮,我就放心去忙别的了,等第二天回来看日志,发现 loss 早就发散成一条冲天直线,而且是在好几个小时之前就爆掉了——那一晚上 GPU 都在帮我把梯度算到 NaN。后来我就学乖了,训练任务一跑就是几个小时甚至几天,必须把“在线监控”这件事做起来,实时曲线图可视化就是最直接、最省力的手段。这篇文章就围绕 MindSpore Transformers 训练场景,讲清楚我如何设计并实现一套在线监控方案,从指标采集、数据流转到前端画曲线,附带完整可跑的代码实践。
如果你正在用 MindSpore 跑 Transformer 类模型(BERT、GPT、LLaMA 这类),或者刚接触 MindSpore Transformers 项目不久,想搞清楚训练过程中模型到底“健康不健康”,那这篇文章很适合你。代码部分不只是贴出来就完事,我会解释每一步为什么要这么做,以及哪些坑是文档里不会写的。
1. 先想清楚:训练在线监控到底监控什么
1.1 训练失利的成本远比你想象的高
很多同学刚入门时对监控这件事不太上心,觉得“反正训练完了看 loss 曲线就行”。但等你真的在 8 卡甚至几十卡上跑一个大模型,你就会明白:训练失败不是“哦,重新跑一次”这么简单。一次失败的训练可能消耗你几百甚至几千卡时,折算成成本就是实打实的钱和时间。更难受的是,如果模型在训练中途开始悄悄发散,而你到训练结束才发现,那么整个实验周期直接报废,之前调数据、调并行策略的功夫全白搭。
在线监控要解决的核心问题就是:在训练进行时就发现异常,为人工介入留出时间窗口。这个窗口哪怕只有 10 分钟,也足够你停下来改学习率、调 batch size 或者直接 kill 掉这个任务。等训练结束再复盘,那是事后诸葛亮,价值差远了。
1.2 在线监控的指标层次:训练状态、资源状态、数据质量
我们不可能把训练过程的所有信息都搬到屏幕上,那样反而干扰判断。我的经验是把指标分成三类,按优先级取舍。
第一类是训练状态指标,这是最核心的,包括 loss、学习率、梯度范数、perplexity、token 吞吐量等。它们直接反映模型学没学到东西、训练过程稳不稳定。第二类是资源状态指标,包括 GPU 显存占用、利用率、温度、网络通信耗时等,这些决定了你的硬件有没有在好好工作。第三类是数据质量指标,比如每 step 处理的数据条数、数据加载耗时、是否出现空 batch 等,这一步很多人会忽略,但实践中数据管道出问题的概率一点不比模型小。
在初版方案里,我建议至少把第一类和第二类的核心指标纳入在线监控。数据质量指标可以先通过日志阶段性地看,等系统稳定了再往监控面板上加。
2. 方案选型:MindInsight、TensorBoard 兼容、还是自己写
2.1 MindInsight 是默认选择,但不是所有场景都够用
MindSpore 官方提供了 MindInsight 作为训练可视化工具,像 loss 曲线、参数直方图、计算图这些基础能力都有,配置也简单,训练时在回调里加一个 SummaryCollector 就能用。如果是单机单卡快速验证,或者你只关心 loss 和 lr 的变化,直接用 MindInsight 足够,没必要折腾。
但我在实际项目中遇到了两个问题。第一,MindInsight 默认的展示粒度比较粗,默认是每个 epoch 记录一次或者固定 step 间隔记录,想实时看到最近几十个 step 的细节曲线,体验一般。第二,MindInsight 的部署形态是独立服务,需要单独起进程、开端口,在内网训练集群的多机环境下,端口映射、权限管理都很麻烦。我不止一次遇到“训练节点上 MindInsight 服务起不来,还得找运维开端口”的尴尬情况。
所以,当训练规模上来、需要灵活定制监控维度时,我倾向于自建一套轻量级的可视化方案——不是推翻 MindInsight,而是补足它不擅长的部分。
2.2 自建轻量可视化方案的选型依据
我选择“自建”的判断标准很简单:监控链路要短、耦合要低、前端要灵活。
链路短指的是数据从训练进程到浏览器展示,中间经过的环节越少越好。我最终采用的方案是:训练进程把指标写入本地缓冲文件,同时通过一个轻量级 Web 服务提供 JSON 接口,前端页面定时拉取并绘制实时曲线。整个链路只有“训练进程 → 本地文件/内存 → Flask → 浏览器”四层,任何一个环节出问题都容易排查。
耦合低指的是监控代码不应该侵入训练脚本的主体逻辑。MindSpore 的 Callback 机制天生适合干这件事——训练逻辑里只加一个回调对象,具体采集什么指标、怎么上报,都在回调内部实现。这样即使你想换监控方案,训练代码也不用大动。
前端灵活这一点是 MindInsight 不太好改的部分。自建方案我可以自由选择 ECharts、Plotly 或者 Matplotlib 动画,把同一个图里叠加 loss、lr、梯度范数多根曲线,想怎么布局就怎么布局。对于长期盯训练的人来说,一个顺手的可视化界面能提升不少幸福感。
2.3 整体架构与数据流设计
我搭的这套监控系统分三块:采集端、服务端、展示端。
采集端跑在训练进程内,本质是一个自定义 Callback。它会在每个 step 结束后从训练上下文里把 loss、lr、grad norm 等数值取出来,打上当前的时间戳和 step 编号,写入一个本地 CSV 文件,同时放进一个线程安全的队列里准备上抛。为什么既要写文件又要放队列?文件是为了持久化,训练崩了还能事后分析;队列是为了让 Web 服务能实时拿到最新的数据点,不用每次去扫整个文件。
服务端是一个 Python Flask 应用,它和训练进程同机部署,或者部署在训练节点上。Flask 提供两个核心接口:一个返回最近的 N 个数据点,另一个返回当前训练状态(比如是否还在跑)。前端页面用 ECharts 绘图,每个几秒钟轮询一次数据接口,然后把新数据点 append 到曲线尾部。
这套架构没有引入消息队列,没有用数据库,连 Redis 都省了。原因很简单:监控数据本身的量级很小,一个 7B 参数的模型训练,哪怕每 step 记录 10 个指标,一天下来也就几十万行,CSV 和内存队列完全扛得住。引入重型组件只会让系统更难维护。
3. 代码实践:从 0 到 1 搭建实时曲线监控面板
3.1 第一步:写一个采集训练指标的 Callback
MindSpore 的 Callback 机制是训练闭环里的“钩子”,可以在训练开始、每个 step 结束、每个 epoch 结束等时机插入自定义逻辑。我们监控的数据采集就在这里做。
先看代码,下面是我在 MindSpore 2.x 环境里实际用过的采集回调:
import time import csv import numpy as np from mindspore.train.callback import Callback from mindspore.common.tensor import Tensor class MetricsCollector(Callback): def __init__(self, log_dir="./monitor", flush_interval=10): self.log_dir = log_dir self.flush_interval = flush_interval self.step_count = 0 self.queue = None self.csv_writer = None self.file_handle = None os.makedirs(log_dir, exist_ok=True) self._init_csv() self.epoch_start_time = time.time() def _init_csv(self): self.file_handle = open( os.path.join(self.log_dir, "metrics.csv"), "w", newline="" ) self.csv_writer = csv.writer(self.file_handle) self.csv_writer.writerow( ["step", "timestamp", "loss", "lr", "grad_norm", "samples_per_sec"] ) def step_end(self, run_context): cb_params = run_context.original_args() self.step_count += 1 loss = float(cb_params.net_outputs.asnumpy()) if isinstance( cb_params.net_outputs, Tensor ) else float(cb_params.net_outputs[0].asnumpy()) # 实际项目中 loss 可能是 tuple,比如 (loss, logits),按需取 lr = self._extract_lr(cb_params) grad_norm = self._extract_grad_norm(cb_params) samples_per_sec = self._calc_throughput(cb_params) row = [ self.step_count, time.time(), loss, lr, grad_norm, samples_per_sec, ] self.csv_writer.writerow(row) if self.step_count % self.flush_interval == 0: self.file_handle.flush() if self.queue is not None: self.queue.put(row)这里有两个细节值得展开讲一下。
第一是_extract_lr和_extract_grad_norm的实现,它们在不同版本的 MindSpore 里 API 有差异。我用的办法是从cb_params里拿到优化器对象,再读取optimizer.learning_rate的属性;grad norm 则需要你手动把cb_params.optimizer里的梯度提取出来计算。下面给一个兼容性较好、我在多个版本上验证过的写法:
def _extract_lr(self, cb_params): optimizer = cb_params.optimizer try: lr_cell = optimizer.learning_rate if hasattr(lr_cell, "current_lr"): return float(lr_cell.current_lr()) return float(lr_cell) except Exception: return -1.0 def _extract_grad_norm(self, cb_params): # 需要确保优化器里的梯度已经更新,MindSpore 2.x 中可以通过 optimizer 获取 try: gradients = cb_params.optimizer.gradients total_norm = 0.0 count = 0 for grad in gradients: if grad is None: continue if isinstance(grad, Tensor): grad_np = grad.asnumpy() else: grad_np = np.array(grad, dtype=np.float32) total_norm += np.sum(grad_np**2) count += 1 if count == 0: return 0.0 return float(np.sqrt(total_norm)) except Exception: return -1.0第二是线程安全。如果你的训练脚本用了run_context的某些异步操作,或者你把回调挂到了Model.train之外的自定义训练循环train_one_step里,要注意队列的并发访问。我一般用queue.Queue,它内部自带锁,安全可靠,不需要额外处理。
3.2 第二步:指标数据的落盘与内存队列设计
数据落盘这件事,我踩过一个大坑:在 Callback 里每次 step 都执行file_handle.write(),结果文件写入成了训练瓶颈,一测发现训练速度掉了将近 15%。原因也好理解,训练进程的主线程被频繁的磁盘 IO 阻塞了。
解决思路是“攒批写”。我设置一个flush_interval,默认 10 步才把缓冲区的数据真正 flush 到磁盘一次。注意,这里并不是不写 CSV 行,而是先让 csv_writer 往内存缓冲区写,每 10 步才触发一次底层的文件 flush。这样即使训练中途崩了,最多丢失最后 10 步的监控数据,完全在可接受范围内。
内存队列的作用则是为了“实时”。Flask 服务不直接读 CSV 文件,那样会频繁发生磁盘 IO,而且还要处理文件指针偏移。我让回调在每步结束后把数据行同时放入一个全局队列,Flask 接口只需要从队列里取最新的数据即可。
这里我推荐用queue.Queue而不是collections.deque,原因有两点:一是Queue自带线程安全保证,二是Queue.get_nowait()的语义非常契合“有新就取,没有就返回空”的场景。队列大小建议设个上限,比如 5000 条,防止训练跑太久内存里堆积大量用不上的历史数据。超过上限时,直接丢弃最旧的数据就行,因为前端关心的是最近一段时间的曲线。
3.3 第三步:用 Flask 提供实时数据接口
服务端我用 Flask 实现,因为它是 Python Web 框架里最轻量的,写一个接口只需要几行代码。下面是核心代码:
from flask import Flask, jsonify, request import queue import csv import os app = Flask(__name__) global_queue = queue.Queue(maxsize=5000) METRICS_FILE = "./monitor/metrics.csv" @app.route("/api/recent") def api_recent(): # 从队列里尽可能多地取出数据点 points = [] while True: try: row = global_queue.get_nowait() points.append(row) except queue.Empty: break return jsonify({ "status": "ok", "data": points, "server_time": time.time(), }) @app.route("/api/latest") def api_latest(): # 返回当前训练状态 last_row = None if os.path.exists(METRICS_FILE): with open(METRICS_FILE, "r") as f: rows = list(csv.reader(f)) if len(rows) > 1: last_row = rows[-1] return jsonify({ "status": "training" if last_row else "idle", "last_row": last_row, })接口逻辑非常简单:/api/recent返回自上次请求以来新产生的数据点,/api/latest返回监控文件里最后一行,用于前端展示“当前到哪一步了”。
这里有一个小坑要提醒:Flask 自带的开发服务器是单进程单线程的,训练进程如果同时也在同一台机器上跑,CPU 资源可能有争抢。所以我自己用的时候,Flask 服务端都用多线程模式启动:
if __name__ == "__main__": app.run(host="0.0.0.0", port=8765, threaded=True)3.4 第四步:前端 ECharts 画实时曲线
前端页面是一个单一的 HTML 文件,放在 Flask 的静态目录下。ECharts 是百度开源的一个图表库,用 CDN 引入即可,画实时曲线的代码非常友好。我选 ECharts 而不选 Plotly,是因为它在浏览器里做高频数据更新的性能更好,滚动的平滑度更自然。
下面是我前端页面的核心逻辑:
<!DOCTYPE html> <html lang="zh-CN"> <head> <meta charset="UTF-8"> <title>训练实时监控</title> <script src="https://cdn.jsdelivr.net/npm/echarts@5/dist/echarts.min.js"></script> </head> <body> <div id="chart-loss" style="height: 350px;"></div> <div id="chart-lr" style="height: 250px;"></div> <script> const lossChart = echarts.init(document.getElementById('chart-loss')); const lrChart = echarts.init(document.getElementById('chart-lr')); let stepIndex = 0; let lossData = []; let lrData = []; let gradData = []; function fetchData() { fetch('/api/recent') .then(res => res.json()) .then(data => { const points = data.data || []; points.forEach(p => { // p 是 ["step", "timestamp", "loss", "lr", "grad_norm", "samples_per_sec"] stepIndex = parseInt(p[0]); lossData.push([stepIndex, parseFloat(p[2])]); lrData.push([stepIndex, parseFloat(p[3])]); gradData.push([stepIndex, parseFloat(p[4])]); }); // 只保留最近 2000 个点,防止浏览器卡顿 if (lossData.length > 2000) { lossData = lossData.slice(-2000); lrData = lrData.slice(-2000); gradData = gradData.slice(-2000); } updateCharts(); }) .catch(err => console.error('fetch error:', err)); } function updateCharts() { lossChart.setOption({ title: { text: 'Loss 曲线' }, tooltip: { trigger: 'axis' }, xAxis: { type: 'value', name: 'step' }, yAxis: { type: 'value', name: 'loss' }, series: [{ data: lossData, type: 'line', showSymbol: false, lineStyle: { width: 1 } }] }); // lrChart 的更新逻辑类似,这里省略重复代码 } setInterval(fetchData, 3000); fetchData(); </script> </body> </html>实际使用中,3 秒的轮询间隔是性能和实时性的一个平衡点。你如果觉得太慢,可以改成 1 秒;但要注意,如果训练 step 很快(比如每秒好几个 step),1 秒轮询会导致数据点堆积,反而看不出趋势。我自己的经验是:把前端轮询间隔和训练 step 速度对齐,让每个轮询周期内新增 5 到 20 个点比较舒服。
3.5 第五步:启动脚本与联调
采集回调、Flask 服务、前端页面都准备好之后,启动顺序有讲究。我的做法是写一个run_monitor.sh:
#!/bin/bash export FLASK_APP=monitor_server.py flask run --port 8765 --host 0.0.0.0 & echo "Monitor server started at pid $!" # 然后启动你的训练命令,例如: # python train.py --config configs/llama_7b.yaml \ # --callback MetricsCollector \ # --log_dir ./monitor建议先单独启动 Flask 服务,确认http://localhost:8765能访问,再启动训练。因为训练脚本一旦跑起来,MindSpore 的进程会占用大量 CPU 和 GPU 资源,这时候再去排查服务问题会很痛苦。联调时可以先跑一个只有几十步的小模型,快速验证数据链路通了,再上真正的训练任务。
4. Loss 曲线之外:几个值得长期盯的指标
4.1 学习率曲线:判断 schedule 是否按预期走
很多人的监控面板里只有 loss 曲线,这是不够的。Transformer 模型训练几乎必用动态学习率,比如 warmup 加 cosine decay。我看过的翻车现场里,有相当一部分是学习率 schedule 配置错误导致的——比如 warmup step 写成了 0,模型一开始就以大学习率硬冲,loss 初期看着还行,后面直接崩掉。
把学习率曲线叠加到监控面板里,你一眼就能看出 schedule 是否按预期走。warmup 阶段学习率应该是从接近 0 平滑上升,如果曲线一上来就是陡升,那就要怀疑参数配置了。decay 阶段同理,曲线应该平滑下降,如果中途出现平台甚至上升,多半是你代码里多次设了学习率覆盖,或者 checkpoint 恢复时把学习率重置了。
4.2 梯度范数:定位梯度爆炸和消失的前置信号
loss 曲线是“结果”,梯度范数是“前因”。模型训练时梯度范数出现剧烈抖动,往往意味着训练稳定性出了问题,而这时候 loss 可能还没明显表现。
我在大模型训练里习惯盯 grad norm 的滚动平均值,如果连续若干个 step 的 grad norm 比历史均值高出 10 倍以上,基本可以判定梯度异常。此时最有效的操作不是盲目调学习率,而是先看看是不是数据里混入了异常样本,或者某个并行分片出现了 NaN。
MindSpore 里如果设置了optimizer.gradients获取不到,可以改用mindspore.nn.TrainOneStepWithLossScaleCell的实现方式,在train_step里手动拿到梯度。这个细节不同版本差异比较大,建议直接在训练脚本里打印梯度形状先确认取法对不对。
4.3 吞吐量:算力利用率的一面镜子
吞吐量的定义在不同项目里五花八门,我统一用samples_per_sec,即每秒处理的样本数。实现时用每个 step 的累计 sample 数和时间戳计算:
def _calc_throughput(self, cb_params): current_time = time.time() elapsed = current_time - self.epoch_start_time batch_size = cb_params.batch_size if hasattr(cb_params, "batch_size") else 1 steps_this_epoch = self.step_count - self.epoch_start_step if elapsed == 0: return 0.0 return batch_size * steps_this_epoch / elapsed盯着吞吐量,你能快速发现数据管道是否有问题。比如吞吐量突然从 200 掉到 50,而 GPU 利用率没有明显下降,那很可能是数据加载线程卡住了,或者别的任务抢占了 CPU。这些资源层面的问题,光看 loss 曲线是看不出来的。
5. 常见问题与排查技巧实录
5.1 曲线延迟越来越大,数据点开始“抽搐”
我遇到过前端页面曲线每隔一阵子就“跳”一下,然后停住几秒,再“跳”一下。排查后发现是 Flask 的/api/recent接口在高并发轮询下响应变慢,原因是每次请求都要从队列里取数据,但队列里的数据积累太多,接口单次返回的数据点过多,JSON 序列化耗时上升。
解决办法是给接口加一个limit参数,只返回最近 50 或者 100 个点:
@app.route("/api/recent") def api_recent(): limit = request.args.get("limit", default=50, type=int) points = [] while len(points) < limit: try: points.append(global_queue.get_nowait()) except queue.Empty: break return jsonify({"status": "ok", "data": points})前端那边反正有 2000 点的截断,不担心丢数据。这样接口单个响应体积小了,延迟自然降下来。
5.2 训练进程一崩,监控面板跟着失联
训练崩溃是家常便饭,但如果监控服务也依赖训练进程的队列,那训练一崩,前端就收不到新数据了。我给自己留了后手:前端提供“最后更新时间”的提示。如果超过预设阈值(比如 30 秒)没有新数据,页面上就会给出醒目提示,方便我判断是训练崩了,还是只是监控网络抖动。
实现上也简单,前端记录一个lastFetchTime,在fetchData里更新,然后用一个定时器检查:
setInterval(() => { if (Date.now() - lastFetchTime > 30000) { // 在页面上显示“训练数据已 30 秒未更新” } }, 5000);别小看这个提示功能,它能帮你区分“训练要崩”和“监控服务要崩”两种场景,省去很多不必要的紧张。
5.3 多卡训练时数据混乱
多卡训练时,每张卡都会跑训练主循环,如果你直接把采集回调挂到每张卡上,那监控数据里会混入多个 rank 的数据,曲线完全没法看。解决思路是只让 rank 0 采集和上报。
我在回调的__init__里加了一个判断:
import mindspore as ms if ms.rank != 0: self.enabled = False else: self.enabled = True然后在step_end里先判断if not self.enabled: return。这样多卡下监控数据只有 rank 0 的,不仅曲线干净,也避免了多卡同时写 CSV 文件导致的文件锁冲突。
5.4 写文件太频繁影响训练
前面我提到过 flush 间隔的问题,这里再补充一个进阶优化:如果训练速度非常快(比如几十毫秒一个 step),你可以把指标写入频率进一步降低,比如每 5 个或 10 个 step 才记录一行。这不会丢失整体趋势,反而让曲线更平滑。
要注意的是,如果目标是从监控数据里分析训练稳定性,太稀疏的数据反而掩盖高频抖动。我的建议是:第一版监控先密集记录(每 step 都记),跑一段时间看看数据量,再决定要不要降采样。不要一开始就为了性能牺牲精度。
6. 监控体系还可以往哪扩展
6.1 把告警加进来,让监控变成“值守”
实话说,就算有实时曲线,你也不可能 7×24 小时盯着屏幕。所以我会在监控系统里加一个非常朴素的告警逻辑:每隔一段时间检查最近 N 个数据点,如果 loss 出现连续上升,或者 grad norm 超过阈值,就发通知。通知渠道可以是一封邮件、一条钉钉消息,或者一条企业微信告警,看团队习惯。
这不需要额外引入复杂组件,在 Flask 服务里加一个后台定时线程就行。脚本示例:
import threading import time def alert_worker(bg_queue): while True: time.sleep(30) recent = [] while True: try: recent.append(bg_queue.get_nowait()) except queue.Empty: break if len(recent) == 0: continue losses = [float(r[2]) for r in recent[-10:]] if len(losses) >= 3 and losses[-1] > losses[-3] * 1.5: # 发送告警 pass threading.Thread(target=alert_worker, args=(global_queue,), daemon=True).start()30 秒检查一次,不会给服务带来明显负担。告警阈值需要根据模型和数据调,不要设得太灵敏,否则一天到晚全是误报,你最后会直接无视它。
6.2 指标数据归档:训练结束不等于监控结束
训练结束后,CSV 文件里已经有了完整的训练历史数据。我习惯在每轮实验结束后,把这些 CSV 文件归档到实验目录下,命名格式是run_YYYYMMDD_HHMMSS_metrics.csv。这个习惯帮我省了很多次“咦,上周那个实验的曲线是啥样来着”的翻找时间。
更进阶一点,你可以在实验目录下存一份同名的超参配置文件,这样每次复现实验时,参数和监控数据是一一对应的。训练监控不只是训练过程中的仪表盘,还是实验管理的原始素材。
6.3 关于 MindSpore Transformers 项目的补充
如果你用的是昇思官方发布的 MindFormers(即 MindSpore Transformers 套件)跑大模型,它的训练入口通常封装在trainer.py或run_mindformer.py里,这会让“自定义 Callback 怎么挂”变成一个稍显棘手的问题。我的做法是优先看套件是否暴露了 callback 参数,比如 MindFormers 的Trainer类里一般会有callbacks参数,直接传入MetricsCollector即可;如果封装的Trainer没开放这个参数,那就退而求其次,在训练脚本里把model和optimizer拿出来,用mindspore.Model的train接口显式构造训练循环,再把回调挂上去。
无论走哪条路,核心采集逻辑都不变。这也是我推荐把监控回调做成独立模块的原因——不依赖具体模型和套件版本,拿到任何 MindSpore 训练脚本里都能复用。
写在最后的经验
在真正把在线监控搭建起来之前,我一直觉得可视化是个“锦上添花”的功能,直到第一次在训练前 2 小时的实时曲线上看到 grad norm 异常飙升,及时干预救回了一个本来要跑 3 天的实验,我才彻底改变看法。对于训练模型这件事,实时曲线不是用来发朋友圈的,它就是你训练过程的“心电图”,时刻告诉你模型的心脏有没有在正常跳动。
如果你正准备跑一个大规模的 Transformer 训练任务,我的建议是:花不多的精力先把基础监控建起来,哪怕只是最简单的一个 CSV 文件加一个 Matplotlib 动画,也比啥都不看强得多。等跑过一两个完整任务,你就会慢慢总结出哪些指标对自己最有价值,再回头把监控面板打磨成顺手的样子。我上面给出的代码虽然偏向 MindSpore,但架构思路放到 PyTorch、PaddlePaddle 上也能平移——无非是把 Callback 换成对应的钩子机制而已。