Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 54 additions & 35 deletions cli/azd/extensions/azure.ai.finetune/internal/cmd/operations.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,14 +7,17 @@ import (
"fmt"
"strings"

"github.com/azure/azure-dev/cli/azd/pkg/azdext"
"github.com/azure/azure-dev/cli/azd/pkg/ux"
"github.com/fatih/color"
"github.com/spf13/cobra"

"github.com/azure/azure-dev/cli/azd/pkg/azdext"
"github.com/azure/azure-dev/cli/azd/pkg/ux"

FTYaml "azure.ai.finetune/internal/fine_tuning_yaml"
"azure.ai.finetune/internal/services"
JobWrapper "azure.ai.finetune/internal/tools"
"azure.ai.finetune/internal/utils"
"azure.ai.finetune/pkg/models"
)

func newOperationCommand() *cobra.Command {
Expand All @@ -35,8 +38,8 @@ func newOperationCommand() *cobra.Command {
return cmd
}

// getStatusSymbol returns a symbol representation for job status
func getStatusSymbol(status string) string {
// getStatusSymbolFromString returns a symbol representation for job status
func getStatusSymbolFromString(status string) string {
switch status {
case "pending":
return "⌛"
Expand Down Expand Up @@ -139,44 +142,52 @@ func newOperationSubmitCommand() *cobra.Command {
return cmd
}

// newOperationShowCommand creates a command to show the fine-tuning job details
func newOperationShowCommand() *cobra.Command {
var jobID string

cmd := &cobra.Command{
Use: "show",
Short: "Show the fine tuning job details",
Short: "Show fine-tuning job details.",
RunE: func(cmd *cobra.Command, args []string) error {
ctx := azdext.WithAccessToken(cmd.Context())
azdClient, err := azdext.NewAzdClient()
if err != nil {
return fmt.Errorf("failed to create azd client: %w", err)
}
defer azdClient.Close()
// Show spinner while fetching jobs

// Show spinner while fetching job
spinner := ux.NewSpinner(&ux.SpinnerOptions{
Text: fmt.Sprintf("Fetching fine-tuning job %s...", jobID),
})
if err := spinner.Start(ctx); err != nil {
fmt.Printf("Failed to start spinner: %v\n", err)
fmt.Printf("failed to start spinner: %v\n", err)
}

// Fetch fine-tuning job details using job wrapper
job, err := JobWrapper.GetJobDetails(ctx, azdClient, jobID)
_ = spinner.Stop(ctx)
fineTuneSvc, err := services.NewFineTuningService(ctx, azdClient, nil)
if err != nil {
_ = spinner.Stop(ctx)
fmt.Println()
return err
}

job, err := fineTuneSvc.GetFineTuningJobDetails(ctx, jobID)
_ = spinner.Stop(ctx)
if err != nil {
return fmt.Errorf("failed to get fine-tuning job details: %w", err)
fmt.Println()
return err
}

// Print job details
color.Green("\nFine-Tuning Job Details\n")
fmt.Printf("Job ID: %s\n", job.Id)
fmt.Printf("Status: %s %s\n", getStatusSymbol(job.Status), job.Status)
// Display job details
color.Green("\nFine-tuning Job Details\n")
fmt.Printf("Job ID: %s\n", job.ID)
fmt.Printf("Status: %s %s\n", utils.GetStatusSymbol(job.Status), job.Status)
fmt.Printf("Model: %s\n", job.Model)
fmt.Printf("Fine-tuned Model: %s\n", formatFineTunedModel(job.FineTunedModel))
fmt.Printf("Created At: %s\n", job.CreatedAt)
if job.FinishedAt != "" {
fmt.Printf("Finished At: %s\n", job.FinishedAt)
fmt.Printf("Created At: %s\n", utils.FormatTime(job.CreatedAt))
if !job.FinishedAt.IsZero() {
fmt.Printf("Finished At: %s\n", utils.FormatTime(job.FinishedAt))
}
fmt.Printf("Method: %s\n", job.Method)
fmt.Printf("Training File: %s\n", job.TrainingFile)
Expand All @@ -197,44 +208,47 @@ func newOperationShowCommand() *cobra.Command {
Text: "Fetching job events...",
})
if err := eventsSpinner.Start(ctx); err != nil {
fmt.Printf("Failed to start spinner: %v\n", err)
fmt.Printf("failed to start spinner: %v\n", err)
}

events, err := JobWrapper.GetJobEvents(ctx, azdClient, jobID)
events, err := fineTuneSvc.GetJobEvents(ctx, jobID)
_ = eventsSpinner.Stop(ctx)

if err != nil {
fmt.Printf("Warning: failed to fetch job events: %v\n", err)
fmt.Println()
return err
} else if events != nil && len(events.Data) > 0 {
fmt.Println("\nJob Events:")
for i, event := range events.Data {
fmt.Printf(" %d. [%s] %s - %s\n", i+1, event.Level, event.CreatedAt, event.Message)
fmt.Printf(" %d. Event ID: %s\n", i+1, event.ID)
fmt.Printf(" [%s] %s - %s\n", event.Level, utils.FormatTime(event.CreatedAt), event.Message)
}
if events.HasMore {
fmt.Println(" ... (more events available)")
}
}

// Fetch and print checkpoints if job is completed
if job.Status == "succeeded" {
if job.Status == models.StatusSucceeded {
checkpointsSpinner := ux.NewSpinner(&ux.SpinnerOptions{
Text: "Fetching job checkpoints...",
})
if err := checkpointsSpinner.Start(ctx); err != nil {
fmt.Printf("Failed to start spinner: %v\n", err)
fmt.Printf("failed to start spinner: %v\n", err)
}

checkpoints, err := JobWrapper.GetJobCheckPoints(ctx, azdClient, jobID)
checkpoints, err := fineTuneSvc.GetJobCheckpoints(ctx, jobID)
_ = checkpointsSpinner.Stop(ctx)

if err != nil {
fmt.Printf("Warning: failed to fetch job checkpoints: %v\n", err)
fmt.Println()
return err
} else if checkpoints != nil && len(checkpoints.Data) > 0 {
fmt.Println("\nJob Checkpoints:")
for i, checkpoint := range checkpoints.Data {
fmt.Printf(" %d. Checkpoint ID: %s\n", i+1, checkpoint.ID)
fmt.Printf(" Checkpoint Name: %s\n", checkpoint.FineTunedModelCheckpoint)
fmt.Printf(" Created On: %s\n", checkpoint.CreatedAt)
fmt.Printf(" Created On: %s\n", utils.FormatTime(checkpoint.CreatedAt))
fmt.Printf(" Step Number: %d\n", checkpoint.StepNumber)
if checkpoint.Metrics != nil {
fmt.Printf(" Full Validation Loss: %.6f\n", checkpoint.Metrics.FullValidLoss)
Expand All @@ -251,8 +265,10 @@ func newOperationShowCommand() *cobra.Command {
return nil
},
}

cmd.Flags().StringVarP(&jobID, "job-id", "i", "", "Fine-tuning job ID")
cmd.MarkFlagRequired("job-id")

return cmd
}

Expand All @@ -262,7 +278,7 @@ func newOperationListCommand() *cobra.Command {
var after string
cmd := &cobra.Command{
Use: "list",
Short: "list the fine tuning jobs",
Short: "List fine-tuning jobs.",
RunE: func(cmd *cobra.Command, args []string) error {
ctx := azdext.WithAccessToken(cmd.Context())
azdClient, err := azdext.NewAzdClient()
Expand All @@ -273,7 +289,7 @@ func newOperationListCommand() *cobra.Command {

// Show spinner while fetching jobs
spinner := ux.NewSpinner(&ux.SpinnerOptions{
Text: "fetching fine-tuning jobs...",
Text: "Fetching fine-tuning jobs...",
})
if err := spinner.Start(ctx); err != nil {
fmt.Printf("failed to start spinner: %v\n", err)
Expand All @@ -289,22 +305,25 @@ func newOperationListCommand() *cobra.Command {
jobs, err := fineTuneSvc.ListFineTuningJobs(ctx, limit, after)
_ = spinner.Stop(ctx)
if err != nil {
fmt.Println()
fmt.Println()
return err
}

// Display job list
for i, job := range jobs {
fmt.Printf("\n%d. Job ID: %s | Status: %s %s | Model: %s | Fine-tuned: %s | Created: %s",
i+1, job.ID, getStatusSymbol(string(job.Status)), job.Status, job.BaseModel, formatFineTunedModel(job.FineTunedModel), job.CreatedAt)
i+1, job.ID, utils.GetStatusSymbol(job.Status), job.Status, job.BaseModel,
formatFineTunedModel(job.FineTunedModel), utils.FormatTime(job.CreatedAt))
}

fmt.Printf("\ntotal jobs: %d\n", len(jobs))
fmt.Printf("\nTotal jobs: %d\n", len(jobs))

return nil
},
}
cmd.Flags().IntVarP(&limit, "top", "t", 50, "number of fine-tuning jobs to list")
cmd.Flags().StringVarP(&after, "after", "a", "", "cursor for pagination")

cmd.Flags().IntVarP(&limit, "top", "t", 50, "Number of fine-tuning jobs to list")
cmd.Flags().StringVarP(&after, "after", "a", "", "Cursor for pagination")
return cmd
}

Expand Down Expand Up @@ -361,7 +380,7 @@ func newOperationActionCommand() *cobra.Command {
color.Green(fmt.Sprintf("\nSuccessfully %sd fine-tuning Job!\n", action))
fmt.Printf("Job ID: %s\n", job.Id)
fmt.Printf("Model: %s\n", job.Model)
fmt.Printf("Status: %s %s\n", getStatusSymbol(job.Status), job.Status)
fmt.Printf("Status: %s %s\n", getStatusSymbolFromString(job.Status), job.Status)
fmt.Printf("Created: %s\n", job.CreatedAt)
if job.FineTunedModel != "" {
fmt.Printf("Fine-tuned: %s\n", job.FineTunedModel)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -50,13 +50,13 @@ func (p *AzureProvider) GetFineTuningJobDetails(ctx context.Context, jobID strin
}

// GetJobEvents retrieves events for a fine-tuning job
func (p *AzureProvider) GetJobEvents(ctx context.Context, jobID string, limit int, after string) ([]*models.JobEvent, error) {
func (p *AzureProvider) GetJobEvents(ctx context.Context, jobID string, limit int, after string) (*models.JobEventsList, error) {
// TODO: Implement
return nil, nil
}

// GetJobCheckpoints retrieves checkpoints for a fine-tuning job
func (p *AzureProvider) GetJobCheckpoints(ctx context.Context, jobID string, limit int, after string) ([]*models.JobCheckpoint, error) {
Comment thread
saanikaguptamicrosoft marked this conversation as resolved.
func (p *AzureProvider) GetJobCheckpoints(ctx context.Context, jobID string, limit int, after string) (*models.JobCheckpointsList, error) {
// TODO: Implement
return nil, nil
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,10 +25,10 @@ type FineTuningProvider interface {
GetFineTuningJobDetails(ctx context.Context, jobID string) (*models.FineTuningJobDetail, error)

// GetJobEvents retrieves events for a fine-tuning job
GetJobEvents(ctx context.Context, jobID string, limit int, after string) ([]*models.JobEvent, error)
GetJobEvents(ctx context.Context, jobID string) (*models.JobEventsList, error)

// GetJobCheckpoints retrieves checkpoints for a fine-tuning job
GetJobCheckpoints(ctx context.Context, jobID string, limit int, after string) ([]*models.JobCheckpoint, error)
GetJobCheckpoints(ctx context.Context, jobID string) (*models.JobCheckpointsList, error)

// PauseJob pauses a fine-tuning job
PauseJob(ctx context.Context, jobID string) (*models.FineTuningJob, error)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,11 @@
package openai

import (
"github.com/openai/openai-go/v3"
"github.com/openai/openai-go/v3/packages/pagination"

"azure.ai.finetune/internal/utils"
"azure.ai.finetune/pkg/models"
"github.com/openai/openai-go/v3"
)

// OpenAI Status Constants - matches OpenAI SDK values
Expand Down Expand Up @@ -47,3 +49,75 @@ func convertOpenAIJobToModel(openaiJob openai.FineTuningJob) *models.FineTuningJ
CreatedAt: utils.UnixTimestampToUTC(openaiJob.CreatedAt),
}
}

// convertOpenAIJobToDetailModel converts OpenAI SDK job to detailed domain model
func convertOpenAIJobToDetailModel(openaiJob *openai.FineTuningJob) *models.FineTuningJobDetail {
// Extract hyperparameters from OpenAI job
hyperparameters := &models.Hyperparameters{}
hyperparameters.BatchSize = openaiJob.Hyperparameters.BatchSize.OfInt
hyperparameters.LearningRateMultiplier = openaiJob.Hyperparameters.LearningRateMultiplier.OfFloat
hyperparameters.NEpochs = openaiJob.Hyperparameters.NEpochs.OfInt

jobDetail := &models.FineTuningJobDetail{
ID: openaiJob.ID,
Status: mapOpenAIStatusToJobStatus(openaiJob.Status),
Model: openaiJob.Model,
FineTunedModel: openaiJob.FineTunedModel,
CreatedAt: utils.UnixTimestampToUTC(openaiJob.CreatedAt),
FinishedAt: utils.UnixTimestampToUTC(openaiJob.FinishedAt),
Method: openaiJob.Method.Type,
TrainingFile: openaiJob.TrainingFile,
ValidationFile: openaiJob.ValidationFile,
Hyperparameters: hyperparameters,
}

return jobDetail
}

// convertOpenAIJobEventsToModel converts OpenAI SDK job events to domain model
func convertOpenAIJobEventsToModel(eventsPage *pagination.CursorPage[openai.FineTuningJobEvent]) *models.JobEventsList {
var events []models.JobEvent
for _, event := range eventsPage.Data {
jobEvent := models.JobEvent{
ID: event.ID,
CreatedAt: utils.UnixTimestampToUTC(event.CreatedAt),
Level: string(event.Level),
Message: event.Message,
Data: event.Data,
Type: string(event.Type),
}
events = append(events, jobEvent)
}

return &models.JobEventsList{
Data: events,
HasMore: eventsPage.HasMore,
}
}

// convertOpenAIJobCheckpointsToModel converts OpenAI SDK job checkpoints to domain model
func convertOpenAIJobCheckpointsToModel(checkpointsPage *pagination.CursorPage[openai.FineTuningJobCheckpoint]) *models.JobCheckpointsList {
var checkpoints []models.JobCheckpoint

for _, checkpoint := range checkpointsPage.Data {
metrics := &models.CheckpointMetrics{
FullValidLoss: checkpoint.Metrics.FullValidLoss,
FullValidMeanTokenAccuracy: checkpoint.Metrics.FullValidMeanTokenAccuracy,
}

jobCheckpoint := models.JobCheckpoint{
ID: checkpoint.ID,
CreatedAt: utils.UnixTimestampToUTC(checkpoint.CreatedAt),
FineTunedModelCheckpoint: checkpoint.FineTunedModelCheckpoint,
Metrics: metrics,
FineTuningJobID: checkpoint.FineTuningJobID,
StepNumber: checkpoint.StepNumber,
}
checkpoints = append(checkpoints, jobCheckpoint)
}

return &models.JobCheckpointsList{
Data: checkpoints,
HasMore: checkpointsPage.HasMore,
}
}
Loading