summaryrefslogtreecommitdiff
path: root/app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java
diff options
context:
space:
mode:
Diffstat (limited to 'app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java')
-rw-r--r--app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java418
1 files changed, 0 insertions, 418 deletions
diff --git a/app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java b/app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java
deleted file mode 100644
index 3579ab5..0000000
--- a/app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java
+++ /dev/null
@@ -1,418 +0,0 @@
-package com.wuhei.reccoon;
-
-import android.content.Context;
-import android.net.Uri;
-
-import com.k2fsa.sherpa.onnx.FeatureConfig;
-import com.k2fsa.sherpa.onnx.OfflineModelConfig;
-import com.k2fsa.sherpa.onnx.OfflineRecognizer;
-import com.k2fsa.sherpa.onnx.OfflineRecognizerConfig;
-import com.k2fsa.sherpa.onnx.OfflineRecognizerResult;
-import com.k2fsa.sherpa.onnx.OfflineStream;
-import com.k2fsa.sherpa.onnx.OfflineWhisperModelConfig;
-import com.k2fsa.sherpa.onnx.SileroVadModelConfig;
-import com.k2fsa.sherpa.onnx.SpeechSegment;
-import com.k2fsa.sherpa.onnx.Vad;
-import com.k2fsa.sherpa.onnx.VadModelConfig;
-
-import java.io.BufferedInputStream;
-import java.io.DataInputStream;
-import java.io.File;
-import java.io.IOException;
-import java.io.InputStream;
-import java.nio.charset.StandardCharsets;
-import java.util.ArrayList;
-import java.util.List;
-
-/**
- * On-device transcription using sherpa-onnx with a multilingual Whisper model.
- *
- * <p>Long recordings are streamed, resampled to 16 kHz and segmented with
- * Silero VAD, then each detected speech segment is decoded. This keeps memory
- * bounded regardless of the recording length (offline Whisper itself only
- * handles up to 30 seconds at a time).
- */
-public final class TranscriptionEngine {
-
- public static final String LANG_AUTO = "";
- public static final String LANG_EN = "en";
- public static final String LANG_IT = "it";
-
- private static final int TARGET_RATE = 16_000;
-
- public static final class Word {
- public final String text;
- public final float startSeconds;
- public final float endSeconds;
-
- public Word(String text, float startSeconds, float endSeconds) {
- this.text = text;
- this.startSeconds = startSeconds;
- this.endSeconds = endSeconds;
- }
- }
-
- public static final class Segment {
- public final float startSeconds;
- public final float endSeconds;
- public final String text;
- public final List<Word> words = new ArrayList<>();
-
- public Segment(float startSeconds, float endSeconds, String text) {
- this.startSeconds = startSeconds;
- this.endSeconds = endSeconds;
- this.text = text;
- }
- }
-
- public static final class Result {
- public final List<Segment> segments = new ArrayList<>();
- public String language = "";
-
- /** Plain transcript with one segment per paragraph. */
- public String asText() {
- StringBuilder sb = new StringBuilder();
- for (Segment segment : segments) {
- if (sb.length() > 0) {
- sb.append('\n');
- }
- sb.append(segment.text);
- }
- return sb.toString();
- }
- }
-
- public interface Callback {
- /** @param percent 0..100 */
- void onProgress(int percent, String message);
-
- boolean isCancelled();
- }
-
- private final Context context;
- private final ModelRepository models;
-
- public TranscriptionEngine(Context context) {
- this.context = context.getApplicationContext();
- this.models = new ModelRepository(this.context);
- }
-
- public Result transcribe(Uri wavUri,
- ModelRepository.WhisperModel model,
- String language,
- Callback callback) throws IOException {
- if (!models.isReady(model)) {
- throw new IOException("Model is not downloaded");
- }
-
- OfflineRecognizer recognizer = buildRecognizer(model, language);
- Vad vad = buildVad();
- Result result = new Result();
- result.language = language;
-
- try (InputStream raw = context.getContentResolver().openInputStream(wavUri)) {
- if (raw == null) {
- throw new IOException("Cannot open " + wavUri);
- }
- BufferedInputStream buffered = new BufferedInputStream(raw, 256 * 1024);
- WavInfo info = readWavHeader(buffered);
- if (info.bitsPerSample != 16 || info.audioFormat != 1) {
- throw new IOException("Only 16-bit PCM WAV files are supported");
- }
-
- StreamingResampler resampler = new StreamingResampler(info.sampleRate, TARGET_RATE);
- long totalPcm16k = estimatePcm16kSamples(info);
- long processed = 0;
-
- int frameBytes = info.channels * 2;
- byte[] readBuffer = new byte[frameBytes * 8192];
- long remaining = info.dataSize > 0 ? info.dataSize : Long.MAX_VALUE;
-
- while (remaining > 0) {
- if (callback != null && callback.isCancelled()) {
- return null;
- }
- int toRead = (int) Math.min(readBuffer.length, remaining);
- int read = buffered.read(readBuffer, 0, toRead);
- if (read <= 0) {
- break;
- }
- remaining -= read;
-
- int frames = read / frameBytes;
- float[] mono = new float[frames];
- for (int i = 0; i < frames; i++) {
- int offset = i * frameBytes;
- if (info.channels == 1) {
- mono[i] = pcm16(readBuffer, offset) / 32768f;
- } else {
- int sum = 0;
- for (int c = 0; c < info.channels; c++) {
- sum += pcm16(readBuffer, offset + c * 2);
- }
- mono[i] = sum / (float) info.channels / 32768f;
- }
- }
-
- float[] resampled = resampler.process(mono);
- processed += resampled.length;
- if (resampled.length > 0) {
- vad.acceptWaveform(resampled);
- drainSegments(vad, recognizer, result);
- }
- if (callback != null && totalPcm16k > 0) {
- int percent = (int) Math.min(99, processed * 100 / totalPcm16k);
- callback.onProgress(percent, "Transcribing…");
- }
- }
-
- vad.flush();
- drainSegments(vad, recognizer, result);
- } finally {
- vad.release();
- recognizer.release();
- }
-
- if (callback != null) {
- callback.onProgress(100, "Done");
- }
- return result;
- }
-
- private void drainSegments(Vad vad, OfflineRecognizer recognizer, Result result) {
- while (!vad.empty()) {
- SpeechSegment speech = vad.front();
- vad.pop();
- float[] samples = speech.getSamples();
- if (samples == null || samples.length < TARGET_RATE / 10) {
- continue; // ignore blips shorter than 100 ms
- }
- OfflineStream stream = recognizer.createStream();
- try {
- stream.acceptWaveform(samples, TARGET_RATE);
- recognizer.decode(stream);
- OfflineRecognizerResult decoded = recognizer.getResult(stream);
- String text = decoded.getText().trim();
- if (!text.isEmpty()) {
- float start = speech.getStart() / (float) TARGET_RATE;
- float end = start + samples.length / (float) TARGET_RATE;
- Segment segment = new Segment(start, end, text);
- buildWords(segment, decoded.getTokens(), decoded.getTimestamps(),
- decoded.getDurations());
- result.segments.add(segment);
- }
- } finally {
- stream.release();
- }
- }
- }
-
- /**
- * Fills in word timings. Uses token timestamps when the model provides
- * them, otherwise distributes the words across the segment by length.
- */
- private static void buildWords(Segment segment, String[] tokens,
- float[] timestamps, float[] durations) {
- if (tokens == null || tokens.length == 0) {
- return;
- }
- boolean haveTimestamps = timestamps != null && timestamps.length >= tokens.length;
- if (haveTimestamps) {
- for (int i = 0; i < tokens.length; i++) {
- if (tokens[i] == null || tokens[i].trim().isEmpty()) {
- continue;
- }
- float start = segment.startSeconds + timestamps[i];
- float end = durations != null && i < durations.length && durations[i] > 0f
- ? start + durations[i] : start + 0.2f;
- segment.words.add(new Word(tokens[i], start, end));
- }
- return;
- }
- int total = 0;
- for (String token : tokens) {
- if (token != null && !token.trim().isEmpty()) {
- total += token.trim().length();
- }
- }
- if (total <= 0) {
- return;
- }
- float span = Math.max(0.001f, segment.endSeconds - segment.startSeconds);
- int consumed = 0;
- for (String token : tokens) {
- if (token == null || token.trim().isEmpty()) {
- continue;
- }
- int weight = token.trim().length();
- float start = segment.startSeconds + span * consumed / total;
- consumed += weight;
- float end = segment.startSeconds + span * consumed / total;
- segment.words.add(new Word(token, start, end));
- }
- }
-
- private OfflineRecognizer buildRecognizer(ModelRepository.WhisperModel model, String language) {
- FeatureConfig feat = new FeatureConfig();
- feat.setSampleRate(TARGET_RATE);
- feat.setFeatureDim(80);
- feat.setDither(0f);
-
- OfflineWhisperModelConfig whisper = new OfflineWhisperModelConfig();
- whisper.setEncoder(models.encoder(model).getAbsolutePath());
- whisper.setDecoder(models.decoder(model).getAbsolutePath());
- whisper.setLanguage(language == null ? "" : language);
- whisper.setTask("transcribe");
- whisper.setTailPaddings(1000);
- whisper.setEnableTokenTimestamps(true);
- whisper.setEnableSegmentTimestamps(false);
-
- OfflineModelConfig modelConfig = new OfflineModelConfig();
- modelConfig.setWhisper(whisper);
- modelConfig.setTokens(models.tokens(model).getAbsolutePath());
- modelConfig.setNumThreads(Math.max(2, Runtime.getRuntime().availableProcessors() / 2));
- modelConfig.setDebug(false);
- modelConfig.setProvider("cpu");
- modelConfig.setModelType("whisper");
-
- OfflineRecognizerConfig config = new OfflineRecognizerConfig();
- config.setFeatConfig(feat);
- config.setModelConfig(modelConfig);
- config.setDecodingMethod("greedy_search");
-
- // Passing a null AssetManager makes sherpa-onnx load the model from
- // the absolute file paths configured above.
- return new OfflineRecognizer(null, config);
- }
-
- private Vad buildVad() {
- SileroVadModelConfig silero = new SileroVadModelConfig();
- silero.setModel(models.vadModel().getAbsolutePath());
- silero.setThreshold(0.5f);
- silero.setMinSilenceDuration(0.5f);
- silero.setMinSpeechDuration(0.25f);
- silero.setWindowSize(512);
- silero.setMaxSpeechDuration(20.0f);
-
- VadModelConfig config = new VadModelConfig();
- config.setSileroVadModelConfig(silero);
- config.setSampleRate(TARGET_RATE);
- config.setNumThreads(1);
- config.setProvider("cpu");
- config.setDebug(false);
-
- return new Vad(null, config);
- }
-
- // --- WAV handling -----------------------------------------------------
-
- public static final class WavInfo {
- public int audioFormat;
- public int channels;
- public int sampleRate;
- public int bitsPerSample;
- public long dataSize;
-
- public long durationMillis() {
- int byteRate = sampleRate * channels * (bitsPerSample / 8);
- if (byteRate <= 0 || dataSize <= 0) {
- return 0;
- }
- return dataSize * 1000L / byteRate;
- }
- }
-
- public static WavInfo readWavHeader(InputStream stream) throws IOException {
- DataInputStream in = stream instanceof DataInputStream
- ? (DataInputStream) stream
- : new DataInputStream(stream);
-
- byte[] header = new byte[12];
- in.readFully(header);
- if (header[0] != 'R' || header[1] != 'I' || header[2] != 'F' || header[3] != 'F'
- || header[8] != 'W' || header[9] != 'A' || header[10] != 'V' || header[11] != 'E') {
- throw new IOException("Not a RIFF/WAVE file");
- }
-
- WavInfo info = new WavInfo();
- while (true) {
- byte[] chunkHeader = new byte[8];
- in.readFully(chunkHeader);
- String id = new String(chunkHeader, 0, 4, StandardCharsets.US_ASCII);
- long size = le32(chunkHeader, 4) & 0xffffffffL;
-
- if ("fmt ".equals(id)) {
- if (size < 16 || size > 1024) {
- throw new IOException("Unsupported fmt chunk");
- }
- byte[] fmt = new byte[(int) size];
- in.readFully(fmt);
- info.audioFormat = (int) le16(fmt, 0);
- info.channels = (int) le16(fmt, 2);
- info.sampleRate = (int) le32(fmt, 4);
- info.bitsPerSample = (int) le16(fmt, 14);
- if ((size & 1) == 1) {
- in.readFully(new byte[1]);
- }
- } else if ("data".equals(id)) {
- info.dataSize = size;
- return info;
- } else {
- skipFully(in, size + (size & 1));
- }
- }
- }
-
- public static long durationMillis(Context context, Uri uri) {
- try (InputStream in = context.getContentResolver().openInputStream(uri)) {
- if (in == null) {
- return 0;
- }
- WavInfo info = readWavHeader(new BufferedInputStream(in));
- return info.durationMillis();
- } catch (Exception e) {
- return 0;
- }
- }
-
- private static long estimatePcm16kSamples(WavInfo info) {
- if (info.dataSize <= 0 || info.sampleRate <= 0) {
- return 0;
- }
- int frameBytes = info.channels * (info.bitsPerSample / 8);
- if (frameBytes <= 0) {
- return 0;
- }
- long frames = info.dataSize / frameBytes;
- return frames * TARGET_RATE / info.sampleRate;
- }
-
- private static int pcm16(byte[] data, int offset) {
- return (short) ((data[offset] & 0xff) | (data[offset + 1] << 8));
- }
-
- private static long le32(byte[] data, int offset) {
- return (data[offset] & 0xffL)
- | ((data[offset + 1] & 0xffL) << 8)
- | ((data[offset + 2] & 0xffL) << 16)
- | ((data[offset + 3] & 0xffL) << 24);
- }
-
- private static long le16(byte[] data, int offset) {
- return (data[offset] & 0xffL) | ((data[offset + 1] & 0xffL) << 8);
- }
-
- private static void skipFully(DataInputStream in, long count) throws IOException {
- long remaining = count;
- while (remaining > 0) {
- long skipped = in.skip(remaining);
- if (skipped <= 0) {
- if (in.read() < 0) {
- throw new IOException("Unexpected end of WAV file");
- }
- skipped = 1;
- }
- remaining -= skipped;
- }
- }
-}