Llm/TinyStoriesLanguageModelService.cs

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.

Native Interop
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" );
	}
}