多模态AI后端架构:统一推理调度与资源管理实战解析

0 阅读

多模态AI后端的调度挑战

现代AI应用早已不再局限于单一模态。一个典型的智能助手可能需要同时处理文本对话、图像识别、语音转写和语音合成。这意味着后端服务需要部署多个大型模型,例如用于文本生成的LLaMA 70B、用于图像理解的LLaVA 13B、用于语音识别的Whisper Large-v3,以及用于语音合成的Bark。这些模型在资源需求上差异巨大:LLaMA 70B的FP16权重就高达140GB,需要至少两张A100(每张80GB)进行张量并行推理;而Whisper Large-v3的权重仅3GB,甚至可以在CPU上运行。

将这些模型部署到同一个GPU集群时,我们面临三个核心问题:

显存碎片化:当清除一个70B模型后,虽然释放了140GB显存,但由于CUDA内存分配器的碎片化,这些零散的空闲块可能无法满足下一个需要50GB连续显存的模型。这要求我们实现显存池管理,而不是简单依赖系统分配。

调度优先级倒置:低优先级的TTS任务可能长时间占用GPU,导致高优先级的文本生成请求排队等待。我们需要一个优先级感知的调度器,确保关键任务及时获得资源。

多模态请求的关联性:用户上传一张图片并询问“这是什么?”时,图像理解模型(LLaVA)和文本生成模型(LLaMA)需要串联执行,并且共享上下文。这种多步推理管线要求调度器能够协调步骤间的数据传递和资源分配。

统一推理调度器的架构设计

为了解决上述问题,我们设计了一个统一推理调度器,它由以下核心模块组成:

路由模块:根据请求类型(如/v1/chat对应文本模型,/v1/image/describe对应LLaVA)选择目标模型。对于多步请求,如“图像理解+总结”,则将其分解为管线步骤。

优先级队列:按模型分组维护请求队列,每个队列内部按QoS等级排序。分发器从各队列中拉取请求,并根据GPU负载和显存情况决定优先执行哪个模型的请求。

显存分配器:维护GPU集群的显存状态图。当需要加载模型但显存不足时,驱逐最久未使用(LRU)的空闲模型。使用LRU而非FIFO,是因为高频模型(如Embedding服务)应被保留,不应被低频模型挤占。

管线编排:多模型串联调用时,第一步的输出作为第二步的输入。调度时需将连续步骤分配到同一GPU节点(或通过NVLink连接的低延迟节点),以减少中间数据在GPU间的传输开销。

统一调度器的Rust实现

以下是一个基于Rust和Tokio的简化实现,展示了调度器的核心逻辑。

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 {
    pub min_vram: usize,
    pub recommended_gpus: usize,
    pub latency_ms: (u64, u64),
    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,
    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,
}

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 {
    model_placement: RwLock<HashMap<ModelId, Vec<String>>>,
    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, 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;
        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)?;
            let available_gpus: Vec<_> = 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
                })
                .collect();
            if !available_gpus.is_empty() {
                if let Some(req) = queue.front() {
                    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));
                }
            }
        }

        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() {
                let gpu_id = self.select_best_gpu(&model_id, &gpu_nodes)?;
                return Some((req, gpu_id));
            }
        }
        None
    }

    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())
    }

    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() {
            node.loaded_models.sort_by_key(|m| m.last_used);
            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 {
    pub async fn describe_and_summarize(&self, image_data: Vec<u8>, user_id: &str) -> Result<PipelineResult, PipelineError> {
        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;
        // 实际通过回调获取结果
        let description = "Image shows a cat sitting on a sofa.".to_string();
        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,
}

关键设计决策解析

模型加载状态优先:调度时优先选择已加载目标模型的GPU节点,因为模型加载的冷启动延迟可能高达10-30秒,而推理本身仅需几秒。通过model_loaded || free_vram >= min_vram条件,我们确保请求要么被调度到已加载的节点,要么被调度到有足够显存可加载的节点。

优先级分数综合计算QoS权重 + 等待时间惩罚。等待惩罚随排队时间线性增长,防止低优先级请求被无限饥饿。权重系数的选择需要根据实际流量分布校准,例如Critical请求权重设为10,而Low请求设为1。

多因素GPU选择select_best_gpu中,首先比较模型是否已加载,其次比较当前负载,最后比较空闲显存。这是一个多目标排序问题,当前实现使用优先级链式比较,简单有效。

显存GC的保守策略current_load < 0.9的节点才驱逐模型,避免在高负载时驱逐导致正在排队的请求陷入冷启动。同时,通过max_idle_secs参数控制最小驻留时间,防止模型频繁加载/卸载。

适用场景与权衡

适用场景

  • 多模型推理平台:需要同时服务3种以上不同模型类型的推理平台,如智能助手、多模态搜索等。
  • 显存受限的GPU集群:当GPU集群显存总量有限(如4张A100),需要通过动态加载/卸载来服务超过显存总量的模型集合。
  • 多模态应用:图像+文本、语音+文本等请求之间有关联关系,需要管线编排。

不适用场景

  • 模型数量少:只有1-2个模型时,简单的负载均衡即可,无需复杂调度。
  • 显存充足:所有模型能同时加载到显存中,显存管理逻辑多余。
  • 纯CPU推理:无GPU时,显存感知调度的核心逻辑无用。

主要权衡

  1. 模型加载/卸载的抖动:如果模型A和B交替调用,显存GC会频繁加载/卸载,产生抖动。引入“最小驻留时间”可有效缓解,即模型加载后至少保留N秒。
  2. 管线步骤的数据传递:步骤间的大数据(如224×224图像的embedding约150KB)在内存中传递。如果步骤调度到不同节点,需要网络传输,应选择网络延迟最低的节点对(如NVSwitch域)。
  3. 调度延迟 vs 批处理效率:立即执行单条请求降低延迟,但批次推理的吞吐是逐条推理的2-5倍。连续批处理(Continuous Batching)是折中方案,可在保证延迟的同时提高吞吐。

总结

多模态推理的统一调度需要同时管理显存分配、优先级队列和跨模型管线编排。模型加载状态是调度决策的首要因素,已加载模型的节点优先级最高,以消除冷启动延迟。等待时间惩罚机制防止低优先级请求被高频高优请求永久饥饿。显存GC需要平衡释放空间和抖动,引入最小驻留时间可有效防止模型频繁加载/卸载。管线步骤间优先调度到同一节点(或NVLink直连节点),可减少中间数据跨GPU传输的延迟。

通过合理的架构设计和实现,我们可以构建一个高效、稳定的多模态AI后端,满足日益复杂的应用需求。