51 changed files with 155505 additions and 138 deletions
@ -1,33 +0,0 @@
|
||||
Ниже — подтверждённые по коду проблемы; отдельно указал, что воспроизведено запуском. |
||||
|
||||
[P1] Exit смешивает TCP-потоки разных клиентов. Поиск соединения использует только stream_id, хотя каждый клиент начинает нумерацию заново. После подключения второго клиента с таким же ID данные первого могут уйти в соединение второго; FIN/CLOSE также закрывают чужой поток. Нужен ключ (src_node_id, stream_id) во всех обработчиках, включая проверку повторного CONNECT. src/proxy/tcp_proxy_server.c:342 |
||||
|
||||
[P1] Клиент не проверяет отправителя ответов. DATA/FIN/CLOSE принимаются по одному stream_id, без сравнения источника с via_node_id. Обработка CLOSE_ALL читает peer_id, но затем уничтожает вообще все клиентские соединения. Пакеты другого узла, доставленные этому сервису, могут повредить или закрыть существующие потоки. src/proxy/tcp_proxy_client.c:668 |
||||
|
||||
[P1, воспроизведено] UDP-ответы в TUN имеют неправильную checksum. Обнуляется только IP-заголовок; байты UDP checksum остаются заполненными аллокатором значением 0xAAAA. Проверка сформированного пакета дала сумму 0xC1E2 вместо 0xFFFF. Такие ответы будут отбрасываться принимающим стеком. Нужно рассчитывать checksum либо явно устанавливать ноль для IPv4. src/proxy/udp_proxy.c:188 |
||||
|
||||
[P1, воспроизведено] ICMP с нечётной длиной payload получает неверную checksum. Цикл читает последнее 16-битное слово целиком, захватывая байт за пределами выделенного пользовательского буфера. Ошибка есть и на exit, и при восстановлении ответа клиенту. Для payload длиной 2 проверочная сумма получилась FFFF, длиной 3 — FF10. Отправка (src/proxy/icmp_proxy.c:80), ответ в TUN (src/proxy/icmp_proxy.c:285) |
||||
|
||||
[P1] UDP/ICMP-контексты глобальные, а ядро поддерживает несколько instances. Следующий init() перезаписывает g_udp_ctx/g_icmp_ctx; обработчики старого экземпляра используют новый контекст, а destroy(inst) освобождает его без проверки владельца. Дополнительно UDP различает запрос и ответ только по глобальному is_exit: узел, одновременно работающий клиентом и exit, трактует входящий ответ как новый исходящий запрос. src/proxy/udp_proxy.c:137, инициализация (src/proxy/udp_proxy.c:237), src/proxy/icmp_proxy.c:321 |
||||
|
||||
[P1] Нет сквозного ограничения потока при медленном получателе. Принятые ETCP-данные без ограничения добавляются в write_queue сокета или to_lwip. Заполненность этих очередей не останавливает доставку и подтверждение данных маршрутизатором. Медленный destination или локальный клиент приводит к росту памяти; при отказе выделения TCP-данные просто теряются. src/proxy/tcp_proxy_server.c:424, src/proxy/tcp_proxy_client.c:584 |
||||
|
||||
[P1] TUN-клиент может уничтожить ещё не отправленные данные при FIN. Если локальный FIN уже получен, но tx_queue остаётся заблокированной, входящий FIN от exit вызывает conn_finish() без проверки этой очереди. Она освобождается вместе с данными. Завершать поток нужно после опустошения очередей обоих направлений. src/proxy/tcp_proxy_client.c:642 |
||||
|
||||
[P1] Exit может отправить CLOSE раньше остатка ответа. В on_fin_cb() проверка tc->fin_local имеет приоритет над pend_r. Если клиент уже закрыл свою половину, а ответ destination остался в tx_buf из-за backpressure, сервер отправляет CLOSE и освобождает остаток. Ветка on_flushed_cb() также не проверяет ожидающие отправки данные ответа. src/proxy/tcp_proxy_server.c:108 |
||||
|
||||
[P1] HTTP CONNECT зависит от границ TCP-чтений. Парсер запускает туннель после первой \r\n, не дожидаясь \r\n\r\n. Если заголовки приходят следующим чтением, они пересылаются destination как содержимое туннеля — например, перед TLS ClientHello. Если вместе с заголовками уже пришли данные туннеля, они уничтожаются обнулением buf_len. src/proxy/socks_proxy.c:195 |
||||
|
||||
[P1] HTTP POST во время DNS может потерять весь накопленный запрос. dns_pending накапливается до 65535 байт, затем отправляется одним DATA. Exit принимает максимум 8192 байта в одном сообщении и отбрасывает превышение. При переполнении самого dns_pending очередные данные также теряются без прекращения потока. Нужны ограниченная очередь и отправка частями. src/proxy/socks_proxy.c:339, накопление (src/proxy/socks_proxy.c:430), ограничение exit (src/proxy/tcp_proxy_server.c:438) |
||||
|
||||
[P2] FIN теряется в двух штатных сценариях. SOCKS/HTTP при непустой очереди ответа вызывает tcp_conn_set_flushed(tc, NULL) — продолжение, которое должно переслать FIN, отсутствует. Exit молча игнорирует FIN, пришедший до завершения TCP connect. Протоколы, ожидающие EOF перед ответом, могут зависнуть. src/proxy/socks_proxy.c:538, src/proxy/tcp_proxy_server.c:492 |
||||
|
||||
[P2] SOCKS5-парсер некорректно обрабатывает запросы. После greeting/CONNECT сбрасывается весь буфер, включая следующие байты; всегда выбирается метод NO AUTH, даже если клиент его не предлагал. IPv6 вместо отказа превращается в IPv4 из неправильного смещения buf + 12. Доменный CONNECT длиной менее 10 байт бесконечно ожидает продолжения. src/proxy/socks_proxy.c:132 |
||||
|
||||
[P2] SOCKS/HTTP сообщают об успешном соединении до подключения exit. Ответ SOCKS success или HTTP 200 формируется даже до send_connect(). Подтверждения успешного TCP connect от exit в протоколе нет. При отказе подключения клиент сначала получает успех, затем закрытие вместо корректной ошибки. src/proxy/socks_proxy.c:316 |
||||
|
||||
[P2] После RST TUN-соединение остаётся в списке. tcp_proxy_client_err_cb() обнуляет pcb и выставляет error, но не освобождает pc. Очистка предусмотрена в poll callback уже уничтоженного PCB, который больше не вызовется. src/proxy/tcp_proxy_client.c:446 |
||||
|
||||
[P2] Неправильно разбираются IP options и фрагменты в TUN. UDP/ICMP используют фиксированные смещения от 20-байтового IPv4-заголовка, игнорируя IHL и fragment offset. Пакет с options или последующий IP-фрагмент превращается в запрос с неверными портами/данными. src/proxy/tcp_proxy_client.c:158 |
||||
|
||||
[P2] ICMP-ответы разных клиентов могут перепутаться. Exit сохраняет исходные echo_id/echo_seq и ищет ответ только по этой паре, без уникального преобразования ID и проверки адреса отправителя. Совпадающие ping-запросы разных клиентов получают чужие ответы. src/proxy/icmp_proxy.c:56 |
||||
@ -1,89 +0,0 @@
|
||||
# Задача: BGP NODEINFO теряется при флапающем реконнекте (chatgui member online не возвращается) |
||||
|
||||
## Кратко |
||||
|
||||
После рестарта Android-пира и реконнекта локальный chatgui показывает мембера **offline** и не возвращает **online**. |
||||
Корневая причина — **не в chat/member коде**, а в ETCP: BGP-пакет NODEINFO (347 байт) расшифровывается, но **не доходит до BGP-обработчика** `topo_group_receive_cbk()`. |
||||
|
||||
Нужно: **точно локализовать точку потери NODEINFO в receive-пути ETCP и починить её.** |
||||
|
||||
## Что уже сделано (исправлено и работает) |
||||
|
||||
1. `tools/chatgui/src/mainwindow.cpp` — **исправлен** `onMemberUpdatedCallback`/`onMemberRemovedCallback`: проверка `len < 1+64+1+79` отбрасывала все события с channel_id короче 64 симв. Теперь `len < 1` + точная проверка. (Это была причина «не обновляется вообще».) |
||||
2. `src/chat/chat_member.c/h` — добавлен `chat_core_member_online(ch_id,node_id)` = присутствие в `topo_group` (`topo_node_find_by_id`). Используется в `chat_core_get_member_list`/`get_single_member`. Убран `nodes.online` из SQL и блок `connected→2`. |
||||
3. `src/chat/chat_core.c`, `src/chat/chat_sync.c`, `src/chat/member_sync.c/h` — удалён онлайн-статус из БД + merkle-online gossip (`member_sync_set_online`, `_ms_apply_update`, `push_update`). |
||||
4. `src/chat/chat_status.c` + `tools/chatgui/src/accountlist.cpp` — флаг «bgp present» (0x08) для detail-панели. |
||||
5. `tools/chatgui/src/nodespage.cpp` — убрана колонка Online из диагностики. |
||||
|
||||
Эти правки корректны; статус работает на чистом connect/disconnect/стабильном реконнекте. |
||||
|
||||
## Найденная причина (подтверждено логами) |
||||
|
||||
- Android **шлёт** свой NODEINFO: `send_nodeinfo: node 1657b3ea8281c0a8 ver=2` (bgp-debug на Android). |
||||
- Локально пакет **расшифровывается**: `decrypt: code=01 dlen=347 plen=350` (347-байтный BGP-пакет, `code=0x01` = `ETCP_ID_TOPO_ENTRY`). |
||||
- Но `BGP recv NODEINFO nid=1657` (строка 193 в `topo_group.c`) **не появляется** — до обработчика пакет не доходит. |
||||
- При этом маленький `TABLE_COMPLETE` (10 байт) **доходит** (виден в логе как «mislabeled» `nid=aaaaaaaaaadeadbe` — это строка 193 логирует не-NODEINFO пакет как NODEINFO, garbage-поля). |
||||
- Признак: в стабильных рестартах (18:18/18:20) узел **добавляется** каждый раз, в флапающем (18:07) — нет. Т.е. потеря привязана к флапу (TCP-линки + ETCP reinit). |
||||
|
||||
## Receive-путь ETCP (где искать потерю) |
||||
|
||||
``` |
||||
decrypt (etcp_connections.c:2262, лог "decrypt: code=.. dlen=..") |
||||
→ if (link_state==CONNECTED) etcp_conn_input(pkt) |
||||
else memory_pool_free(pkt) ← добавлен WARN "decrypted pkt DROPPED" |
||||
→ etcp_conn_input (etcp.c:1506, "RX pkt dlen=%d") |
||||
→ нормализатор (reassembly) → int_queue |
||||
→ etcp_int_recv (etcp_api.c) ← добавлен WARN "int_recv BGP id=0x01 len=.." |
||||
→ dispatch по route id (0x01) → topo_group_receive_cbk (topo_group.c:193 "BGP recv NODEINFO") |
||||
``` |
||||
|
||||
## Уже добавленный дебаг (для локализации) |
||||
|
||||
1. `src/transport_layer/etcp_connections.c` (после дешифровки): |
||||
```c |
||||
} else { |
||||
DEBUG_WARN(DEBUG_CATEGORY_ETCP, "[%s] decrypted pkt DROPPED: link_state=%d code=0x%02x dlen=%u", |
||||
link->etcp->log_name, link->link_state, pkt_code, pkt->data_len); |
||||
memory_pool_free(e_sock->instance->pkt_pool, pkt); |
||||
} |
||||
``` |
||||
2. `src/transport_layer/etcp_api.c` (`etcp_int_recv`, после чтения `id`): |
||||
```c |
||||
uint8_t id = e->dgram[0]; |
||||
if (id == 0x01) DEBUG_WARN(DEBUG_CATEGORY_ETCP, "int_recv BGP id=0x01 conn=%s len=%zu", conn->log_name, e->len); |
||||
``` |
||||
3. Android: `tools/chatgui-android/libutun_lite/instance_lite.c` добавлено |
||||
`debug_set_category_level(DEBUG_CATEGORY_BGP, DEBUG_LEVEL_DEBUG);` (пересобрать APK: `cd tools/chatgui-android && ./build.sh`). |
||||
4. chatgui конфиг `tools/chatgui/build/vibechat.cfg` — включены `bgp=debug`, `chat_sync=debug`, `member_sync=debug`, `debug=debug`. |
||||
|
||||
Как читать (по таймстампу флапа): |
||||
- есть `decrypt code=01 dlen=..` + `DROPPED` → потеря в проверке `link_state==CONNECTED`; |
||||
- есть `decrypt` + нет `DROPPED` + нет `int_recv BGP` → потеря в `etcp_conn_input`/нормализаторе; |
||||
- есть `int_recv BGP` + нет `BGP recv NODEINFO` → потеря в BGP-обработчике. |
||||
|
||||
## Задача для агента |
||||
|
||||
1. Пересобрать chatgui (`cd tools/chatgui && cmake --build build -j4`) и Android APK (`cd tools/chatgui-android && ./build.sh`). |
||||
2. Запустить chatgui (DISPLAY=:0, `./vibechat`), воспроизвести флап реконнекта: |
||||
```bash |
||||
adb shell am force-stop com.utun.chat && adb shell am start -n com.utun.chat/.MainActivity |
||||
``` |
||||
(повторить несколько раз; флап ловится не каждый раз). |
||||
3. В логах (`tools/chatgui/build/chatgui.log`) найти потерю NODEINFO и **точно определить точку** по таблице выше. |
||||
4. Починить первопричину (варианты): |
||||
- буферизовать расшифрованный пакет до `LINK_STATE_CONNECTED` вместо `memory_pool_free`; |
||||
- либо гарантированный BGP-retry (после conn UP, если узел пира не в `group->nodes` через N мс — повторно `topo_group_send_table_request`). |
||||
5. Убрать весь диагностический мусор: |
||||
- WARN "DROPPED" (etcp_connections.c), WARN "int_recv BGP" (etcp_api.c); |
||||
- DEBUG-логи в `chat_member.c` (`on_member_props_changed`), `chat_sync.c` (`cs_on_peer_status_changed`), `memberlistmodel.cpp`; |
||||
- `instance_lite.c` BGP-debug; |
||||
- `vibechat.cfg` — вернуть категории в закомментированное состояние. |
||||
6. Прогнать: connect→online, disconnect→offline, реконнект (флап)→online. `./check.sh` (или `cd src && make -j4`). |
||||
|
||||
## Воспроизведение / ключевые артефакты |
||||
|
||||
- Пир: Android SM-A525F (`com.utun.chat`), node_id `0x1657b3ea8281c0a8`. |
||||
- Локальный: chatgui (`tools/chatgui/build/vibechat`), node_id `0x24cd036a6e659b9a`. |
||||
- Канал: `ch_id=7206723622466219923`, `group_id=0x64036cf3a783ef93`. |
||||
- Лог chatgui: `tools/chatgui/build/chatgui.log` (debug_file из конфига). |
||||
- Лог Android: `adb logcat -d | grep -iE "send_nodeinfo|NOT in registry|NODEINFO|DROPPED|int_recv"`. |
||||
@ -0,0 +1,301 @@
|
||||
#include "silero_vad.h" |
||||
#include "debug_config.h" |
||||
#include "mem.h" |
||||
|
||||
#ifdef HAVE_SILERO_VAD |
||||
|
||||
#include <stdio.h> |
||||
#include <string.h> |
||||
|
||||
#include <onnxruntime/onnxruntime_c_api.h> |
||||
|
||||
#include "silero_vad_model.inc" |
||||
|
||||
/* Число float-элементов рекуррентного состояния модели: 2 * 1 * 128 */ |
||||
#define SILERO_VAD_STATE_COUNT 256 |
||||
|
||||
struct silero_vad { |
||||
const OrtApi* ort; |
||||
OrtEnv* env; |
||||
OrtSession* session; |
||||
OrtMemoryInfo* mem_info; |
||||
float state[SILERO_VAD_STATE_COUNT]; /* рекуррентное состояние между вызовами */ |
||||
uint8_t* model_data; /* байты .onnx, загруженные в память (для file-версии) */ |
||||
size_t model_len; |
||||
}; |
||||
|
||||
/* Освободить OrtStatus и залогировать его сообщение как ошибку */ |
||||
static void vad_report_ort_error(const OrtApi* ort, OrtStatus* status, const char* what) { |
||||
const char* msg = ort->GetErrorMessage(status); |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "%s: %s", what, msg ? msg : "unknown onnxruntime error"); |
||||
ort->ReleaseStatus(status); |
||||
} |
||||
|
||||
/* Прочитать файл целиком в буфер. 0 при успехе: буфер пишется в out/out_len (u_malloc). */ |
||||
static int vad_read_file(const char* path, uint8_t** out, size_t* out_len) { |
||||
FILE* f; |
||||
long size; |
||||
uint8_t* buf; |
||||
|
||||
*out = NULL; |
||||
*out_len = 0; |
||||
|
||||
f = fopen(path, "rb"); |
||||
if (!f) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "vad_read_file: cannot open model '%s'", path); |
||||
return -1; |
||||
} |
||||
if (fseek(f, 0, SEEK_END) != 0) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "vad_read_file: fseek failed on '%s'", path); |
||||
fclose(f); |
||||
return -1; |
||||
} |
||||
size = ftell(f); |
||||
if (size <= 0) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "vad_read_file: empty or unreadable model '%s'", path); |
||||
fclose(f); |
||||
return -1; |
||||
} |
||||
rewind(f); |
||||
|
||||
buf = (uint8_t*)u_malloc((size_t)size); |
||||
if (!buf) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "vad_read_file: OOM for %ld bytes", size); |
||||
fclose(f); |
||||
return -1; |
||||
} |
||||
if (fread(buf, 1, (size_t)size, f) != (size_t)size) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "vad_read_file: short read on '%s'", path); |
||||
u_free(buf); |
||||
fclose(f); |
||||
return -1; |
||||
} |
||||
fclose(f); |
||||
|
||||
*out = buf; |
||||
*out_len = (size_t)size; |
||||
return 0; |
||||
} |
||||
|
||||
/* Общая инициализация детектора из байтов модели (файл или вшитая модель).
|
||||
* data может указывать на внешний буфер (вшитая модель) — тогда не освобождается. */ |
||||
static silero_vad_t* silero_vad_create_from_bytes(const uint8_t* data, size_t len, |
||||
const uint8_t* owned_buf, const char* label) { |
||||
silero_vad_t* vad = NULL; |
||||
const OrtApi* ort; |
||||
OrtStatus* status; |
||||
OrtSessionOptions* opts = NULL; |
||||
|
||||
if (!data || len == 0) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_create(%s): empty model", label); |
||||
return NULL; |
||||
} |
||||
|
||||
ort = OrtGetApiBase() ? OrtGetApiBase()->GetApi(ORT_API_VERSION) : NULL; |
||||
if (!ort) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_create(%s): onnxruntime C API unavailable", label); |
||||
return NULL; |
||||
} |
||||
|
||||
vad = (silero_vad_t*)u_calloc(1, sizeof(*vad)); |
||||
if (!vad) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_create(%s): OOM", label); |
||||
return NULL; |
||||
} |
||||
vad->ort = ort; |
||||
vad->model_data = (uint8_t*)owned_buf; |
||||
vad->model_len = len; |
||||
|
||||
status = ort->CreateEnv(ORT_LOGGING_LEVEL_WARNING, "utun_silero_vad", &vad->env); |
||||
if (status) { vad_report_ort_error(ort, status, "CreateEnv"); goto fail; } |
||||
|
||||
status = ort->CreateSessionOptions(&opts); |
||||
if (status) { vad_report_ort_error(ort, status, "CreateSessionOptions"); goto fail; } |
||||
status = ort->SetIntraOpNumThreads(opts, 1); |
||||
if (status) { vad_report_ort_error(ort, status, "SetIntraOpNumThreads"); } |
||||
status = ort->SetSessionGraphOptimizationLevel(opts, ORT_ENABLE_ALL); |
||||
if (status) { vad_report_ort_error(ort, status, "SetSessionGraphOptimizationLevel"); } |
||||
|
||||
status = ort->CreateSessionFromArray(vad->env, data, len, opts, &vad->session); |
||||
if (status) { vad_report_ort_error(ort, status, "CreateSessionFromArray"); goto fail; } |
||||
|
||||
status = ort->CreateCpuMemoryInfo(OrtArenaAllocator, OrtMemTypeDefault, &vad->mem_info); |
||||
if (status) { vad_report_ort_error(ort, status, "CreateCpuMemoryInfo"); goto fail; } |
||||
|
||||
ort->ReleaseSessionOptions(opts); |
||||
|
||||
DEBUG_INFO(DEBUG_CATEGORY_VAD, "silero_vad_create(%s): ok bytes=%zu", label, len); |
||||
return vad; |
||||
|
||||
fail: |
||||
if (opts) ort->ReleaseSessionOptions(opts); |
||||
silero_vad_destroy(vad); |
||||
return NULL; |
||||
} |
||||
|
||||
silero_vad_t* silero_vad_create(const char* model_path) { |
||||
uint8_t* buf = NULL; |
||||
size_t len = 0; |
||||
silero_vad_t* vad; |
||||
|
||||
if (!model_path || !*model_path) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_create: empty model path"); |
||||
return NULL; |
||||
} |
||||
if (vad_read_file(model_path, &buf, &len) != 0) return NULL; |
||||
vad = silero_vad_create_from_bytes(buf, len, buf, model_path); |
||||
if (!vad) u_free(buf); |
||||
return vad; |
||||
} |
||||
|
||||
silero_vad_t* silero_vad_create_default(void) { |
||||
return silero_vad_create_from_bytes(silero_vad_model_data, silero_vad_model_len, NULL, "embedded"); |
||||
} |
||||
|
||||
void silero_vad_destroy(silero_vad_t* vad) { |
||||
if (!vad) return; |
||||
const OrtApi* ort = vad->ort; |
||||
if (ort) { |
||||
if (vad->mem_info) ort->ReleaseMemoryInfo(vad->mem_info); |
||||
if (vad->session) ort->ReleaseSession(vad->session); |
||||
if (vad->env) ort->ReleaseEnv(vad->env); |
||||
} |
||||
if (vad->model_data) u_free(vad->model_data); |
||||
DEBUG_DEBUG(DEBUG_CATEGORY_VAD, "silero_vad_destroy"); |
||||
u_free(vad); |
||||
} |
||||
|
||||
void silero_vad_reset(silero_vad_t* vad) { |
||||
if (!vad) return; |
||||
memset(vad->state, 0, sizeof(vad->state)); |
||||
DEBUG_DEBUG(DEBUG_CATEGORY_VAD, "silero_vad_reset"); |
||||
} |
||||
|
||||
int silero_vad_process(silero_vad_t* vad, const float* samples, float* prob) { |
||||
const OrtApi* ort; |
||||
OrtStatus* status; |
||||
OrtValue* in_input = NULL; |
||||
OrtValue* in_state = NULL; |
||||
OrtValue* in_sr = NULL; |
||||
OrtValue* out_prob = NULL; |
||||
OrtValue* out_state = NULL; |
||||
const char* input_names[3]; |
||||
const char* output_names[2]; |
||||
const OrtValue* inputs[3]; |
||||
OrtValue* outputs[2]; |
||||
int64_t shape_input[2] = {1, SILERO_VAD_WINDOW_SAMPLES}; |
||||
int64_t shape_state[3] = {2, 1, 128}; |
||||
int64_t shape_prob[2] = {1, 1}; |
||||
int64_t sr_value = SILERO_VAD_SAMPLE_RATE; |
||||
float prob_local = 0.0f; |
||||
float next_state[SILERO_VAD_STATE_COUNT]; |
||||
int rc = -1; |
||||
|
||||
if (!vad || !samples || !prob) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_process: invalid args"); |
||||
return -1; |
||||
} |
||||
ort = vad->ort; |
||||
|
||||
status = ort->CreateTensorWithDataAsOrtValue(vad->mem_info, (void*)samples, |
||||
SILERO_VAD_WINDOW_SAMPLES * sizeof(float), shape_input, 2, |
||||
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT, &in_input); |
||||
if (status) { vad_report_ort_error(ort, status, "input tensor"); return -1; } |
||||
|
||||
status = ort->CreateTensorWithDataAsOrtValue(vad->mem_info, vad->state, |
||||
SILERO_VAD_STATE_COUNT * sizeof(float), shape_state, 3, |
||||
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT, &in_state); |
||||
if (status) { vad_report_ort_error(ort, status, "state tensor"); goto done; } |
||||
|
||||
status = ort->CreateTensorWithDataAsOrtValue(vad->mem_info, &sr_value, |
||||
sizeof(sr_value), NULL, 0, |
||||
ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64, &in_sr); |
||||
if (status) { vad_report_ort_error(ort, status, "sr tensor"); goto done; } |
||||
|
||||
status = ort->CreateTensorWithDataAsOrtValue(vad->mem_info, &prob_local, |
||||
sizeof(float), shape_prob, 2, |
||||
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT, &out_prob); |
||||
if (status) { vad_report_ort_error(ort, status, "output tensor"); goto done; } |
||||
|
||||
status = ort->CreateTensorWithDataAsOrtValue(vad->mem_info, next_state, |
||||
SILERO_VAD_STATE_COUNT * sizeof(float), shape_state, 3, |
||||
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT, &out_state); |
||||
if (status) { vad_report_ort_error(ort, status, "stateN tensor"); goto done; } |
||||
|
||||
input_names[0] = "input"; |
||||
input_names[1] = "state"; |
||||
input_names[2] = "sr"; |
||||
inputs[0] = in_input; |
||||
inputs[1] = in_state; |
||||
inputs[2] = in_sr; |
||||
|
||||
output_names[0] = "output"; |
||||
output_names[1] = "stateN"; |
||||
outputs[0] = out_prob; |
||||
outputs[1] = out_state; |
||||
|
||||
status = ort->Run(vad->session, NULL, input_names, inputs, 3, output_names, 2, outputs); |
||||
if (status) { |
||||
vad_report_ort_error(ort, status, "Run"); |
||||
goto done; |
||||
} |
||||
|
||||
memcpy(vad->state, next_state, sizeof(next_state)); |
||||
*prob = prob_local; |
||||
rc = 0; |
||||
|
||||
done: |
||||
if (in_input) ort->ReleaseValue(in_input); |
||||
if (in_state) ort->ReleaseValue(in_state); |
||||
if (in_sr) ort->ReleaseValue(in_sr); |
||||
if (out_prob) ort->ReleaseValue(out_prob); |
||||
if (out_state) ort->ReleaseValue(out_state); |
||||
return rc; |
||||
} |
||||
|
||||
int silero_vad_process_pcm16(silero_vad_t* vad, const int16_t* pcm, float* prob) { |
||||
float samples[SILERO_VAD_WINDOW_SAMPLES]; |
||||
int i; |
||||
|
||||
if (!pcm) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_process_pcm16: pcm == NULL"); |
||||
return -1; |
||||
} |
||||
for (i = 0; i < SILERO_VAD_WINDOW_SAMPLES; i++) { |
||||
samples[i] = (float)pcm[i] / 32768.0f; |
||||
} |
||||
return silero_vad_process(vad, samples, prob); |
||||
} |
||||
|
||||
#else /* !HAVE_SILERO_VAD — пустые стабы (сборки без onnxruntime) */ |
||||
|
||||
silero_vad_t* silero_vad_create(const char* model_path) { |
||||
(void)model_path; |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_create: built without onnxruntime (HAVE_SILERO_VAD not defined)"); |
||||
return NULL; |
||||
} |
||||
|
||||
silero_vad_t* silero_vad_create_default(void) { |
||||
DEBUG_ERROR(DEBUG_CATEGORY_VAD, "silero_vad_create_default: built without onnxruntime (HAVE_SILERO_VAD not defined)"); |
||||
return NULL; |
||||
} |
||||
|
||||
void silero_vad_destroy(silero_vad_t* vad) { |
||||
(void)vad; |
||||
} |
||||
|
||||
void silero_vad_reset(silero_vad_t* vad) { |
||||
(void)vad; |
||||
} |
||||
|
||||
int silero_vad_process(silero_vad_t* vad, const float* samples, float* prob) { |
||||
(void)vad; (void)samples; (void)prob; |
||||
return -1; |
||||
} |
||||
|
||||
int silero_vad_process_pcm16(silero_vad_t* vad, const int16_t* pcm, float* prob) { |
||||
(void)vad; (void)pcm; (void)prob; |
||||
return -1; |
||||
} |
||||
|
||||
#endif /* HAVE_SILERO_VAD */ |
||||
@ -0,0 +1,76 @@
|
||||
/**
|
||||
* Silero VAD (Voice Activity Detector) — C-обёртка над ONNX Runtime. |
||||
* |
||||
* Запускает официальную стриминговую модель silero_vad.onnx (v5, ~2.3 МБ), |
||||
* которая по окну аудио возвращает вероятность наличия речи [0..1]. |
||||
* Рекуррентное состояние (GRU) хранится внутри объекта и переносится между |
||||
* вызовами, поэтому детектор работает в потоковом (real-time) режиме. |
||||
* |
||||
* Требования к входному аудио: |
||||
* - частота дискретизации 16 кГц (модель поддерживает только 16k); |
||||
* - моно; |
||||
* - окно ровно SILERO_VAD_WINDOW_SAMPLES = 512 сэмплов (32 мс); |
||||
* - float-сэмплы в диапазоне примерно [-1;1] (для int16 есть |
||||
* silero_vad_process_pcm16, который сам делит на 32768). |
||||
* |
||||
* Порог и гистерезис (задержки начала/конца речи) намеренно оставлены |
||||
* вызывающему — радио/звонок сами решают, как интерпретировать вероятность. |
||||
* |
||||
* Зависимость: libonnxruntime (C API). Сборка включается опцией |
||||
* `--with-silero-vad` (см. configure.ac) или макросом HAVE_SILERO_VAD в |
||||
* CMake-сборках. Без HAVE_SILERO_VAD модуль компилируется в пустые стабы |
||||
* (все create-функции возвращают NULL, process — ошибку), чтобы исходники lib/ |
||||
* можно было GLOB-ить во всех сборках (Android/chatgui) без onnxruntime. |
||||
* |
||||
* Лицензия модели: MIT (Silero Team, https://github.com/snakers4/silero-vad).
|
||||
*/ |
||||
#ifndef SILERO_VAD_H |
||||
#define SILERO_VAD_H |
||||
|
||||
#include <stdint.h> |
||||
|
||||
#ifdef __cplusplus |
||||
extern "C" { |
||||
#endif |
||||
|
||||
/* Фиксированное окно стриминговой модели: 512 сэмплов @ 16 кГц = 32 мс */ |
||||
#define SILERO_VAD_WINDOW_SAMPLES 512 |
||||
#define SILERO_VAD_SAMPLE_RATE 16000 |
||||
|
||||
typedef struct silero_vad silero_vad_t; |
||||
|
||||
/**
|
||||
* Создать детектор, загрузив ONNX-модель из файла model_path. |
||||
* Возвращает NULL при ошибке (подробности — в лог категории "vad"). |
||||
*/ |
||||
silero_vad_t* silero_vad_create(const char* model_path); |
||||
|
||||
/**
|
||||
* Создать детектор из модели, вшитой в бинарник (silero_vad_model_data в |
||||
* silero_vad.c). Не требует файла на диске. NULL при ошибке. |
||||
*/ |
||||
silero_vad_t* silero_vad_create_default(void); |
||||
|
||||
/* Освободить детектор. NULL безопасен. */ |
||||
void silero_vad_destroy(silero_vad_t* vad); |
||||
|
||||
/* Сбросить рекуррентное состояние (начало нового аудиопотока / после длительной тишины). */ |
||||
void silero_vad_reset(silero_vad_t* vad); |
||||
|
||||
/**
|
||||
* Обработать окно из 512 float-сэмплов @ 16 кГц (примерно [-1;1]). |
||||
* При успехе возвращает 0 и заполняет *prob вероятностью речи [0..1]; |
||||
* при ошибке — отрицательное значение (проб не трогается). |
||||
*/ |
||||
int silero_vad_process(silero_vad_t* vad, const float* samples, float* prob); |
||||
|
||||
/**
|
||||
* То же, но вход — int16 PCM: сэмплы преобразуются к float делением на 32768. |
||||
*/ |
||||
int silero_vad_process_pcm16(silero_vad_t* vad, const int16_t* pcm, float* prob); |
||||
|
||||
#ifdef __cplusplus |
||||
} |
||||
#endif |
||||
|
||||
#endif /* SILERO_VAD_H */ |
||||
Binary file not shown.
@ -0,0 +1,95 @@
|
||||
// radio_vad.h — авто-PTT рации по VAD: чистая логика без I/O (юнит-тестируется).
|
||||
//
|
||||
// Содержит:
|
||||
// - ресемплер 48кГц→16кГц моно (3-tap box + децимация 3:1, стерео микшируется);
|
||||
// - фильтр-гистерезис вероятности речи и машину состояний «передаю/молчу».
|
||||
//
|
||||
// Всё — static inline, без состояния (кроме struct radio_vad_fsm, который владеет
|
||||
// вызывающий). Используется сервисным слоем radio_audio.c и тестом test_radio_vad.
|
||||
//
|
||||
// Время — в timebase-единицах (0.1мс), как get_time_tb() в ядре.
|
||||
#ifndef RADIO_VAD_H |
||||
#define RADIO_VAD_H |
||||
|
||||
#include <stdint.h> |
||||
|
||||
/* Выходов ресемплера на один 20мс кадр 48кГц: 960 сэмплов / 3 = 320 @ 16кГц. */ |
||||
#define RADIO_VAD_FRAME_SAMPLES_16K 320 |
||||
|
||||
/* Число 16кГц-сэмплов в окне Silero: 512 = 32мс. */ |
||||
#define RADIO_VAD_WINDOW_16K 512 |
||||
|
||||
/* Ресемпл 20мс кадра 48кГц (frame_total int16, channels=1|2) → 320 float @16кГц
|
||||
* моно в диапазоне [-1;1]. Стерео миксируется (среднее L/R), затем децимация 3:1. */ |
||||
static inline void radio_vad_resample(const int16_t* pcm, int channels, float* out16k) { |
||||
for (int j = 0; j < RADIO_VAD_FRAME_SAMPLES_16K; j++) { |
||||
float s = 0.0f; |
||||
for (int t = 0; t < 3; t++) { |
||||
int idx = j * 3 + t; |
||||
float v = (channels == 2) |
||||
? ((float)pcm[2 * idx] + (float)pcm[2 * idx + 1]) * 0.5f |
||||
: (float)pcm[idx]; |
||||
s += v; |
||||
} |
||||
out16k[j] = (s / 3.0f) * (1.0f / 32768.0f); |
||||
} |
||||
} |
||||
|
||||
/* ── машина состояний авто-PTT ── */ |
||||
|
||||
struct radio_vad_fsm { |
||||
int talking; /* идёт VAD-передача */ |
||||
int confirm_count; /* подряд окон с prob >= threshold */ |
||||
uint64_t silence_since_tb; /* начало текущей тишины (0 — не в тишине) */ |
||||
}; |
||||
|
||||
static inline void radio_vad_fsm_reset(struct radio_vad_fsm* f) { |
||||
f->talking = 0; |
||||
f->confirm_count = 0; |
||||
f->silence_since_tb = 0; |
||||
} |
||||
|
||||
/* Обработать очередную вероятность речи (окно 32мс).
|
||||
* prob — вероятность [0..1]; |
||||
* threshold — порог срабатывания [0..1]; |
||||
* confirm_windows — сколько окон подряд выше порога нужно для старта; |
||||
* hangover_tb — сколько тишины (tb) держать передачу перед стопом; |
||||
* busy — 1 если сейчас говорит кто-то другой (канал занят); |
||||
* now — текущее время (tb). |
||||
* Возврат: 1 = начать передачу, -1 = закончить передачу, 0 = без изменений. */ |
||||
static inline int radio_vad_fsm_update(struct radio_vad_fsm* f, float prob, float threshold, |
||||
int confirm_windows, uint64_t hangover_tb, |
||||
int busy, uint64_t now) { |
||||
if (busy) { |
||||
/* чужой разговор: старт запрещён, свою передачу гасим сразу */ |
||||
f->confirm_count = 0; |
||||
f->silence_since_tb = 0; |
||||
if (f->talking) { f->talking = 0; return -1; } |
||||
return 0; |
||||
} |
||||
if (prob >= threshold) { |
||||
f->silence_since_tb = 0; |
||||
if (!f->talking) { |
||||
f->confirm_count++; |
||||
if (f->confirm_count >= confirm_windows) { |
||||
f->confirm_count = 0; |
||||
f->talking = 1; |
||||
return 1; /* старт: подтверждение набрано */ |
||||
} |
||||
} |
||||
return 0; |
||||
} |
||||
/* prob < threshold */ |
||||
f->confirm_count = 0; |
||||
if (f->talking) { |
||||
if (f->silence_since_tb == 0) f->silence_since_tb = now; |
||||
else if (now - f->silence_since_tb >= hangover_tb) { |
||||
f->talking = 0; |
||||
f->silence_since_tb = 0; |
||||
return -1; /* стоп: hangover истёк */ |
||||
} |
||||
} |
||||
return 0; |
||||
} |
||||
|
||||
#endif /* RADIO_VAD_H */ |
||||
@ -0,0 +1,134 @@
|
||||
/* Тест чистой логики VAD авто-PTT рации (radio_vad.h): ресемплер 48k→16k и
|
||||
* фильтр-гистерезис + машина состояний. Без ONNX Runtime и без аудио-устройств — |
||||
* детерминированные синтетические последовательности. |
||||
*/ |
||||
#include <stdio.h> |
||||
#include <string.h> |
||||
#include <math.h> |
||||
|
||||
#include "../src/radio/radio_vad.h" |
||||
|
||||
#define PI 3.14159265358979323846 |
||||
|
||||
static int g_failures = 0; |
||||
|
||||
#define CHECK(cond, ...) do { \ |
||||
if (!(cond)) { printf(" FAIL: " __VA_ARGS__); printf("\n"); g_failures++; } \
|
||||
} while (0) |
||||
|
||||
static int near(float a, float b, float eps) { |
||||
return fabsf(a - b) <= eps; |
||||
} |
||||
|
||||
/* ── ресемплер ── */ |
||||
|
||||
static void test_resample(void) { |
||||
printf("[radio_vad] resampler 48k -> 16k\n"); |
||||
|
||||
/* тишина → тишина */ |
||||
{ |
||||
int16_t in[960] = {0}; |
||||
float out[RADIO_VAD_FRAME_SAMPLES_16K]; |
||||
radio_vad_resample(in, 1, out); |
||||
for (int i = 0; i < RADIO_VAD_FRAME_SAMPLES_16K; i++) CHECK(out[i] == 0.0f, "silence out[%d]=%f", i, out[i]); |
||||
} |
||||
|
||||
/* DC 0.5*32768 → выход ~0.5 (box-фильтр сохраняет DC) */ |
||||
{ |
||||
int16_t in[960]; |
||||
float out[RADIO_VAD_FRAME_SAMPLES_16K]; |
||||
for (int i = 0; i < 960; i++) in[i] = 16384; |
||||
radio_vad_resample(in, 1, out); |
||||
for (int i = 0; i < RADIO_VAD_FRAME_SAMPLES_16K; i++) CHECK(near(out[i], 0.5f, 0.001f), "DC out[%d]=%f", i, out[i]); |
||||
} |
||||
|
||||
/* стерео L=R → тот же результат, что моно */ |
||||
{ |
||||
int16_t mono[960]; |
||||
int16_t stereo[1920]; |
||||
float om[RADIO_VAD_FRAME_SAMPLES_16K], os[RADIO_VAD_FRAME_SAMPLES_16K]; |
||||
for (int i = 0; i < 960; i++) { mono[i] = (int16_t)(sin(2.0 * PI * 1000.0 * i / 48000.0) * 20000.0); } |
||||
for (int i = 0; i < 960; i++) { stereo[2*i] = mono[i]; stereo[2*i+1] = mono[i]; } |
||||
radio_vad_resample(mono, 1, om); |
||||
radio_vad_resample(stereo, 2, os); |
||||
for (int i = 0; i < RADIO_VAD_FRAME_SAMPLES_16K; i++) CHECK(near(om[i], os[i], 1e-4f), "stereo!=mono out[%d] %f vs %f", i, om[i], os[i]); |
||||
} |
||||
|
||||
/* синус 1кГц: амплитуда сохраняется (box-фильтр даёт ~0.6% ослабления) */ |
||||
{ |
||||
int16_t in[960]; |
||||
float out[RADIO_VAD_FRAME_SAMPLES_16K]; |
||||
float peak = 0.0f; |
||||
for (int i = 0; i < 960; i++) in[i] = (int16_t)(sin(2.0 * PI * 1000.0 * i / 48000.0) * 30000.0); |
||||
radio_vad_resample(in, 1, out); |
||||
for (int i = 0; i < RADIO_VAD_FRAME_SAMPLES_16K; i++) { float a = fabsf(out[i]); if (a > peak) peak = a; } |
||||
CHECK(peak > 0.85f && peak <= 1.0f, "1kHz peak=%.3f (ожидали ~0.91)", peak); |
||||
} |
||||
} |
||||
|
||||
/* ── машина состояний ── */ |
||||
|
||||
static void test_fsm(void) { |
||||
printf("[radio_vad] fsm (confirm / hangover / busy)\n"); |
||||
struct radio_vad_fsm f; |
||||
const float thr = 0.5f; |
||||
const int confirm = 2; |
||||
const uint64_t hangover = 2000; /* 200мс */ |
||||
uint64_t now = 0; |
||||
|
||||
/* подтверждение: 2 окна выше порога → старт */ |
||||
radio_vad_fsm_reset(&f); |
||||
now += 320; |
||||
CHECK(radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 0, now) == 0, "confirm#1 должен молчать"); |
||||
now += 320; |
||||
CHECK(radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 0, now) == 1, "confirm#2 должен дать START"); |
||||
CHECK(f.talking == 1, "после старта talking=1"); |
||||
|
||||
/* одиночный пик не запускает передачу */ |
||||
radio_vad_fsm_reset(&f); |
||||
now += 320; |
||||
CHECK(radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 0, now) == 0, "одиночный пик не стартует"); |
||||
now += 320; |
||||
CHECK(radio_vad_fsm_update(&f, 0.1f, thr, confirm, hangover, 0, now) == 0, "сброс подтверждения"); |
||||
|
||||
/* hangover: после старта тишина ~200мс → STOP */ |
||||
radio_vad_fsm_reset(&f); |
||||
now += 320; radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 0, now); |
||||
now += 320; CHECK(radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 0, now) == 1, "старт"); |
||||
{ |
||||
int stops = 0; |
||||
for (int i = 0; i < 10; i++) { |
||||
now += 320; |
||||
int a = radio_vad_fsm_update(&f, 0.1f, thr, confirm, hangover, 0, now); |
||||
if (a == -1) { stops++; break; } |
||||
CHECK(a == 0, "hangover: до ~200мс не стопим"); |
||||
} |
||||
CHECK(stops == 1, "hangover должен дать STOP"); |
||||
CHECK(f.talking == 0, "после STOP talking=0"); |
||||
} |
||||
|
||||
/* busy: чужой разговор гасит нашу передачу сразу */ |
||||
radio_vad_fsm_reset(&f); |
||||
now += 320; radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 0, now); |
||||
now += 320; radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 0, now); |
||||
CHECK(f.talking == 1, "старт перед busy"); |
||||
now += 320; |
||||
CHECK(radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 1, now) == -1, "busy должен дать STOP"); |
||||
|
||||
/* busy блокирует старт */ |
||||
radio_vad_fsm_reset(&f); |
||||
for (int i = 0; i < 5; i++) { |
||||
now += 320; |
||||
CHECK(radio_vad_fsm_update(&f, 0.9f, thr, confirm, hangover, 1, now) == 0, "busy блокирует старт"); |
||||
} |
||||
CHECK(f.talking == 0, "busy: не стартуем"); |
||||
} |
||||
|
||||
int main(void) { |
||||
test_resample(); |
||||
test_fsm(); |
||||
|
||||
if (g_failures == 0) { printf("TEST PASSED\n"); return 0; } |
||||
printf("TEST FAILED (%d failures)\n", g_failures); |
||||
return 1; |
||||
} |
||||
@ -0,0 +1,196 @@
|
||||
/* Тест Silero VAD-обёртки (lib/silero_vad) поверх ONNX Runtime.
|
||||
* |
||||
* Проверяет: |
||||
* - загрузку модели и жизненный цикл (create/destroy/reset); |
||||
* - обработку ошибок (пустой/несуществующий путь, NULL-аргументы); |
||||
* - что на тишине вероятность речи близка к нулю (реальный инференс, не мусор); |
||||
* - что путь pcm16 даёт тот же результат, что и float (конвертация /32768); |
||||
* - что вероятность всегда лежит в [0;1]. |
||||
* |
||||
* Модель ищется в argv[1], по умолчанию "../lib/silero_vad.onnx" (запуск из tests/). |
||||
*/ |
||||
#include <stdio.h> |
||||
#include <string.h> |
||||
#include <math.h> |
||||
|
||||
#include "../lib/silero_vad.h" |
||||
#include "../lib/debug_config.h" |
||||
|
||||
#define PI 3.14159265358979323846 |
||||
|
||||
static int g_failures = 0; |
||||
|
||||
#define CHECK(cond, ...) do { \ |
||||
if (!(cond)) { \
|
||||
printf(" FAIL: " __VA_ARGS__); printf("\n"); \
|
||||
g_failures++; \
|
||||
} \
|
||||
} while (0) |
||||
|
||||
/* Детерминированный «гласный»: модулированный гармонический комплекс (f0 ~120 Гц). */ |
||||
static void gen_vowel(float* out, int n, int offset) { |
||||
int i; |
||||
for (i = 0; i < n; i++) { |
||||
double t = (double)(offset + i) / 16000.0; |
||||
double f0 = 120.0 + 20.0 * sin(2.0 * PI * 3.0 * t); |
||||
double env = 0.5 + 0.5 * sin(2.0 * PI * 2.0 * t); |
||||
double s = sin(2.0 * PI * f0 * t) |
||||
+ 0.5 * sin(2.0 * PI * 2.0 * f0 * t) |
||||
+ 0.3 * sin(2.0 * PI * 3.0 * f0 * t) |
||||
+ 0.2 * sin(2.0 * PI * 4.0 * f0 * t); |
||||
out[i] = (float)(s * env * 0.25); |
||||
} |
||||
} |
||||
|
||||
static int test_lifecycle_and_errors(const char* model_path) { |
||||
silero_vad_t* vad; |
||||
|
||||
printf("[silero_vad] lifecycle + error paths\n"); |
||||
|
||||
vad = silero_vad_create(NULL); |
||||
CHECK(vad == NULL, "create(NULL) should fail"); |
||||
if (vad) silero_vad_destroy(vad); |
||||
|
||||
vad = silero_vad_create(""); |
||||
CHECK(vad == NULL, "create(\"\") should fail"); |
||||
|
||||
vad = silero_vad_create("/nonexistent/silero_vad.onnx"); |
||||
CHECK(vad == NULL, "create(bad path) should fail"); |
||||
|
||||
vad = silero_vad_create(model_path); |
||||
CHECK(vad != NULL, "create(%s) should succeed", model_path); |
||||
if (!vad) return 1; |
||||
|
||||
silero_vad_reset(vad); |
||||
silero_vad_reset(vad); /* повторный reset безопасен */ |
||||
silero_vad_destroy(vad); |
||||
silero_vad_destroy(NULL); /* destroy(NULL) безопасен */ |
||||
return 0; |
||||
} |
||||
|
||||
static int test_silence(const char* model_path) { |
||||
silero_vad_t* vad; |
||||
float zeros[SILERO_VAD_WINDOW_SAMPLES]; |
||||
float prob = -1.0f; |
||||
float max_prob = 0.0f; |
||||
int i, rc; |
||||
|
||||
printf("[silero_vad] silence -> prob ~ 0\n"); |
||||
memset(zeros, 0, sizeof(zeros)); |
||||
|
||||
vad = silero_vad_create(model_path); |
||||
CHECK(vad != NULL, "create failed"); |
||||
if (!vad) return 1; |
||||
|
||||
for (i = 0; i < 32; i++) { |
||||
rc = silero_vad_process(vad, zeros, &prob); |
||||
CHECK(rc == 0, "process(%d) rc=%d", i, rc); |
||||
CHECK(prob >= 0.0f && prob <= 1.0f, "prob out of range: %f", prob); |
||||
if (prob > max_prob) max_prob = prob; |
||||
} |
||||
printf(" silence max_prob=%.4f\n", max_prob); |
||||
CHECK(max_prob < 0.2f, "silence probability too high: %f", max_prob); |
||||
|
||||
silero_vad_destroy(vad); |
||||
return 0; |
||||
} |
||||
|
||||
static int test_pcm16_matches_float(const char* model_path) { |
||||
silero_vad_t* vad; |
||||
int16_t pcm[SILERO_VAD_WINDOW_SAMPLES]; |
||||
float samples[SILERO_VAD_WINDOW_SAMPLES]; |
||||
float prob_f = -1.0f, prob_i = -1.0f; |
||||
int rc; |
||||
|
||||
printf("[silero_vad] pcm16 path == float path\n"); |
||||
|
||||
vad = silero_vad_create(model_path); |
||||
CHECK(vad != NULL, "create failed"); |
||||
if (!vad) return 1; |
||||
|
||||
gen_vowel(samples, SILERO_VAD_WINDOW_SAMPLES, 0); |
||||
for (int i = 0; i < SILERO_VAD_WINDOW_SAMPLES; i++) { |
||||
pcm[i] = (int16_t)(samples[i] * 32768.0f); |
||||
} |
||||
|
||||
rc = silero_vad_process(vad, samples, &prob_f); |
||||
CHECK(rc == 0, "process(float) rc=%d", rc); |
||||
|
||||
silero_vad_reset(vad); /* одинаковое состояние перед вторым прогоном */ |
||||
|
||||
rc = silero_vad_process_pcm16(vad, pcm, &prob_i); |
||||
CHECK(rc == 0, "process_pcm16 rc=%d", rc); |
||||
|
||||
printf(" float prob=%.5f pcm16 prob=%.5f\n", prob_f, prob_i); |
||||
CHECK(fabsf(prob_f - prob_i) < 1e-4f, "pcm16 != float: %f vs %f", prob_i, prob_f); |
||||
|
||||
silero_vad_destroy(vad); |
||||
return 0; |
||||
} |
||||
|
||||
static int test_default_model(void) { |
||||
silero_vad_t* vad; |
||||
float zeros[SILERO_VAD_WINDOW_SAMPLES]; |
||||
float prob = -1.0f; |
||||
int rc; |
||||
|
||||
printf("[silero_vad] create_default (embedded model)\n"); |
||||
memset(zeros, 0, sizeof(zeros)); |
||||
|
||||
vad = silero_vad_create_default(); |
||||
CHECK(vad != NULL, "create_default failed"); |
||||
if (!vad) return 1; |
||||
|
||||
rc = silero_vad_process(vad, zeros, &prob); |
||||
CHECK(rc == 0, "process rc=%d", rc); |
||||
CHECK(prob >= 0.0f && prob <= 1.0f, "prob out of range: %f", prob); |
||||
CHECK(prob < 0.2f, "silence prob too high: %f", prob); |
||||
printf(" embedded silence prob=%.4f\n", prob); |
||||
|
||||
silero_vad_destroy(vad); |
||||
return 0; |
||||
} |
||||
|
||||
static int test_null_args(const char* model_path) { |
||||
silero_vad_t* vad; |
||||
float samples[SILERO_VAD_WINDOW_SAMPLES]; |
||||
float prob; |
||||
|
||||
printf("[silero_vad] NULL-arg guards\n"); |
||||
memset(samples, 0, sizeof(samples)); |
||||
|
||||
vad = silero_vad_create(model_path); |
||||
CHECK(vad != NULL, "create failed"); |
||||
if (!vad) return 1; |
||||
|
||||
CHECK(silero_vad_process(NULL, samples, &prob) != 0, "process(NULL vad) should fail"); |
||||
CHECK(silero_vad_process(vad, NULL, &prob) != 0, "process(NULL samples) should fail"); |
||||
CHECK(silero_vad_process(vad, samples, NULL) != 0, "process(NULL prob) should fail"); |
||||
CHECK(silero_vad_process_pcm16(vad, NULL, &prob) != 0, "process_pcm16(NULL) should fail"); |
||||
|
||||
silero_vad_destroy(vad); |
||||
return 0; |
||||
} |
||||
|
||||
int main(int argc, char** argv) { |
||||
const char* model_path = (argc > 1) ? argv[1] : "../lib/silero_vad.onnx"; |
||||
|
||||
debug_config_init(); |
||||
debug_set_level(DEBUG_LEVEL_INFO); |
||||
debug_set_category_level(DEBUG_CATEGORY_VAD, DEBUG_LEVEL_INFO); |
||||
|
||||
printf("Silero VAD test, model: %s\n", model_path); |
||||
|
||||
test_lifecycle_and_errors(model_path); |
||||
test_silence(model_path); |
||||
test_pcm16_matches_float(model_path); |
||||
test_null_args(model_path); |
||||
test_default_model(); |
||||
|
||||
if (g_failures == 0) { |
||||
printf("TEST PASSED\n"); |
||||
return 0; |
||||
} |
||||
printf("TEST FAILED (%d failures)\n", g_failures); |
||||
return 1; |
||||
} |
||||
@ -0,0 +1,19 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "onnxruntime_c_api.h" |
||||
|
||||
#ifdef __cplusplus |
||||
extern "C" { |
||||
#endif |
||||
|
||||
/**
|
||||
* \param use_arena zero: false. non-zero: true. |
||||
*/ |
||||
ORT_EXPORT |
||||
ORT_API_STATUS(OrtSessionOptionsAppendExecutionProvider_CPU, _In_ OrtSessionOptions* options, int use_arena) |
||||
ORT_ALL_ARGS_NONNULL; |
||||
|
||||
#ifdef __cplusplus |
||||
} |
||||
#endif |
||||
@ -0,0 +1,62 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
#pragma once |
||||
|
||||
#include "onnxruntime_c_api.h" |
||||
|
||||
// NNAPIFlags are bool options we want to set for NNAPI EP
|
||||
// This enum is defined as bit flags, and cannot have negative value
|
||||
// To generate an uint32_t nnapi_flags for using with OrtSessionOptionsAppendExecutionProvider_Nnapi below,
|
||||
// uint32_t nnapi_flags = 0;
|
||||
// nnapi_flags |= NNAPI_FLAG_USE_FP16;
|
||||
enum NNAPIFlags { |
||||
NNAPI_FLAG_USE_NONE = 0x000, |
||||
|
||||
// Using fp16 relaxation in NNAPI EP, this may improve perf but may also reduce precision
|
||||
NNAPI_FLAG_USE_FP16 = 0x001, |
||||
|
||||
// Use NCHW layout in NNAPI EP, this is only available after Android API level 29
|
||||
// Please note for now, NNAPI perform worse using NCHW compare to using NHWC
|
||||
NNAPI_FLAG_USE_NCHW = 0x002, |
||||
|
||||
// Prevent NNAPI from using CPU devices.
|
||||
//
|
||||
// NNAPI is more efficient using GPU or NPU for execution, and NNAPI might fall back to its own CPU implementation
|
||||
// for operations not supported by GPU/NPU. The CPU implementation of NNAPI (which is called nnapi-reference)
|
||||
// might be less efficient than the optimized versions of the operation of ORT. It might be advantageous to disable
|
||||
// the NNAPI CPU fallback and handle execution using ORT kernels.
|
||||
//
|
||||
// For some models, if NNAPI would use CPU to execute an operation, and this flag is set, the execution of the
|
||||
// model may fall back to ORT kernels.
|
||||
//
|
||||
// This option is only available after Android API level 29, and will be ignored for Android API level 28-
|
||||
//
|
||||
// For NNAPI device assignments, see https://developer.android.com/ndk/guides/neuralnetworks#device-assignment
|
||||
// For NNAPI CPU fallback, see https://developer.android.com/ndk/guides/neuralnetworks#cpu-fallback
|
||||
//
|
||||
// Please note, the NNAPI EP will return error status if both NNAPI_FLAG_CPU_DISABLED
|
||||
// and NNAPI_FLAG_CPU_ONLY flags are set
|
||||
NNAPI_FLAG_CPU_DISABLED = 0x004, |
||||
|
||||
// Using CPU only in NNAPI EP, this may decrease the perf but will provide
|
||||
// reference output value without precision loss, which is useful for validation
|
||||
//
|
||||
// Please note, the NNAPI EP will return error status if both NNAPI_FLAG_CPU_DISABLED
|
||||
// and NNAPI_FLAG_CPU_ONLY flags are set
|
||||
NNAPI_FLAG_CPU_ONLY = 0x008, |
||||
|
||||
// Keep NNAPI_FLAG_LAST at the end of the enum definition
|
||||
// And assign the last NNAPIFlag to it
|
||||
NNAPI_FLAG_LAST = NNAPI_FLAG_CPU_ONLY, |
||||
}; |
||||
|
||||
#ifdef __cplusplus |
||||
extern "C" { |
||||
#endif |
||||
|
||||
ORT_EXPORT ORT_API_STATUS(OrtSessionOptionsAppendExecutionProvider_Nnapi, |
||||
_In_ OrtSessionOptions* options, uint32_t nnapi_flags); |
||||
|
||||
#ifdef __cplusplus |
||||
} |
||||
#endif |
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@ -0,0 +1,988 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
// Do not include this file directly. Please include "onnxruntime_c_api.h" instead.
|
||||
|
||||
#ifdef __cplusplus |
||||
extern "C" { |
||||
#endif |
||||
|
||||
ORT_RUNTIME_CLASS(Ep); |
||||
ORT_RUNTIME_CLASS(EpFactory); |
||||
ORT_RUNTIME_CLASS(EpGraphSupportInfo); |
||||
ORT_RUNTIME_CLASS(MemoryDevice); // opaque class to wrap onnxruntime::OrtDevice
|
||||
ORT_RUNTIME_CLASS(NodeComputeContext); |
||||
|
||||
ORT_RUNTIME_CLASS(DataTransferImpl); |
||||
ORT_RUNTIME_CLASS(SyncNotificationImpl); |
||||
ORT_RUNTIME_CLASS(SyncStreamImpl); |
||||
|
||||
// struct that an EP implements for IDataTransfer to copy between devices it uses and CPU
|
||||
struct OrtDataTransferImpl { |
||||
uint32_t ort_version_supported; ///< Must be initialized to ORT_API_VERSION
|
||||
|
||||
/** \brief Release the OrtDataTransferImpl instance.
|
||||
* |
||||
* This is called by ORT when the OrtDataTransferImpl instance is no longer needed. |
||||
* The implementation should release any resources held by the instance. |
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtDataTransferImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(void, Release, _In_ OrtDataTransferImpl* this_ptr); |
||||
|
||||
/** \brief Check if the implementation can copy between the source and destination memory devices.
|
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtDataTransferImpl instance. |
||||
* \param[in] src_memory_device Source OrtMemoryDevice to copy from. |
||||
* \param[in] dst_memory_device Destination OrtMemoryDevice to copy to. |
||||
* \return True if the implementation can copy between the devices. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(bool, CanCopy, _In_ const OrtDataTransferImpl* this_ptr, |
||||
_In_ const OrtMemoryDevice* src_memory_device, _In_ const OrtMemoryDevice* dst_memory_device); |
||||
|
||||
/** \brief Copy tensors from src_tensors to dst_tensors using the provided streams.
|
||||
* |
||||
* The implementation can use the provided streams to perform asynchronous copies if supported. |
||||
* If a stream is not available, the copy is performed synchronously. |
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtDataTransferImpl instance. |
||||
* \param[in] src_tensors Array of source OrtValue pointers to copy from. |
||||
* \param[in] dst_tensors Array of destination OrtValue pointers to copy to. |
||||
* \param[in] streams Array of OrtSyncStream pointers for the copy operations, if the execution provider is stream |
||||
* aware. nullptr if it is not. |
||||
* \param[in] num_tensors Number of tensors to copy. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(CopyTensors, _In_ OrtDataTransferImpl* this_ptr, |
||||
_In_reads_(num_tensors) const OrtValue** src_tensors, |
||||
_In_reads_(num_tensors) OrtValue** dst_tensors, |
||||
_In_reads_(num_tensors) OrtSyncStream** streams, |
||||
_In_ size_t num_tensors); |
||||
}; |
||||
|
||||
/** \brief Struct that an EP implements for Stream Notifications.
|
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
struct OrtSyncNotificationImpl { |
||||
uint32_t ort_version_supported; ///< Must be initialized to ORT_API_VERSION
|
||||
|
||||
/** \brief Release the OrtSyncNotificationImpl instance.
|
||||
* |
||||
* This is called by ORT when the OrtSyncNotificationImpl instance is no longer needed. |
||||
* The implementation should release any resources held by the instance. |
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncNotificationImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(void, Release, _In_ OrtSyncNotificationImpl* this_ptr); |
||||
|
||||
/** \brief Called by ORT to activate the notification.
|
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncNotificationImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(Activate, _In_ OrtSyncNotificationImpl* this_ptr); |
||||
|
||||
/** \brief Wait for a device to device operation to complete.
|
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncNotificationImpl instance. |
||||
* \param[in] stream The OrtSyncStream instance that will wait on this notification to be activated. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(WaitOnDevice, _In_ OrtSyncNotificationImpl* this_ptr, _In_ OrtSyncStream* consumer_stream); |
||||
|
||||
/** \brief Wait for a device to host operation to complete.
|
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncNotificationImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(WaitOnHost, _In_ OrtSyncNotificationImpl* this_ptr); |
||||
}; |
||||
|
||||
/** \brief Struct that an EP implements if it wishes to implement Stream support.
|
||||
* |
||||
* This struct provides the overrides for onnxruntime::Stream's virtual methods. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
struct OrtSyncStreamImpl { |
||||
uint32_t ort_version_supported; ///< Must be initialized to ORT_API_VERSION
|
||||
|
||||
/** \brief Release the OrtSyncStreamImpl instance.
|
||||
* |
||||
* This is called by ORT when the OrtSyncStreamImpl instance is no longer needed. |
||||
* The implementation should release any resources held by the instance. |
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncStreamImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(void, Release, _In_ OrtSyncStreamImpl* this_ptr); |
||||
|
||||
/** \brief Get the handle of the stream.
|
||||
* |
||||
* This returns the native handle for the stream. e.g. cudaStream_t for CUDA streams. |
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncStreamImpl instance. |
||||
* \return The handle of the stream. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(void*, GetHandle, _In_ OrtSyncStreamImpl* this_ptr); |
||||
|
||||
/** \brief Create an OrtSyncNotificationImpl for the OrtSyncStreamImpl instance.
|
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncStreamImpl instance |
||||
* \param[out] notification The new OrtSyncNotificationImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(CreateNotification, _In_ OrtSyncStreamImpl* this_ptr, |
||||
_Outptr_ OrtSyncNotificationImpl** notification); |
||||
|
||||
/** \brief Flush the stream.
|
||||
* |
||||
* This is called by ORT to flush the stream, ensuring that all operations submitted to the stream are completed. |
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncStreamImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(Flush, _In_ OrtSyncStreamImpl* this_ptr); |
||||
|
||||
/** \brief Notify the stream that a session run has ended.
|
||||
* |
||||
* This is called by ORT to notify the stream that a session run has ended, allowing the stream to perform any |
||||
* necessary cleanup or finalization. |
||||
* |
||||
* \param[in] this_ptr Pointer to the OrtSyncStreamImpl instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(OnSessionRunEnd, _In_ OrtSyncStreamImpl* this_ptr); |
||||
}; |
||||
|
||||
struct OrtNodeFusionOptions; |
||||
typedef struct OrtNodeFusionOptions OrtNodeFusionOptions; |
||||
|
||||
struct OrtNodeComputeInfo; |
||||
typedef struct OrtNodeComputeInfo OrtNodeComputeInfo; |
||||
|
||||
/**
|
||||
* \brief The OrtNodeFusionOptions struct specifies options for fusing nodes supported by an execution provider. |
||||
* |
||||
* Refer to OrtEpApi::EpGraphSupportInfo_AddNodesToFuse. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
struct OrtNodeFusionOptions { |
||||
/** \brief The ONNX Runtime version the OrtNodeFusionOptions was compiled with.
|
||||
* |
||||
* Implementation should set to ORT_API_VERSION. |
||||
* ORT will use this to ensure it does not use members that were not available when the EP library was compiled. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
uint32_t ort_version_supported; |
||||
|
||||
/** \brief If set to true, specify that the execution provider does not require ONNX Runtime to provide constant
|
||||
* initializers as inputs to the fused node during model inference. This is used when the execution |
||||
* provider saves a copy of constant initializers, and allows ONNX Runtime to release constant initializers that |
||||
* are not used by any execution provider. |
||||
* |
||||
* If not specified, defaults to false. That is, ONNX Runtime provides constant initializers as inputs to |
||||
* the fused node by default. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
bool drop_constant_initializers; |
||||
|
||||
// const OrtNode* fused_node_schema;
|
||||
}; |
||||
|
||||
/**
|
||||
* \brief The OrtNodeComputeInfo struct provides functions that an OrtEp implements to specify the compute |
||||
* function for a compiled OrtGraph instance. |
||||
* \since Version 1.23. |
||||
*/ |
||||
struct OrtNodeComputeInfo { |
||||
/** \brief The ONNX Runtime version the OrtNodeComputeInfo was compiled with.
|
||||
* |
||||
* Implementation should set to ORT_API_VERSION. |
||||
* ORT will use this to ensure it does not call functions that were not available when the EP library was compiled. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
uint32_t ort_version_supported; |
||||
|
||||
/** \brief Creates an opaque compute state object that is then passed to the Compute() function during inference.
|
||||
* \param[in] this_ptr The OrtNodeComputeInfo instance. |
||||
* \param[in] compute_context OrtNodeComputeContext instance that contains compiled/fused node's name and host |
||||
* memory allocation functions. Can optionally be used to build the compute state. |
||||
* \param[out] compute_state Output parameter that is assigned the opaque computation state. ONNX Runtime calls |
||||
* ReleaseState() (after calling Compute()) to allow the implementer to release the |
||||
* compute state. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
OrtStatus*(ORT_API_CALL* CreateState)(_In_ OrtNodeComputeInfo* this_ptr, |
||||
_In_ OrtNodeComputeContext* compute_context, |
||||
_Outptr_ void** compute_state); |
||||
|
||||
/** \brief Computation function called to execute the fused node compiled by an OrtEp instance.
|
||||
* \param[in] this_ptr The OrtNodeComputeInfo instance. |
||||
* \param[in] compute_state The opaque computation state returned by CreateState(). |
||||
* \param[in] kernel_context The OrtKernelContext instance used to access inputs/outputs. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
OrtStatus*(ORT_API_CALL* Compute)(_In_ OrtNodeComputeInfo* this_ptr, _In_ void* compute_state, |
||||
_In_ OrtKernelContext* kernel_context); |
||||
|
||||
/** \brief Releases the compute state returned by CreateState().
|
||||
* \param[in] this_ptr The OrtNodeComputeInfo instance. |
||||
* \param[inout] compute_state The opaque compute state returned by CreateState(). |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
void(ORT_API_CALL* ReleaseState)(_In_ OrtNodeComputeInfo* this_ptr, _Frees_ptr_opt_ void* compute_state); |
||||
}; |
||||
|
||||
struct OrtEpApi { |
||||
/** \brief Create an OrtEpDevice for the EP and an OrtHardwareDevice.
|
||||
* \param[in] ep_factory Execution provider factory that is creating the instance. |
||||
* \param[in] hardware_device Hardware device that the EP can utilize. |
||||
* \param[in] ep_metadata Optional OrtKeyValuePairs instance for execution provider metadata that may be used |
||||
* during execution provider selection and passed to CreateEp. |
||||
* ep_device will copy this instance and the user should call ReleaseKeyValuePairs. |
||||
* \param[in] ep_options Optional OrtKeyValuePairs instance for execution provider options that will be added |
||||
* to the Session configuration options if the execution provider is selected. |
||||
* ep_device will copy this instance and the user should call ReleaseKeyValuePairs. |
||||
* \param ep_device OrtExecutionDevice that is created. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
ORT_API2_STATUS(CreateEpDevice, _In_ OrtEpFactory* ep_factory, |
||||
_In_ const OrtHardwareDevice* hardware_device, |
||||
_In_opt_ const OrtKeyValuePairs* ep_metadata, |
||||
_In_opt_ const OrtKeyValuePairs* ep_options, |
||||
_Out_ OrtEpDevice** ep_device); |
||||
|
||||
ORT_CLASS_RELEASE(EpDevice); |
||||
|
||||
/** \brief Specify nodes that are supported by an OrtEp and should be fused into one node.
|
||||
* |
||||
* Because the nodes will be fused into one "fused node", there must not exist an unsupported node in |
||||
* a path between two of the provided nodes. Otherwise, the graph will become invalid. |
||||
* |
||||
* This function can be called multiple times. A subsequent call to this function will force the next set of |
||||
* nodes to be fused into a different node. |
||||
* |
||||
* \param[in] graph_support_info OrtEpGraphSupportInfo instance to which to add the supported nodes. |
||||
* \param[in] nodes Array of nodes supported by the EP that should be fused/compiled. |
||||
* \param[in] num_nodes The number of supported nodes. |
||||
* \param[in] node_fusion_options Optional node fusion options. Ignored if set to NULL. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(EpGraphSupportInfo_AddNodesToFuse, _In_ OrtEpGraphSupportInfo* graph_support_info, |
||||
_In_reads_(num_nodes) const OrtNode* const* nodes, _In_ size_t num_nodes, |
||||
_In_opt_ const OrtNodeFusionOptions* node_fusion_options); |
||||
|
||||
/** \brief Specify a node that is supported by an OrtEp and should be run with a registered EP kernel.
|
||||
* |
||||
* \param[in] graph_support_info OrtEpGraphSupportInfo instance to which to add the supported node. |
||||
* \param[in] node The supported OrtNode instance. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(EpGraphSupportInfo_AddSingleNode, _In_ OrtEpGraphSupportInfo* graph_support_info, |
||||
_In_ const OrtNode* node); |
||||
|
||||
/** \brief Query a OrtNodeComputeContext for the name of the node that encapsulates the compiled/fused node.
|
||||
* |
||||
* Used in OrtNodeComputeInfo::CreateComputeState(). |
||||
* |
||||
* \param[in] context The OrtNodeComputeContext instance to query. |
||||
* \return The node's name. |
||||
* |
||||
* \note Returned string is owned by ORT and valid only while OrtNodeComputeInfo::CreateComputeState() is called. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(const char*, NodeComputeContext_NodeName, _In_ const OrtNodeComputeContext* context); |
||||
|
||||
/** \brief Register an allocator with the OrtEpDevice.
|
||||
* |
||||
* This allows an EP to provide OrtMemoryInfo for DEFAULT and HOST_ACCESSIBLE memory type as needed. |
||||
* The registered values will be used in calls to OrtEpFactory::CreateAllocator to ensure the required allocator/s |
||||
* are available for EP usage. |
||||
* |
||||
* Multiple calls for the same entry type will replace a previous entry. |
||||
* |
||||
* Available entries: |
||||
* - OrtDeviceAllocator with type of OrtDeviceMemoryType_DEFAULT |
||||
* - OrtDeviceAllocator with type of OrtDeviceMemoryType_HOST_ACCESSIBLE |
||||
* - OrtReadOnlyAllocator with type of OrtDeviceMemoryType_DEFAULT |
||||
* - if provided this allocator will only be used to copy initializers to the device the EP uses. |
||||
* ORT will use the OrtDeviceAllocator if not provided. |
||||
* |
||||
* \param[in] ep_device The OrtEpDevice instance to register the OrtMemoryInfo with. |
||||
* \param[in] allocator_memory_info The OrtMemoryInfo information for the allocator. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(EpDevice_AddAllocatorInfo, _In_ OrtEpDevice* ep_device, |
||||
_In_ const OrtMemoryInfo* allocator_memory_info); |
||||
|
||||
/** \brief Get the OrtMemoryDevice from an OrtMemoryInfo instance.
|
||||
* |
||||
* This is required for OrtDataTransferImpl (which implements onnxruntime::IDataTransfer) where the OrtMemoryDevice |
||||
* is used in the CanCopy and CopyTensors functions. |
||||
* |
||||
* \param[in] memory_info The OrtMemoryInfo instance to get the memory device from. |
||||
* \return The OrtMemoryDevice associated with the OrtMemoryInfo instance. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(const OrtMemoryDevice*, MemoryInfo_GetMemoryDevice, _In_ const OrtMemoryInfo* memory_info); |
||||
|
||||
/** \brief Get the OrtMemoryDevice from an OrtValue instance if it contains a Tensor.
|
||||
* |
||||
* \param[in] value The OrtValue instance to get the memory device from. |
||||
* \return Memory device if OrtValue contains a Tensor, nullptr otherwise. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(const OrtMemoryDevice*, Value_GetMemoryDevice, _In_ const OrtValue* value); |
||||
|
||||
/** \brief Compare two OrtMemoryDevice instances for equality.
|
||||
* |
||||
* This is used to check if two memory devices are the same. |
||||
* Used to implement DataTransferImpl::CanCopy. |
||||
* |
||||
* \param[in] a The first OrtMemoryDevice instance to compare. |
||||
* \param[in] b The second OrtMemoryDevice instance to compare. |
||||
* \return True if the two OrtMemoryDevice instances are equal, false otherwise. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(bool, MemoryDevice_AreEqual, _In_ const OrtMemoryDevice* a, _In_ const OrtMemoryDevice* b); |
||||
|
||||
/** \brief Get the OrtMemoryInfoDeviceType value from an OrtMemoryDevice instance.
|
||||
* |
||||
* \param[in] memory_device OrtMemoryDevice instance. |
||||
* \return The OrtMemoryInfoDeviceType value. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(OrtMemoryInfoDeviceType, MemoryDevice_GetDeviceType, _In_ const OrtMemoryDevice* memory_device); |
||||
|
||||
/** \brief Get the OrtDeviceMemoryType value from an OrtMemoryDevice instance.
|
||||
* |
||||
* \param[in] memory_device OrtMemoryDevice instance. |
||||
* \return The OrtDeviceMemoryType value. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(OrtDeviceMemoryType, MemoryDevice_GetMemoryType, _In_ const OrtMemoryDevice* memory_device); |
||||
|
||||
/** \brief Get the vendor ID from an OrtMemoryDevice instance.
|
||||
* |
||||
* The vendor ID is used to identify the vendor of the device, and is typically set to the PCI vendor ID. |
||||
* |
||||
* If the device is not vendor specific (e.g. CPU memory) the vendor ID is set to 0. |
||||
* |
||||
* \param[in] memory_device OrtMemoryDevice instance. |
||||
* \return The vendor ID value. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(uint32_t, MemoryDevice_GetVendorId, _In_ const OrtMemoryDevice* memory_device); |
||||
|
||||
/** \brief Get the device ID from an OrtMemoryDevice instance.
|
||||
* |
||||
* \param[in] memory_device OrtMemoryDevice instance. |
||||
* \return The device ID. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(uint32_t, MemoryDevice_GetDeviceId, _In_ const OrtMemoryDevice* memory_device); |
||||
|
||||
/** \brief Get the OrtSyncStreamImpl associated with an OrtSyncStream instance.
|
||||
* |
||||
* This allows an the plugin library to connect its OrtSyncStreamImpl instance with an OrtSyncStream if needed. |
||||
* |
||||
* \param[in] stream The OrtSyncStream instance to find an OrtSyncStreamImpl for. |
||||
* \return The associated OrtSyncStreamImpl if found. nullptr otherwise. |
||||
* |
||||
* \since Version 1.23. |
||||
* |
||||
* \remarks There should always be an OrtSyncStreamImpl associated with an OrtSyncStream instance that the EP gets. |
||||
*/ |
||||
ORT_API_T(const OrtSyncStreamImpl*, SyncStream_GetImpl, _In_ const OrtSyncStream* stream); |
||||
|
||||
/** \brief Get the current sync ID for a stream.
|
||||
* |
||||
* \param[in] stream The OrtSyncStream to get the sync ID for. |
||||
* \return Current sync ID. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(uint64_t, SyncStream_GetSyncId, _In_ const OrtSyncStream* stream); |
||||
|
||||
/** \brief Get the sync ID for the last time the consumer_stream waited on the producer_stream.
|
||||
* |
||||
* When two streams are synchronized, the sync id represents the event used in that synchronization. |
||||
* |
||||
* \param[in] producer_stream The OrtSyncStream that produced the data. |
||||
* \param[in] consumer_stream The OrtSyncStream that waited on the producer_stream. |
||||
* \return ID for last sync. 0 if no sync has occurred between the two streams. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(uint64_t, GetSyncIdForLastWaitOnSyncStream, |
||||
_In_ const OrtSyncStream* producer_stream, _In_ const OrtSyncStream* consumer_stream); |
||||
}; |
||||
|
||||
/**
|
||||
* \brief The data layout type. |
||||
* |
||||
* EPs may specify a preferred data layout type. ORT's default layout type is OrtEpDataLayout_NCHW, or |
||||
* OrtEpDataLayout_Default. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
typedef enum OrtEpDataLayout { |
||||
OrtEpDataLayout_NCHW = 0, |
||||
OrtEpDataLayout_NHWC, |
||||
|
||||
OrtEpDataLayout_Default = OrtEpDataLayout_NCHW, |
||||
} OrtEpDataLayout; |
||||
|
||||
/**
|
||||
* \brief The OrtEp struct provides functions to implement for an execution provider. |
||||
* \since Version 1.22. |
||||
*/ |
||||
struct OrtEp { |
||||
/** \brief The ONNX Runtime version the execution provider was compiled with.
|
||||
* |
||||
* Implementation should set to ORT_API_VERSION. |
||||
* ORT will use this to ensure it does not call functions that were not available when the library was compiled. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
uint32_t ort_version_supported; |
||||
|
||||
/** \brief Get the execution provider name.
|
||||
* |
||||
* The returned string should be a null-terminated, UTF-8 encoded string. ORT will copy it. |
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \return The execution provider name. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
ORT_API_T(const char*, GetName, _In_ const OrtEp* this_ptr); |
||||
|
||||
/** \brief Get information about the nodes supported by the OrtEp instance.
|
||||
* |
||||
* IMPORTANT: This is not the final version of this API function. This is currently experimental but will |
||||
* be stabilized by the ONNX Runtime 1.23 release. |
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[in] graph The OrtGraph instance for which to populate node support. The OrtGraph could be a nested subgraph |
||||
* contained by a node (e.g., an If or Loop node). ONNX Runtime calls this function separately |
||||
* for each nested subgraph. |
||||
* \param[inout] graph_support_info OrtEpGraphSupportInfo instance that the implementer must fill out in order to |
||||
* specify the supported nodes. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(GetCapability, _In_ OrtEp* this_ptr, _In_ const OrtGraph* graph, |
||||
_Inout_ OrtEpGraphSupportInfo* graph_support_info); |
||||
|
||||
/** \brief Compile OrtGraph instances assigned to the OrtEp. Implementer must set a OrtNodeComputeInfo instance
|
||||
* for each OrtGraph in order to define its computation function. |
||||
* |
||||
* If the session is configured to generate a pre-compiled model, the execution provider must return EPContext nodes, |
||||
* as OrtNode instances, that ONNX Runtime uses to create a pre-compiled model, known as an "EPContext model". |
||||
* An EPContext model contains EPContext nodes. Each EPContext node encapsulates the pre-compiled binary data for a |
||||
* OrtGraph compiled for a specific execution provider. For more details about the EPContext design, refer to: |
||||
* \htmlonly |
||||
* <a href="https://onnxruntime.ai/docs/execution-providers/EP-Context-Design.html">EPContext design document.</a> |
||||
* \endhtmlonly |
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[in] graphs Array of `count` OrtGraph instances to compile. Each graph contains only the nodes for |
||||
* which the execution provider indicated support. Nested subgraphs contained by a |
||||
* node, such as an If or Loop, have separate OrtGraph instances. |
||||
* \param[in] fused_nodes Array of `count` fused nodes that will replace the compiled graphs. |
||||
* Each fused node is an OrtNode initialized with the intended fused node name and |
||||
* input/output information. |
||||
* \param[in] count The number of OrtGraph instances to compile. |
||||
* \param[out] node_compute_infos Array of `count` OrtNodeComputeInfo instances that define each OrtGraph instance's |
||||
* computation function. The implementer allocates the OrtNodeComputeInfo instances. |
||||
* ORT calls ReleaseNodeComputeInfos() to release multiple instances in a batch. |
||||
* \param[out] ep_context_nodes Output array of `count` OrtNode instances, each representing an EPContext |
||||
* node for a compiled OrtGraph. The execution provider must use |
||||
* OrtModelEditorApi::CreateNode to create the OrtNode instances. ONNX Runtime takes |
||||
* ownership of the OrtNode instances, so the execution provider must NOT call |
||||
* OrtApi::ReleaseNode. Should be ignored if the session is not configured to generate an |
||||
* EPContext model. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \note Do NOT cache the provided OrtGraph instances in any of the OrtNodeComputeInfo functions because the |
||||
* graphs are only valid for the duration of the call to Compile. Any graph/node/input/output |
||||
* names that are needed by the OrtNodeComputeInfo functions must be copied and stored by the OrtEp. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(Compile, _In_ OrtEp* this_ptr, _In_ const OrtGraph** graphs, |
||||
_In_ const OrtNode** fused_nodes, _In_ size_t count, |
||||
_Out_writes_all_(count) OrtNodeComputeInfo** node_compute_infos, |
||||
_Out_writes_(count) OrtNode** ep_context_nodes); |
||||
|
||||
/** \brief Release OrtNodeComputeInfo instances.
|
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[inout] node_compute_infos The OrtNodeComputeInfo instances to release. |
||||
* \param[in] num_node_compute_infos The number of OrtNodeComputeInfo instances. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(void, ReleaseNodeComputeInfos, _In_ OrtEp* this_ptr, |
||||
OrtNodeComputeInfo** node_compute_infos, |
||||
_In_ size_t num_node_compute_infos); |
||||
|
||||
/** \brief Get the EP's preferred data layout.
|
||||
* |
||||
* \note Implementation of this function is optional. |
||||
* If not implemented, ORT will assume that this EP prefers the data layout `OrtEpDataLayout::NCHW`. |
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[out] preferred_data_layout The EP's preferred data layout. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(GetPreferredDataLayout, _In_ OrtEp* this_ptr, _Out_ OrtEpDataLayout* preferred_data_layout); |
||||
|
||||
/** \brief Given an op with domain `domain` and type `op_type`, determine whether an associated node's data layout
|
||||
* should be converted to `target_data_layout`. |
||||
* If the EP prefers a non-default data layout (see `GetPreferredDataLayout()`), this function will be called |
||||
* during layout transformation with `target_data_layout` set to the EP's preferred data layout. |
||||
* |
||||
* \note Implementation of this function is optional. |
||||
* If an EP prefers a non-default data layout, it may implement this to customize the specific op data layout |
||||
* preferences at a finer granularity. |
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[in] domain The op domain. An empty string means the ONNX domain. |
||||
* \param[in] op_type The op type. |
||||
* \param[in] target_data_layout The target data layout. |
||||
* \param[out] should_convert Whether the associated node's data layout should be converted to `target_data_layout`. |
||||
* If greater than 0, convert. |
||||
* If 0, don't convert. |
||||
* Otherwise, if less than 0, leave the decision to ORT. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(ShouldConvertDataLayoutForOp, _In_ OrtEp* this_ptr, |
||||
_In_z_ const char* domain, _In_z_ const char* op_type, |
||||
_In_ OrtEpDataLayout target_data_layout, |
||||
_Outptr_ int* should_convert); |
||||
|
||||
/** \brief Set dynamic options on this EP.
|
||||
* |
||||
* Dynamic options can be set by the user at any time after session creation with `OrtApi::SetEpDynamicOptions()`. |
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[in] option_keys The dynamic option keys. |
||||
* \param[in] option_values The dynamic option values. |
||||
* \param[in] num_options The number of dynamic options. |
||||
* |
||||
* \note Implementation of this function is optional. |
||||
* An EP should only implement this if it needs to handle any dynamic options. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(SetDynamicOptions, _In_ OrtEp* this_ptr, |
||||
_In_reads_(num_options) const char* const* option_keys, |
||||
_In_reads_(num_options) const char* const* option_values, |
||||
_In_ size_t num_options); |
||||
|
||||
/** \brief Called by ORT to notify the EP of the start of a run.
|
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[in] run_options The run options for this run. |
||||
* |
||||
* \note Implementation of this function is optional. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(OnRunStart, _In_ OrtEp* this_ptr, _In_ const OrtRunOptions* run_options); |
||||
|
||||
/** \brief Called by ORT to notify the EP of the end of a run.
|
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[in] run_options The run options for this run. |
||||
* \param[in] sync_stream Whether any associated stream should be synchronized during this call. |
||||
* Only applicable if there is such a stream. |
||||
* |
||||
* \note Implementation of this function is optional. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(OnRunEnd, _In_ OrtEp* this_ptr, _In_ const OrtRunOptions* run_options, _In_ bool sync_stream); |
||||
|
||||
/** \brief Create an OrtAllocator for the given OrtMemoryInfo for an OrtSession.
|
||||
* |
||||
* The OrtMemoryInfo instance will match one of the values set in the OrtEpDevice using EpDevice_AddAllocatorInfo. |
||||
* Any allocator specific options should be read from the session options. |
||||
* |
||||
* If nullptr OrtEpFactory::CreateAllocator will be used. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[in] memory_info The OrtMemoryInfo to create the allocator for. May be nullptr. |
||||
* \param[out] allocator The created OrtAllocator instance. Set to nullptr if the default CPU allocator is used. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(CreateAllocator, _In_ OrtEp* this_ptr, |
||||
_In_ const OrtMemoryInfo* memory_info, |
||||
_Outptr_result_maybenull_ OrtAllocator** allocator); |
||||
|
||||
/** \brief Create a synchronization stream for the given memory device for an OrtSession.
|
||||
* |
||||
* This is used to create a synchronization stream for the execution provider and is used to synchronize |
||||
* operations on the device during model execution. |
||||
* Any stream specific options should be read from the session options. |
||||
* |
||||
* If nullptr OrtEpFactory::CreateSyncStreamForDevice will be used. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[in] memory_device The OrtMemoryDevice to create the synchronization stream for. |
||||
* \param[out] stream The created OrtSyncStreamImpl instance. nullptr if the execution provider is not stream aware. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(CreateSyncStreamForDevice, _In_ OrtEp* this_ptr, |
||||
_In_ const OrtMemoryDevice* memory_device, |
||||
_Outptr_ OrtSyncStreamImpl** stream); |
||||
|
||||
/** \brief Get a string with details about the EP stack used to produce a compiled model.
|
||||
* |
||||
* This function gets a compatibility information string that contains details about the execution provider |
||||
* used to compile a given model. This string can later be used with ValidateCompiledModelCompatibilityInfo |
||||
* to determine if a compiled model is compatible with the EP. |
||||
* |
||||
* The returned string should be a null-terminated, UTF-8 encoded string. ORT will copy it. |
||||
* |
||||
* \param[in] this_ptr The OrtEp instance. |
||||
* \param[in] graph The OrtGraph instance for which to generate compatibility information. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(const char*, GetCompiledModelCompatibilityInfo, _In_ OrtEp* this_ptr, |
||||
_In_ const OrtGraph* graph); |
||||
}; |
||||
|
||||
/** \brief The function signature that ORT will call to create OrtEpFactory instances.
|
||||
* |
||||
* This must be available in a function called 'CreateEpFactories' in the execution provider library. |
||||
* |
||||
* \param[in] registered_name The name the execution library is registered with by RegisterExecutionProviderLibrary |
||||
* \param[in] ort_api_base The OrtApiBase instance that is used by the factory to get the OrtApi instance for the |
||||
* version of ORT that the library was compiled against. |
||||
* \param[in] default_logger The default ORT logger that can be used for logging outside of an inference session. |
||||
* \param[in,out] factories The implementation should create and add OrtEpFactory instances to this |
||||
* pre-allocated array. |
||||
* i.e. usage is `factories[0] = new MyEpFactory();` |
||||
* \param[in] max_factories The maximum number of OrtEpFactory instances that can be added to `factories`. |
||||
* Current default is to allow 4 factories. This can be increased in the future if needed. |
||||
* \param[out] num_factories The number of OrtEpFactory instances created by the factory and added to `factories`. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
typedef OrtStatus* (*CreateEpApiFactoriesFn)(_In_ const char* registered_name, _In_ const OrtApiBase* ort_api_base, |
||||
_In_ const OrtLogger* default_logger, |
||||
_Inout_ OrtEpFactory** factories, _In_ size_t max_factories, |
||||
_Out_ size_t* num_factories); |
||||
|
||||
/** \brief The function signature that ORT will call to release an OrtEpFactory instance.
|
||||
* |
||||
* This must be available in a function called 'ReleaseEpFactory' in the execution provider library. |
||||
* |
||||
* \param[in] factory The OrtEpFactory instance to release. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
typedef OrtStatus* (*ReleaseEpApiFactoryFn)(_In_ OrtEpFactory* factory); |
||||
|
||||
/**
|
||||
* \brief The OrtEpFactory provides functions to create and manage execution providers. |
||||
* \since Version 1.22. |
||||
*/ |
||||
struct OrtEpFactory { |
||||
/** \brief The ONNX Runtime version the execution provider was compiled with.
|
||||
* |
||||
* Implementation should set to ORT_API_VERSION. |
||||
* ORT will use this to ensure it does not call functions that were not available when the library was compiled. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
uint32_t ort_version_supported; |
||||
|
||||
/** \brief Get the name of the execution provider that the factory creates.
|
||||
* |
||||
* The returned string should be a null-terminated, UTF-8 encoded string. ORT will copy it. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \return The name of the execution provider the factory creates. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
ORT_API_T(const char*, GetName, const OrtEpFactory* this_ptr); |
||||
|
||||
/** \brief Get the name of vendor who owns the execution provider that the factory creates.
|
||||
* |
||||
* The returned string should be a null-terminated, UTF-8 encoded string. ORT will copy it. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \return vendor The vendor name of the execution provider the factory creates. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
ORT_API_T(const char*, GetVendor, const OrtEpFactory* this_ptr); // return EP vendor
|
||||
|
||||
/** \brief Get information from the execution provider about OrtHardwareDevice support.
|
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* Non-const as the factory is passed through to the CreateEp call via the OrtEpDevice. |
||||
* \param[in] devices The OrtHardwareDevice instances that are available. |
||||
* \param[in] num_devices The number of OrtHardwareDevice instances. |
||||
* \param[out] ep_devices OrtEpDevice instances for each OrtHardwareDevice that the EP can use. |
||||
* The implementation should call OrtEpApi::CreateEpDevice to create, and add the OrtEpDevice |
||||
* instances to this pre-allocated array. ORT will take ownership of the values returned. |
||||
* i.e. usage is `ep_devices[0] = <ptr to OrtEpDevice created with OrtEpApi::CreateEpDevice>;` |
||||
* \param[in] max_ep_devices The maximum number of OrtEpDevices that can be added to ep_devices. |
||||
* Current default is 8. This can be increased if needed. |
||||
* \param[out] num_ep_devices The number of EP devices added to ep_devices. |
||||
* \return true if the factory can create an execution provider that uses `device`. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
ORT_API2_STATUS(GetSupportedDevices, _In_ OrtEpFactory* this_ptr, |
||||
_In_reads_(num_devices) const OrtHardwareDevice* const* devices, |
||||
_In_ size_t num_devices, |
||||
_Inout_ OrtEpDevice** ep_devices, |
||||
_In_ size_t max_ep_devices, |
||||
_Out_ size_t* num_ep_devices); |
||||
|
||||
/** \brief Function to create an OrtEp instance for use in a Session.
|
||||
* |
||||
* ORT will call ReleaseEp to release the instance when it is no longer needed. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[in] devices The OrtHardwareDevice instances that the execution provider was selected to use. |
||||
* May be a subset of the OrtHardwareDevice instances that the execution provider's factory |
||||
* set as supported in the call to OrtEpFactory::GetSupportedDevices. |
||||
* \param[in] ep_metadata_pairs Execution provider metadata that was provided to OrtEpApi::CreateEpDevice, for each |
||||
* device. |
||||
* \param[in] num_devices The number of devices the execution provider was selected for. |
||||
* \param[in] session_options The OrtSessionOptions instance that contains the configuration options for the |
||||
* session. This will include ep_options from GetSupportedDevices as well as any |
||||
* user provided overrides. |
||||
* Execution provider options will have been added with a prefix of 'ep.[ep name].'. |
||||
* The OrtSessionOptions instance will NOT be valid after this call and should not be |
||||
* stored for later use. |
||||
* \param[in] logger The OrtLogger instance for the session that the execution provider should use for logging. |
||||
* \param[out] ep The OrtEp instance created by the factory. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
ORT_API2_STATUS(CreateEp, _In_ OrtEpFactory* this_ptr, |
||||
_In_reads_(num_devices) const OrtHardwareDevice* const* devices, |
||||
_In_reads_(num_devices) const OrtKeyValuePairs* const* ep_metadata_pairs, |
||||
_In_ size_t num_devices, |
||||
_In_ const OrtSessionOptions* session_options, |
||||
_In_ const OrtLogger* logger, _Outptr_ OrtEp** ep); |
||||
|
||||
/** \brief Release the OrtEp instance.
|
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[in] ep The OrtEp instance to release. |
||||
* |
||||
* \since Version 1.22. |
||||
*/ |
||||
ORT_API_T(void, ReleaseEp, OrtEpFactory* this_ptr, struct OrtEp* ep); |
||||
|
||||
/** \brief Get the vendor id who owns the execution provider that the factory creates.
|
||||
* |
||||
* This is typically the PCI vendor ID. See https://pcisig.com/membership/member-companies
|
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \return vendor_id The vendor ID of the execution provider the factory creates. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(uint32_t, GetVendorId, const OrtEpFactory* this_ptr); |
||||
|
||||
/** \brief Get the version of the execution provider that the factory creates.
|
||||
* |
||||
* The version string should adhere to the Semantic Versioning 2.0 specification |
||||
* (https://github.com/semver/semver/blob/v2.0.0/semver.md).
|
||||
* |
||||
* The returned string should be a null-terminated, UTF-8 encoded string. ORT will copy it. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \return The execution provider version string. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(const char*, GetVersion, _In_ const OrtEpFactory* this_ptr); |
||||
|
||||
/** \brief Validate the compatibility of a compiled model with the execution provider factory for one or more devices.
|
||||
* |
||||
* Given a compatibility info string produced during model compilation, the EP factory should determine whether the |
||||
* compiled model is compatible with the EP factory when targeting the provided hardware devices. All devices provided |
||||
* must belong to the same execution provider instance that this factory creates. |
||||
* |
||||
* The EP factory implementation should consider the set of devices (e.g., multi-adapter or multi-GPU scenarios) when |
||||
* evaluating compatibility and set `model_compatibility` accordingly. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[in] devices Array of OrtHardwareDevice pointers that the EP would run on. All must map to this EP. |
||||
* \param[in] num_devices Number of entries in `devices`. |
||||
* \param[in] compatibility_info The compatibility information string produced when the model was compiled. |
||||
* \param[out] model_compatibility OrtCompiledModelCompatibility value describing the compatibility of the model with the EP. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(ValidateCompiledModelCompatibilityInfo, _In_ OrtEpFactory* this_ptr, |
||||
_In_reads_(num_devices) const OrtHardwareDevice* const* devices, |
||||
_In_ size_t num_devices, |
||||
_In_ const char* compatibility_info, |
||||
_Out_ OrtCompiledModelCompatibility* model_compatibility); |
||||
|
||||
/** \brief Create an OrtAllocator that can be shared across sessions for the given OrtMemoryInfo.
|
||||
* |
||||
* The factory that creates the EP is responsible for providing the allocators required by the EP. |
||||
* The OrtMemoryInfo instance will match one of the values set in the OrtEpDevice using EpDevice_AddAllocatorInfo. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[in] memory_info The OrtMemoryInfo to create the allocator for. May be nullptr. |
||||
* \param[in] allocator_options Optional key-value pairs for allocator options, can be nullptr. |
||||
* \param[out] allocator The created OrtAllocator instance. Set to nullptr if the default CPU allocator is used. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(CreateAllocator, _In_ OrtEpFactory* this_ptr, |
||||
_In_ const OrtMemoryInfo* memory_info, |
||||
_In_opt_ const OrtKeyValuePairs* allocator_options, |
||||
_Outptr_result_maybenull_ OrtAllocator** allocator); |
||||
|
||||
/** \brief Release an OrtAllocator created by the factory.
|
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(void, ReleaseAllocator, _In_ OrtEpFactory* this_ptr, _In_ OrtAllocator* allocator); |
||||
|
||||
/** \brief Create an OrtDataTransferImpl instance for the factory.
|
||||
* |
||||
* This is used to create an IDataTransfer implementation that can be used to copy data between devices |
||||
* that the execution provider supports. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[out] data_transfer The created OrtDataTransferImpl instance. Set to nullptr if not required. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(CreateDataTransfer, _In_ OrtEpFactory* this_ptr, |
||||
_Outptr_result_maybenull_ OrtDataTransferImpl** data_transfer); |
||||
|
||||
/** \brief Check if execution providers created by the factory are stream aware.
|
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \return True if the factory creates execution providers that are stream aware and it implements CreateSyncStreamForDevice. |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API_T(bool, IsStreamAware, _In_ const OrtEpFactory* this_ptr); |
||||
|
||||
/** \brief Create a synchronization stream for the given memory device.
|
||||
* |
||||
* This is used to create a synchronization stream for the memory device that can be used for operations outside of |
||||
* a session. |
||||
* |
||||
* \param[in] this_ptr The OrtEpFactory instance. |
||||
* \param[in] memory_device The OrtMemoryDevice to create the synchronization stream for. |
||||
* \param[in] stream_options Options for stream creation. May be nullptr. |
||||
* \param[out] stream The created OrtSyncStreamImpl instance. nullptr if the execution provider is not stream aware. |
||||
* |
||||
* \snippet{doc} snippets.dox OrtStatus Return Value |
||||
* |
||||
* \since Version 1.23. |
||||
*/ |
||||
ORT_API2_STATUS(CreateSyncStreamForDevice, _In_ OrtEpFactory* this_ptr, |
||||
_In_ const OrtMemoryDevice* memory_device, |
||||
_In_opt_ const OrtKeyValuePairs* stream_options, |
||||
_Outptr_ OrtSyncStreamImpl** stream); |
||||
}; |
||||
|
||||
#ifdef __cplusplus |
||||
} |
||||
#endif |
||||
@ -0,0 +1,18 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once |
||||
|
||||
// This file contains well-known keys for OrtEpDevice EP metadata entries.
|
||||
// It does NOT specify all available metadata keys.
|
||||
|
||||
// Key for the execution provider version string. This should be available for all plugin EPs.
|
||||
static const char* const kOrtEpDevice_EpMetadataKey_Version = "version"; |
||||
|
||||
// Prefix for execution provider compatibility information stored in model metadata.
|
||||
// Used when generating EP context models to store compatibility strings for each EP.
|
||||
// Full key format: "ep_compatibility_info.<EP_TYPE>"
|
||||
static const char* const kOrtModelMetadata_EpCompatibilityInfoPrefix = "ep_compatibility_info."; |
||||
|
||||
// Key for the execution provider library path (for dynamically loaded EPs)
|
||||
static const char* const kOrtEpDevice_EpMetadataKey_LibraryPath = "library_path"; |
||||
@ -0,0 +1,535 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once |
||||
|
||||
#include <stdint.h> |
||||
#include <cmath> |
||||
#include <cstring> |
||||
#include <limits> |
||||
|
||||
namespace onnxruntime_float16 { |
||||
|
||||
namespace detail { |
||||
|
||||
enum class endian { |
||||
#if defined(_WIN32) |
||||
little = 0, |
||||
big = 1, |
||||
native = little, |
||||
#elif defined(__GNUC__) || defined(__clang__) |
||||
little = __ORDER_LITTLE_ENDIAN__, |
||||
big = __ORDER_BIG_ENDIAN__, |
||||
native = __BYTE_ORDER__, |
||||
#else |
||||
#error onnxruntime_float16::detail::endian is not implemented in this environment. |
||||
#endif |
||||
}; |
||||
|
||||
static_assert( |
||||
endian::native == endian::little || endian::native == endian::big, |
||||
"Only little-endian or big-endian native byte orders are supported."); |
||||
|
||||
} // namespace detail
|
||||
|
||||
/// <summary>
|
||||
/// Shared implementation between public and internal classes. CRTP pattern.
|
||||
/// </summary>
|
||||
template <class Derived> |
||||
struct Float16Impl { |
||||
protected: |
||||
/// <summary>
|
||||
/// Converts from float to uint16_t float16 representation
|
||||
/// </summary>
|
||||
/// <param name="v"></param>
|
||||
/// <returns></returns>
|
||||
constexpr static uint16_t ToUint16Impl(float v) noexcept; |
||||
|
||||
/// <summary>
|
||||
/// Converts float16 to float
|
||||
/// </summary>
|
||||
/// <returns>float representation of float16 value</returns>
|
||||
float ToFloatImpl() const noexcept; |
||||
|
||||
/// <summary>
|
||||
/// Creates an instance that represents absolute value.
|
||||
/// </summary>
|
||||
/// <returns>Absolute value</returns>
|
||||
uint16_t AbsImpl() const noexcept { |
||||
return static_cast<uint16_t>(val & ~kSignMask); |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Creates a new instance with the sign flipped.
|
||||
/// </summary>
|
||||
/// <returns>Flipped sign instance</returns>
|
||||
uint16_t NegateImpl() const noexcept { |
||||
return IsNaN() ? val : static_cast<uint16_t>(val ^ kSignMask); |
||||
} |
||||
|
||||
public: |
||||
// uint16_t special values
|
||||
static constexpr uint16_t kSignMask = 0x8000U; |
||||
static constexpr uint16_t kBiasedExponentMask = 0x7C00U; |
||||
static constexpr uint16_t kPositiveInfinityBits = 0x7C00U; |
||||
static constexpr uint16_t kNegativeInfinityBits = 0xFC00U; |
||||
static constexpr uint16_t kPositiveQNaNBits = 0x7E00U; |
||||
static constexpr uint16_t kNegativeQNaNBits = 0xFE00U; |
||||
static constexpr uint16_t kMaxValueBits = 0x7BFFU; // Largest normal number
|
||||
static constexpr uint16_t kOneBits = 0x3C00U; |
||||
static constexpr uint16_t kMinusOneBits = 0xBC00U; |
||||
|
||||
uint16_t val{0}; |
||||
|
||||
Float16Impl() = default; |
||||
|
||||
/// <summary>
|
||||
/// Checks if the value is negative
|
||||
/// </summary>
|
||||
/// <returns>true if negative</returns>
|
||||
bool IsNegative() const noexcept { |
||||
return static_cast<int16_t>(val) < 0; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is NaN
|
||||
/// </summary>
|
||||
/// <returns>true if NaN</returns>
|
||||
bool IsNaN() const noexcept { |
||||
return AbsImpl() > kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is finite
|
||||
/// </summary>
|
||||
/// <returns>true if finite</returns>
|
||||
bool IsFinite() const noexcept { |
||||
return AbsImpl() < kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value represents positive infinity.
|
||||
/// </summary>
|
||||
/// <returns>true if positive infinity</returns>
|
||||
bool IsPositiveInfinity() const noexcept { |
||||
return val == kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value represents negative infinity
|
||||
/// </summary>
|
||||
/// <returns>true if negative infinity</returns>
|
||||
bool IsNegativeInfinity() const noexcept { |
||||
return val == kNegativeInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is either positive or negative infinity.
|
||||
/// </summary>
|
||||
/// <returns>True if absolute value is infinity</returns>
|
||||
bool IsInfinity() const noexcept { |
||||
return AbsImpl() == kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is NaN or zero. Useful for comparisons.
|
||||
/// </summary>
|
||||
/// <returns>True if NaN or zero.</returns>
|
||||
bool IsNaNOrZero() const noexcept { |
||||
auto abs = AbsImpl(); |
||||
return (abs == 0 || abs > kPositiveInfinityBits); |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is normal (not zero, subnormal, infinite, or NaN).
|
||||
/// </summary>
|
||||
/// <returns>True if so</returns>
|
||||
bool IsNormal() const noexcept { |
||||
auto abs = AbsImpl(); |
||||
return (abs < kPositiveInfinityBits) // is finite
|
||||
&& (abs != 0) // is not zero
|
||||
&& ((abs & kBiasedExponentMask) != 0); // is not subnormal (has a non-zero exponent)
|
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is subnormal (denormal).
|
||||
/// </summary>
|
||||
/// <returns>True if so</returns>
|
||||
bool IsSubnormal() const noexcept { |
||||
auto abs = AbsImpl(); |
||||
return (abs < kPositiveInfinityBits) // is finite
|
||||
&& (abs != 0) // is not zero
|
||||
&& ((abs & kBiasedExponentMask) == 0); // is subnormal (has a zero exponent)
|
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Creates an instance that represents absolute value.
|
||||
/// </summary>
|
||||
/// <returns>Absolute value</returns>
|
||||
Derived Abs() const noexcept { return Derived::FromBits(AbsImpl()); } |
||||
|
||||
/// <summary>
|
||||
/// Creates a new instance with the sign flipped.
|
||||
/// </summary>
|
||||
/// <returns>Flipped sign instance</returns>
|
||||
Derived Negate() const noexcept { return Derived::FromBits(NegateImpl()); } |
||||
|
||||
/// <summary>
|
||||
/// IEEE defines that positive and negative zero are equal, this gives us a quick equality check
|
||||
/// for two values by or'ing the private bits together and stripping the sign. They are both zero,
|
||||
/// and therefore equivalent, if the resulting value is still zero.
|
||||
/// </summary>
|
||||
/// <param name="lhs">first value</param>
|
||||
/// <param name="rhs">second value</param>
|
||||
/// <returns>True if both arguments represent zero</returns>
|
||||
static bool AreZero(const Float16Impl& lhs, const Float16Impl& rhs) noexcept { |
||||
return static_cast<uint16_t>((lhs.val | rhs.val) & ~kSignMask) == 0; |
||||
} |
||||
|
||||
bool operator==(const Float16Impl& rhs) const noexcept { |
||||
if (IsNaN() || rhs.IsNaN()) { |
||||
// IEEE defines that NaN is not equal to anything, including itself.
|
||||
return false; |
||||
} |
||||
return val == rhs.val; |
||||
} |
||||
|
||||
bool operator!=(const Float16Impl& rhs) const noexcept { return !(*this == rhs); } |
||||
|
||||
bool operator<(const Float16Impl& rhs) const noexcept { |
||||
if (IsNaN() || rhs.IsNaN()) { |
||||
// IEEE defines that NaN is unordered with respect to everything, including itself.
|
||||
return false; |
||||
} |
||||
|
||||
const bool left_is_negative = IsNegative(); |
||||
if (left_is_negative != rhs.IsNegative()) { |
||||
// When the signs of left and right differ, we know that left is less than right if it is
|
||||
// the negative value. The exception to this is if both values are zero, in which case IEEE
|
||||
// says they should be equal, even if the signs differ.
|
||||
return left_is_negative && !AreZero(*this, rhs); |
||||
} |
||||
return (val != rhs.val) && ((val < rhs.val) ^ left_is_negative); |
||||
} |
||||
}; |
||||
|
||||
// The following Float16_t conversions are based on the code from
|
||||
// Eigen library.
|
||||
|
||||
// The conversion routines are Copyright (c) Fabian Giesen, 2016.
|
||||
// The original license follows:
|
||||
//
|
||||
// Copyright (c) Fabian Giesen, 2016
|
||||
// All rights reserved.
|
||||
// Redistribution and use in source and binary forms, with or without
|
||||
// modification, are permitted.
|
||||
// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
|
||||
// "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
|
||||
// LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR
|
||||
// A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT
|
||||
// HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL,
|
||||
// SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT
|
||||
// LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
|
||||
// DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY
|
||||
// THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
|
||||
// (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
// OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
|
||||
namespace detail { |
||||
union float32_bits { |
||||
unsigned int u; |
||||
float f; |
||||
}; |
||||
} // namespace detail
|
||||
|
||||
template <class Derived> |
||||
inline constexpr uint16_t Float16Impl<Derived>::ToUint16Impl(float v) noexcept { |
||||
detail::float32_bits f{}; |
||||
f.f = v; |
||||
|
||||
constexpr detail::float32_bits f32infty = {255 << 23}; |
||||
constexpr detail::float32_bits f16max = {(127 + 16) << 23}; |
||||
constexpr detail::float32_bits denorm_magic = {((127 - 15) + (23 - 10) + 1) << 23}; |
||||
constexpr unsigned int sign_mask = 0x80000000u; |
||||
uint16_t val = static_cast<uint16_t>(0x0u); |
||||
|
||||
unsigned int sign = f.u & sign_mask; |
||||
f.u ^= sign; |
||||
|
||||
// NOTE all the integer compares in this function can be safely
|
||||
// compiled into signed compares since all operands are below
|
||||
// 0x80000000. Important if you want fast straight SSE2 code
|
||||
// (since there's no unsigned PCMPGTD).
|
||||
|
||||
if (f.u >= f16max.u) { // result is Inf or NaN (all exponent bits set)
|
||||
val = (f.u > f32infty.u) ? 0x7e00 : 0x7c00; // NaN->qNaN and Inf->Inf
|
||||
} else { // (De)normalized number or zero
|
||||
if (f.u < (113 << 23)) { // resulting FP16 is subnormal or zero
|
||||
// use a magic value to align our 10 mantissa bits at the bottom of
|
||||
// the float. as long as FP addition is round-to-nearest-even this
|
||||
// just works.
|
||||
f.f += denorm_magic.f; |
||||
|
||||
// and one integer subtract of the bias later, we have our final float!
|
||||
val = static_cast<uint16_t>(f.u - denorm_magic.u); |
||||
} else { |
||||
unsigned int mant_odd = (f.u >> 13) & 1; // resulting mantissa is odd
|
||||
|
||||
// update exponent, rounding bias part 1
|
||||
// Equivalent to `f.u += ((unsigned int)(15 - 127) << 23) + 0xfff`, but
|
||||
// without arithmetic overflow.
|
||||
f.u += 0xc8000fffU; |
||||
// rounding bias part 2
|
||||
f.u += mant_odd; |
||||
// take the bits!
|
||||
val = static_cast<uint16_t>(f.u >> 13); |
||||
} |
||||
} |
||||
|
||||
val |= static_cast<uint16_t>(sign >> 16); |
||||
return val; |
||||
} |
||||
|
||||
template <class Derived> |
||||
inline float Float16Impl<Derived>::ToFloatImpl() const noexcept { |
||||
constexpr detail::float32_bits magic = {113 << 23}; |
||||
constexpr unsigned int shifted_exp = 0x7c00 << 13; // exponent mask after shift
|
||||
detail::float32_bits o{}; |
||||
|
||||
o.u = (val & 0x7fff) << 13; // exponent/mantissa bits
|
||||
unsigned int exp = shifted_exp & o.u; // just the exponent
|
||||
o.u += (127 - 15) << 23; // exponent adjust
|
||||
|
||||
// handle exponent special cases
|
||||
if (exp == shifted_exp) { // Inf/NaN?
|
||||
o.u += (128 - 16) << 23; // extra exp adjust
|
||||
} else if (exp == 0) { // Zero/Denormal?
|
||||
o.u += 1 << 23; // extra exp adjust
|
||||
o.f -= magic.f; // re-normalize
|
||||
} |
||||
|
||||
// Attempt to workaround the Internal Compiler Error on ARM64
|
||||
// for bitwise | operator, including std::bitset
|
||||
#if (defined _MSC_VER) && (defined _M_ARM || defined _M_ARM64 || defined _M_ARM64EC) |
||||
if (IsNegative()) { |
||||
return -o.f; |
||||
} |
||||
#else |
||||
// original code:
|
||||
o.u |= (val & 0x8000U) << 16U; // sign bit
|
||||
#endif |
||||
return o.f; |
||||
} |
||||
|
||||
/// Shared implementation between public and internal classes. CRTP pattern.
|
||||
template <class Derived> |
||||
struct BFloat16Impl { |
||||
protected: |
||||
/// <summary>
|
||||
/// Converts from float to uint16_t float16 representation
|
||||
/// </summary>
|
||||
/// <param name="v"></param>
|
||||
/// <returns></returns>
|
||||
static uint16_t ToUint16Impl(float v) noexcept; |
||||
|
||||
/// <summary>
|
||||
/// Converts bfloat16 to float
|
||||
/// </summary>
|
||||
/// <returns>float representation of bfloat16 value</returns>
|
||||
float ToFloatImpl() const noexcept; |
||||
|
||||
/// <summary>
|
||||
/// Creates an instance that represents absolute value.
|
||||
/// </summary>
|
||||
/// <returns>Absolute value</returns>
|
||||
uint16_t AbsImpl() const noexcept { |
||||
return static_cast<uint16_t>(val & ~kSignMask); |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Creates a new instance with the sign flipped.
|
||||
/// </summary>
|
||||
/// <returns>Flipped sign instance</returns>
|
||||
uint16_t NegateImpl() const noexcept { |
||||
return IsNaN() ? val : static_cast<uint16_t>(val ^ kSignMask); |
||||
} |
||||
|
||||
public: |
||||
// uint16_t special values
|
||||
static constexpr uint16_t kSignMask = 0x8000U; |
||||
static constexpr uint16_t kBiasedExponentMask = 0x7F80U; |
||||
static constexpr uint16_t kPositiveInfinityBits = 0x7F80U; |
||||
static constexpr uint16_t kNegativeInfinityBits = 0xFF80U; |
||||
static constexpr uint16_t kPositiveQNaNBits = 0x7FC1U; |
||||
static constexpr uint16_t kNegativeQNaNBits = 0xFFC1U; |
||||
static constexpr uint16_t kMaxValueBits = 0x7F7FU; |
||||
static constexpr uint16_t kRoundToNearest = 0x7FFFU; |
||||
static constexpr uint16_t kOneBits = 0x3F80U; |
||||
static constexpr uint16_t kMinusOneBits = 0xBF80U; |
||||
|
||||
uint16_t val{0}; |
||||
|
||||
BFloat16Impl() = default; |
||||
|
||||
/// <summary>
|
||||
/// Checks if the value is negative
|
||||
/// </summary>
|
||||
/// <returns>true if negative</returns>
|
||||
bool IsNegative() const noexcept { |
||||
return static_cast<int16_t>(val) < 0; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is NaN
|
||||
/// </summary>
|
||||
/// <returns>true if NaN</returns>
|
||||
bool IsNaN() const noexcept { |
||||
return AbsImpl() > kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is finite
|
||||
/// </summary>
|
||||
/// <returns>true if finite</returns>
|
||||
bool IsFinite() const noexcept { |
||||
return AbsImpl() < kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value represents positive infinity.
|
||||
/// </summary>
|
||||
/// <returns>true if positive infinity</returns>
|
||||
bool IsPositiveInfinity() const noexcept { |
||||
return val == kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value represents negative infinity
|
||||
/// </summary>
|
||||
/// <returns>true if negative infinity</returns>
|
||||
bool IsNegativeInfinity() const noexcept { |
||||
return val == kNegativeInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is either positive or negative infinity.
|
||||
/// </summary>
|
||||
/// <returns>True if absolute value is infinity</returns>
|
||||
bool IsInfinity() const noexcept { |
||||
return AbsImpl() == kPositiveInfinityBits; |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is NaN or zero. Useful for comparisons.
|
||||
/// </summary>
|
||||
/// <returns>True if NaN or zero.</returns>
|
||||
bool IsNaNOrZero() const noexcept { |
||||
auto abs = AbsImpl(); |
||||
return (abs == 0 || abs > kPositiveInfinityBits); |
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is normal (not zero, subnormal, infinite, or NaN).
|
||||
/// </summary>
|
||||
/// <returns>True if so</returns>
|
||||
bool IsNormal() const noexcept { |
||||
auto abs = AbsImpl(); |
||||
return (abs < kPositiveInfinityBits) // is finite
|
||||
&& (abs != 0) // is not zero
|
||||
&& ((abs & kBiasedExponentMask) != 0); // is not subnormal (has a non-zero exponent)
|
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Tests if the value is subnormal (denormal).
|
||||
/// </summary>
|
||||
/// <returns>True if so</returns>
|
||||
bool IsSubnormal() const noexcept { |
||||
auto abs = AbsImpl(); |
||||
return (abs < kPositiveInfinityBits) // is finite
|
||||
&& (abs != 0) // is not zero
|
||||
&& ((abs & kBiasedExponentMask) == 0); // is subnormal (has a zero exponent)
|
||||
} |
||||
|
||||
/// <summary>
|
||||
/// Creates an instance that represents absolute value.
|
||||
/// </summary>
|
||||
/// <returns>Absolute value</returns>
|
||||
Derived Abs() const noexcept { return Derived::FromBits(AbsImpl()); } |
||||
|
||||
/// <summary>
|
||||
/// Creates a new instance with the sign flipped.
|
||||
/// </summary>
|
||||
/// <returns>Flipped sign instance</returns>
|
||||
Derived Negate() const noexcept { return Derived::FromBits(NegateImpl()); } |
||||
|
||||
/// <summary>
|
||||
/// IEEE defines that positive and negative zero are equal, this gives us a quick equality check
|
||||
/// for two values by or'ing the private bits together and stripping the sign. They are both zero,
|
||||
/// and therefore equivalent, if the resulting value is still zero.
|
||||
/// </summary>
|
||||
/// <param name="lhs">first value</param>
|
||||
/// <param name="rhs">second value</param>
|
||||
/// <returns>True if both arguments represent zero</returns>
|
||||
static bool AreZero(const BFloat16Impl& lhs, const BFloat16Impl& rhs) noexcept { |
||||
// IEEE defines that positive and negative zero are equal, this gives us a quick equality check
|
||||
// for two values by or'ing the private bits together and stripping the sign. They are both zero,
|
||||
// and therefore equivalent, if the resulting value is still zero.
|
||||
return static_cast<uint16_t>((lhs.val | rhs.val) & ~kSignMask) == 0; |
||||
} |
||||
}; |
||||
|
||||
template <class Derived> |
||||
inline uint16_t BFloat16Impl<Derived>::ToUint16Impl(float v) noexcept { |
||||
uint16_t result; |
||||
if (std::isnan(v)) { |
||||
result = kPositiveQNaNBits; |
||||
} else { |
||||
auto get_msb_half = [](float fl) { |
||||
uint16_t result; |
||||
#ifdef __cpp_if_constexpr |
||||
if constexpr (detail::endian::native == detail::endian::little) { |
||||
#else |
||||
if (detail::endian::native == detail::endian::little) { |
||||
#endif |
||||
std::memcpy(&result, reinterpret_cast<char*>(&fl) + sizeof(uint16_t), sizeof(uint16_t)); |
||||
} else { |
||||
std::memcpy(&result, &fl, sizeof(uint16_t)); |
||||
} |
||||
return result; |
||||
}; |
||||
|
||||
uint16_t upper_bits = get_msb_half(v); |
||||
union { |
||||
uint32_t U32; |
||||
float F32; |
||||
}; |
||||
F32 = v; |
||||
U32 += (upper_bits & 1) + kRoundToNearest; |
||||
result = get_msb_half(F32); |
||||
} |
||||
return result; |
||||
} |
||||
|
||||
template <class Derived> |
||||
inline float BFloat16Impl<Derived>::ToFloatImpl() const noexcept { |
||||
if (IsNaN()) { |
||||
return std::numeric_limits<float>::quiet_NaN(); |
||||
} |
||||
float result; |
||||
char* const first = reinterpret_cast<char*>(&result); |
||||
char* const second = first + sizeof(uint16_t); |
||||
#ifdef __cpp_if_constexpr |
||||
if constexpr (detail::endian::native == detail::endian::little) { |
||||
#else |
||||
if (detail::endian::native == detail::endian::little) { |
||||
#endif |
||||
std::memset(first, 0, sizeof(uint16_t)); |
||||
std::memcpy(second, &val, sizeof(uint16_t)); |
||||
} else { |
||||
std::memcpy(first, &val, sizeof(uint16_t)); |
||||
std::memset(second, 0, sizeof(uint16_t)); |
||||
} |
||||
return result; |
||||
} |
||||
|
||||
} // namespace onnxruntime_float16
|
||||
File diff suppressed because it is too large
Load Diff
@ -0,0 +1,54 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once |
||||
|
||||
/*
|
||||
* This file defines RunOptions Config Keys and format of the Config Values. |
||||
* |
||||
* The Naming Convention for a RunOptions Config Key, |
||||
* "[Area][.[SubArea1].[SubArea2]...].[Keyname]" |
||||
* Such as "ep.cuda.use_arena" |
||||
* The Config Key cannot be empty |
||||
* The maximum length of the Config Key is 128 |
||||
* |
||||
* The string format of a RunOptions Config Value is defined individually for each Config. |
||||
* The maximum length of the Config Value is 1024 |
||||
*/ |
||||
|
||||
// Key for enabling shrinkages of user listed device memory arenas.
|
||||
// Expects a list of semi-colon separated key value pairs separated by colon in the following format:
|
||||
// "device_0:device_id_0;device_1:device_id_1"
|
||||
// No white-spaces allowed in the provided list string.
|
||||
// Currently, the only supported devices are : "cpu", "gpu" (case sensitive).
|
||||
// If "cpu" is included in the list, DisableCpuMemArena() API must not be called (i.e.) arena for cpu should be enabled.
|
||||
// Example usage: "cpu:0;gpu:0" (or) "gpu:0"
|
||||
// By default, the value for this key is empty (i.e.) no memory arenas are shrunk
|
||||
static const char* const kOrtRunOptionsConfigEnableMemoryArenaShrinkage = "memory.enable_memory_arena_shrinkage"; |
||||
|
||||
// Set to '1' to not synchronize execution providers with CPU at the end of session run.
|
||||
// Per default it will be set to '0'
|
||||
// Taking CUDA EP as an example, it omit triggering cudaStreamSynchronize on the compute stream.
|
||||
static const char* const kOrtRunOptionsConfigDisableSynchronizeExecutionProviders = "disable_synchronize_execution_providers"; |
||||
|
||||
// Set HTP performance mode for QNN HTP backend before session run.
|
||||
// options for HTP performance mode: "burst", "balanced", "default", "high_performance",
|
||||
// "high_power_saver", "low_balanced", "extreme_power_saver", "low_power_saver", "power_saver",
|
||||
// "sustained_high_performance". Default to "default".
|
||||
static const char* const kOrtRunOptionsConfigQnnPerfMode = "qnn.htp_perf_mode"; |
||||
|
||||
// Set HTP performance mode for QNN HTP backend post session run.
|
||||
static const char* const kOrtRunOptionsConfigQnnPerfModePostRun = "qnn.htp_perf_mode_post_run"; |
||||
|
||||
// Set RPC control latency for QNN HTP backend
|
||||
static const char* const kOrtRunOptionsConfigQnnRpcControlLatency = "qnn.rpc_control_latency"; |
||||
|
||||
// Set QNN Lora Config File for apply Lora in QNN context binary
|
||||
static const char* const kOrtRunOptionsConfigQnnLoraConfig = "qnn.lora_config"; |
||||
|
||||
// Set graph annotation id for CUDA EP. Use with enable_cuda_graph=true.
|
||||
// The value should be an integer. If the value is not set, the default value is 0 and
|
||||
// ORT session only captures one cuda graph before another capture is requested.
|
||||
// If the value is set to -1, cuda graph capture/replay is disabled in that run.
|
||||
// User are not expected to set the value to 0 as it is reserved for internal use.
|
||||
static const char* const kOrtRunOptionsConfigCudaGraphAnnotation = "gpu_graph_id"; |
||||
@ -0,0 +1,417 @@
|
||||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once |
||||
|
||||
/*
|
||||
* This file defines SessionOptions Config Keys and format of the Config Values. |
||||
* |
||||
* The Naming Convention for a SessionOptions Config Key, |
||||
* "[Area][.[SubArea1].[SubArea2]...].[Keyname]" |
||||
* Such as "ep.cuda.use_arena" |
||||
* The Config Key cannot be empty |
||||
* The maximum length of the Config Key is 1024 |
||||
* |
||||
* The string format of a SessionOptions Config Value is defined individually for each Config. |
||||
* The maximum length of the Config Value is 2048 |
||||
*/ |
||||
|
||||
// Key for disable PrePacking,
|
||||
// If the config value is set to "1" then the prepacking is disabled, otherwise prepacking is enabled (default value)
|
||||
static const char* const kOrtSessionOptionsConfigDisablePrepacking = "session.disable_prepacking"; |
||||
|
||||
// A value of "1" means allocators registered in the env will be used. "0" means the allocators created in the session
|
||||
// will be used. Use this to override the usage of env allocators on a per session level.
|
||||
static const char* const kOrtSessionOptionsConfigUseEnvAllocators = "session.use_env_allocators"; |
||||
|
||||
// Set to 'ORT' (case sensitive) to load an ORT format model.
|
||||
// If unset, model type will default to ONNX unless inferred from filename ('.ort' == ORT format) or bytes to be ORT
|
||||
static const char* const kOrtSessionOptionsConfigLoadModelFormat = "session.load_model_format"; |
||||
|
||||
// Set to 'ORT' (case sensitive) to save optimized model in ORT format when SessionOptions.optimized_model_path is set.
|
||||
// If unset, format will default to ONNX unless optimized_model_filepath ends in '.ort'.
|
||||
static const char* const kOrtSessionOptionsConfigSaveModelFormat = "session.save_model_format"; |
||||
|
||||
// If a value is "1", flush-to-zero and denormal-as-zero are applied. The default is "0".
|
||||
// When multiple sessions are created, a main thread doesn't override changes from succeeding session options,
|
||||
// but threads in session thread pools follow option changes.
|
||||
// When ORT runs with OpenMP, the same rule is applied, i.e. the first session option to flush-to-zero and
|
||||
// denormal-as-zero is only applied to global OpenMP thread pool, which doesn't support per-session thread pool.
|
||||
// Note that an alternative way not using this option at runtime is to train and export a model without denormals
|
||||
// and that's recommended because turning this option on may hurt model accuracy.
|
||||
static const char* const kOrtSessionOptionsConfigSetDenormalAsZero = "session.set_denormal_as_zero"; |
||||
|
||||
// It controls to run quantization model in QDQ (QuantizelinearDeQuantizelinear) format or not.
|
||||
// "0": enable. ORT does fusion logic for QDQ format.
|
||||
// "1": disable. ORT doesn't do fusion logic for QDQ format.
|
||||
// Its default value is "0" unless the DirectML execution provider is registered, in which case it defaults to "1".
|
||||
static const char* const kOrtSessionOptionsDisableQuantQDQ = "session.disable_quant_qdq"; |
||||
|
||||
// It controls whether to enable Double QDQ remover and Identical Children Consolidation
|
||||
// "0": not to disable. ORT does remove the middle 2 Nodes from a Q->(QD->Q)->QD pairs
|
||||
// "1": disable. ORT doesn't remove the middle 2 Nodes from a Q->(QD->Q)->QD pairs
|
||||
// Its default value is "0"
|
||||
static const char* const kOrtSessionOptionsDisableDoubleQDQRemover = "session.disable_double_qdq_remover"; |
||||
|
||||
// If set to "1", enables the removal of QuantizeLinear/DequantizeLinear node pairs once all QDQ handling has been
|
||||
// completed. e.g. If after all QDQ handling has completed and we have -> FloatOp -> Q -> DQ -> FloatOp -> the
|
||||
// Q -> DQ could potentially be removed. This will provide a performance benefit by avoiding going from float to
|
||||
// 8-bit and back to float, but could impact accuracy. The impact on accuracy will be model specific and depend on
|
||||
// other factors like whether the model was created using Quantization Aware Training or Post Training Quantization.
|
||||
// As such, it's best to test to determine if enabling this works well for your scenario.
|
||||
// The default value is "0"
|
||||
// Available since version 1.11.
|
||||
static const char* const kOrtSessionOptionsEnableQuantQDQCleanup = "session.enable_quant_qdq_cleanup"; |
||||
|
||||
// Enable or disable gelu approximation in graph optimization. "0": disable; "1": enable. The default is "0".
|
||||
// GeluApproximation has side effects which may change the inference results. It is disabled by default due to this.
|
||||
static const char* const kOrtSessionOptionsEnableGeluApproximation = "optimization.enable_gelu_approximation"; |
||||
|
||||
// Enable or disable Cast chain elimination in graph optimization. "0": disable; "1": enable. The default is "0".
|
||||
// CastElimination with chain elimination has side effects which may change the inference results. It is disabled by default due to this.
|
||||
static const char* const kOrtSessionOptionsEnableCastChainElimination = "optimization.enable_cast_chain_elimination"; |
||||
|
||||
// This setting controls whether to enable AheadOfTime function inlining.
|
||||
// AOT function inlining examines the graph and attempts to inline as many locally defined functions in the model
|
||||
// as possible with the help of enabled execution providers.
|
||||
// This can reduce the number of function calls and improve performance because it is done before
|
||||
// Level1 optimizers and constant folding. However, under some circumstances, when the EPs are not available,
|
||||
// one can disable the AOT inlining, produce an optimized model and postpone AOT until run time.
|
||||
// "0": enable; "1": disable.
|
||||
// Its default value is "0".
|
||||
static const char* const kOrtSessionOptionsDisableAheadOfTimeFunctionInlining = "session.disable_aot_function_inlining"; |
||||
|
||||
#ifdef ENABLE_TRAINING |
||||
// Specifies a path of the file containing a list of memory optimization configurations.
|
||||
// The value should be a string indicating the file path of the config file.
|
||||
// The content of the config file is a JSON struct like this:
|
||||
// [
|
||||
// "Gelu+Cast+:1:0",
|
||||
// "Dropout+:1:1"
|
||||
// ]
|
||||
// Taking the example of "Gelu+Cast+:1:0",
|
||||
// > "Gelu+Cast+" is the subgraph string, a valid "subgraph string" should be one subgraph representation
|
||||
// output by ORT graph transformations.
|
||||
// > "1" is "optimization strategy", valid values: 0 - disabled, 1 - recompute.
|
||||
// > "0" is "number of subgraph to apply" which is used to control how many subgraphs to apply optimization,
|
||||
// to avoid "oversaving" the memory.
|
||||
static const char* const kOrtSessionOptionsMemoryOptimizerApplyConfig = "optimization.memory_optimizer_config"; |
||||
|
||||
// Specifies the config for detecting subgraphs for memory footprint reduction.
|
||||
// The value should be a string contains int separated using commas. The default value is "0:0".
|
||||
static const char* const kOrtSessionOptionsMemoryOptimizerProbeConfig = "optimization.enable_memory_probe_recompute_config"; |
||||
#endif |
||||
|
||||
// This setting if set should contain a comma separated list of optimizers names that should be disabled.
|
||||
// Optimizers may take time to execute and affect model loading time. If you feel that a specific optimizer
|
||||
// does not provider runtime benefits, but affects your model loading time you may disable it using this config
|
||||
// entry. This option is not enabled in ORT_MINIMAL_BUILD build.
|
||||
// A list of optimizes is available in onnxruntime/core/optimizer/graph_transformer_utils.cc
|
||||
//
|
||||
// Default is an empty string which means no optimizers are disabled.
|
||||
static const char* const kOrtSessionOptionsDisableSpecifiedOptimizers = "optimization.disable_specified_optimizers"; |
||||
|
||||
// It controls whether to run graph optimizations in loop or not.
|
||||
//
|
||||
// "0": disable. Graph Optimization Loop is disabled.
|
||||
// ```
|
||||
// Level 2 --> Level 3 --> InsertCastTransforms --> Level 4
|
||||
// ^ |
|
||||
// | "No Loop" |
|
||||
// | |
|
||||
// X xxxxxxxxxxx X
|
||||
// ```
|
||||
// "1": enable. Graph Optimization Loop is enabled, such that, if optimizations at Level 4 are applied then
|
||||
// the loop will check for any other valid optimization that can happen.
|
||||
// ```
|
||||
// Level 2 --> Level 3 --> InsertCastTransforms --> Level 4
|
||||
// ^ |
|
||||
// | "Loop only depending on Level 4" |
|
||||
// | |
|
||||
// ---------------------------------------------------
|
||||
// ```
|
||||
// "2": enable. Graph Optimization Loop is enabled, such that, if optimizations at Level 2 or above are applied then
|
||||
// The loop will check for any other valid optimization that can happen.
|
||||
// ```
|
||||
// Level 2 --> Level 3 --> InsertCastTransforms --> Level 4
|
||||
// ^ |
|
||||
// | "Loop" |
|
||||
// | |
|
||||
// ---------------------------------------------------
|
||||
// ```
|
||||
// Default value is set to "1".
|
||||
static const char* const kOrtSessionOptionsGraphOptimizationsLoopLevel = "session.graph_optimizations_loop_level"; |
||||
|
||||
// Enable or disable using device allocator for allocating initialized tensor memory. "1": enable; "0": disable. The default is "0".
|
||||
// Using device allocators means the memory allocation is made using malloc/new.
|
||||
static const char* const kOrtSessionOptionsUseDeviceAllocatorForInitializers = "session.use_device_allocator_for_initializers"; |
||||
|
||||
// Configure whether to allow the inter_op/intra_op threads spinning a number of times before blocking
|
||||
// "0": thread will block if found no job to run
|
||||
// "1": thread will spin a number of times before blocking
|
||||
// The default is "0" when ORT is built with "ORT_CLIENT_PACKAGE_BUILD" and "1" otherwise.
|
||||
// Thread spinning is disabled by default for client/on-device workloads to reduce cpu utilization and improve power efficiency.
|
||||
static const char* const kOrtSessionOptionsConfigAllowInterOpSpinning = "session.inter_op.allow_spinning"; |
||||
static const char* const kOrtSessionOptionsConfigAllowIntraOpSpinning = "session.intra_op.allow_spinning"; |
||||
|
||||
// Key for using model bytes directly for ORT format
|
||||
// If a session is created using an input byte array contains the ORT format model data,
|
||||
// By default we will copy the model bytes at the time of session creation to ensure the model bytes
|
||||
// buffer is valid.
|
||||
// Setting this option to "1" will disable copy the model bytes, and use the model bytes directly. The caller
|
||||
// has to guarantee that the model bytes are valid until the ORT session using the model bytes is destroyed.
|
||||
static const char* const kOrtSessionOptionsConfigUseORTModelBytesDirectly = "session.use_ort_model_bytes_directly"; |
||||
|
||||
/// <summary>
|
||||
/// Key for using the ORT format model flatbuffer bytes directly for initializers.
|
||||
/// This avoids copying the bytes and reduces peak memory usage during model loading and initialization.
|
||||
/// Requires `session.use_ort_model_bytes_directly` to be true.
|
||||
/// If set, the flatbuffer bytes provided when creating the InferenceSession MUST remain valid for the entire
|
||||
/// duration of the InferenceSession.
|
||||
/// </summary>
|
||||
static const char* const kOrtSessionOptionsConfigUseORTModelBytesForInitializers = |
||||
"session.use_ort_model_bytes_for_initializers"; |
||||
|
||||
// This should only be specified when exporting an ORT format model for use on a different platform.
|
||||
// If the ORT format model will be used on ARM platforms set to "1". For other platforms set to "0"
|
||||
// Available since version 1.11.
|
||||
static const char* const kOrtSessionOptionsQDQIsInt8Allowed = "session.qdqisint8allowed"; |
||||
|
||||
// x64 SSE4.1/AVX2/AVX512(with no VNNI) has overflow problem with quantizied matrix multiplication with U8S8.
|
||||
// To avoid this we need to use slower U8U8 matrix multiplication instead. This option, if
|
||||
// turned on, use slower U8U8 matrix multiplications. Only effective with AVX2 or AVX512
|
||||
// platforms.
|
||||
static const char* const kOrtSessionOptionsAvx2PrecisionMode = "session.x64quantprecision"; |
||||
|
||||
// Specifies how minimal build graph optimizations are handled in a full build.
|
||||
// These optimizations are at the extended level or higher.
|
||||
// Possible values and their effects are:
|
||||
// "save": Save runtime optimizations when saving an ORT format model.
|
||||
// "apply": Only apply optimizations available in a minimal build.
|
||||
// ""/<unspecified>: Apply optimizations available in a full build.
|
||||
// Available since version 1.11.
|
||||
static const char* const kOrtSessionOptionsConfigMinimalBuildOptimizations = |
||||
"optimization.minimal_build_optimizations"; |
||||
|
||||
// Note: The options specific to an EP should be specified prior to appending that EP to the session options object in
|
||||
// order for them to take effect.
|
||||
|
||||
// Specifies a list of stop op types. Nodes of a type in the stop op types and nodes downstream from them will not be
|
||||
// run by the NNAPI EP.
|
||||
// The value should be a ","-delimited list of op types. For example, "Add,Sub".
|
||||
// If not specified, the default set of stop ops is used. To specify an empty stop ops types list and disable stop op
|
||||
// exclusion, set the value to "".
|
||||
static const char* const kOrtSessionOptionsConfigNnapiEpPartitioningStopOps = "ep.nnapi.partitioning_stop_ops"; |
||||
|
||||
// Enabling dynamic block-sizing for multithreading.
|
||||
// With a positive value, thread pool will split a task of N iterations to blocks of size starting from:
|
||||
// N / (num_of_threads * dynamic_block_base)
|
||||
// As execution progresses, the size will decrease according to the diminishing residual of N,
|
||||
// meaning the task will be distributed in smaller granularity for better parallelism.
|
||||
// For some models, it helps to reduce the variance of E2E inference latency and boost performance.
|
||||
// The feature will not function by default, specify any positive integer, e.g. "4", to enable it.
|
||||
// Available since version 1.11.
|
||||
static const char* const kOrtSessionOptionsConfigDynamicBlockBase = "session.dynamic_block_base"; |
||||
|
||||
// This option allows to decrease CPU usage between infrequent
|
||||
// requests and forces any TP threads spinning stop immediately when the last of
|
||||
// concurrent Run() call returns.
|
||||
// Spinning is restarted on the next Run() call.
|
||||
// Applies only to internal thread-pools
|
||||
static const char* const kOrtSessionOptionsConfigForceSpinningStop = "session.force_spinning_stop"; |
||||
|
||||
// "1": all inconsistencies encountered during shape and type inference
|
||||
// will result in failures.
|
||||
// "0": in some cases warnings will be logged but processing will continue. The default.
|
||||
// May be useful to expose bugs in models.
|
||||
static const char* const kOrtSessionOptionsConfigStrictShapeTypeInference = "session.strict_shape_type_inference"; |
||||
|
||||
// "1": every model using a more recent opset than the latest released one will fail
|
||||
// "0": the model may or may not work if onnxruntime cannot find an implementation, this option
|
||||
// is used for development purpose.
|
||||
static const char* const kOrtSessionOptionsConfigStrictAllowReleasedOpsetsOnly = "session.allow_released_opsets_only"; |
||||
|
||||
// The file saves configuration for partitioning node among logic streams
|
||||
static const char* const kNodePartitionConfigFile = "session.node_partition_config_file"; |
||||
|
||||
// This Option allows setting affinities for intra op threads.
|
||||
// Affinity string follows format:
|
||||
// logical_processor_id,logical_processor_id;logical_processor_id,logical_processor_id
|
||||
// Semicolon isolates configurations among threads, while comma split processors where ith thread expected to attach to.
|
||||
// e.g.1,2,3;4,5
|
||||
// specifies affinities for two threads, with the 1st thread attach to the 1st, 2nd, and 3rd processor, and 2nd thread to the 4th and 5th.
|
||||
// To ease the configuration, an "interval" is also allowed:
|
||||
// e.g. 1-8;8-16;17-24
|
||||
// orders that the 1st thread runs on first eight processors, 2nd thread runs on next eight processors, and so forth.
|
||||
// Note:
|
||||
// 1. Once set, the number of thread affinities must equal to intra_op_num_threads - 1, since ort does not set affinity on the main thread which
|
||||
// is started and managed by the calling app;
|
||||
// 2. For windows, ort will infer the group id from a logical processor id, for example, assuming there are two groups with each has 64 logical processors,
|
||||
// an id of 64 will be inferred as the last processor of the 1st group, while 65 will be interpreted as the 1st processor of the second group.
|
||||
// Hence 64-65 is an invalid configuration, because a windows thread cannot be attached to processors across group boundary.
|
||||
static const char* const kOrtSessionOptionsConfigIntraOpThreadAffinities = "session.intra_op_thread_affinities"; |
||||
|
||||
// This option will dump out the model to assist debugging any issues with layout transformation,
|
||||
// and is primarily intended for developer usage. It is only relevant if an execution provider that requests
|
||||
// NHWC layout is enabled such as NNAPI, XNNPACK or QNN.
|
||||
//
|
||||
// Default is off. Set to "1" to enable.
|
||||
//
|
||||
// If modified by layout transformation the model will be dumped after these steps:
|
||||
// 1) insertion of the layout transformation Transpose nodes
|
||||
// 2) after those are optimized using the transpose optimizer,
|
||||
// 3) after the L1 transformers are applied to the updated graph.
|
||||
// The model will be saved to filename post_layout_transform_step_<step_number>.onnx.
|
||||
static const char* const kDebugLayoutTransformation = "session.debug_layout_transformation"; |
||||
|
||||
// Graph nodes that are not supported by the execution providers (EPs) explicitly added to the session are
|
||||
// assigned (i.e., "fallback") to the CPU EP by default.
|
||||
//
|
||||
// This option allows the user to disable the fallback of unsupported graph nodes to the CPU EP.
|
||||
// If this option is set to "1", session creation will fail if the execution providers other than the CPU EP cannot
|
||||
// fully support all of the nodes in the graph.
|
||||
//
|
||||
// It is invalid to set this option and explicitly add the CPU EP to the session. In this case, session creation
|
||||
// will also fail with an error.
|
||||
//
|
||||
// Option values:
|
||||
// - "0": CPU EP fallback is not disabled. [DEFAULT]
|
||||
// - "1": CPU EP fallback is disabled.
|
||||
static const char* const kOrtSessionOptionsDisableCPUEPFallback = "session.disable_cpu_ep_fallback"; |
||||
|
||||
// Use this config when serializing a large model after optimization to specify an external initializers file
|
||||
static const char* const kOrtSessionOptionsOptimizedModelExternalInitializersFileName = |
||||
"session.optimized_model_external_initializers_file_name"; |
||||
|
||||
// Use this config to control the minimum size of the initializer when externalizing it during serialization
|
||||
static const char* const kOrtSessionOptionsOptimizedModelExternalInitializersMinSizeInBytes = |
||||
"session.optimized_model_external_initializers_min_size_in_bytes"; |
||||
|
||||
// When loading model from memory buffer and the model has external initializers
|
||||
// Use this config to set the external data file folder path
|
||||
// All external data files should be in the same folder
|
||||
static const char* const kOrtSessionOptionsModelExternalInitializersFileFolderPath = |
||||
"session.model_external_initializers_file_folder_path"; |
||||
|
||||
// Use this config when saving pre-packed constant initializers to an external data file.
|
||||
// This allows you to memory map pre-packed initializers on model load and leave it to
|
||||
// to the OS the amount of memory consumed by the pre-packed initializers. Otherwise,
|
||||
// pre-packed data resides on the heap.
|
||||
//
|
||||
// - "0": Default is not save pre-packed initializers to a data file.
|
||||
// - "1": Save pre-packed constant initializers to an external data file.
|
||||
// Sample usage: sess_options.add_session_config_entry(kOrtSessionOptionsSavePrePackedConstantInitializers, "1")
|
||||
static const char* const kOrtSessionOptionsSavePrePackedConstantInitializers = |
||||
"session.save_external_prepacked_constant_initializers"; |
||||
|
||||
// Use this config when you want to collect memory stats for each node in the graph.
|
||||
// The file format is a CSV file with the following columns:
|
||||
// The file will be created if it does not exist, and will be overwritten if it does.
|
||||
//
|
||||
// The content of the file can be used to estimate memory requirements at run time including
|
||||
// the temporary allocations. This operation is preferably done on a CPU device, as the model may exceed
|
||||
// device memory limits in constrained environments. When enabling this option, it is important to disable
|
||||
// memory patterns, as they tend to allocate large blocks to avoid fragmentation and accommodate needs of multiple
|
||||
// kernels. Memory patterns may make it difficult to allocate on a device with limited memory.
|
||||
//
|
||||
// The collected stats then can be used to partition the graph among the devices in a way that only the
|
||||
// required memory is allocated on each device.
|
||||
//
|
||||
// node_name, initializers_memory, dynamic_outputs_sizes, temp_allocations_size
|
||||
//
|
||||
// - "full path to file": there is not a default for this option. If the file can not be opened for writing, an error will be returned.
|
||||
static const char* const kOrtSessionOptionsCollectNodeMemoryStatsToFile = "session.collect_node_memory_stats_to_file"; |
||||
|
||||
/// This is a composite CSV setting formatted as "memory limit in kb,file name for collected stats"
|
||||
/// "limit > 0": enables Capacity Aware Partitioning for Cuda EP. `limit` is optional and when absent
|
||||
/// the provider may attempt to figure out the memory available automatically.
|
||||
/// The setting with no limit is expected to look like: ",file name for collected stats"
|
||||
/// The EP will place nodes on device "file name" :
|
||||
/// this file is expected to be found at the same folder with the model. The file contains
|
||||
/// pre-recorded stats collected when running with kOrtSessionOptionsCollectNodeMemoryStatsToFile enforce (see above)
|
||||
static const char* const kOrtSessionOptionsResourceCudaPartitioningSettings = |
||||
"session.resource_cuda_partitioning_settings"; |
||||
|
||||
// Enable EP context feature to dump the partitioned graph which includes the EP context into Onnx file.
|
||||
// The dumped Onnx model with EP context can be used for future inference to avoid the EP graph partitioning/compile overhead.
|
||||
// "0": disable. (default)
|
||||
// "1": enable.
|
||||
static const char* const kOrtSessionOptionEpContextEnable = "ep.context_enable"; |
||||
|
||||
// Specify the file path for the Onnx model which has EP context.
|
||||
// Default to original_file_name_ctx.onnx if not specified
|
||||
// Folder is not a valid option
|
||||
static const char* const kOrtSessionOptionEpContextFilePath = "ep.context_file_path"; |
||||
|
||||
// Flag to specify whether to dump the EP context into the Onnx model.
|
||||
// "0": dump the EP context into separate file, keep the file name in the Onnx model. (default).
|
||||
// "1": dump the EP context into the Onnx model.
|
||||
static const char* const kOrtSessionOptionEpContextEmbedMode = "ep.context_embed_mode"; |
||||
|
||||
// Specify the EPContext node name prefix to make it unique
|
||||
// in case user need to merge/connect multiple EPContext nodes in one model
|
||||
static const char* const kOrtSessionOptionEpContextNodeNamePrefix = "ep.context_node_name_prefix"; |
||||
|
||||
// Share EP related resources across sessions
|
||||
static const char* const kOrtSessionOptionShareEpContexts = "ep.share_ep_contexts"; |
||||
|
||||
// Stop to share EP related resources across sessions from then on
|
||||
static const char* const kOrtSessionOptionStopShareEpContexts = "ep.stop_share_ep_contexts"; |
||||
|
||||
// Used only for context model generation.
|
||||
// This configuration is used when some nodes are partitioned on the CPU EP and those nodes have external initializers.
|
||||
// When generating the EP context model, the new model should not rely on the old external data file used by the source ONNX model.
|
||||
// Use this setting when dumping the EP context model with an external initializers file.
|
||||
// If specified, all initializers will be placed inside the external data file.
|
||||
// Otherwise, all initializers will be embedded inside the generated ONNX file.
|
||||
// By default, this option is not set, meaning all initializers will be included within the ONNX file.
|
||||
static const char* const kOrtSessionOptionsEpContextModelExternalInitializersFileName = |
||||
"ep.context_model_external_initializers_file_name"; |
||||
|
||||
// Gemm fastmath mode provides fp32 gemm acceleration with bfloat16 based matmul.
|
||||
// Option values:
|
||||
// - "0": Gemm FastMath mode is not enabled. [DEFAULT]
|
||||
// - "1": Gemm FastMath mode is enabled.
|
||||
static const char* const kOrtSessionOptionsMlasGemmFastMathArm64Bfloat16 = "mlas.enable_gemm_fastmath_arm64_bfloat16"; |
||||
|
||||
// When converting DQ + MatMul -> MatMulNBits, the accuracy level of the MatMulNBits is controlled by this option.
|
||||
// Refer to MatMulNBits op schema for more details.
|
||||
// If not provided, default is 4.
|
||||
static const char* const kOrtSessionOptionsQDQMatMulNBitsAccuracyLevel = "session.qdq_matmulnbits_accuracy_level"; |
||||
|
||||
// THIS OPTION IS NOT A REGULAR SESSION OPTION SINCE IT CAN BE MODIFIED AT ANY TIME
|
||||
// Meant to be used with SetEpDynamicOptions
|
||||
// Specify the type of workload for this session.
|
||||
// "Default": OS determines the scheduling priority and processor performance to service this workload. [Default]
|
||||
// "Efficient": OS treats this workload is efficiency oriented with low scheduling priority and efficient processor performance.
|
||||
static const char* const kOrtEpDynamicOptionsWorkloadType = "ep.dynamic.workload_type"; |
||||
|
||||
// Disables model compilation during session initialization.
|
||||
//
|
||||
// If this option is set to "1", inference session creation will fail with error code ORT_MODEL_REQUIRES_COMPILATION
|
||||
// if compilation is required to run the model on any Execution Provider added to the session.
|
||||
// Only the following kinds of models are valid when this option is set to "1":
|
||||
// - Pre-compiled models that have EPContext nodes for the compiling Execution Providers in the session.
|
||||
// - Non-compiled models that run only on non-compiling Execution Providers, like CPU EP.
|
||||
//
|
||||
// See \href https://onnxruntime.ai/docs/execution-providers/EP-Context-Design.html for details about
|
||||
// compiled models with EPContext nodes.
|
||||
//
|
||||
// Option values:
|
||||
// - "0": EP compile is not disabled. [DEFAULT]
|
||||
// - "1": EP compile is disabled.
|
||||
static const char* const kOrtSessionOptionsDisableModelCompile = "session.disable_model_compile"; |
||||
|
||||
// Controls behavior when compiled model compatibility is SUPPORTED_PREFER_RECOMPILATION.
|
||||
// "0": Allow execution with suboptimal performance. [DEFAULT]
|
||||
// "1": Fail session creation to require recompilation for optimal performance.
|
||||
// Note: UNSUPPORTED models always fail regardless of this setting.
|
||||
static const char* const kOrtSessionOptionsFailOnSuboptimalCompiledModel = |
||||
"session.fail_on_suboptimal_compiled_model"; |
||||
|
||||
// THIS OPTION IS NOT A REGULAR SESSION OPTION SINCE IT CAN BE MODIFIED AT ANY TIME
|
||||
// Meant to be used with SetEpDynamicOptions
|
||||
// options for HTP performance mode: "burst", "balanced", "default", "high_performance",
|
||||
// "high_power_saver", "low_balanced", "extreme_power_saver", "low_power_saver", "power_saver",
|
||||
// "sustained_high_performance". Default to "default".
|
||||
static const char* const kOrtEpDynamicOptionsQnnHtpPerformanceMode = "ep.dynamic.qnn_htp_performance_mode"; |
||||
Binary file not shown.
Binary file not shown.
|
After Width: | Height: | Size: 1.2 MiB |
|
Before Width: | Height: | Size: 84 KiB |
|
After Width: | Height: | Size: 2.0 MiB |
|
After Width: | Height: | Size: 1.2 MiB |
Loading…
Reference in new issue