多模态 AI 系统的后端架构:文本、图像、音频的统一推理调度与资源管理
一、多模型后端的调度复杂度
一个典型的多模态 AI 后端需要同时服务:文本生成(LLaMA 70B)、图像理解(LLaVA 13B)、语音识别(Whisper Large-v3)、文本到语音(Bark/VITS)。四个模型有截然不同的资源需求:
- LLaMA 70B:FP16 权重 140GB,至少需要 2×A100(每张 80GB)做张量并行推理。单次推理 2-10 秒。
- LLaVA 13B:文本部分 26GB + 视觉编码器 2GB,1×A100 足够。单次推理 1-3 秒。
- Whisper Large-v3:权重 3GB,CPU 即可推理。编码器(30秒音频)1-2 秒,解码器 0.5-2 秒。
- Bark TTS:权重 3GB,GPU 推理。20 秒输出音频生成 5-15 秒。
将这些模型部署到同一个 GPU 集群上,面临三个核心问题:
- 显存碎片:清除 70B 模型后的 140GB 空闲显存,由于 CUDA 内存分配器的碎片化,可能无法加载下一个需要 50GB 的模型。需要显存池管理。
- 调度优先级倒置:低优先级的 TTS 任务占据了 GPU,导致高优先级的文本生成排队等待。需要优先级感知调度。
- 多模态请求的关联:用户上传一张图片并问"这是什么?"——图像理解模型 (LLaVA) 和文本模型 (LLaMA) 需要串联执行,且共享上下文。
统一调度器的设计需要解决:显存感知的资源分配、模型间的优先级调度、以及多步推理管线中步骤间的数据传递。
二、统一推理调度的架构设计
统一调度器由以下模块组成:
路由模块:根据请求类型(/v1/chat → 文本模型、/v1/image/describe → LLaVA、/v1/speech/transcribe → Whisper)选择目标模型。多步请求(如图像理解 + 总结)被分解为管线步骤。
优先级队列:按模型分组维护请求队列。每个模型队列内部按 QoS 等级排序。分发器从各队列中拉取请求,根据 GPU 的负载和显存情况决定优先执行哪个模型的请求。
显存分配器:维护 GPU 集群的显存状态图。当需要加载模型但显存不足时,驱逐最久未使用(LRU)的空闲模型。使用 LRU 而非 FIFO——因为高频模型(如 Embedding 服务)被频繁调用,不应被低频模型挤占。
管线编排:多模型串联调用时,第一步输出作为第二步输入。管线的调度需要将连续步骤分配到同一 GPU 节点(或通过 NVLink 连接的低延迟节点),减少中间数据在 GPU 间的传输。
三、统一调度器的 Rust 实现
use std::collections::{HashMap, BinaryHeap, VecDeque};
use std::sync::Arc;
use tokio::sync::{Mutex, RwLock, Semaphore};
use std::cmp::Ordering;
/// 模型标识
#[derive(Clone, PartialEq, Eq, Hash, Debug)]
pub enum ModelId {
Llama70B,
Llava13B,
WhisperLargeV3,
BarkTTS,
EmbeddingBge,
}
/// 模型资源需求
pub struct ModelResourceProfile {
/// 最小显存需求 (bytes)
pub min_vram: usize,
/// 推荐 GPU 数量(张量并行)
pub recommended_gpus: usize,
/// 单次推理预估延迟范围 (ms): (min, p95)
pub latency_ms: (u64, u64),
/// 支持的 GPU 架构列表
pub gpu_arch: Vec<GpuArch>,
}
/// 推理请求
#[derive(Clone)]
pub struct InferenceRequest {
pub id: String,
pub model: ModelId,
pub qos: QosLevel,
/// 请求到达时间戳
pub arrived_at: std::time::Instant,
/// 后续步骤(管线场景)
pub next_steps: Vec<InferenceRequest>,
}
/// QoS 优先级
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Debug)]
pub enum QosLevel {
Critical = 0, // 实时交互
High = 1, // API 调用
Normal = 2, // 批处理
Low = 3, // 离线任务
}
/// GPU 节点状态
pub struct GpuNode {
pub id: String,
pub total_vram: usize,
pub free_vram: usize,
pub loaded_models: Vec<LoadedModelInfo>,
pub current_load: f64, // 0.0 ~ 1.0
}
pub struct LoadedModelInfo {
pub model: ModelId,
pub loaded_at: std::time::Instant,
pub last_used: std::time::Instant,
pub vram_used: usize,
}
/// 统一推理调度器
pub struct UnifiedScheduler {
/// 模型到 GPU 节点的分配
model_placement: RwLock<HashMap<ModelId, Vec<String>>>,
/// GPU 集群状态
gpu_nodes: RwLock<HashMap<String, GpuNode>>,
/// 按模型分组的请求队列
request_queues: RwLock<HashMap<ModelId, VecDeque<InferenceRequest>>>,
/// 每个模型的并发限制
model_concurrency: HashMap<ModelId, Arc<Semaphore>>,
/// 模型资源需求配置
profiles: HashMap<ModelId, ModelResourceProfile>,
}
impl UnifiedScheduler {
/// 注册模型及其资源需求
pub fn register_model(
&mut self,
model: ModelId,
profile: ModelResourceProfile,
) {
self.profiles.insert(model.clone(), profile);
}
/// 入队推理请求
pub async fn enqueue(&self, request: InferenceRequest) {
let mut queues = self.request_queues.write().await;
queues.entry(request.model.clone())
.or_insert_with(VecDeque::new)
.push_back(request);
}
/// 主调度循环:选择下一个要执行的请求
pub async fn schedule_next(&self) -> Option<(InferenceRequest, String)> {
let gpu_nodes = self.gpu_nodes.read().await;
let mut queues = self.request_queues.write().await;
// 1. 收集所有可执行的请求
let mut candidates = Vec::new();
for (model_id, queue) in queues.iter_mut() {
if queue.is_empty() {
continue;
}
let profile = self.profiles.get(model_id)?;
// 检查是否有 GPU 节点能满足该模型的显存需求
let available_gpus: Vec<_> = gpu_nodes.iter()
.filter(|(id, node)| {
// 显存检查: 要么模型已加载,要么有足够空闲显存
let model_loaded = node.loaded_models.iter()
.any(|m| m.model == *model_id);
model_loaded || node.free_vram >= profile.min_vram
})
.collect();
if !available_gpus.is_empty() {
if let Some(req) = queue.front() {
// 计算该请求的调度优先级分数
// 分数 = QoS权重 × 等待时间惩罚 × 显存效率
let qos_weight = match req.qos {
QosLevel::Critical => 10.0,
QosLevel::High => 5.0,
QosLevel::Normal => 2.0,
QosLevel::Low => 1.0,
};
let wait_penalty = req.arrived_at.elapsed().as_secs_f64() / 10.0;
let score = qos_weight + wait_penalty;
candidates.push((model_id.clone(), score));
}
}
}
// 2. 按优先级分数排序,选择最高分的请求
candidates.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal));
for (model_id, _) in candidates {
if let Some(req) = queues.get_mut(&model_id)?.pop_front() {
// 选择最优 GPU 节点
let gpu_id = self.select_best_gpu(&model_id, &gpu_nodes)?;
return Some((req, gpu_id));
}
}
None
}
/// 选择最优 GPU 节点 —— 考虑因素:
/// 1. 模型是否已加载(零启动延迟)
/// 2. 空闲显存量
/// 3. 当前负载
/// 4. 是否为推荐 GPU 架构
fn select_best_gpu(
&self,
model_id: &ModelId,
gpu_nodes: &HashMap<String, GpuNode>,
) -> Option<String> {
let profile = self.profiles.get(model_id)?;
gpu_nodes.iter()
.filter(|(_, node)| {
// 基本过滤:是否有足够显存
let model_loaded = node.loaded_models.iter()
.any(|m| m.model == *model_id);
model_loaded || node.free_vram >= profile.min_vram
})
.max_by(|(_, a), (_, b)| {
// 排序:模型已加载 > 空闲显存多 > 负载低
let a_loaded = a.loaded_models.iter().any(|m| m.model == *model_id);
let b_loaded = b.loaded_models.iter().any(|m| m.model == *model_id);
if a_loaded != b_loaded {
return a_loaded.cmp(&b_loaded);
}
// 次选:负载低的节点
a.current_load.partial_cmp(&b.current_load).unwrap_or(Ordering::Equal).reverse()
})
.map(|(id, _)| id.clone())
}
/// 显存 GC:驱逐空闲模型以释放显存
pub async fn evict_idle_models(&self, max_idle_secs: u64) -> usize {
let mut gpu_nodes = self.gpu_nodes.write().await;
let mut evicted = 0;
let now = std::time::Instant::now();
for (_, node) in gpu_nodes.iter_mut() {
// 按 last_used 排序,驱逐最久未使用的模型
node.loaded_models.sort_by_key(|m| m.last_used);
// 注意:不驱逐正在处理的模型(通过 current_load > 0 判断)
// 被驱逐的模型从 GPU 节点移除,释放显存
node.loaded_models.retain(|m| {
let idle_duration = now.duration_since(m.last_used).as_secs();
if idle_duration > max_idle_secs && node.current_load < 0.9 {
node.free_vram += m.vram_used;
evicted += 1;
false // 不保留
} else {
true
}
});
}
evicted
}
}
/// 多模态管线:图像理解 → 文本总结
pub struct MultimodalPipeline {
scheduler: Arc<UnifiedScheduler>,
}
impl MultimodalPipeline {
/// 执行多模态管线
/// 1. LLaVA: 图像 → 描述文本
/// 2. LLaMA: 描述文本 → 结构化摘要
pub async fn describe_and_summarize(
&self,
image_data: Vec<u8>,
user_id: &str,
) -> Result<PipelineResult, PipelineError> {
// 步骤 1: 图像理解
let img_req = InferenceRequest {
id: format!("img-{}-1", user_id),
model: ModelId::Llava13B,
qos: QosLevel::High,
arrived_at: std::time::Instant::now(),
next_steps: vec![],
};
self.scheduler.enqueue(img_req).await;
// 等待步骤 1 结果...(实际代码通过回调/通道获取)
let description = "Image shows a cat sitting on a sofa.".to_string();
// 步骤 2: 文本总结(使用步骤 1 的输出作为输入)
let summary_req = InferenceRequest {
id: format!("img-{}-2", user_id),
model: ModelId::Llama70B,
qos: QosLevel::High,
arrived_at: std::time::Instant::now(),
next_steps: vec![],
};
self.scheduler.enqueue(summary_req).await;
Ok(PipelineResult {
description,
summary: "A cat resting on furniture.".to_string(),
})
}
}
pub struct PipelineResult {
pub description: String,
pub summary: String,
}
#[derive(Debug)]
pub enum PipelineError {
Timeout,
ModelUnavailable,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum GpuArch {
A100,
H100,
A10,
T4,
}
关键设计决策:
- 模型加载检查(
model_loaded || free_vram >= min_vram):优先将请求调度到已加载该模型的 GPU 节点——消除模型重载的冷启动延迟(10-30秒)。 - 优先级分数的综合计算:
QoS权重 + 等待时间惩罚。等待惩罚随排队时间线性增长,防止低优先级请求被无限饥饿。权重系数的选择需要根据实际流量分布校准。 select_best_gpu中的多因素比较:模型已加载(最高优先级)、当前负载低(次优先)、空闲显存多。这是一个多目标排序问题——当前实现使用优先级链式比较。- 显存 GC 的保守策略:
current_load < 0.9的节点才驱逐——避免在高负载时驱逐模型导致正在排队的请求陷入冷启动。
四、多模态统一调度的适用边界与权衡
适用场景:
- 需要同时服务 3 种以上不同模型类型的推理平台。
- GPU 集群显存总量有限(如 4 张 A100),需要通过动态加载/卸载来服务超过显存总量的模型集合。
- 多模态应用(图像+文本、语音+文本),请求之间有关联关系。
不适用场景:
- 只有 1-2 个模型——不需要复杂的调度和显存管理,简单的负载均衡即可。
- 所有模型能同时加载到显存中——显存管理逻辑是多余的。
- 无 GPU 的纯 CPU 推理——显存感知调度的核心逻辑在此无用。
主要权衡:
- 模型加载/卸载的抖动:如果模型 A 和模型 B 交替调用,显存 GC 会频繁加载/卸载两个模型,产生抖动。需要引入"最小驻留时间"——模型加载后至少保留 N 秒,即使被标为空闲也不立即驱逐。
- 管线步骤的数据传递:步骤间的大数据(如 224×224 图像的 embedding = 150KB)直接在内存中传递。如果步骤调度到不同节点,需要网络传输——选择网络延迟最低的节点对(如同 NVSwitch 域)。
- 调度延迟 vs 批处理效率:立即执行单条请求降低延迟,但批次推理(batch inference)的吞吐是逐条推理的 2-5 倍。连续批处理(Continuous Batching)是折中方案。
五、总结
- 多模态推理的统一调度需同时管理显存分配、优先级队列和跨模型管线编排。
- 模型加载状态是调度决策的首要因素——已加载模型的节点优先级最高,消除冷启动延迟。
- 等待时间惩罚机制防止低优先级请求被高频高优请求永久饥饿。
- 显存 GC 需要平衡释放空间和抖动——引入最小驻留时间可有效防止模型频繁加载/卸载。
- 管线步骤间优先调度到同一节点(或 NVLink 直连节点),可减少中间数据跨 GPU 传输的延迟。
转载自 CSDN-专业IT技术社区
原文链接:https://blog.csdn.net/2301_81410839/article/details/163119576



