基于 EMR Serverless Ray 实现 Qwen 模型批量推理实践
1. 背景与目标
随着大语言模型在业务场景中的深入应用,批量推理成为常见需求。相比在线推理,批量推理更关注吞吐量、资源利用率和成本控制。EMR Serverless Ray 提供无服务器化的 Ray 计算环境,能够按需弹性伸缩,适合承载大规模、可并行的推理任务。
本文以 Qwen 系列模型为例,介绍如何在 EMR Serverless Ray 上实现批量推理。全文包含环境准备、代码实现、任务提交、性能调优和常见问题排查,所有代码均可直接复制运行。
2. 技术选型与架构
整体架构由三个核心部分组成:
- EMR Serverless Ray:负责弹性调度 Ray 集群,按任务自动拉起和释放计算资源。
- Qwen 模型:使用 Hugging Face Transformers 加载,支持 Qwen2、Qwen2.5 等系列模型。
- 对象存储 OSS:存放输入数据、模型权重和推理结果。
批量推理任务的数据流如下:
flowchart LR
A[OSS 输入数据] --> B[Ray Driver]
B --> C[Ray Worker 并行推理]
C --> D[OSS 输出结果]
D --> E[下游消费]
每个 Ray Worker 加载一份模型副本,对分片数据进行推理,最后将结果写回 OSS。通过调整 Worker 数量和每 Worker 的并发度,可以灵活控制吞吐量。
3. 环境准备
3.1 创建 EMR Serverless Ray 应用
在阿里云 EMR Serverless 控制台创建 Ray 应用,选择运行时版本和计算规格。建议按以下参数配置:
| 配置项 | 推荐值 | 说明 |
|---|---|---|
| 运行时版本 | EMR 5.x Ray 版本 | 包含 Ray 2.x 和 Python 3.10 |
| Driver 规格 | 4 vCPU / 16 GB | 负责任务调度和结果汇总 |
| Worker 规格 | 8 vCPU / 32 GB | 根据模型大小调整 |
| Worker 数量 | 4 - 16 | 按数据量和预算弹性调整 |
3.2 准备 Python 依赖
推理任务需要以下核心依赖,建议通过 requirements.txt 管理:
transformers>=4.40.0
torch>=2.1.0
accelerate>=0.30.0
ray[default]>=2.9.0
oss2>=2.18.0
sentencepiece>=0.1.99
protobuf>=4.25.0
在提交任务时,通过 --pip 参数或自定义镜像安装依赖。推荐使用自定义镜像,将依赖预置到镜像中,减少任务启动时间。
4. 数据准备
批量推理的输入数据通常以 JSON Lines 格式存储在 OSS 中,每行一条推理请求。示例输入文件 input.jsonl:
{"id": "001", "prompt": "请用一句话介绍杭州。"}
{"id": "002", "prompt": "解释什么是大语言模型。"}
{"id": "003", "prompt": "写一首关于秋天的五言绝句。"}
每条记录包含唯一 ID 和推理提示词。输出结果同样以 JSON Lines 格式写回 OSS,便于下游任务消费。
5. 核心代码实现
5.1 模型加载与推理函数
首先定义模型加载函数和单条推理函数。模型加载在 Worker 初始化时执行一次,推理函数对每条数据进行处理:
import os
import json
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2-1.5B-Instruct")
def load_model():
"""每个 Ray Worker 启动时加载一次模型"""
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
MODEL_NAME,
torch_dtype=torch.float16,
device_map="auto",
trust_remote_code=True
)
model.eval()
return model, tokenizer
def generate_text(model, tokenizer, prompt, max_new_tokens=256):
"""对单条 prompt 执行推理"""
messages = [{"role": "user", "content": prompt}]
text = tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
inputs = tokenizer(text, return_tensors="pt").to(model.device)
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=0.7,
top_p=0.9
)
response = tokenizer.decode(outputs[0][inputs.input_ids.shape[1]:], skip_special_tokens=True)
return response
5.2 Ray 并行推理主程序
主程序负责读取 OSS 数据、分发任务到 Ray Worker 并收集结果。使用 ray.remote 将推理函数封装为远程任务:
import ray
import oss2
import io
import json
from typing import List, Dict
@ray.remote(num_gpus=1, max_calls=100)
class InferenceWorker:
"""每个 Worker 持有一份模型副本"""
def __init__(self):
self.model, self.tokenizer = load_model()
def infer_batch(self, batch: List[Dict]) -> List[Dict]:
results = []
for item in batch:
try:
response = generate_text(
self.model, self.tokenizer, item["prompt"]
)
results.append({
"id": item["id"],
"prompt": item["prompt"],
"response": response,
"status": "success"
})
except Exception as e:
results.append({
"id": item["id"],
"prompt": item["prompt"],
"response": str(e),
"status": "failed"
})
return results
def read_input_from_oss(bucket, key: str) -> List[Dict]:
"""从 OSS 读取 JSONL 输入文件"""
content = bucket.get_object(key).read().decode("utf-8")
return [json.loads(line) for line in content.strip().split("\n") if line]
def write_output_to_oss(bucket, key: str, results: List[Dict]):
"""将推理结果写回 OSS"""
lines = "\n".join(json.dumps(r, ensure_ascii=False) for r in results)
bucket.put_object(key, lines.encode("utf-8"))
def main():
ray.init(address="auto", ignore_reinit_error=True)
# OSS 配置
auth = oss2.Auth(os.getenv("OSS_ACCESS_KEY_ID"), os.getenv("OSS_ACCESS_KEY_SECRET"))
bucket = oss2.Bucket(auth, os.getenv("OSS_ENDPOINT"), os.getenv("OSS_BUCKET"))
input_key = os.getenv("INPUT_KEY", "data/input.jsonl")
output_key = os.getenv("OUTPUT_KEY", "data/output.jsonl")
batch_size = int(os.getenv("BATCH_SIZE", "8"))
num_workers = int(os.getenv("NUM_WORKERS", "4"))
# 读取输入
data = read_input_from_oss(bucket, input_key)
print(f"Loaded {len(data)} records from {input_key}")
# 创建 Worker 池
workers = [InferenceWorker.remote() for _ in range(num_workers)]
# 分批分发任务
batches = [data[i:i + batch_size] for i in range(0, len(data), batch_size)]
futures = [workers[i % num_workers].infer_batch.remote(batch)
for i, batch in enumerate(batches)]
# 收集结果
all_results = []
for future in ray.get(futures):
all_results.extend(future)
# 写回 OSS
write_output_to_oss(bucket, output_key, all_results)
print(f"Completed {len(all_results)} records, output to {output_key}")
ray.shutdown()
if __name__ == "__main__":
main()
5.3 使用 Ray Data 的流式处理方式
对于超大规模数据,推荐使用 Ray Data 进行流式处理,避免一次性加载全部数据到内存。以下示例展示基于 Ray Data 的推理流程:
import ray
import ray.data
from ray.data import ActorPoolStrategy
def main_with_ray_data():
ray.init(address="auto", ignore_reinit_error=True)
# 从 OSS 读取数据
ds = ray.data.read_json(
"oss://your-bucket/data/input.jsonl",
include_paths=False
)
# 使用 Actor 池并行推理
ds = ds.map_batches(
InferenceWorker,
concurrency=4,
batch_size=8,
num_gpus=1,
max_concurrency=2
)
# 写回 OSS
ds.write_json("oss://your-bucket/data/output_raydata/")
print("Ray Data inference completed")
if __name__ == "__main__":
main_with_ray_data()
Ray Data 自动处理数据分片、任务调度和容错重试,适合 TB 级数据的批量推理场景。
6. 任务提交与运行
6.1 通过控制台提交
在 EMR Serverless Ray 控制台创建作业,上传 Python 脚本,配置环境变量和资源参数后提交运行。关键环境变量如下:
| 环境变量 | 示例值 | 说明 |
|---|---|---|
| MODEL_NAME | Qwen/Qwen2-1.5B-Instruct | 模型名称或 OSS 路径 |
| OSS_ENDPOINT | oss-cn-hangzhou.aliyuncs.com | OSS 地域节点 |
| OSS_BUCKET | my-bucket | 存储桶名称 |
| INPUT_KEY | data/input.jsonl | 输入文件路径 |
| OUTPUT_KEY | data/output.jsonl | 输出文件路径 |
| NUM_WORKERS | 4 | Ray Worker 数量 |
| BATCH_SIZE | 8 | 每个批次的数据条数 |
6.2 通过命令行提交
也可以使用 EMR Serverless 提供的 CLI 工具提交任务:
emr-serverless submit-job \
--application-id app-xxxx \
--job-name qwen-batch-inference \
--job-type RAY \
--python-file oss://my-bucket/scripts/inference.py \
--pip "transformers torch accelerate ray oss2 sentencepiece protobuf" \
--env "MODEL_NAME=Qwen/Qwen2-1.5B-Instruct" \
--env "OSS_ENDPOINT=oss-cn-hangzhou.aliyuncs.com" \
--env "OSS_BUCKET=my-bucket" \
--env "INPUT_KEY=data/input.jsonl" \
--env "OUTPUT_KEY=data/output.jsonl" \
--env "NUM_WORKERS=4" \
--env "BATCH_SIZE=8" \
--driver-cpu 4 --driver-memory 16G \
--worker-cpu 8 --worker-memory 32G \
--worker-count 4
7. 性能调优
7.1 吞吐量优化
批量推理的吞吐量受多个因素影响,可以从以下方向调优:
- 增大 Batch Size:在显存允许范围内增大每批数据量,提高 GPU 利用率。
- 使用 vLLM 加速:将推理后端替换为 vLLM,可显著提升吞吐量。示例代码如下:
from vllm import LLM, SamplingParams
def load_model_vllm():
llm = LLM(
model=MODEL_NAME,
tensor_parallel_size=1,
dtype="float16",
max_model_len=4096
)
return llm
def generate_batch_vllm(llm, prompts: List[str], max_tokens=256):
params = SamplingParams(
temperature=0.7,
top_p=0.9,
max_tokens=max_tokens
)
outputs = llm.generate(prompts, params)
return [o.outputs[0].text for o in outputs]
7.2 资源规划建议
| 模型规模 | 单 Worker 规格 | 推荐 Worker 数 | 说明 |
|---|---|---|---|
| 0.5B - 1.5B | 8 vCPU / 32 GB / 1 GPU | 4 - 8 | 适合快速验证和小规模任务 |
| 7B - 14B | 16 vCPU / 64 GB / 1 GPU | 4 - 16 | 需要较大显存,建议使用 A10 或 A100 |
| 32B 以上 | 32 vCPU / 128 GB / 多 GPU | 2 - 8 | 建议使用张量并行和模型并行 |
7.3 减少冷启动时间
模型加载是 Worker 启动的主要耗时点。建议将模型权重提前下载到 OSS,并通过环境变量指定本地路径,避免每次任务重复下载。同时可以使用自定义镜像预装依赖,减少 pip 安装时间。
8. 容错与重试机制
批量推理任务可能因单条数据异常或 Worker 故障而中断。Ray 提供了多种容错机制:
- 任务级重试:通过
max_retries参数设置远程任务的最大重试次数。 - Worker 自动恢复:Ray 会自动检测 Worker 故障并重新调度任务。
- 结果持久化:建议分批写回 OSS,避免全部完成后一次性写入导致数据丢失。
以下代码展示带重试机制的推理调用:
@ray.remote(num_gpus=1, max_retries=3)
def infer_with_retry(worker, batch):
try:
return ray.get(worker.infer_batch.remote(batch))
except ray.exceptions.RayTaskError:
# 重试时重新创建 Worker
new_worker = InferenceWorker.remote()
return ray.get(new_worker.infer_batch.remote(batch))
9. 结果验证与质量检查
推理完成后,需要对输出结果进行质量检查。建议从以下维度验证:
- 完整性:输出记录数是否与输入一致,是否有缺失或重复。
- 成功率:统计 status 为 success 的记录占比,排查失败原因。
- 内容质量:抽样检查生成文本是否符合预期,是否存在明显错误。
以下脚本用于统计输出结果的成功率和抽样展示:
import json
def validate_output(output_path):
with open(output_path, "r", encoding="utf-8") as f:
lines = [json.loads(line) for line in f if line.strip()]
total = len(lines)
success = sum(1 for r in lines if r["status"] == "success")
failed = total - success
print(f"Total: {total}, Success: {success}, Failed: {failed}")
print(f"Success rate: {success / total * 100:.2f}%")
# 抽样展示前 3 条结果
for r in lines[:3]:
print(f"ID: {r['id']}")
print(f"Prompt: {r['prompt']}")
print(f"Response: {r['response'][:200]}")
print("-" * 50)
if __name__ == "__main__":
validate_output("output.jsonl")
10. 常见问题排查
10.1 模型加载失败
如果模型加载时报错,优先检查网络连通性和模型路径。建议将模型下载到 OSS 后使用本地路径加载,避免每次任务从 Hugging Face 下载。同时确认 trust_remote_code=True 参数已设置,部分 Qwen 模型需要该参数。
10.2 GPU 显存不足
显存不足通常表现为 CUDA Out of Memory 错误。解决方案包括:减小 Batch Size、使用 torch.float16 半精度加载、启用梯度检查点或使用更小的模型版本。
10.3 Worker 频繁重启
Worker 频繁重启通常由内存溢出或 OOM 导致。建议检查 Worker 内存规格是否充足,适当增大内存配置,并减少单 Worker 的并发推理数。
10.4 数据倾斜
如果部分 Worker 处理时间明显长于其他 Worker,可能是数据倾斜导致。建议按 ID 哈希或随机方式打散数据,确保各分片数据量均衡。
11. 总结
本文详细介绍了基于 EMR Serverless Ray 实现 Qwen 模型批量推理的完整流程,包括环境准备、数据准备、核心代码实现、任务提交、性能调优和容错机制。通过 Ray 的弹性调度能力,可以按需伸缩计算资源,在保证吞吐量的同时有效控制成本。
实际生产环境中,建议结合 vLLM 加速推理、使用 Ray Data 处理超大规模数据,并建立完善的结果质量检查机制。希望本文的代码和调优经验能为你的批量推理实践提供参考。
更多推荐


所有评论(0)