diff --git a/LLama.Unittest/StatelessExecutorTest.cs b/LLama.Unittest/StatelessExecutorTest.cs
index 19f618af..07439390 100644
--- a/LLama.Unittest/StatelessExecutorTest.cs
+++ b/LLama.Unittest/StatelessExecutorTest.cs
@@ -43,7 +43,8 @@ namespace LLama.Unittest
Assert.Equal(result1, result2);
}
- [Fact(Skip = "Very very slow in CI")]
+ //[Fact(Skip = "Very very slow in CI")]
+ [Fact]
public async Task OutOfContext()
{
var executor = new StatelessExecutor(_weights, _params);
diff --git a/LLama/LLamaStatelessExecutor.cs b/LLama/LLamaStatelessExecutor.cs
index ab1e9bbc..16401566 100644
--- a/LLama/LLamaStatelessExecutor.cs
+++ b/LLama/LLamaStatelessExecutor.cs
@@ -6,7 +6,6 @@ using System.Linq;
using System.Runtime.CompilerServices;
using System.Threading;
using System.Threading.Tasks;
-using LLama.Extensions;
using LLama.Native;
using Microsoft.Extensions.Logging;
@@ -47,61 +46,58 @@ namespace LLama
}
///
- public async IAsyncEnumerable InferAsync(string text, IInferenceParams? inferenceParams = null, [EnumeratorCancellation] CancellationToken cancellationToken = default)
+ public async IAsyncEnumerable InferAsync(string prompt, IInferenceParams? inferenceParams = null, [EnumeratorCancellation] CancellationToken cancellationToken = default)
{
- using var context = _weights.CreateContext(_params, _logger);
- Context = context;
-
+ // Ensure the context from last time is disposed (it always hould be)
if (!Context.NativeHandle.IsClosed)
Context.Dispose();
- Context = _weights.CreateContext(Context.Params, _logger);
-
- var decoder = new StreamingTokenDecoder(Context);
- var antiprocessor = new AntipromptProcessor(inferenceParams?.AntiPrompts ?? Array.Empty());
-
- if (inferenceParams != null)
- {
- if (inferenceParams.TokensKeep > Context.ContextSize)
- throw new ArgumentOutOfRangeException(nameof(inferenceParams), $"TokensKeep ({inferenceParams.TokensKeep}) cannot be larger than ContextSize ({Context.ContextSize})");
- }
- cancellationToken.ThrowIfCancellationRequested();
+ // Create an inference context which will be disposed when this method exits
+ using var context = _weights.CreateContext(_params, _logger);
+ Context = context;
+ // Sanity check inference params
inferenceParams ??= new InferenceParams();
+ if (inferenceParams.TokensKeep > Context.ContextSize)
+ throw new ArgumentOutOfRangeException(nameof(inferenceParams), $"TokensKeep ({inferenceParams.TokensKeep}) cannot be larger than ContextSize ({Context.ContextSize})");
+
+ // Create decoders for the token stream
+ var decoder = new StreamingTokenDecoder(Context);
+ var antiprocessor = new AntipromptProcessor(inferenceParams.AntiPrompts);
- var lastTokens = new List(inferenceParams.RepeatLastTokensCount);
- for (var i = 0; i < inferenceParams.RepeatLastTokensCount; i++)
+ // Keep track of the last N tokens emitted
+ var repeat_last_n = Math.Max(0, inferenceParams.RepeatLastTokensCount <0 ? _weights.ContextSize : inferenceParams.RepeatLastTokensCount);
+ var lastTokens = new List(repeat_last_n);
+ for (var i = 0; i < repeat_last_n; i++)
lastTokens.Add(0);
- var tokens = Context.Tokenize(text).ToList();
+ // Tokenize the prompt
+ var tokens = Context.Tokenize(prompt).ToList();
+ lastTokens.AddRange(tokens);
+ var n_past = 1 + tokens.Count;
+ // Evaluate the prompt
await Task.Run(() => { Context.Eval(tokens, 1); }, cancellationToken)
.ConfigureAwait(false);
- lastTokens.AddRange(tokens);
- var n_past = 1 + tokens.Count;
-
+ // Begin loop, evaluating one token at a time
var mu = (float?)null;
var max_tokens = inferenceParams.MaxTokens < 0 ? int.MaxValue : inferenceParams.MaxTokens;
- for(var i = 0; i < max_tokens; i++)
+ for(var i = 0; i < max_tokens && !cancellationToken.IsCancellationRequested; i++)
{
- if (cancellationToken.IsCancellationRequested)
- break;
-
- var repeat_last_n = inferenceParams.RepeatLastTokensCount < 0 ? Context.ContextSize : inferenceParams.RepeatLastTokensCount;
-
+ // Penalize the generated tokens by various penalties
var tokenDataArray = Context.ApplyPenalty(lastTokens, inferenceParams.LogitBias, repeat_last_n,
inferenceParams.RepeatPenalty, inferenceParams.FrequencyPenalty, inferenceParams.PresencePenalty, inferenceParams.PenalizeNL);
+ // Sample a single token
var id = Context.Sample(tokenDataArray, ref mu, inferenceParams.Temperature, inferenceParams.Mirostat, inferenceParams.MirostatTau,
inferenceParams.MirostatEta, inferenceParams.TopK, inferenceParams.TopP, inferenceParams.TfsZ, inferenceParams.TypicalP, inferenceParams.Grammar);
- lastTokens.Add(id);
-
decoder.Add(id);
var decoded = decoder.Read();
yield return decoded;
+ lastTokens.Add(id);
tokens.Clear();
tokens.Add(id);