一句话总结
推理服务跑起来了,现在要解决的问题是: 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; string datatype = 2; repeated int64 shape = 3; 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; string datatype = 2; repeated int64 shape = 3; InferTensorContents contents = 4; }
|
为什么不自定义 Proto
Triton 的 gRPC 接口是标准协议,直接用官方 Proto 生成 Go 代码即可. 但在实际项目中,通常会在外面再包一层业务 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; 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, Timeout: 3 * time.Second, PermitWithoutStream: true, }), grpc.WithDefaultCallOptions( grpc.MaxCallRecvMsgSize(4*1024*1024), ), ) 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 }
maxLen := int64(0) for _, ids := range batch { if int64(len(ids)) > maxLen { maxLen = int64(len(ids)) } }
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])
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) { 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 }
switch st.Code() { case codes.Unavailable, codes.ResourceExhausted: time.Sleep(time.Duration(attempt+1) * 10 * time.Millisecond) continue case codes.DeadlineExceeded: return nil, err default: return nil, err } } return nil, fmt.Errorf("max retries exceeded: %w", lastErr) }
|
推理重试的陷阱
推理服务过载时盲目重试会加重雪崩. 必须配合:
- 重试预算(一个请求最多重试 2 次)
- 退避间隔(指数退避或固定间隔)
- 熔断器(连续失败 N 次后停止调用,走降级路径)
连接池与负载均衡
gRPC 连接池
单个 gRPC 连接基于 HTTP/2 多路复用,理论上可以承载大量并发流. 但实践中,单连接的吞吐有上限(受 flow control window 和服务端处理能力限制). 建议维护一个小型连接池:
type ConnPool struct { conns []*grpc.ClientConn clients []pb.GRPCInferenceServiceClient next uint64 }
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) { 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 被删除后,连接断开,客户端需要重连. 确保:
- 客户端有重连逻辑(gRPC-Go 默认有)
- 服务端设置
MaxConnectionAge 让旧连接优雅关闭
- 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 只发文本
动手练习
- 生成 Go 代码: 用
protoc 从 Triton 官方 Proto 生成 Go gRPC 客户端代码
- 单条推理: 实现
Predict 函数,向本地 Triton 发送一条文本的 input_ids,验证返回 logits
- 批量推理: 构造 10 条消息的 batch,验证
BatchPredict 返回 10 个结果
- 超时测试: 把推理超时设为 1ms,验证超时错误被正确捕获和处理
- BatchCollector: 实现客户端聚合器,用 10 个 goroutine 并发调用
Infer,观察实际 gRPC 请求数是否小于 10
- 负载均衡: 部署两个 Triton 副本,配置 round_robin LB,验证请求分散到两个 Pod