VerseFlow-Studio / src /verseflow /core /gpu_manager.py
julius119's picture
Upload 23 files
63ca2a0 verified
Raw History Blame Contribute Delete
2.71 kB
"""
显存池化、分片与 Zero-Bubble 张量并行调度器
Author: XiaoZhe (Commercial Contact: janejulius119@gmail.com / WeChat: julius119)
"""
import asyncio
import logging
from dataclasses import dataclass
from typing import Dict, List
logger = logging.getLogger("VerseFlow.GPUManager")
@dataclass
class GPUAllocation:
gpu_ids: List[int]
memory_allocated_mb: float
model_id: str
task_id: str
tp_size: int
class ProductionGPUMemoryManager:
"""分布式 GPU 显存管理器,支持多卡张量并行与防死锁并发申请"""
def __init__(self, gpu_count: int = 4, memory_per_gpu_mb: float = 24576.0):
self.gpu_count = gpu_count
self.memory_per_gpu_mb = memory_per_gpu_mb
self.gpu_used_mb = [0.0] * gpu_count
self.lock = asyncio.Lock()
self.model_vram_map = {
"minimax_h3": (38000.0, 2),[cite: 1]
"wan2.2_14b": (28000.0, 2),[cite: 1]
"hunyuan_video_1_5": (24000.0, 1),[cite: 1]
"ltx_2_3": (16000.0, 1),[cite: 1]
"wan2.2_5b": (8000.0, 1)[cite: 1]
}
async def allocate(self, model_id: str, task_id: str, timeout: float = 60.0) -> GPUAllocation:
req_mb, req_tp = self.model_vram_map.get(model_id, (16000.0, 1))
per_gpu_req = req_mb / req_tp
start_time = asyncio.get_event_loop().time()
while True:
async with self.lock:
available_gpus = []
for i in range(self.gpu_count):
free_mem = self.memory_per_gpu_mb - self.gpu_used_mb[i] - 2000.0
if free_mem >= per_gpu_req:
available_gpus.append(i)
if len(available_gpus) >= req_tp:
selected_gpus = available_gpus[:req_tp]
for gpu_id in selected_gpus:
self.gpu_used_mb[gpu_id] += per_gpu_req
logger.info(f"显存分配成功: Task={task_id}, Model={model_id}, GPUs={selected_gpus}")
return GPUAllocation(selected_gpus, req_mb, model_id, task_id, req_tp)
if asyncio.get_event_loop().time() - start_time > timeout:
raise TimeoutError(f"显存申请超时: Task={task_id}, Model={model_id}")
await asyncio.sleep(0.5)
async def release(self, alloc: GPUAllocation):
async with self.lock:
per_gpu_req = alloc.memory_allocated_mb / alloc.tp_size
for gpu_id in alloc.gpu_ids:
self.gpu_used_mb[gpu_id] = max(0.0, self.gpu_used_mb[gpu_id] - per_gpu_req)
logger.info(f"显存成功释放: Task={alloc.task_id}, GPUs={alloc.gpu_ids}")