Repository navigation
Expand file tree
/
Copy pathmain.cpp
More file actions
110 lines (95 loc) · 3.43 KB
/
Copy pathmain.cpp
File metadata and controls
110 lines (95 loc) · 3.43 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
/*
* Copyright (c) Meta Platforms, Inc. and affiliates.
* All rights reserved.
*
* This source code is licensed under the BSD-style license found in the
* LICENSE file in the root directory of this source tree.
*/
#include <gflags/gflags.h>
#include <exception>
#include <iostream>
#include <string>
#include "parakeet_transcriber.h"
#include <executorch/runtime/platform/log.h>
#ifdef ET_BUILD_METAL
#include <executorch/backends/apple/metal/runtime/stats.h>
#endif
DEFINE_string(model_path, "parakeet.pte", "Path to Parakeet model (.pte).");
DEFINE_string(audio_path, "", "Path to input audio file (.wav).");
DEFINE_string(
tokenizer_path,
"tokenizer.model",
"Path to SentencePiece tokenizer model file.");
DEFINE_string(
data_path,
"",
"Path to data file (.ptd) for delegate data (optional, required for CUDA).");
DEFINE_string(
timestamps,
"segment",
"Timestamp output mode: none|token|word|segment|all");
DEFINE_bool(
runtime_profile,
false,
"Print a detailed runtime profile for preprocessor, encoder, and decode-loop execution.");
int main(int argc, char** argv) {
gflags::ParseCommandLineFlags(&argc, &argv, true);
parakeet::TimestampOutputMode timestamp_mode;
try {
timestamp_mode = parakeet::parse_timestamp_output_mode(FLAGS_timestamps);
} catch (const std::invalid_argument& e) {
ET_LOG(Error, "%s", e.what());
return 1;
}
if (FLAGS_audio_path.empty()) {
ET_LOG(Error, "audio_path flag must be provided.");
return 1;
}
try {
parakeet::ParakeetTranscriber transcriber(
FLAGS_model_path, FLAGS_tokenizer_path, FLAGS_data_path);
const auto result = transcriber.transcribe_wav_path(
FLAGS_audio_path,
parakeet::TranscribeConfig{timestamp_mode, FLAGS_runtime_profile});
std::cout << "Transcribed text: " << result.text << std::endl;
if (!result.stats_json.empty()) {
std::cout << "PyTorchObserver " << result.stats_json << std::endl;
}
if (result.runtime_profile_report.has_value()) {
std::cout << *result.runtime_profile_report;
}
#ifdef ET_BUILD_METAL
executorch::backends::metal::print_metal_backend_stats();
#endif
if (timestamp_mode.segment) {
std::cout << "\nSegment timestamps:" << std::endl;
for (const auto& segment : result.segment_offsets) {
const double start = segment.start_offset * result.frame_to_seconds;
const double end = segment.end_offset * result.frame_to_seconds;
std::cout << start << "s - " << end << "s : " << segment.text
<< std::endl;
}
}
if (timestamp_mode.word) {
std::cout << "\nWord timestamps:" << std::endl;
for (const auto& word : result.word_offsets) {
const double start = word.start_offset * result.frame_to_seconds;
const double end = word.end_offset * result.frame_to_seconds;
std::cout << start << "s - " << end << "s : " << word.text << std::endl;
}
}
if (timestamp_mode.token) {
std::cout << "\nToken timestamps:" << std::endl;
for (const auto& token : result.token_offsets) {
const double start = token.start_offset * result.frame_to_seconds;
const double end = token.end_offset * result.frame_to_seconds;
std::cout << start << "s - " << end << "s : " << token.decoded_text
<< std::endl;
}
}
return 0;
} catch (const std::exception& e) {
ET_LOG(Error, "%s", e.what());
return 1;
}
}