Gliner/Gpu/GlinerGpuRuntime.cs
using System;
using System.Collections.Generic;
using System.Diagnostics;
using Sandbox;

namespace GlinerPoc.Gpu;

/// <summary>
/// Phase 8B.3 — narrow metadata descriptor for one GLiNER GPU FP32 buffer.
/// Logical tensor identity for debugging; NOT a generic tensor framework.
/// FP32 storage only.
/// </summary>
public sealed class GlinerGpuBufferDesc
{
	public string Name;
	public GpuBuffer<float> Buffer;
	public int Rows;
	public int Cols;
	/// <summary>Logical lifetime: persistent (weights, ping-pong activations) vs request-local scratch.</summary>
	public bool Persistent;
	/// <summary>Free-form ownership/usage note ("weight: layer0 Q", "activation ping-pong A", ...).</summary>
	public string Usage;
	/// <summary>
	/// CPU-side staging copy for weight uploads (P8B.38): SetData needs a
	/// managed float[], which this references. Kept because the CPU oracle
	/// also needs the weights; a GPU-only backend could drop it after upload.
	/// </summary>
	public float[] CpuStaging;

	public int ElementCount => Buffer.ElementCount;
	public long Bytes => (long)Buffer.ElementCount * 4;
}

/// <summary>
/// Phase 8B.2 — GPU resource ownership for the GLiNER backend.
///
/// Owns:
///   - the GlinerGpuExecutor (render-context dispatch)
///   - shader instances, constructed ONCE at initialization (P8A discovered
///     Material.FromShader static caches that survive hotload — never
///     construct shaders in hot paths; a poisoned cache needs an editor
///     restart, so these paths are only touched here)
///   - persistent weight buffers (P8B.5 residency)
///   - activation/scratch buffers (P8B.24/25 ping-pong + FFN scratch)
///
/// Intended eventual dependency chain (P8E): UI → service → inference
/// backend → GlinerGpuRuntime → executor. Nothing UI-side ever owns a raw
/// GpuBuffer. Logical lifetime (weights vs request activations) is distinct
/// from allocation lifetime (persistent allocations reused across requests).
/// </summary>
public sealed class GlinerGpuRuntime : IDisposable
{
	public GlinerGpuExecutor Executor { get; } = new();

	public bool Initialized { get; private set; }

	// shader set — one instance each, created once (see class doc)
	public ComputeShader VecAdd;      // P8A.1 canary kernel
	public ComputeShader GemmNaive;   // P8A.1 correctness reference
	public ComputeShader GemmTiled16; // P8B.8 optimized (16x16 groupshared tiles)
	public ComputeShader GemmTiled8;  // P8B.8 variant (8x8)
	public ComputeShader Add;
	public ComputeShader Relu;
	public ComputeShader Gelu;
	public ComputeShader LayerNorm;
	public ComputeShader Softmax;
	public ComputeShader Copy;
	public ComputeShader Gather;
	public ComputeShader C2C;
	public ComputeShader C2P;
	public ComputeShader P2C;
	public ComputeShader Combine3;
	public ComputeShader MaskKernel;
	public ComputeShader Context;

	private readonly Dictionary<string, GlinerGpuBufferDesc> _buffers = new( StringComparer.Ordinal );

	public IReadOnlyDictionary<string, GlinerGpuBufferDesc> Buffers => _buffers;

	public long PersistentBytes
	{
		get
		{
			long sum = 0;
			foreach ( var kv in _buffers )
			{
				if ( kv.Value.Persistent )
				{
					sum += kv.Value.Bytes;
				}
			}
			return sum;
		}
	}

	public long TotalBytes
	{
		get
		{
			long sum = 0;
			foreach ( var kv in _buffers )
			{
				sum += kv.Value.Bytes;
			}
			return sum;
		}
	}

	public void Initialize( Scene scene )
	{
		if ( Initialized )
		{
			return;
		}
		Executor.Initialize( scene.SceneWorld );

		VecAdd = new ComputeShader( "Shaders/gliner/gliner_cl_vec_add.shader" );
		GemmNaive = new ComputeShader( "Shaders/gliner/gliner_cl_gemm.shader" );
		GemmTiled16 = new ComputeShader( "Shaders/gliner/gliner_gpu_gemm_tiled16.shader" );
		GemmTiled8 = new ComputeShader( "Shaders/gliner/gliner_gpu_gemm_tiled8.shader" );
		Add = new ComputeShader( "Shaders/gliner/gliner_gpu_add.shader" );
		Relu = new ComputeShader( "Shaders/gliner/gliner_gpu_relu.shader" );
		Gelu = new ComputeShader( "Shaders/gliner/gliner_gpu_gelu.shader" );
		LayerNorm = new ComputeShader( "Shaders/gliner/gliner_gpu_layernorm.shader" );
		Softmax = new ComputeShader( "Shaders/gliner/gliner_gpu_softmax.shader" );
		Copy = new ComputeShader( "Shaders/gliner/gliner_gpu_copy.shader" );
		Gather = new ComputeShader( "Shaders/gliner/gliner_gpu_gather.shader" );
		C2C = new ComputeShader( "Shaders/gliner/gliner_gpu_c2c.shader" );
		C2P = new ComputeShader( "Shaders/gliner/gliner_gpu_c2p.shader" );
		P2C = new ComputeShader( "Shaders/gliner/gliner_gpu_p2c.shader" );
		Combine3 = new ComputeShader( "Shaders/gliner/gliner_gpu_combine3.shader" );
		MaskKernel = new ComputeShader( "Shaders/gliner/gliner_gpu_mask.shader" );
		Context = new ComputeShader( "Shaders/gliner/gliner_gpu_context.shader" );

		Initialized = true;
	}

	/// <summary>Allocate (or resize) a named buffer. Disposes a previous buffer of the same name.</summary>
	public GlinerGpuBufferDesc CreateBuffer( string name, int rows, int cols, bool persistent, string usage )
	{
		if ( _buffers.TryGetValue( name, out var old ) )
		{
			old.Buffer.Dispose();
		}
		var desc = new GlinerGpuBufferDesc
		{
			Name = name,
			Buffer = new GpuBuffer<float>( rows * cols ),
			Rows = rows,
			Cols = cols,
			Persistent = persistent,
			Usage = usage,
		};
		_buffers[name] = desc;
		return desc;
	}

	public GlinerGpuBufferDesc GetBuffer( string name ) =>
		_buffers.TryGetValue( name, out var d ) ? d : null;

	/// <summary>
	/// P8B.5 controlled weight-residency upload: managed float[] → persistent
	/// GpuBuffer. The float[] is a temporary staging copy required by SetData
	/// (P8B.38 documents this; callers must not retain it solely for the GPU's
	/// sake). Upload time is measured separately and returned.
	/// </summary>
	public (GlinerGpuBufferDesc Desc, double UploadMs) UploadPersistent( string name, float[] data, int rows, int cols, string usage )
	{
		var desc = CreateBuffer( name, rows, cols, persistent: true, usage );
		var sw = Stopwatch.StartNew();
		desc.Buffer.SetData( data );
		sw.Stop();
		desc.CpuStaging = data;
		return (desc, sw.Elapsed.TotalMilliseconds);
	}

	public void Dispose()
	{
		foreach ( var kv in _buffers )
		{
			kv.Value.Buffer.Dispose();
		}
		_buffers.Clear();
		Executor.Dispose();
	}
}