A sandboxed game component that implements ILanguageModelService for the TinyStories model. It loads tokenizer, model and config artifacts, exposes Initialize and Generate methods, runs heavy work on a worker thread, and includes a regression validator with expected tokens and text.
using Sandbox.Diagnostics;
namespace LlmPoc.Llm;
[Title( "TinyStories Language Model Service" )]
[Category( "LLM" )]
public sealed class TinyStoriesLanguageModelService : Component, ILanguageModelService
{
public const int InteractiveDefaultMaxNewTokens = 24;
[Property]
public bool InitializeOnStart { get; set; } = true;
[Property]
public int DefaultMaxNewTokens { get; set; } = InteractiveDefaultMaxNewTokens;
[Property]
public TinyStoriesModelResource RuntimeModelResource { get; set; }
public string DisplayName => "TinyStories-Instruct-1M";
public bool IsRealInference => true;
public LanguageModelServiceState State { get; private set; } =
LanguageModelServiceState.NotLoaded;
public string ErrorMessage { get; private set; } = "";
public bool CanGenerate =>
State == LanguageModelServiceState.Ready ||
(State == LanguageModelServiceState.Error && _model is not null);
public double InitializationMilliseconds { get; private set; }
public LanguageModelGenerationResult LastResult { get; private set; }
private SboxLlmModel _model;
private TinyStoriesConfig _config;
private Gpt2ByteBpeTokenizer _tokenizer;
private Task _initializationTask;
protected override void OnStart()
{
if ( InitializeOnStart )
{
_ = InitializeAsync();
}
}
public async Task InitializeAsync()
{
if ( _model is not null )
{
return;
}
if ( State == LanguageModelServiceState.Loading && _initializationTask is not null )
{
await _initializationTask;
return;
}
if ( State == LanguageModelServiceState.Error )
{
throw new InvalidOperationException(
string.IsNullOrWhiteSpace( ErrorMessage )
? "Model failed to load."
: ErrorMessage );
}
State = LanguageModelServiceState.Loading;
ErrorMessage = "";
_initializationTask = InitializeCoreAsync();
await _initializationTask;
}
public async Task<LanguageModelGenerationResult> GenerateAsync(
string prompt,
int? maxNewTokens = null )
{
if ( string.IsNullOrWhiteSpace( prompt ) )
{
throw new ArgumentException(
"[LLM:ERROR] Prompt cannot be empty.", nameof( prompt ) );
}
await InitializeAsync();
if ( State == LanguageModelServiceState.Generating )
{
throw new InvalidOperationException(
"[LLM:ERROR] A generation request is already active." );
}
if ( _model is null || _config is null || _tokenizer is null )
{
throw new InvalidOperationException(
"[LLM:ERROR] Model service is not initialized." );
}
int requestedTokens = maxNewTokens ?? DefaultMaxNewTokens;
if ( requestedTokens <= 0 )
{
throw new ArgumentOutOfRangeException(
nameof( maxNewTokens ), requestedTokens,
"[LLM:ERROR] maxNewTokens must be positive." );
}
// A previous inference error is recoverable because the immutable loaded
// artifacts remain valid. A load failure never reaches this point.
State = LanguageModelServiceState.Generating;
ErrorMessage = "";
LlmLog.Info(
"GEN",
$"request start chars={prompt.Length} max_new_tokens={requestedTokens} " +
"worker_thread=true strategy=greedy kv_cache=false" );
try
{
LanguageModelGenerationResult result = await Task.RunInThreadAsync(
() => GenerateOnWorker( prompt, requestedTokens ) );
LastResult = result;
State = LanguageModelServiceState.Ready;
LlmLog.Info(
"GEN",
$"complete input_tokens={result.InputTokenIds.Length} " +
$"generated_tokens={result.GeneratedTokenIds.Length} " +
$"first_token_ms={result.FirstTokenMilliseconds:N2} " +
$"elapsed_ms={result.TotalGenerationMilliseconds:N2} " +
$"request_ms={result.TotalRequestMilliseconds:N2} " +
$"tokens_per_second={result.TokensPerSecond:N2} " +
$"stop={result.StopReason} worker_thread=true" );
return result;
}
catch ( Exception error )
{
State = LanguageModelServiceState.Error;
ErrorMessage = BuildFriendlyError( "Generation failed", error );
LlmLog.Error( $"Runtime generation failed: {error}" );
throw;
}
}
private async Task InitializeCoreAsync()
{
FastTimer timer = FastTimer.StartNew();
LlmLog.Info(
"INIT",
"Loading runtime-only TinyStories config, model weights, and tokenizer." );
try
{
TinyStoriesModelResource resource = RuntimeModelResource;
TinyStoriesRuntimeArtifacts artifacts = await Task.RunInThreadAsync(
() => resource is not null
? TinyStoriesRuntimeArtifactLoader.Load( resource, LlmPaths.RuntimeModelResource )
: TinyStoriesRuntimeArtifactLoader.LoadFromResourceLibrary(
LlmPaths.RuntimeModelResource ) );
_model = artifacts.Model;
_config = artifacts.Config;
_tokenizer = artifacts.Tokenizer;
InitializationMilliseconds = timer.ElapsedMilliSeconds;
State = LanguageModelServiceState.Ready;
LlmLog.Info(
"INIT",
$"Runtime model ready in {InitializationMilliseconds:N2} ms; " +
$"tensors={_model.TensorCount} parameters={_model.TotalParameterCount:N0} " +
"reference_artifacts_loaded=false." );
}
catch ( Exception error )
{
State = LanguageModelServiceState.Error;
ErrorMessage = BuildFriendlyError( "Model failed to load", error );
LlmLog.Error( $"Runtime model initialization failed: {error}" );
throw;
}
}
private LanguageModelGenerationResult GenerateOnWorker(
string prompt,
int maxNewTokens )
{
FastTimer requestTimer = FastTimer.StartNew();
FastTimer tokenizerTimer = FastTimer.StartNew();
int[] promptTokenIds = _tokenizer.Encode( prompt );
double tokenizationMilliseconds = tokenizerTimer.ElapsedMilliSeconds;
if ( promptTokenIds.Length == 0 )
{
throw new InvalidOperationException(
"[LLM:ERROR] Tokenizer produced an empty prompt sequence." );
}
if ( promptTokenIds.Length > _config.MaximumPositions - maxNewTokens )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Prompt encodes to {promptTokenIds.Length} tokens and " +
$"requesting {maxNewTokens} more would exceed the " +
$"{_config.MaximumPositions}-token context. Truncation is disabled." );
}
GreedyGenerationResult generation = TinyStoriesGreedyGenerator.Generate(
_model,
_config,
_tokenizer,
promptTokenIds,
maxNewTokens,
observer: null,
logLifecycle: false );
return new LanguageModelGenerationResult
{
Prompt = prompt,
InputTokenIds = promptTokenIds,
GeneratedTokenIds = generation.GeneratedTokenIds,
GeneratedText = generation.GeneratedText,
StopReason = generation.StopReason,
EosReached = generation.EosReached,
MaximumNewTokens = maxNewTokens,
TokenizationMilliseconds = tokenizationMilliseconds,
FirstTokenMilliseconds = generation.Steps.Length > 0
? generation.Steps[0].StepMilliseconds
: 0,
TotalForwardMilliseconds = generation.TotalForwardMilliseconds,
TotalGenerationMilliseconds = generation.TotalGenerationMilliseconds,
TotalRequestMilliseconds = requestTimer.ElapsedMilliSeconds
};
}
private static string BuildFriendlyError( string prefix, Exception error )
{
string detail = error?.Message ?? "Unknown error.";
if ( detail.StartsWith( "[LLM:ERROR] " ) )
{
detail = detail[12..];
}
if ( detail.Length > 220 )
{
detail = detail[..217] + "...";
}
return $"{prefix}: {detail}";
}
}
public static class TinyStoriesLanguageModelServiceRegression
{
public const string GoldenPrompt = "Once upon a time";
public const string GoldenText =
", there was a little girl named Lily. She loved to";
public static readonly int[] GoldenTokenIds =
{
11, 612, 373, 257, 1310, 2576, 3706, 20037, 13, 1375, 6151, 284
};
public static void ValidateGolden(
LanguageModelGenerationResult result,
LanguageModelServiceState serviceState )
{
if ( result is null )
{
throw new ArgumentNullException( nameof( result ) );
}
if ( result.Prompt != GoldenPrompt )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Service golden prompt expected '{GoldenPrompt}', " +
$"found '{result.Prompt}'." );
}
if ( result.MaximumNewTokens != GoldenTokenIds.Length )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Service golden request expected " +
$"{GoldenTokenIds.Length} new tokens, found {result.MaximumNewTokens}." );
}
if ( result.GeneratedTokenIds.Length != GoldenTokenIds.Length )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Service golden generation expected " +
$"{GoldenTokenIds.Length} tokens, found " +
$"{result.GeneratedTokenIds.Length}." );
}
for ( int index = 0; index < GoldenTokenIds.Length; index++ )
{
if ( result.GeneratedTokenIds[index] != GoldenTokenIds[index] )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Service golden token mismatch at {index}: " +
$"expected={GoldenTokenIds[index]} " +
$"actual={result.GeneratedTokenIds[index]}." );
}
}
if ( result.GeneratedText != GoldenText )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Service golden decode expected '{GoldenText}', " +
$"found '{result.GeneratedText}'." );
}
if ( result.StopReason != "max_new_tokens" || result.EosReached )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Service golden stop expected max_new_tokens/non-EOS, " +
$"found {result.StopReason}/eos={result.EosReached}." );
}
if ( serviceState != LanguageModelServiceState.Ready )
{
throw new InvalidOperationException(
$"[LLM:ERROR] Service golden completion expected Ready state, " +
$"found {serviceState}." );
}
LlmLog.Info(
"PARITY",
$"service_golden generated_ids=[{string.Join( ",", result.GeneratedTokenIds )}] " +
$"decoded='{result.GeneratedText}' stop={result.StopReason} " +
"state=Ready PASS" );
}
}