unity/inference-engine-whisper-tiny
5881
1using System.Collections.Generic;2using UnityEngine;3using Unity.InferenceEngine;4using System.Text;5using Unity.Collections;6using Newtonsoft.Json;7 8public class RunWhisper : MonoBehaviour9{10 Worker decoder1, decoder2, encoder, spectrogram;11 Worker argmax;12 13 public AudioClip audioClip;14 15 // This is how many tokens you want. It can be adjusted.16 const int maxTokens = 100;17 18 // Special tokens see added tokens file for details19 const int END_OF_TEXT = 50257;20 const int START_OF_TRANSCRIPT = 50258;21 const int ENGLISH = 50259;22 const int GERMAN = 50261;23 const int FRENCH = 50265;24 const int TRANSCRIBE = 50359; //for speech-to-text in specified language25 const int TRANSLATE = 50358; //for speech-to-text then translate to English26 const int NO_TIME_STAMPS = 50363;27 const int START_TIME = 50364;28 29 int numSamples;30 string[] tokens;31 32 int tokenCount = 0;33 NativeArray<int> outputTokens;34 35 // Used for special character decoding36 int[] whiteSpaceCharacters = new int[256];37 38 Tensor<float> encodedAudio;39 40 bool transcribe = false;41 string outputString = "";42 43 // Maximum size of audioClip (30s at 16kHz)44 const int maxSamples = 30 * 16000;45 46 public ModelAsset audioDecoder1, audioDecoder2;47 public ModelAsset audioEncoder;48 public ModelAsset logMelSpectro;49 50 public async void Start()51 {52 SetupWhiteSpaceShifts();53 GetTokens();54 55 decoder1 = new Worker(ModelLoader.Load(audioDecoder1), BackendType.GPUCompute);56 decoder2 = new Worker(ModelLoader.Load(audioDecoder2), BackendType.GPUCompute);57 58 FunctionalGraph graph = new FunctionalGraph();59 var input = graph.AddInput(DataType.Float, new DynamicTensorShape(1, 1, 51865));60 var amax = Functional.ArgMax(input, -1, false);61 var selectTokenModel = graph.Compile(amax);62 argmax = new Worker(selectTokenModel, BackendType.GPUCompute);63 64 encoder = new Worker(ModelLoader.Load(audioEncoder), BackendType.GPUCompute);65 spectrogram = new Worker(ModelLoader.Load(logMelSpectro), BackendType.GPUCompute);66 67 outputTokens = new NativeArray<int>(maxTokens, Allocator.Persistent);68 69 outputTokens[0] = START_OF_TRANSCRIPT;70 outputTokens[1] = ENGLISH;// GERMAN;//FRENCH;//71 outputTokens[2] = TRANSCRIBE; //TRANSLATE;//72 //outputTokens[3] = NO_TIME_STAMPS;// START_TIME;//73 tokenCount = 3;74 75 LoadAudio();76 EncodeAudio();77 transcribe = true;78 79 tokensTensor = new Tensor<int>(new TensorShape(1, maxTokens));80 ComputeTensorData.Pin(tokensTensor);81 tokensTensor.Reshape(new TensorShape(1, tokenCount));82 tokensTensor.dataOnBackend.Upload<int>(outputTokens, tokenCount);83 84 lastToken = new NativeArray<int>(1, Allocator.Persistent); lastToken[0] = NO_TIME_STAMPS;85 lastTokenTensor = new Tensor<int>(new TensorShape(1, 1), new[] { NO_TIME_STAMPS });86 87 while (true)88 {89 if (!transcribe || tokenCount >= (outputTokens.Length - 1))90 return;91 m_Awaitable = InferenceStep();92 await m_Awaitable;93 }94 }95 Awaitable m_Awaitable;96 97 NativeArray<int> lastToken;98 Tensor<int> lastTokenTensor;99 Tensor<int> tokensTensor;100 Tensor<float> audioInput;101 102 void LoadAudio()103 {104 numSamples = audioClip.samples;105 var data = new float[maxSamples];106 107 // Handle stereo to mono conversion108 if (audioClip.channels == 2)109 {110 var stereoData = new float[numSamples * 2];111 audioClip.GetData(stereoData, 0);112 113 int monoSamples = Mathf.Min(numSamples, maxSamples);114 for (int i = 0; i < monoSamples; i++)115 {116 data[i] = (stereoData[i * 2] + stereoData[i * 2 + 1]) / 2f;117 }118 }119 else120 {121 numSamples = Mathf.Min(numSamples, maxSamples);122 audioClip.GetData(data, 0);123 }124 125 numSamples = maxSamples;126 audioInput = new Tensor<float>(new TensorShape(1, numSamples), data);127 }128 129 void EncodeAudio()130 {131 spectrogram.Schedule(audioInput);132 var logmel = spectrogram.PeekOutput() as Tensor<float>;133 encoder.Schedule(logmel);134 encodedAudio = encoder.PeekOutput() as Tensor<float>;135 }136 async Awaitable InferenceStep()137 {138 decoder1.SetInput("input_ids", tokensTensor);139 decoder1.SetInput("encoder_hidden_states", encodedAudio);140 decoder1.Schedule();141 142 var past_key_values_0_decoder_key = decoder1.PeekOutput("present.0.decoder.key") as Tensor<float>;143 var past_key_values_0_decoder_value = decoder1.PeekOutput("present.0.decoder.value") as Tensor<float>;144 var past_key_values_1_decoder_key = decoder1.PeekOutput("present.1.decoder.key") as Tensor<float>;145 var past_key_values_1_decoder_value = decoder1.PeekOutput("present.1.decoder.value") as Tensor<float>;146 var past_key_values_2_decoder_key = decoder1.PeekOutput("present.2.decoder.key") as Tensor<float>;147 var past_key_values_2_decoder_value = decoder1.PeekOutput("present.2.decoder.value") as Tensor<float>;148 var past_key_values_3_decoder_key = decoder1.PeekOutput("present.3.decoder.key") as Tensor<float>;149 var past_key_values_3_decoder_value = decoder1.PeekOutput("present.3.decoder.value") as Tensor<float>;150 151 var past_key_values_0_encoder_key = decoder1.PeekOutput("present.0.encoder.key") as Tensor<float>;152 var past_key_values_0_encoder_value = decoder1.PeekOutput("present.0.encoder.value") as Tensor<float>;153 var past_key_values_1_encoder_key = decoder1.PeekOutput("present.1.encoder.key") as Tensor<float>;154 var past_key_values_1_encoder_value = decoder1.PeekOutput("present.1.encoder.value") as Tensor<float>;155 var past_key_values_2_encoder_key = decoder1.PeekOutput("present.2.encoder.key") as Tensor<float>;156 var past_key_values_2_encoder_value = decoder1.PeekOutput("present.2.encoder.value") as Tensor<float>;157 var past_key_values_3_encoder_key = decoder1.PeekOutput("present.3.encoder.key") as Tensor<float>;158 var past_key_values_3_encoder_value = decoder1.PeekOutput("present.3.encoder.value") as Tensor<float>;159 160 decoder2.SetInput("input_ids", lastTokenTensor);161 decoder2.SetInput("past_key_values.0.decoder.key", past_key_values_0_decoder_key);162 decoder2.SetInput("past_key_values.0.decoder.value", past_key_values_0_decoder_value);163 decoder2.SetInput("past_key_values.1.decoder.key", past_key_values_1_decoder_key);164 decoder2.SetInput("past_key_values.1.decoder.value", past_key_values_1_decoder_value);165 decoder2.SetInput("past_key_values.2.decoder.key", past_key_values_2_decoder_key);166 decoder2.SetInput("past_key_values.2.decoder.value", past_key_values_2_decoder_value);167 decoder2.SetInput("past_key_values.3.decoder.key", past_key_values_3_decoder_key);168 decoder2.SetInput("past_key_values.3.decoder.value", past_key_values_3_decoder_value);169 170 decoder2.SetInput("past_key_values.0.encoder.key", past_key_values_0_encoder_key);171 decoder2.SetInput("past_key_values.0.encoder.value", past_key_values_0_encoder_value);172 decoder2.SetInput("past_key_values.1.encoder.key", past_key_values_1_encoder_key);173 decoder2.SetInput("past_key_values.1.encoder.value", past_key_values_1_encoder_value);174 decoder2.SetInput("past_key_values.2.encoder.key", past_key_values_2_encoder_key);175 decoder2.SetInput("past_key_values.2.encoder.value", past_key_values_2_encoder_value);176 decoder2.SetInput("past_key_values.3.encoder.key", past_key_values_3_encoder_key);177 decoder2.SetInput("past_key_values.3.encoder.value", past_key_values_3_encoder_value);178 179 decoder2.Schedule();180 181 var logits = decoder2.PeekOutput("logits") as Tensor<float>;182 argmax.Schedule(logits);183 using var t_Token = await argmax.PeekOutput().ReadbackAndCloneAsync() as Tensor<int>;184 int index = t_Token[0];185 186 outputTokens[tokenCount] = lastToken[0];187 lastToken[0] = index;188 tokenCount++;189 tokensTensor.Reshape(new TensorShape(1, tokenCount));190 tokensTensor.dataOnBackend.Upload<int>(outputTokens, tokenCount);191 lastTokenTensor.dataOnBackend.Upload<int>(lastToken, 1);192 193 if (index == END_OF_TEXT)194 {195 transcribe = false;196 }197 else if (index < tokens.Length)198 {199 outputString += GetUnicodeText(tokens[index]);200 }201 202 Debug.Log(outputString);203 }204 205 // Tokenizer206 public TextAsset vocabAsset;207 void GetTokens()208 {209 var vocab = JsonConvert.DeserializeObject<Dictionary<string, int>>(vocabAsset.text);210 tokens = new string[vocab.Count];211 foreach (var item in vocab)212 {213 tokens[item.Value] = item.Key;214 }215 }216 217 string GetUnicodeText(string text)218 {219 var bytes = Encoding.GetEncoding("ISO-8859-1").GetBytes(ShiftCharacterDown(text));220 return Encoding.UTF8.GetString(bytes);221 }222 223 string ShiftCharacterDown(string text)224 {225 string outText = "";226 foreach (char letter in text)227 {228 outText += ((int)letter <= 256) ? letter : (char)whiteSpaceCharacters[(int)(letter - 256)];229 }230 return outText;231 }232 233 void SetupWhiteSpaceShifts()234 {235 for (int i = 0, n = 0; i < 256; i++)236 {237 if (IsWhiteSpace((char)i)) whiteSpaceCharacters[n++] = i;238 }239 }240 241 bool IsWhiteSpace(char c)242 {243 return !(('!' <= c && c <= '~') || ('�' <= c && c <= '�') || ('�' <= c && c <= '�'));244 }245 246 private void OnDestroy()247 {248 decoder1.Dispose();249 decoder2.Dispose();250 encoder.Dispose();251 spectrogram.Dispose();252 argmax.Dispose();253 audioInput.Dispose();254 lastTokenTensor.Dispose();255 tokensTensor.Dispose();256 }257}258 