WebSocket无法实时获取Keras训练stdout日志问题求助
问题描述
我只用过两次WebSocket,肯定漏了些简单细节,或者当前实现太复杂了(另外我需要跟踪连接,所以用了连接管理器)。
我想把FastAPI上运行的机器学习程序的所有stdout输出通过WebSocket流式传输,在网页上展示Keras的stdout日志。普通字符串(比如"starting training...")流式传输正常,但运行model.fit时,直到训练完成才会把内容发去WebSocket,所有Keras日志会一次性发送,像是被缓冲区卡住了。
WebSocket路由
@router.websocket('/ws') async def start_websocket_logging(websocket: WebSocket): await websocket_helper.start_websocket_logging(websocket)
websocket_helper.py
connection_manager = ConnectionManager() socket_sleep = 0.5 async def redirect_std_out(websocket): """ redirects the std output to the websocket """ stdout_buffer = StringIO() async def send_stdout(): try: while True: await asyncio.sleep(socket_sleep) data = stdout_buffer.getvalue() if data: await websocket.send_text(data.rstrip('\n')) stdout_buffer.seek(0) stdout_buffer.truncate() except (asyncio.CancelledError, websockets.ConnectionClosedError): # nothing needs to happen here pass sys.stdout = stdout_buffer return asyncio.create_task(send_stdout()) async def stop_std_redirect(task: asyncio.Task): """ resets the stdout to its original function and cancels the task sending data to the websocket """ sys.stdout = sys.__stdout__ task.cancel() await asyncio.gather(task, return_exceptions=True) async def start_websocket_logging(websocket): """ starts the websocket between this service and the UI """ socket_id = str(uuid.uuid4()) await connection_manager.connect(websocket, socket_id) redirect_task = await redirect_std_out(websocket) try: while True: # force the socket to sleep to prevent it from crashing await asyncio.sleep(socket_sleep) except WebSocketDisconnect: # ignore since we're disconnecting in the finally-block pass except Exception as e: print(e) finally: print('websocket disconnected') await stop_std_redirect(redirect_task) await connection_manager.disconnect(socket_id)
ConnectionManager
class ConnectionManager: def __init__(self): self.open_sockets = {} async def connect(self, websocket: WebSocket, socket_id: str): await websocket.accept() self.open_sockets[socket_id] = websocket async def disconnect(self, socket_id: str): websocket = self.open_sockets[socket_id] del self.open_sockets[socket_id] await websocket.close()
尝试过的自定义回调(无效)
from keras.callbacks import Callback from sys import stderr, stdout class FlushStdIOCallback(Callback): def on_epoch_begin(self, epoch, logs=None): print(f'starting epoch {epoch}') stderr.flush() stdout.flush() def on_epoch_end(self, epoch, logs=None): print(f'finished epoch {epoch}') stderr.flush() stdout.flush()
解决方案
问题根源在两个核心点:
model.fit是同步阻塞操作:FastAPI的WebSocket运行在异步事件循环中,但model.fit是CPU密集型同步函数,会完全阻塞事件循环,导致send_stdout任务无法定期执行,只能等训练结束后批量发送缓冲区内容。- StringIO无有效刷新机制:你替换的
sys.stdout是StringIO,它的flush()方法是空实现,手动调用也不会触发任何缓冲区同步动作。
具体修复步骤:
1. 用线程隔离训练任务,避免阻塞事件循环
把训练代码放到单独线程中执行,让异步的WebSocket发送任务可以正常运行:
# 定义同步训练函数 def run_training(): model = ... # 你的模型定义 model.fit(...) # 执行训练逻辑 # 在WebSocket处理函数中,用asyncio.to_thread包装执行 await asyncio.to_thread(run_training)
2. 替换StringIO为线程安全的自定义缓冲区
实现支持实时读取、线程安全的输出缓冲区,替代StringIO:
class WebSocketBuffer: def __init__(self): self.buffer = [] self.lock = asyncio.Lock() def write(self, data): with self.lock: self.buffer.append(data) def flush(self): # 无需额外操作,锁已保证线程安全 pass async def get_and_clear(self): async with self.lock: data = ''.join(self.buffer) self.buffer = [] return data
修改redirect_std_out函数使用新缓冲区:
async def redirect_std_out(websocket): """ redirects the std output to the websocket """ stdout_buffer = WebSocketBuffer() async def send_stdout(): try: while True: await asyncio.sleep(socket_sleep) data = await stdout_buffer.get_and_clear() if data: await websocket.send_text(data.rstrip('\n')) except (asyncio.CancelledError, websockets.ConnectionClosedError): pass sys.stdout = stdout_buffer return asyncio.create_task(send_stdout())
3. 优化WebSocket处理逻辑
去掉无限sleep循环,改为并行执行训练任务和WebSocket维护:
async def start_websocket_logging(websocket): """ starts the websocket between this service and the UI """ socket_id = str(uuid.uuid4()) await connection_manager.connect(websocket, socket_id) redirect_task = await redirect_std_out(websocket) try: # 并行执行训练任务和WebSocket发送任务 await asyncio.to_thread(run_training) except WebSocketDisconnect: pass except Exception as e: print(e) finally: print('websocket disconnected') await stop_std_redirect(redirect_task) await connection_manager.disconnect(socket_id)
4. 直接通过Keras回调推送日志(可选优化)
绕过stdout重定向,在Keras回调中直接通过WebSocket发送日志,保证实时性:
class WebSocketLoggingCallback(Callback): def __init__(self, websocket): self.websocket = websocket async def send_log(self, message): try: await self.websocket.send_text(message) except Exception: pass def on_epoch_begin(self, epoch, logs=None): msg = f'starting epoch {epoch}' print(msg) # 在同步回调中调用异步函数,用线程安全的方式提交到事件循环 asyncio.run_coroutine_threadsafe(self.send_log(msg), asyncio.get_event_loop()) def on_epoch_end(self, epoch, logs=None): msg = f'finished epoch {epoch}, logs: {logs}' print(msg) asyncio.run_coroutine_threadsafe(self.send_log(msg), asyncio.get_event_loop())
训练时传入该回调:
def run_training(websocket): model = ... callbacks = [WebSocketLoggingCallback(websocket)] model.fit(callbacks=callbacks, ...)
内容的提问来源于stack exchange,提问作者Drew Adams
相关产品推荐
相关产品推荐

