You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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()
解决方案

问题根源在两个核心点:

  1. model.fit是同步阻塞操作:FastAPI的WebSocket运行在异步事件循环中,但model.fit是CPU密集型同步函数,会完全阻塞事件循环,导致send_stdout任务无法定期执行,只能等训练结束后批量发送缓冲区内容。
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.11 14:22:16