diff --git a/README.md b/README.md index afdaf1c..ad9fe02 100644 --- a/README.md +++ b/README.md @@ -43,6 +43,7 @@ Core commands: Alias commands: run Alias for "prediction create" + stream Alias for "prediction create --stream" train Alias for "training create" Additional Commands: @@ -67,6 +68,20 @@ $ replicate run stability-ai/sdxl \ Prediction created: https://replicate.com/p/jpgp263bdekvxileu2ppsy46v4 ``` +### Stream prediction output + +Run [LLaMA 2] and stream output tokens to your terminal. + +```console +$ replicate stream meta/llama-2-70b-chat \ + prompt="Tell me a joke about llamas" +Sure, here's a joke about llamas for you: + +Why did the llama refuse to play poker? + +Because he always got fleeced! +``` + ### Create a local development environment from a prediction Create a Node.js or Python project from a prediction. diff --git a/cmd/replicate/main.go b/cmd/replicate/main.go index 0b9fae6..d183d76 100644 --- a/cmd/replicate/main.go +++ b/cmd/replicate/main.go @@ -50,6 +50,7 @@ func init() { for _, cmd := range []*cobra.Command{ cmd.RunCmd, cmd.TrainCmd, + cmd.StreamCmd, } { rootCmd.AddCommand(cmd) cmd.GroupID = "alias" diff --git a/demo.gif b/demo.gif index ea6f059..6915ccf 100644 Binary files a/demo.gif and b/demo.gif differ diff --git a/demo.tape b/demo.tape index 266520f..17a2d97 100644 --- a/demo.tape +++ b/demo.tape @@ -17,6 +17,16 @@ Sleep 100ms Ctrl+C # Don't actually set the API key Sleep 1s +Type 'replicate stream meta/llama-2-70b-chat \' +Enter +Type@50ms ' prompt="write a haiku about corgis"' +Enter + +Sleep 2s + +Enter + +Enter Type 'replicate run stability-ai/sdxl \' Enter Type@50ms ' prompt="a studio photo of a rainbow colored corgi" \' @@ -24,4 +34,4 @@ Enter Type@50ms ' width=512 height=512 seed=42069' Enter -Sleep 10s +Sleep 15s diff --git a/go.mod b/go.mod index 8694f56..8e588e3 100644 --- a/go.mod +++ b/go.mod @@ -14,7 +14,7 @@ require ( github.com/getkin/kin-openapi v0.120.0 github.com/golangci/golangci-lint v1.55.2 github.com/mattn/go-isatty v0.0.20 - github.com/replicate/replicate-go v0.12.0 + github.com/replicate/replicate-go v0.13.2 github.com/schollz/progressbar/v3 v3.13.1 github.com/spf13/cobra v1.8.0 github.com/stretchr/testify v1.8.4 diff --git a/go.sum b/go.sum index e7f6834..63d1d57 100644 --- a/go.sum +++ b/go.sum @@ -522,8 +522,12 @@ github.com/quasilyte/regex/syntax v0.0.0-20210819130434-b3f0c404a727 h1:TCg2WBOl github.com/quasilyte/regex/syntax v0.0.0-20210819130434-b3f0c404a727/go.mod h1:rlzQ04UMyJXu/aOvhd8qT+hvDrFpiwqp8MRXDY9szc0= github.com/quasilyte/stdinfo v0.0.0-20220114132959-f7386bf02567 h1:M8mH9eK4OUR4lu7Gd+PU1fV2/qnDNfzT635KRSObncs= github.com/quasilyte/stdinfo v0.0.0-20220114132959-f7386bf02567/go.mod h1:DWNGW8A4Y+GyBgPuaQJuWiy0XYftx4Xm/y5Jqk9I6VQ= -github.com/replicate/replicate-go v0.12.0 h1:gd/hw4hCBO5G4M3Fezb3zdKYSbe9NEfRLzGoktFk3Ks= -github.com/replicate/replicate-go v0.12.0/go.mod h1:k9C4+PaYa9+hMRjn4D7ZPHOCUFb8P4jhytsCqcGa2vU= +github.com/replicate/replicate-go v0.13.0 h1:DWpSw8ck+dVK79jcVbg0iWJt4/ajcDaYX7FmiqSh2iI= +github.com/replicate/replicate-go v0.13.0/go.mod h1:k9C4+PaYa9+hMRjn4D7ZPHOCUFb8P4jhytsCqcGa2vU= +github.com/replicate/replicate-go v0.13.1 h1:+WgP8hoWuw8e0ZCA1RlVQZZrDkaksMZJCk8C1i+icp0= +github.com/replicate/replicate-go v0.13.1/go.mod h1:k9C4+PaYa9+hMRjn4D7ZPHOCUFb8P4jhytsCqcGa2vU= +github.com/replicate/replicate-go v0.13.2 h1:S+ENs0cKMlizZzh9Ht/Diy66FCPKSdXNRh/9QvKyf+8= +github.com/replicate/replicate-go v0.13.2/go.mod h1:k9C4+PaYa9+hMRjn4D7ZPHOCUFb8P4jhytsCqcGa2vU= github.com/rivo/uniseg v0.1.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/rivo/uniseg v0.4.4 h1:8TfxU8dW6PdqD27gjM8MVNuicgxIjxpm4K7x4jp8sis= diff --git a/go.work.sum b/go.work.sum index 6f86d97..aaa6194 100644 --- a/go.work.sum +++ b/go.work.sum @@ -33,6 +33,8 @@ github.com/quasilyte/go-ruleguard/dsl v0.3.22/go.mod h1:KeCP03KrjuSO0H1kTuZQCWlQ github.com/quasilyte/go-ruleguard/rules v0.0.0-20211022131956-028d6511ab71/go.mod h1:4cgAphtvu7Ftv7vOT2ZOYhC6CvBxZixcasr8qIOTA50= github.com/replicate/replicate-go v0.8.1 h1:Mza5hWR/R1akZRKwXtA/CQJ2pY4/B9fSCYX+2nTb8zo= github.com/replicate/replicate-go v0.8.1/go.mod h1:k9C4+PaYa9+hMRjn4D7ZPHOCUFb8P4jhytsCqcGa2vU= +github.com/replicate/replicate-go v0.13.0 h1:DWpSw8ck+dVK79jcVbg0iWJt4/ajcDaYX7FmiqSh2iI= +github.com/replicate/replicate-go v0.13.0/go.mod h1:k9C4+PaYa9+hMRjn4D7ZPHOCUFb8P4jhytsCqcGa2vU= github.com/sahilm/fuzzy v0.1.0/go.mod h1:VFvziUEIMCrT6A6tw2RFIXPXXmzXbOsSHF0DOI8ZK9Y= github.com/shoenig/go-m1cpu v0.1.6/go.mod h1:1JJMcUBvfNwpq05QDQVAnx3gUHr9IYF7GNg9SUEw2VQ= github.com/valyala/bytebufferpool v1.0.0/go.mod h1:6bBcMArwyJ5K/AmCkWv1jt77kVWyCJ6HpOuEn7z0Csc= diff --git a/internal/cmd/model/schema.go b/internal/cmd/model/schema.go index fbe1647..8ab0502 100644 --- a/internal/cmd/model/schema.go +++ b/internal/cmd/model/schema.go @@ -71,8 +71,8 @@ func printModelVersionSchema(version *replicate.ModelVersion) error { if inputSchema != nil { fmt.Println("Inputs:") - for _, propName := range util.SortedKeys(inputSchema.Value.Properties) { - prop, ok := inputSchema.Value.Properties[propName] + for _, propName := range util.SortedKeys(inputSchema.Properties) { + prop, ok := inputSchema.Properties[propName] if !ok { continue } @@ -91,9 +91,9 @@ func printModelVersionSchema(version *replicate.ModelVersion) error { if outputSchema != nil { fmt.Println("Output:") - fmt.Printf("- type: %s\n", outputSchema.Value.Type) - if outputSchema.Value.Type == "array" { - fmt.Printf("- items: %s %s\n", outputSchema.Value.Items.Value.Type, outputSchema.Value.Items.Value.Format) + fmt.Printf("- type: %s\n", outputSchema.Type) + if outputSchema.Type == "array" { + fmt.Printf("- items: %s %s\n", outputSchema.Items.Value.Type, outputSchema.Items.Value.Format) } fmt.Println() } diff --git a/internal/cmd/prediction/create.go b/internal/cmd/prediction/create.go index d32d381..fd8b89b 100644 --- a/internal/cmd/prediction/create.go +++ b/internal/cmd/prediction/create.go @@ -3,11 +3,14 @@ package prediction import ( "encoding/json" "fmt" + "os" "path/filepath" + "strings" "time" "github.com/briandowns/spinner" "github.com/cli/browser" + "github.com/getkin/kin-openapi/openapi3" "github.com/replicate/cli/internal/identifier" "github.com/replicate/cli/internal/util" "github.com/replicate/replicate-go" @@ -40,20 +43,12 @@ var CreateCmd = &cobra.Command{ var version *replicate.ModelVersion if id.Version == "" { - model, err := client.GetModel(ctx, id.Owner, id.Name) - if err != nil { - return fmt.Errorf("failed to get model: %w", err) - } - - if model.LatestVersion == nil { - return fmt.Errorf("no versions found for model %s", args[0]) + if model, err := client.GetModel(ctx, id.Owner, id.Name); err == nil { + version = model.LatestVersion } - - version = model.LatestVersion } else { - version, err = client.GetModelVersion(ctx, id.Owner, id.Name, id.Version) - if err != nil { - return fmt.Errorf("failed to get model version: %w", err) + if v, err := client.GetModelVersion(ctx, id.Owner, id.Name, id.Version); err == nil { + version = v } } @@ -68,25 +63,68 @@ var CreateCmd = &cobra.Command{ return fmt.Errorf("failed to parse inputs: %w", err) } - inputSchema, _, err := util.GetSchemas(*version) - if err != nil { - return fmt.Errorf("failed to get input schema for version: %w", err) + var inputSchema *openapi3.Schema + var outputSchema *openapi3.Schema + if version != nil { + inputSchema, outputSchema, err = util.GetSchemas(*version) + if err != nil { + return fmt.Errorf("failed to get input schema for version: %w", err) + } } - coercedInputs, err := util.CoerceTypes(inputs, inputSchema.Value) + coercedInputs, err := util.CoerceTypes(inputs, inputSchema) if err != nil { return fmt.Errorf("failed to coerce inputs: %w", err) } + shouldWait := cmd.Flags().Changed("wait") || !cmd.Flags().Changed("no-wait") + shouldStream := !cmd.Flags().Changed("wait") && cmd.Flags().Changed("stream") || (outputSchema != nil && outputSchema.Type == "array" && outputSchema.Items.Value.Type == "string" && outputSchema.Items.Value.Format != "uri") + s.Start() - prediction, err := client.CreatePrediction(ctx, version.ID, coercedInputs, nil, false) + var prediction *replicate.Prediction + if id.Version == "" { + prediction, err = client.CreatePredictionWithModel(ctx, id.Owner, id.Name, coercedInputs, nil, shouldStream) + // TODO: check status code + if err != nil { + if version != nil { + prediction, err = client.CreatePrediction(ctx, version.ID, coercedInputs, nil, shouldStream) + } + } + } else { + prediction, err = client.CreatePrediction(ctx, id.Version, coercedInputs, nil, shouldStream) + } if err != nil { return fmt.Errorf("failed to create prediction: %w", err) } s.Stop() - shouldWait := cmd.Flags().Changed("wait") || !cmd.Flags().Changed("no-wait") - if cmd.Flags().Changed("json") || !util.IsTTY() { + hasStream := prediction.URLs["stream"] != "" + + if !util.IsTTY() || cmd.Flags().Changed("json") { + if hasStream { + events, _ := client.StreamPrediction(ctx, prediction) + + if cmd.Flags().Changed("json") { + fmt.Print("[") + defer fmt.Print("]") + } + + for event := range events { + if cmd.Flags().Changed("json") { + b, err := json.Marshal(event.Data) + if err != nil { + return fmt.Errorf("failed to marshal event: %w", err) + } + fmt.Printf("%s,", string(b)) + } else { + fmt.Print(event.Data) + } + } + fmt.Println("") + + return nil + } + if shouldWait { err = client.Wait(ctx, prediction) if err != nil { @@ -98,91 +136,148 @@ var CreateCmd = &cobra.Command{ if err != nil { return fmt.Errorf("failed to marshal prediction: %w", err) } - fmt.Println(string(b)) + return nil - } else { - url := fmt.Sprintf("https://replicate.com/p/%s", prediction.ID) - fmt.Printf("Prediction created: %s\n", url) + } - if cmd.Flags().Changed("web") { - if util.IsTTY() { - fmt.Println("Opening in browser...") - } + url := fmt.Sprintf("https://replicate.com/p/%s", prediction.ID) + if !hasStream { + fmt.Printf("Prediction created: %s\n", url) + } - err = browser.OpenURL(url) - if err != nil { - return fmt.Errorf("failed to open browser: %w", err) - } + if cmd.Flags().Changed("web") { + if util.IsTTY() { + fmt.Println("Opening in browser...") + } - return nil + err = browser.OpenURL(url) + if err != nil { + return fmt.Errorf("failed to open browser: %w", err) } - if shouldWait { - bar := progressbar.Default(100) - bar.Describe("processing") - - predChan, errChan := client.WaitAsync(ctx, prediction) - for pred := range predChan { - progress := pred.Progress() - if progress != nil { - bar.ChangeMax(progress.Total) - _ = bar.Set(progress.Current) + return nil + } + + if hasStream { + sseChan, errChan := client.StreamPrediction(ctx, prediction) + + tokens := []string{} + for { + select { + case event, ok := <-sseChan: + if !ok { + return nil } - if pred.Status.Terminated() { - _ = bar.Finish() - break + switch event.Type { + case "output": + token := event.Data + tokens = append(tokens, token) + fmt.Print(token) + case "logs": + // TODO: print logs to stderr + default: + // ignore + } + case err, ok := <-errChan: + if !ok { + return nil } - } - if err := <-errChan; err != nil { - return fmt.Errorf("failed to wait for prediction: %w", err) + return fmt.Errorf("streaming error: %w", err) } - switch prediction.Status { - case replicate.Succeeded: - fmt.Println("✅ Succeeded") - bytes, err := json.MarshalIndent(prediction.Output, "", " ") - if err != nil { - return fmt.Errorf("failed to marshal output: %w", err) + if cmd.Flags().Changed("save") { + var dirname string + if cmd.Flags().Changed("output-directory") { + dirname = cmd.Flag("output-directory").Value.String() + } else { + dirname = fmt.Sprintf("./%s", prediction.ID) } - fmt.Println(string(bytes)) - - if cmd.Flags().Changed("save") { - var dirname string - if cmd.Flags().Changed("output-directory") { - dirname = cmd.Flag("output-directory").Value.String() - } else { - dirname = fmt.Sprintf("./%s", prediction.ID) - } - dir, err := filepath.Abs(dirname) - if err != nil { - return fmt.Errorf("failed to create output directory: %w", err) - } + dir, err := filepath.Abs(dirname) + if err != nil { + return fmt.Errorf("failed to create output directory: %w", err) + } - err = util.DownloadPrediction(ctx, *prediction, dir) - if err != nil { - return fmt.Errorf("failed to save output: %w", err) - } + // write tokens to file + err = os.MkdirAll(dir, 0o755) + if err != nil { + return fmt.Errorf("failed to create directory: %w", err) } - case replicate.Failed: - fmt.Println("❌ Failed") - fmt.Println(*prediction.Logs) - bytes, err := json.MarshalIndent(prediction.Error, "", " ") + + err = os.WriteFile(filepath.Join(dir, "output.txt"), []byte(strings.Join(tokens, "")), 0o644) if err != nil { - return fmt.Errorf("error: %v", prediction.Error) + return fmt.Errorf("failed to write output: %w", err) } - fmt.Println(string(bytes)) - case replicate.Canceled: - fmt.Println("🚫 Canceled") - fmt.Println(prediction.Logs) } } + } else if shouldWait { + bar := progressbar.Default(100) + bar.Describe("processing") + + predChan, errChan := client.WaitAsync(ctx, prediction) + for pred := range predChan { + progress := pred.Progress() + if progress != nil { + bar.ChangeMax(progress.Total) + _ = bar.Set(progress.Current) + } - return nil + if pred.Status.Terminated() { + _ = bar.Finish() + break + } + } + + if err := <-errChan; err != nil { + return fmt.Errorf("failed to wait for prediction: %w", err) + } + + switch prediction.Status { + case replicate.Succeeded: + fmt.Println("✅ Succeeded") + bytes, err := json.MarshalIndent(prediction.Output, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal output: %w", err) + } + fmt.Println(string(bytes)) + case replicate.Failed: + fmt.Println("❌ Failed") + fmt.Println(*prediction.Logs) + bytes, err := json.MarshalIndent(prediction.Error, "", " ") + if err != nil { + return fmt.Errorf("error: %v", prediction.Error) + } + fmt.Println(string(bytes)) + case replicate.Canceled: + fmt.Println("🚫 Canceled") + fmt.Println(prediction.Logs) + } + + if cmd.Flags().Changed("save") && prediction.Status == replicate.Succeeded { + var dirname string + if cmd.Flags().Changed("output-directory") { + dirname = cmd.Flag("output-directory").Value.String() + } else { + dirname = fmt.Sprintf("./%s", prediction.ID) + } + + dir, err := filepath.Abs(dirname) + if err != nil { + return fmt.Errorf("failed to create output directory: %w", err) + } + + err = util.DownloadPrediction(ctx, *prediction, dir) + if err != nil { + return fmt.Errorf("failed to save output: %w", err) + } + } } + + return nil + }, } @@ -194,11 +289,14 @@ func AddCreateFlags(cmd *cobra.Command) { cmd.Flags().Bool("json", false, "Emit JSON") cmd.Flags().Bool("no-wait", false, "Don't wait for prediction to complete") cmd.Flags().BoolP("wait", "w", true, "Wait for prediction to complete") + cmd.Flags().Bool("stream", false, "Stream prediction output") cmd.Flags().Bool("web", false, "View on web") cmd.Flags().String("separator", "=", "Separator between input key and value") cmd.Flags().Bool("save", false, "Save prediction outputs to directory") cmd.Flags().String("output-directory", "", "Output directory, defaults to ./{prediction-id}") cmd.MarkFlagsMutuallyExclusive("json", "web") + cmd.MarkFlagsMutuallyExclusive("json", "stream") + cmd.MarkFlagsMutuallyExclusive("stream", "wait") cmd.MarkFlagsMutuallyExclusive("wait", "no-wait") } diff --git a/internal/cmd/stream.go b/internal/cmd/stream.go new file mode 100644 index 0000000..b150971 --- /dev/null +++ b/internal/cmd/stream.go @@ -0,0 +1,25 @@ +package cmd + +import ( + "github.com/spf13/cobra" + + "github.com/replicate/cli/internal/cmd/prediction" +) + +var StreamCmd = &cobra.Command{ + Use: "stream [input=value] ... [flags]", + Short: `Alias for "prediction create --stream"`, + Args: cobra.MinimumNArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + err := cmd.Flags().Set("stream", "true") + if err != nil { + return err + } + + return prediction.CreateCmd.RunE(cmd, args) + }, +} + +func init() { + prediction.AddCreateFlags(StreamCmd) +} diff --git a/internal/util/schema.go b/internal/util/schema.go index 0d1d461..ba36ed9 100644 --- a/internal/util/schema.go +++ b/internal/util/schema.go @@ -11,7 +11,7 @@ import ( ) // GetSchemas returns the input and output schemas for a model version -func GetSchemas(version replicate.ModelVersion) (input *openapi3.SchemaRef, output *openapi3.SchemaRef, err error) { +func GetSchemas(version replicate.ModelVersion) (input *openapi3.Schema, output *openapi3.Schema, err error) { bytes, err := json.Marshal(version.OpenAPISchema) if err != nil { return nil, nil, fmt.Errorf("failed to serialize schema: %w", err) @@ -23,10 +23,18 @@ func GetSchemas(version replicate.ModelVersion) (input *openapi3.SchemaRef, outp } schemas := spec.Components.Schemas - inputSchema, _ := schemas["Input"] - outputSchema, _ := schemas["Output"] + inputSchemaRef, _ := schemas["Input"] + outputSchemaRef, _ := schemas["Output"] - return inputSchema, outputSchema, nil + if inputSchemaRef != nil { + input = inputSchemaRef.Value + } + + if outputSchemaRef != nil { + output = outputSchemaRef.Value + } + + return input, output, nil } // SortedKeys returns the keys of the properties in the order they should be displayed