""" 显存池化、分片与 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}")