diff options
Diffstat (limited to 'app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java')
| -rw-r--r-- | app/src/main/java/com/wuhei/reccoon/TranscriptionEngine.java | 418 |
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; - } - } -} |
