阶段四 · 模型部署与 Go 集成

Go 服务集成 gRPC 推理

一句话总结

推理服务跑起来了,现在要解决的问题是: Go 业务层怎么高效地调用它.
本篇用 gRPC 打通 Go 服务与 Triton 推理服务,重点关注批量请求, 超时控制和连接管理.

前置回顾

第 16 篇选定了 Triton Inference Server 作为推理服务,它提供标准 gRPC 接口.
现在从 Go 后端开发者的角度,把这个接口接进来.

Proto 定义:推理请求的契约

Triton 使用 KFServing v2 推理协议,核心 Proto 定义:

syntax = "proto3";
package inference;

service GRPCInferenceService {
rpc ModelInfer(ModelInferRequest) returns (ModelInferResponse) {}
rpc ModelReady(ModelReadyRequest) returns (ModelReadyResponse) {}
}

message ModelInferRequest {
string model_name = 1;
string model_version = 2; // 空字符串 = 用最新版本
repeated InferInputTensor inputs = 3;
repeated InferRequestedOutputTensor outputs = 4;
}

message InferInputTensor {
string name = 1; // "input_ids" / "attention_mask"
string datatype = 2; // "INT64"
repeated int64 shape = 3; // [batch_size, seq_len]
InferTensorContents contents = 4;
}

message InferTensorContents {
repeated int64 int64_contents = 5;
repeated float fp32_contents = 6;
}

message ModelInferResponse {
string model_name = 1;
string model_version = 2;
repeated InferOutputTensor outputs = 3;
}

message InferOutputTensor {
string name = 1; // "logits"
string datatype = 2; // "FP32"
repeated int64 shape = 3; // [batch_size, num_classes]
InferTensorContents contents = 4;
}

为什么不自定义 Proto

Triton 的 gRPC 接口是标准协议,直接用官方 Proto 生成 Go 代码即可. 但在实际项目中,通常会在外面再包一层业务 Proto,隐藏推理协议的复杂性:

// 业务层的审核服务 Proto
service ModerationService {
rpc CheckMessage(CheckMessageRequest) returns (CheckMessageResponse) {}
}
message CheckMessageRequest {
string message_text = 1;
string group_id = 2;
string sender_id = 3;
}
message CheckMessageResponse {
bool is_violation = 1;
string violation_type = 2; // "spam" / "toxic" / "ad"
float confidence = 3;
}

Go 业务服务实现这个 Proto,内部调用 Triton 做推理. 上游调用方不需要知道底层用了什么模型.

Go gRPC 客户端实现

基础客户端

package inference

import (
"context"
"fmt"
"time"

pb "your-project/proto/triton"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/keepalive"
)

type Client struct {
conn *grpc.ClientConn
client pb.GRPCInferenceServiceClient
}

func NewClient(addr string) (*Client, error) {
conn, err := grpc.NewClient(addr,
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithKeepaliveParams(keepalive.ClientParameters{
Time: 10 * time.Second, // 每 10s 发 keepalive ping
Timeout: 3 * time.Second, // ping 超时
PermitWithoutStream: true, // 无活跃流也保持连接
}),
grpc.WithDefaultCallOptions(
grpc.MaxCallRecvMsgSize(4*1024*1024), // 4MB, 批量推理结果可能较大
),
)
if err != nil {
return nil, fmt.Errorf("dial triton: %w", err)
}
return &Client{
conn: conn,
client: pb.NewGRPCInferenceServiceClient(conn),
}, nil
}

func (c *Client) Close() error {
return c.conn.Close()
}

单条推理

func (c *Client) Predict(ctx context.Context, inputIDs, attentionMask []int64) ([]float32, error) {
seqLen := int64(len(inputIDs))

req := &pb.ModelInferRequest{
ModelName: "sentiment_model",
Inputs: []*pb.InferInputTensor{
{
Name: "input_ids",
Datatype: "INT64",
Shape: []int64{1, seqLen},
Contents: &pb.InferTensorContents{Int64Contents: inputIDs},
},
{
Name: "attention_mask",
Datatype: "INT64",
Shape: []int64{1, seqLen},
Contents: &pb.InferTensorContents{Int64Contents: attentionMask},
},
},
Outputs: []*pb.InferRequestedOutputTensor{
{Name: "logits"},
},
}

resp, err := c.client.ModelInfer(ctx, req)
if err != nil {
return nil, fmt.Errorf("model infer: %w", err)
}

if len(resp.Outputs) == 0 {
return nil, fmt.Errorf("empty output from model")
}
return resp.Outputs[0].Contents.Fp32Contents, nil
}

批量推理

群聊场景下消息是突发的. 一个群 50 条消息几乎同时到达,逐条调用推理服务效率极低. 批量推理把多条消息打包成一个请求:

type PredictResult struct {
Scores []float32
Label int
}

func (c *Client) BatchPredict(ctx context.Context, batch [][]int64, masks [][]int64) ([]PredictResult, error) {
batchSize := int64(len(batch))
if batchSize == 0 {
return nil, nil
}

// 找到最长序列,做 padding
maxLen := int64(0)
for _, ids := range batch {
if int64(len(ids)) > maxLen {
maxLen = int64(len(ids))
}
}

// 展平为一维数组(Triton 要求)
flatIDs := make([]int64, 0, batchSize*maxLen)
flatMask := make([]int64, 0, batchSize*maxLen)

for i := range batch {
padded := make([]int64, maxLen)
mask := make([]int64, maxLen)
copy(padded, batch[i])
copy(mask, masks[i])
flatIDs = append(flatIDs, padded...)
flatMask = append(flatMask, mask...)
}

req := &pb.ModelInferRequest{
ModelName: "sentiment_model",
Inputs: []*pb.InferInputTensor{
{
Name: "input_ids",
Datatype: "INT64",
Shape: []int64{batchSize, maxLen},
Contents: &pb.InferTensorContents{Int64Contents: flatIDs},
},
{
Name: "attention_mask",
Datatype: "INT64",
Shape: []int64{batchSize, maxLen},
Contents: &pb.InferTensorContents{Int64Contents: flatMask},
},
},
Outputs: []*pb.InferRequestedOutputTensor{
{Name: "logits"},
},
}

resp, err := c.client.ModelInfer(ctx, req)
if err != nil {
return nil, fmt.Errorf("batch infer: %w", err)
}

scores := resp.Outputs[0].Contents.Fp32Contents
numClasses := int(resp.Outputs[0].Shape[1]) // shape: [batch_size, num_classes]

results := make([]PredictResult, batchSize)
for i := range results {
start := i * numClasses
end := start + numClasses
results[i].Scores = scores[start:end]
results[i].Label = argmax(scores[start:end])
}
return results, nil
}

func argmax(scores []float32) int {
maxIdx := 0
for i := 1; i < len(scores); i++ {
if scores[i] > scores[maxIdx] {
maxIdx = i
}
}
return maxIdx
}

客户端侧 Batching:请求聚合

Triton 的 dynamic batching 是服务端聚合. 但客户端也可以做一层聚合,减少 gRPC 调用次数:

type BatchCollector struct {
client *Client
maxBatch int
maxWait time.Duration
pending chan *inferRequest
}

type inferRequest struct {
inputIDs []int64
attentionMask []int64
result chan<- PredictResult
err chan<- error
}

func NewBatchCollector(client *Client, maxBatch int, maxWait time.Duration) *BatchCollector {
bc := &BatchCollector{
client: client,
maxBatch: maxBatch,
maxWait: maxWait,
pending: make(chan *inferRequest, maxBatch*4),
}
go bc.loop()
return bc
}

func (bc *BatchCollector) Infer(ctx context.Context, ids, mask []int64) (PredictResult, error) {
resultCh := make(chan PredictResult, 1)
errCh := make(chan error, 1)

select {
case bc.pending <- &inferRequest{
inputIDs: ids, attentionMask: mask,
result: resultCh, err: errCh,
}:
case <-ctx.Done():
return PredictResult{}, ctx.Err()
}

select {
case r := <-resultCh:
return r, nil
case e := <-errCh:
return PredictResult{}, e
case <-ctx.Done():
return PredictResult{}, ctx.Err()
}
}

func (bc *BatchCollector) loop() {
for {
// 等第一个请求到达
first := <-bc.pending
batch := []*inferRequest{first}

// 开始计时,尽量凑满一批
timer := time.NewTimer(bc.maxWait)
collect:
for len(batch) < bc.maxBatch {
select {
case req := <-bc.pending:
batch = append(batch, req)
case <-timer.C:
break collect
}
}
timer.Stop()

bc.executeBatch(batch)
}
}

func (bc *BatchCollector) executeBatch(batch []*inferRequest) {
ids := make([][]int64, len(batch))
masks := make([][]int64, len(batch))
for i, req := range batch {
ids[i] = req.inputIDs
masks[i] = req.attentionMask
}

ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()

results, err := bc.client.BatchPredict(ctx, ids, masks)
if err != nil {
for _, req := range batch {
req.err <- err
}
return
}
for i, req := range batch {
req.result <- results[i]
}
}

客户端 Batching = 电梯调度. 电梯不会每来一个人就跑一趟,而是等几秒或等满载再出发. 同理,推理请求不必逐条发送,聚合成 batch 一次发出,GPU 利用率更高,整体吞吐更大.

超时与重试策略

分层超时

func (c *Client) PredictWithTimeout(ctx context.Context, ids, mask []int64) ([]float32, error) {
// 外层 context 是业务超时(比如 HTTP handler 的 deadline)
// 内层加一个推理专用超时,防止推理卡住拖垮整个请求链路
inferCtx, cancel := context.WithTimeout(ctx, 30*time.Millisecond)
defer cancel()

scores, err := c.Predict(inferCtx, ids, mask)
if err != nil {
return nil, fmt.Errorf("inference timeout or error: %w", err)
}
return scores, nil
}

超时分层逻辑:

HTTP 请求超时: 200ms (整个审核流程)
└── 分词处理: 5ms
└── 推理调用: 30ms (p99 目标)
└── 后处理 + 规则: 5ms
└── 缓冲: 160ms (网络波动, GC 等)

重试策略

推理调用的重试需要区分错误类型:

import (
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)

func (c *Client) PredictWithRetry(ctx context.Context, ids, mask []int64) ([]float32, error) {
var lastErr error
for attempt := 0; attempt < 3; attempt++ {
scores, err := c.Predict(ctx, ids, mask)
if err == nil {
return scores, nil
}
lastErr = err

st, ok := status.FromError(err)
if !ok {
return nil, err // 非 gRPC 错误,不重试
}

switch st.Code() {
case codes.Unavailable, codes.ResourceExhausted:
// 服务暂时不可用或过载,可以重试
time.Sleep(time.Duration(attempt+1) * 10 * time.Millisecond)
continue
case codes.DeadlineExceeded:
// 超时了,但剩余 budget 不够再试
return nil, err
default:
// InvalidArgument, Internal 等,重试无意义
return nil, err
}
}
return nil, fmt.Errorf("max retries exceeded: %w", lastErr)
}

推理重试的陷阱

推理服务过载时盲目重试会加重雪崩. 必须配合:

  1. 重试预算(一个请求最多重试 2 次)
  2. 退避间隔(指数退避或固定间隔)
  3. 熔断器(连续失败 N 次后停止调用,走降级路径)

连接池与负载均衡

gRPC 连接池

单个 gRPC 连接基于 HTTP/2 多路复用,理论上可以承载大量并发流. 但实践中,单连接的吞吐有上限(受 flow control window 和服务端处理能力限制). 建议维护一个小型连接池:

type ConnPool struct {
conns []*grpc.ClientConn
clients []pb.GRPCInferenceServiceClient
next uint64 // atomic round-robin counter
}

func NewConnPool(addr string, size int) (*ConnPool, error) {
pool := &ConnPool{
conns: make([]*grpc.ClientConn, size),
clients: make([]pb.GRPCInferenceServiceClient, size),
}
for i := 0; i < size; i++ {
conn, err := grpc.NewClient(addr,
grpc.WithTransportCredentials(insecure.NewCredentials()),
)
if err != nil {
pool.Close()
return nil, err
}
pool.conns[i] = conn
pool.clients[i] = pb.NewGRPCInferenceServiceClient(conn)
}
return pool, nil
}

func (p *ConnPool) Get() pb.GRPCInferenceServiceClient {
idx := atomic.AddUint64(&p.next, 1) % uint64(len(p.clients))
return p.clients[idx]
}

func (p *ConnPool) Close() {
for _, conn := range p.conns {
if conn != nil {
conn.Close()
}
}
}

客户端负载均衡

如果 Triton 有多个副本,Go 客户端有两种负载均衡方式:

方式 实现 优缺点
DNS 轮询 K8s Service 的 ClusterIP 简单,但连接复用导致流量不均
客户端 LB grpc.WithDefaultServiceConfig 配置 round_robin 连接级均衡,推荐
外部 LB Envoy/Istio sidecar 功能最全,但引入额外延迟和复杂度

推荐方案:使用 K8s headless Service + gRPC 客户端 round_robin:

import "google.golang.org/grpc/resolver"

func NewClientWithLB(serviceName string) (*Client, error) {
// dns:///triton-inference.default.svc.cluster.local:8001
target := fmt.Sprintf("dns:///%s:8001", serviceName)

conn, err := grpc.NewClient(target,
grpc.WithTransportCredentials(insecure.NewCredentials()),
grpc.WithDefaultServiceConfig(`{"loadBalancingConfig": [{"round_robin":{}}]}`),
)
if err != nil {
return nil, err
}
return &Client{conn: conn, client: pb.NewGRPCInferenceServiceClient(conn)}, nil
}

Tokenizer 的位置问题

一个容易忽略的细节: Triton 接收的是 input_ids(整数数组),不是原始文本. 分词(Tokenization)在哪里做?

方案 A: Go 侧分词                    方案 B: 推理服务侧分词

Go 服务 Go 服务
│ │
├── tokenize(text) → input_ids ├── 直接发送原始 text
│ (需要 Go 分词库或 cgo) │
└── gRPC(input_ids) → Triton └── gRPC(text) → Triton
└── Python ensemble:
tokenize → model → postprocess
方案 优点 缺点
Go 侧分词 减少网络传输, 推理服务更纯粹 Go 生态缺成熟 Tokenizer, cgo 引入复杂度
服务侧分词 Go 只需发文本, Tokenizer 与模型版本绑定 推理服务承担额外计算, 接口不够通用

推荐: 服务侧分词. Triton 支持 Ensemble 模式,把 Tokenizer 和模型串成 pipeline. Go 侧只需发送原始文本,由 Triton 内部完成 tokenize → infer → postprocess 全流程.

Triton Ensemble Pipeline:

┌──────────────────────────────────────────────────────┐
│ ensemble_model │
│ │
│ step 1: tokenizer (Python backend) │
│ input: raw_text (string) │
│ output: input_ids, attention_mask │
│ │
│ step 2: sentiment_model (ONNX backend) │
│ input: input_ids, attention_mask │
│ output: logits │
│ │
│ step 3: postprocess (Python backend) │
│ input: logits │
│ output: label, confidence │
└──────────────────────────────────────────────────────┘

这样 Go 客户端代码简化为:

func (c *Client) PredictText(ctx context.Context, text string) (PredictResult, error) {
req := &pb.ModelInferRequest{
ModelName: "ensemble_model",
Inputs: []*pb.InferInputTensor{
{
Name: "raw_text",
Datatype: "BYTES",
Shape: []int64{1, 1},
Contents: &pb.InferTensorContents{
BytesContents: [][]byte{[]byte(text)},
},
},
},
Outputs: []*pb.InferRequestedOutputTensor{
{Name: "label"},
{Name: "confidence"},
},
}
resp, err := c.client.ModelInfer(ctx, req)
if err != nil {
return PredictResult{}, err
}
// 解析响应...
return parseEnsembleResponse(resp), nil
}

生产注意事项

gRPC 长连接与 K8s 滚动更新

gRPC 长连接在 Pod 滚动更新时不会自动迁移. 旧 Pod 被删除后,连接断开,客户端需要重连. 确保:

  1. 客户端有重连逻辑(gRPC-Go 默认有)
  2. 服务端设置 MaxConnectionAge 让旧连接优雅关闭
  3. K8s 的 terminationGracePeriodSeconds 给足时间处理 in-flight 请求

Proto 版本管理

推理服务的 Proto 由 Triton 项目维护,升级 Triton 版本时 Proto 可能变化. 建议在 Go 项目中 vendor Triton 的 Proto 文件,锁定版本,避免上游更新导致编译失败.

快速回顾

  • 标准 Proto: 使用 KFServing v2 协议,Go 直接生成客户端代码
  • 批量推理: 多条消息打包成 batch 发送,GPU 利用率倍增
  • 客户端聚合: BatchCollector 模式,等待 N 条或 T 时间后统一发送
  • 超时分层: 推理调用设独立超时(30ms),不拖垮整个请求链路
  • 重试区分错误: Unavailable 重试, InvalidArgument 不重试, 过载时配合熔断
  • Tokenizer 在服务侧: Triton Ensemble 把分词和推理串成 pipeline, Go 只发文本

动手练习

  1. 生成 Go 代码: 用 protoc 从 Triton 官方 Proto 生成 Go gRPC 客户端代码
  2. 单条推理: 实现 Predict 函数,向本地 Triton 发送一条文本的 input_ids,验证返回 logits
  3. 批量推理: 构造 10 条消息的 batch,验证 BatchPredict 返回 10 个结果
  4. 超时测试: 把推理超时设为 1ms,验证超时错误被正确捕获和处理
  5. BatchCollector: 实现客户端聚合器,用 10 个 goroutine 并发调用 Infer,观察实际 gRPC 请求数是否小于 10
  6. 负载均衡: 部署两个 Triton 副本,配置 round_robin LB,验证请求分散到两个 Pod