Skip to content

Commit 9ff8bfe

Browse files
committed
feat: add gRPC protocol for fine-tuning with TrainStream RPC
- Add TrainRequest and TrainResponse messages to backend.proto - Add TrainStream server-streaming RPC to Backend service - Update Go gRPC interfaces (interface.go, server.go, client.go, backend.go, embed.go) - Add stub implementation in base.go returning unimplemented error - Regenerate protobuf bindings This implements Phase 1 of the Unsloth fine-tuning backend feature. Subsequent phases will implement the Python backend and Go service layer. Signed-off-by: localai-bot <localai-bot@users.noreply.github.com>
1 parent e832efe commit 9ff8bfe

7 files changed

Lines changed: 170 additions & 0 deletions

File tree

backend/backend.proto

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,8 @@ service Backend {
3939
rpc AudioDecode(AudioDecodeRequest) returns (AudioDecodeResult) {}
4040

4141
rpc ModelMetadata(ModelOptions) returns (ModelMetadataResponse) {}
42+
43+
rpc TrainStream(TrainRequest) returns (stream TrainResponse) {}
4244
}
4345

4446
// Define the empty request
@@ -528,3 +530,49 @@ message ModelMetadataResponse {
528530
string rendered_template = 2; // The rendered chat template with enable_thinking=true (empty if not applicable)
529531
ToolFormatMarkers tool_format = 3; // Auto-detected tool format markers from differential template analysis
530532
}
533+
534+
message TrainRequest {
535+
string model = 1; // Base model name or HuggingFace model ID
536+
string dataset = 2; // Path or HuggingFace dataset ID
537+
string output_dir = 3; // Output directory for the fine-tuned model
538+
539+
// Hyperparameters
540+
int32 epochs = 4; // Number of training epochs
541+
int32 batch_size = 5; // Training batch size
542+
float learning_rate = 6; // Learning rate
543+
int32 max_seq_length = 7; // Maximum sequence length
544+
545+
// LoRA parameters
546+
int32 lora_rank = 8; // LoRA rank (r)
547+
int32 lora_alpha = 9; // LoRA alpha scaling factor
548+
float lora_dropout = 10; // LoRA dropout rate
549+
repeated string target_modules = 11; // LoRA target modules (e.g., "q_proj", "v_proj")
550+
551+
// Quantization
552+
string quantization = 12; // Quantization method (e.g., "4bit", "8bit", "none")
553+
554+
// Additional options
555+
map<string, string> options = 13; // Generic key-value options for backend-specific settings
556+
}
557+
558+
message TrainResponse {
559+
enum TrainStatus {
560+
UNKNOWN = 0;
561+
STARTING = 1;
562+
RUNNING = 2;
563+
COMPLETED = 3;
564+
FAILED = 4;
565+
}
566+
567+
TrainStatus status = 1; // Current training status
568+
float progress = 2; // Overall progress percentage (0-100)
569+
int32 current_epoch = 3; // Current epoch number
570+
int32 total_epochs = 4; // Total number of epochs
571+
int32 current_step = 5; // Current training step
572+
int32 total_steps = 6; // Total number of training steps
573+
float loss = 7; // Current training loss
574+
float learning_rate_current = 8; // Current learning rate (may change with schedulers)
575+
string message = 9; // Human-readable status message
576+
string error = 10; // Error message if status is FAILED
577+
string output_path = 11; // Path to the output model (set on completion)
578+
}

pkg/grpc/backend.go

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -63,4 +63,6 @@ type Backend interface {
6363
AudioDecode(ctx context.Context, in *pb.AudioDecodeRequest, opts ...grpc.CallOption) (*pb.AudioDecodeResult, error)
6464

6565
ModelMetadata(ctx context.Context, in *pb.ModelOptions, opts ...grpc.CallOption) (*pb.ModelMetadataResponse, error)
66+
67+
TrainStream(ctx context.Context, in *pb.TrainRequest, f func(resp *pb.TrainResponse), opts ...grpc.CallOption) error
6668
}

pkg/grpc/base/base.go

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -120,6 +120,10 @@ func (llm *Base) AudioDecode(*pb.AudioDecodeRequest) (*pb.AudioDecodeResult, err
120120
return nil, fmt.Errorf("unimplemented")
121121
}
122122

123+
func (llm *Base) TrainStream(*pb.TrainRequest, chan *pb.TrainResponse) error {
124+
return fmt.Errorf("unimplemented")
125+
}
126+
123127
func memoryUsage() *pb.MemoryUsageData {
124128
mud := pb.MemoryUsageData{
125129
Breakdown: make(map[string]uint64),

pkg/grpc/client.go

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -632,6 +632,54 @@ func (c *Client) AudioDecode(ctx context.Context, in *pb.AudioDecodeRequest, opt
632632
return client.AudioDecode(ctx, in, opts...)
633633
}
634634

635+
func (c *Client) TrainStream(ctx context.Context, in *pb.TrainRequest, f func(resp *pb.TrainResponse), opts ...grpc.CallOption) error {
636+
if !c.parallel {
637+
c.opMutex.Lock()
638+
defer c.opMutex.Unlock()
639+
}
640+
c.setBusy(true)
641+
defer c.setBusy(false)
642+
c.wdMark()
643+
defer c.wdUnMark()
644+
conn, err := grpc.Dial(c.address, grpc.WithTransportCredentials(insecure.NewCredentials()),
645+
grpc.WithDefaultCallOptions(
646+
grpc.MaxCallRecvMsgSize(50*1024*1024), // 50MB
647+
grpc.MaxCallSendMsgSize(50*1024*1024), // 50MB
648+
))
649+
if err != nil {
650+
return err
651+
}
652+
defer conn.Close()
653+
client := pb.NewBackendClient(conn)
654+
655+
stream, err := client.TrainStream(ctx, in, opts...)
656+
if err != nil {
657+
return err
658+
}
659+
660+
for {
661+
select {
662+
case <-ctx.Done():
663+
return ctx.Err()
664+
default:
665+
}
666+
667+
resp, err := stream.Recv()
668+
if err == io.EOF {
669+
break
670+
}
671+
if err != nil {
672+
if ctx.Err() != nil {
673+
return ctx.Err()
674+
}
675+
return err
676+
}
677+
f(resp)
678+
}
679+
680+
return nil
681+
}
682+
635683
func (c *Client) ModelMetadata(ctx context.Context, in *pb.ModelOptions, opts ...grpc.CallOption) (*pb.ModelMetadataResponse, error) {
636684
if !c.parallel {
637685
c.opMutex.Lock()

pkg/grpc/embed.go

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ import (
1010

1111
var _ Backend = new(embedBackend)
1212
var _ pb.Backend_PredictStreamServer = new(embedBackendServerStream)
13+
var _ pb.Backend_TrainStreamServer = new(embedBackendTrainStream)
1314

1415
type embedBackend struct {
1516
s *server
@@ -119,6 +120,14 @@ func (e *embedBackend) ModelMetadata(ctx context.Context, in *pb.ModelOptions, o
119120
return e.s.ModelMetadata(ctx, in)
120121
}
121122

123+
func (e *embedBackend) TrainStream(ctx context.Context, in *pb.TrainRequest, f func(resp *pb.TrainResponse), opts ...grpc.CallOption) error {
124+
bs := &embedBackendTrainStream{
125+
ctx: ctx,
126+
fn: f,
127+
}
128+
return e.s.TrainStream(in, bs)
129+
}
130+
122131
func (e *embedBackend) GetTokenMetrics(ctx context.Context, in *pb.MetricsRequest, opts ...grpc.CallOption) (*pb.MetricsResponse, error) {
123132
return e.s.GetMetrics(ctx, in)
124133
}
@@ -158,3 +167,39 @@ func (e *embedBackendServerStream) SendMsg(m any) error {
158167
func (e *embedBackendServerStream) RecvMsg(m any) error {
159168
return nil
160169
}
170+
171+
type embedBackendTrainStream struct {
172+
ctx context.Context
173+
fn func(resp *pb.TrainResponse)
174+
}
175+
176+
func (e *embedBackendTrainStream) Send(resp *pb.TrainResponse) error {
177+
e.fn(resp)
178+
return nil
179+
}
180+
181+
func (e *embedBackendTrainStream) SetHeader(md metadata.MD) error {
182+
return nil
183+
}
184+
185+
func (e *embedBackendTrainStream) SendHeader(md metadata.MD) error {
186+
return nil
187+
}
188+
189+
func (e *embedBackendTrainStream) SetTrailer(md metadata.MD) {
190+
}
191+
192+
func (e *embedBackendTrainStream) Context() context.Context {
193+
return e.ctx
194+
}
195+
196+
func (e *embedBackendTrainStream) SendMsg(m any) error {
197+
if x, ok := m.(*pb.TrainResponse); ok {
198+
return e.Send(x)
199+
}
200+
return nil
201+
}
202+
203+
func (e *embedBackendTrainStream) RecvMsg(m any) error {
204+
return nil
205+
}

pkg/grpc/interface.go

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,8 @@ type AIModel interface {
3535
AudioDecode(*pb.AudioDecodeRequest) (*pb.AudioDecodeResult, error)
3636

3737
ModelMetadata(*pb.ModelOptions) (*pb.ModelMetadataResponse, error)
38+
39+
TrainStream(*pb.TrainRequest, chan *pb.TrainResponse) error
3840
}
3941

4042
func newReply(s string) *pb.Reply {

pkg/grpc/server.go

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -320,6 +320,27 @@ func (s *server) ModelMetadata(ctx context.Context, in *pb.ModelOptions) (*pb.Mo
320320
return res, nil
321321
}
322322

323+
func (s *server) TrainStream(in *pb.TrainRequest, stream pb.Backend_TrainStreamServer) error {
324+
if s.llm.Locking() {
325+
s.llm.Lock()
326+
defer s.llm.Unlock()
327+
}
328+
resultChan := make(chan *pb.TrainResponse)
329+
330+
done := make(chan bool)
331+
go func() {
332+
for result := range resultChan {
333+
stream.Send(result)
334+
}
335+
done <- true
336+
}()
337+
338+
err := s.llm.TrainStream(in, resultChan)
339+
<-done
340+
341+
return err
342+
}
343+
323344
func StartServer(address string, model AIModel) error {
324345
lis, err := net.Listen("tcp", address)
325346
if err != nil {

0 commit comments

Comments
 (0)