异步 Hogwild 推理
优化并发推理,提高吞吐量
类型: 学习 | 语言: Python | 🏷 前置:《Jamba》(本系列第 20 篇)
学习目标
- 理解 Hogwild 推理的原理
- 实现异步推理管道
- 优化并发处理
- 减少推理延迟
- 提高系统吞吐量
Hogwild 推理概述
Hogwild 推理是一种无锁的并发推理方法,允许多个线程同时访问共享模型。
传统推理的问题
- 锁竞争:多线程访问共享资源时的锁开销
- 序列化:请求被序列化处理
- 资源浪费:GPU 利用率低
Hogwild 的优势
- 无锁设计:避免锁竞争
- 并行处理:同时处理多个请求
- 高吞吐量:提高 GPU 利用率
实现
基本 Hogwild 推理
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107
| import torch import torch.nn as nn from concurrent.futures import ThreadPoolExecutor import threading import queue
class HogwildInference: def __init__(self, model, n_workers=4): self.model = model self.n_workers = n_workers self.request_queue = queue.Queue() self.result_queue = queue.Queue() self.workers = [] self.start_workers() def start_workers(self): """启动工作线程""" for i in range(self.n_workers): worker = threading.Thread( target=self.worker_loop, args=(i,), daemon=True ) worker.start() self.workers.append(worker) def worker_loop(self, worker_id): """工作线程循环""" while True: try: request = self.request_queue.get(timeout=1) if request is None: break result = self.process_request(request) self.result_queue.put(result) self.request_queue.task_done() except queue.Empty: continue def process_request(self, request): """处理请求""" input_ids = request["input_ids"] request_id = request["request_id"] with torch.no_grad(): outputs = self.model(input_ids) return { "request_id": request_id, "outputs": outputs } def infer(self, input_ids, request_id): """异步推理""" request = { "input_ids": input_ids, "request_id": request_id } self.request_queue.put(request) return request_id def get_result(self, request_id, timeout=10): """获取结果""" start_time = time.time() while time.time() - start_time < timeout: try: result = self.result_queue.get(timeout=0.1) if result["request_id"] == request_id: return result else: self.result_queue.put(result) except queue.Empty: continue return None def shutdown(self): """关闭工作线程""" for _ in range(self.n_workers): self.request_queue.put(None) for worker in self.workers: worker.join()
|
批量 Hogwild 推理
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61
| class BatchHogwildInference: def __init__(self, model, batch_size=32, n_workers=4): self.model = model self.batch_size = batch_size self.n_workers = n_workers self.batch_buffer = [] self.buffer_lock = threading.Lock() self.request_queue = queue.Queue() self.batch_thread = threading.Thread( target=self.batch_loop, daemon=True ) self.batch_thread.start() def batch_loop(self): """批量处理循环""" while True: try: requests = [] request = self.request_queue.get(timeout=1) requests.append(request) while len(requests) < self.batch_size: try: request = self.request_queue.get(timeout=0.01) requests.append(request) except queue.Empty: break if requests: self.process_batch(requests) except queue.Empty: continue def process_batch(self, requests): """批量处理""" input_ids = torch.cat([r["input_ids"] for r in requests], dim=0) with torch.no_grad(): outputs = self.model(input_ids) for request, output in zip(requests, outputs): self.result_queue.put({ "request_id": request["request_id"], "outputs": output.unsqueeze(0) })
|
无锁数据结构
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19
| class LockFreeQueue: """无锁队列""" def __init__(self): self.queue = queue.Queue() def push(self, item): """压入元素""" self.queue.put(item) def pop(self): """弹出元素""" try: return self.queue.get_nowait() except queue.Empty: return None def is_empty(self): """检查是否为空""" return self.queue.empty()
|
原子操作
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22
| class AtomicCounter: """原子计数器""" def __init__(self, value=0): self.value = value self.lock = threading.Lock() def increment(self): """原子递增""" with self.lock: self.value += 1 return self.value def decrement(self): """原子递减""" with self.lock: self.value -= 1 return self.value def get(self): """获取值""" with self.lock: return self.value
|
优化技术
请求调度
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22
| class RequestScheduler: def __init__(self, n_workers): self.n_workers = n_workers self.worker_loads = [0] * n_workers self.lock = threading.Lock() def schedule(self, request): """调度请求到负载最低的工作线程""" with self.lock: min_load = min(self.worker_loads) min_worker = self.worker_loads.index(min_load) self.worker_loads[min_worker] += 1 return min_worker def complete(self, worker_id): """标记工作线程完成""" with self.lock: self.worker_loads[worker_id] -= 1
|
动态批处理
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46
| class DynamicBatcher: def __init__(self, max_batch_size=32, max_wait_time=0.1): self.max_batch_size = max_batch_size self.max_wait_time = max_wait_time self.batch = [] self.batch_start_time = None def add_request(self, request): """添加请求到批次""" if not self.batch: self.batch_start_time = time.time() self.batch.append(request) if (len(self.batch) >= self.max_batch_size or time.time() - self.batch_start_time >= self.max_wait_time): return self.process_batch() return None def process_batch(self): """处理当前批次""" if not self.batch: return None input_ids = torch.cat([r["input_ids"] for r in self.batch], dim=0) outputs = model(input_ids) results = [] for request, output in zip(self.batch, outputs): results.append({ "request_id": request["request_id"], "outputs": output.unsqueeze(0) }) self.batch = [] self.batch_start_time = None return results
|
负载均衡
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25
| class LoadBalancer: def __init__(self, workers): self.workers = workers self.worker_stats = {w.id: {"requests": 0, "latency": 0} for w in workers} def select_worker(self, request): """选择工作线程""" scores = {} for worker_id, stats in self.worker_stats.items(): load = stats["requests"] * (stats["latency"] + 1) scores[worker_id] = 1.0 / (load + 1) selected = max(scores, key=scores.get) return selected def update_stats(self, worker_id, latency): """更新统计信息""" self.worker_stats[worker_id]["requests"] += 1 self.worker_stats[worker_id]["latency"] = ( self.worker_stats[worker_id]["latency"] * 0.9 + latency * 0.1 )
|
性能监控
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40
| class HogwildMonitor: def __init__(self): self.metrics = { "total_requests": 0, "completed_requests": 0, "avg_latency": 0, "throughput": 0, "queue_size": 0, } self.start_time = time.time() def record_request(self): """记录请求""" self.metrics["total_requests"] += 1 def record_completion(self, latency): """记录完成""" self.metrics["completed_requests"] += 1 n = self.metrics["completed_requests"] self.metrics["avg_latency"] = ( self.metrics["avg_latency"] * (n - 1) + latency ) / n elapsed = time.time() - self.start_time self.metrics["throughput"] = self.metrics["completed_requests"] / elapsed def update_queue_size(self, size): """更新队列大小""" self.metrics["queue_size"] = size def print_stats(self): """打印统计信息""" print(f"总请求数: {self.metrics['total_requests']}") print(f"完成请求数: {self.metrics['completed_requests']}") print(f"平均延迟: {self.metrics['avg_latency']:.4f}s") print(f"吞吐量: {self.metrics['throughput']:.2f} req/s") print(f"队列大小: {self.metrics['queue_size']}")
|
实际应用
Web 服务
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51
| from fastapi import FastAPI, HTTPException from pydantic import BaseModel import uvicorn
app = FastAPI()
class InferenceRequest(BaseModel): prompt: str max_new_tokens: int = 512
class InferenceResponse(BaseModel): response: str request_id: str latency: float
hogwild = HogwildInference(model, n_workers=4)
@app.post("/infer", response_model=InferenceResponse) async def infer(request: InferenceRequest): """推理接口""" start_time = time.time() inputs = tokenizer(request.prompt, return_tensors="pt") request_id = hogwild.infer(inputs["input_ids"], "req_" + str(time.time())) result = hogwild.get_result(request_id, timeout=30) if result is None: raise HTTPException(status_code=500, detail="推理超时") latency = time.time() - start_time response_text = tokenizer.decode( result["outputs"][0].argmax(dim=-1), skip_special_tokens=True ) return InferenceResponse( response=response_text, request_id=request_id, latency=latency )
if __name__ == "__main__": uvicorn.run(app, host="0.0.0.0", port=8000)
|
批量处理服务
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36
| class BatchInferenceService: def __init__(self, model, batch_size=32): self.model = model self.batcher = DynamicBatcher(max_batch_size=batch_size) self.process_thread = threading.Thread( target=self.process_loop, daemon=True ) self.process_thread.start() def process_loop(self): """处理循环""" while True: if self.batcher.batch: results = self.batcher.process_batch() for result in results: self.result_queue.put(result) time.sleep(0.01) async def infer(self, request): """异步推理""" result = self.batcher.add_request(request) if result is not None: return result return await self.wait_for_result(request["request_id"])
|
总结
Hogwild 推理通过无锁设计和并发处理,显著提高了推理吞吐量。关键组件包括请求队列、工作线程、批量处理和负载均衡。适用于高并发、低延迟的推理场景。
下一步
下一课将介绍梯度检查点,优化内存使用。
📚 本文改编自 AI Engineering from Scratch(MIT License · 作者 Rohit Ghumare),中文内容来自官方中文镜像。原课程共 503 课 · 20 阶段 · 免费开源,教程网站见 aiengineeringfromscratch.com。