You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

BatchedExecutorSaveAndLoad.cs 2.8 kB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980
  1. using LLama.Batched;
  2. using LLama.Common;
  3. using LLama.Native;
  4. using LLama.Sampling;
  5. using Spectre.Console;
  6. namespace LLama.Examples.Examples;
  7. /// <summary>
  8. /// This demonstrates generating multiple replies to the same prompt, with a shared cache
  9. /// </summary>
  10. public class BatchedExecutorSaveAndLoad
  11. {
  12. private const int n_len = 18;
  13. public static async Task Run()
  14. {
  15. string modelPath = UserSettings.GetModelPath();
  16. var parameters = new ModelParams(modelPath);
  17. using var model = LLamaWeights.LoadFromFile(parameters);
  18. var prompt = AnsiConsole.Ask("Prompt (or ENTER for default):", "Not many people know that");
  19. // Create an executor that can evaluate a batch of conversations together
  20. using var executor = new BatchedExecutor(model, parameters);
  21. // Print some info
  22. var name = executor.Model.Metadata.GetValueOrDefault("general.name", "unknown model name");
  23. Console.WriteLine($"Created executor with model: {name}");
  24. // Create a conversation
  25. var conversation = executor.Create();
  26. conversation.Prompt(prompt);
  27. // Run inference loop
  28. var decoder = new StreamingTokenDecoder(executor.Context);
  29. var sampler = new DefaultSamplingPipeline();
  30. var lastToken = (LLamaToken)0;
  31. for (var i = 0; i < n_len; i++)
  32. {
  33. await executor.Infer();
  34. var token = sampler.Sample(executor.Context.NativeHandle, conversation.Sample(), ReadOnlySpan<LLamaToken>.Empty);
  35. lastToken = token;
  36. decoder.Add(token);
  37. conversation.Prompt(token);
  38. }
  39. // Can't save a conversation while RequiresInference is true
  40. if (conversation.RequiresInference)
  41. await executor.Infer();
  42. // Save this conversation and dispose it
  43. conversation.Save("demo_conversation.state");
  44. conversation.Dispose();
  45. AnsiConsole.WriteLine($"Saved state: {new FileInfo("demo_conversation.state").Length} bytes");
  46. // Now create a new conversation by loading that state
  47. conversation = executor.Load("demo_conversation.state");
  48. AnsiConsole.WriteLine("Loaded state");
  49. // Prompt it again with the last token, so we can continue generating
  50. conversation.Rewind(1);
  51. conversation.Prompt(lastToken);
  52. // Continue generating text
  53. for (var i = 0; i < n_len; i++)
  54. {
  55. await executor.Infer();
  56. var token = sampler.Sample(executor.Context.NativeHandle, conversation.Sample(), ReadOnlySpan<LLamaToken>.Empty);
  57. decoder.Add(token);
  58. conversation.Prompt(token);
  59. }
  60. // Display final ouput
  61. AnsiConsole.MarkupLine($"[red]{prompt}{decoder.Read()}[/]");
  62. }
  63. }