// C# example: Run Cohere Transcribe with baked-in feature extraction. // // The encoder ONNX takes raw audio waveform directly -- no external // mel spectrogram computation needed. // // NuGet packages required: // dotnet add package Microsoft.ML.OnnxRuntime (or Microsoft.ML.OnnxRuntime.Gpu) // // Usage: // dotnet run -- audio.wav en // // Model files expected in ./onnx-baked/: // cohere-encoder.int8.onnx (+ .data file for external weights) // cohere-decoder.int8.onnx // tokens.txt using System; using System.Collections.Generic; using System.IO; using System.Linq; using Microsoft.ML.OnnxRuntime; using Microsoft.ML.OnnxRuntime.Tensors; namespace CohereTranscribeExample { class Program { // --- Model constants from config.json --- const int NumDecoderLayers = 8; const int DecoderHiddenSize = 1024; const int NumHeads = 8; const int HeadDim = 128; // 1024 / 8 const int MaxSeqLen = 1024; const int VocabSize = 16384; const int SampleRate = 16000; static void Main(string[] args) { if (args.Length < 1) { Console.WriteLine("Usage: CohereTranscribe [language]"); Console.WriteLine(" language: en, de, fr, es, it, pt, nl, pl, el, ar, ja, zh, vi, ko"); return; } string audioPath = args[0]; string language = args.Length > 1 ? args[1] : "en"; string modelDir = "./onnx-baked"; // Load tokens var tokens = LoadTokens(Path.Combine(modelDir, "tokens.txt")); var tokenToId = tokens.ToDictionary(kv => kv.Value, kv => kv.Key); // Build decoder prompt tokens var promptTokenIds = BuildPromptTokens(language, tokenToId); Console.WriteLine($"Prompt: {string.Join(" ", promptTokenIds.Select(id => tokens[id]))}"); // Read audio var wav = new WaveReader(audioPath); float[] audio = wav.Samples; if (wav.SampleRate != SampleRate) { Console.WriteLine($"Warning: audio is {wav.SampleRate}Hz, model expects {SampleRate}Hz. Resample first!"); return; } Console.WriteLine($"Audio: {audio.Length} samples ({audio.Length / (float)SampleRate:F1}s)"); // Create ONNX sessions var sessionOptions = new SessionOptions(); sessionOptions.InterOpNumThreads = 4; sessionOptions.IntraOpNumThreads = 4; string encoderPath = Path.Combine(modelDir, "cohere-encoder.int8.onnx"); string decoderPath = Path.Combine(modelDir, "cohere-decoder.int8.onnx"); Console.WriteLine("Loading encoder..."); using var encoder = new InferenceSession(encoderPath, sessionOptions); Console.WriteLine("Loading decoder..."); using var decoder = new InferenceSession(decoderPath, sessionOptions); // --- Run encoder --- Console.WriteLine("Running encoder..."); var audioTensor = new DenseTensor(audio, new[] { 1, audio.Length }); var encoderInputs = new List { NamedOnnxValue.CreateFromTensor("audio", audioTensor) }; float[] crossK, crossV; int[] crossKShape, crossVShape; using (var encoderResults = encoder.Run(encoderInputs)) { var crossKTensor = encoderResults.First(r => r.Name == "n_layer_cross_k").AsTensor(); var crossVTensor = encoderResults.First(r => r.Name == "n_layer_cross_v").AsTensor(); crossK = crossKTensor.ToArray(); crossV = crossVTensor.ToArray(); crossKShape = crossKTensor.Dimensions.ToArray(); crossVShape = crossVTensor.Dimensions.ToArray(); } int T_enc = crossKShape[2]; // encoder output time steps Console.WriteLine($"Encoder output: T_enc={T_enc}"); // --- Autoregressive decoding --- Console.WriteLine("Decoding..."); var generatedIds = new List(promptTokenIds); int eosId = tokenToId.GetValueOrDefault("<|endoftext|>", -1); int maxNewTokens = 256; // Initialize self-attention KV cache: (n_layers, batch, n_heads, max_ctx, head_dim) int cacheSize = NumDecoderLayers * 1 * NumHeads * MaxSeqLen * HeadDim; float[] selfKCache = new float[cacheSize]; float[] selfVCache = new float[cacheSize]; int[] cacheShape = new[] { NumDecoderLayers, 1, NumHeads, MaxSeqLen, HeadDim }; // First decoder call: process all prompt tokens at once int offset = 0; var currentTokens = promptTokenIds.ToArray(); for (int step = 0; step < maxNewTokens; step++) { int nTokens = currentTokens.Length; var tokensTensor = new DenseTensor( currentTokens.Select(t => (long)t).ToArray(), new[] { 1, nTokens }); var decoderInputs = new List { NamedOnnxValue.CreateFromTensor("tokens", tokensTensor), NamedOnnxValue.CreateFromTensor("in_n_layer_self_k_cache", new DenseTensor(selfKCache, cacheShape)), NamedOnnxValue.CreateFromTensor("in_n_layer_self_v_cache", new DenseTensor(selfVCache, cacheShape)), NamedOnnxValue.CreateFromTensor("n_layer_cross_k", new DenseTensor(crossK, crossKShape)), NamedOnnxValue.CreateFromTensor("n_layer_cross_v", new DenseTensor(crossV, crossVShape)), NamedOnnxValue.CreateFromTensor("offset", new DenseTensor(new[] { (long)offset }, Array.Empty())), }; using var decoderResults = decoder.Run(decoderInputs); // Get logits for last token position var logitsTensor = decoderResults.First(r => r.Name == "logits").AsTensor(); int lastPos = nTokens - 1; // Greedy: argmax over vocab for last token int bestId = 0; float bestScore = float.NegativeInfinity; for (int v = 0; v < VocabSize; v++) { float score = logitsTensor[0, lastPos, v]; if (score > bestScore) { bestScore = score; bestId = v; } } // Check EOS if (bestId == eosId) break; generatedIds.Add(bestId); // Update KV cache var outKCache = decoderResults.First(r => r.Name == "out_n_layer_self_k_cache").AsTensor(); var outVCache = decoderResults.First(r => r.Name == "out_n_layer_self_v_cache").AsTensor(); selfKCache = outKCache.ToArray(); selfVCache = outVCache.ToArray(); // Next step: single token, advance offset offset += nTokens; currentTokens = new[] { bestId }; } // --- Decode tokens to text --- // Skip prompt tokens, decode the rest var outputIds = generatedIds.Skip(promptTokenIds.Length).ToList(); string text = string.Join("", outputIds .Where(id => tokens.ContainsKey(id)) .Select(id => tokens[id]) .Select(t => t.StartsWith("<|") ? "" : t.Replace("\u2581", " "))); Console.WriteLine($"\nLanguage: {language}"); Console.WriteLine($"Text: {text.Trim()}"); Console.WriteLine($"Generated {outputIds.Count} tokens"); } static int[] BuildPromptTokens(string language, Dictionary tokenToId) { // Cohere prompt format: <|startofcontext|><|startoftranscript|><|emo:undefined|> // <|lang|><|lang|><|pnc|><|noitn|><|notimestamp|><|nodiarize|> var promptParts = new[] { "<|startofcontext|>", "<|startoftranscript|>", "<|emo:undefined|>", $"<|{language}|>", $"<|{language}|>", "<|pnc|>", "<|noitn|>", "<|notimestamp|>", "<|nodiarize|>", }; return promptParts .Where(t => tokenToId.ContainsKey(t)) .Select(t => tokenToId[t]) .ToArray(); } static Dictionary LoadTokens(string path) { var tokens = new Dictionary(); foreach (var line in File.ReadAllLines(path)) { int lastSpace = line.LastIndexOf(' '); if (lastSpace < 0) continue; string token = line.Substring(0, lastSpace); if (int.TryParse(line.Substring(lastSpace + 1), out int id)) tokens[id] = token; } return tokens; } } class WaveReader { public int SampleRate { get; } public float[] Samples { get; } public WaveReader(string path) { using var reader = new BinaryReader(File.OpenRead(path)); reader.ReadBytes(4); // "RIFF" reader.ReadInt32(); // file size reader.ReadBytes(4); // "WAVE" while (true) { string chunkId = new string(reader.ReadChars(4)); int chunkSize = reader.ReadInt32(); if (chunkId == "fmt ") { int audioFormat = reader.ReadInt16(); int numChannels = reader.ReadInt16(); SampleRate = reader.ReadInt32(); reader.ReadInt32(); reader.ReadInt16(); int bitsPerSample = reader.ReadInt16(); if (chunkSize > 16) reader.ReadBytes(chunkSize - 16); while (true) { string dataId = new string(reader.ReadChars(4)); int dataSize = reader.ReadInt32(); if (dataId == "data") { int numSamples = dataSize / (bitsPerSample / 8) / numChannels; Samples = new float[numSamples]; for (int i = 0; i < numSamples; i++) { float sample = 0; for (int ch = 0; ch < numChannels; ch++) { short s = reader.ReadInt16(); sample += s / 32768.0f; } Samples[i] = sample / numChannels; } return; } else { reader.ReadBytes(dataSize); } } } else { reader.ReadBytes(chunkSize); } } } } }