【大模型】- 异步 Hogwild 推理

算法

异步 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 推理
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

  • 标题: 【大模型】- 异步 Hogwild 推理
  • 作者:
  • 创建于 : 2026-08-19 09:22:00
  • 更新于 : 2026-08-21 16:20:12
  • 链接: https://sxl-space.tk/2026/08/19/010_LLM/010_LLM-22-AsyncHogwild/
  • 版权声明: 版权所有 © 宋,禁止转载。