Compare commits
17 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 1766a0f2e3 | |||
| 1fb0454ac3 | |||
| 0c065a0931 | |||
| 0d1bc4d936 | |||
| 515a434928 | |||
| 66048bff05 | |||
| 9787b69588 | |||
| f0c1d23121 | |||
| ecf33beb01 | |||
| 9e0f21023d | |||
| 87b45a1484 | |||
| 14ae41d2db | |||
| ecef2f41af | |||
| a081d4d7e8 | |||
| 0dd3e46b6f | |||
| cf005944b8 | |||
| 35dfb8d3a5 |
@ -11,6 +11,7 @@ else()
|
||||
set(CMAKE_CXX_FLAGS_RELEASE "-O3 -ffast-math -funroll-loops")
|
||||
endif()
|
||||
|
||||
|
||||
add_compile_options(-Wall -Wextra -Wno-unused-parameter)
|
||||
|
||||
set(SOURCES
|
||||
@ -27,8 +28,13 @@ set(SOURCES
|
||||
xenith/converter.cpp
|
||||
help.h
|
||||
xenith/ollama_api/handler.cpp
|
||||
xenith/ollama_api/handler.h
|
||||
xenith/ollama_api/server.h
|
||||
xenith/ollama_api/handler.hpp
|
||||
xenith/ollama_api/server.hpp
|
||||
xenith/ollama_api/server.cpp
|
||||
xenith/ui/tui/tui.cpp
|
||||
xenith/ui/tui.hpp
|
||||
xenith/ui/tui/mouse.cpp
|
||||
xenith/ui/tui/message.cpp
|
||||
)
|
||||
|
||||
# Заголовочные файлы (для удобства в IDE)
|
||||
@ -38,9 +44,9 @@ set(CMAKE_RUNTIME_OUTPUT_DIRECTORY ${CMAKE_SOURCE_DIR}/bin)
|
||||
|
||||
add_executable(Xenith ${SOURCES} ${HEADER_FILES})
|
||||
|
||||
# Настройка include директорий
|
||||
target_include_directories(Xenith SYSTEM PRIVATE ./xenith/libs/cpp-httplib-0.57.1)
|
||||
|
||||
target_include_directories(Xenith PRIVATE ${CMAKE_CURRENT_SOURCE_DIR})
|
||||
|
||||
# Потоки (pthread)
|
||||
find_package(Threads REQUIRED)
|
||||
target_link_libraries(Xenith PRIVATE Threads::Threads)
|
||||
50
README.md
50
README.md
@ -15,7 +15,7 @@
|
||||
|
||||
| файл | что делает |
|
||||
|---|---|
|
||||
| `xenith/core/tensor.h` | матрицы, три формы GEMM, RMSNorm, RNG, пул потоков |
|
||||
| `xenith/core/tensor.hpp` | матрицы, три формы GEMM, RMSNorm, RNG, пул потоков |
|
||||
| `xenith/core/model.cpp` | конфигурация, инициализация, прямой проход, RoPE, Adam |
|
||||
| `xenith/core/backward.cpp` | обратный проход (ручной, без автомиффов) |
|
||||
| `xenith/core/trainer.cpp` | цикл обучения, расписание lr, чекпоинты, валидация |
|
||||
@ -26,16 +26,33 @@
|
||||
## Быстрый старт
|
||||
|
||||
```bash
|
||||
# 1) собрать модель из своего текста
|
||||
./bin/Xenith new models/my.xnh --corpus data.txt \
|
||||
# 1) активация окружения
|
||||
source activate.sh
|
||||
|
||||
# 2) собрать модель из своего текста
|
||||
Xenith new models/my.xnh --corpus data.txt \
|
||||
--vocab 512 --embd 128 --layers 6 --heads 8 --ctx 256
|
||||
|
||||
# 2) обучить
|
||||
./bin/Xenith train models/my.xnh --corpus data.txt \
|
||||
# 3) обучить
|
||||
Xenith train models/my.xnh --corpus data.txt \
|
||||
--out models/my_trained.xnh --steps 5000 --batch 16 --lr 3e-4
|
||||
|
||||
# 3) сгенерировать
|
||||
./bin/Xenith gen models/my_trained.xnh --prompt "Привет" --n 200 --temp 0.8
|
||||
# 4) сгенерировать
|
||||
Xenith gen models/my_trained.xnh --prompt "Привет" --n 200 --temp 0.8
|
||||
```
|
||||
|
||||
## Перед использованием ***Xenith*** желательно активировать окружение
|
||||
|
||||
### Команда активации окружения
|
||||
```bash
|
||||
# для bash
|
||||
source activate.sh
|
||||
|
||||
# для zsh
|
||||
source activate.zsh
|
||||
|
||||
# для fish
|
||||
source activate.fish
|
||||
```
|
||||
|
||||
## Команды
|
||||
@ -109,15 +126,24 @@
|
||||
--conf NAME.conf конфиг с настройками и списком моделей
|
||||
```
|
||||
|
||||
### update - проверка обновления
|
||||
### gradcheck - проверить градиенты численно
|
||||
|
||||
### info / gradcheck / bench
|
||||
```
|
||||
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
||||
```
|
||||
|
||||
### bench - замерить скорость форварда и шага обучения
|
||||
|
||||
```
|
||||
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
||||
```
|
||||
|
||||
### info
|
||||
|
||||
```bash
|
||||
./bin/Xenith info models/my.xnh # конфигурация и словарь
|
||||
./bin/Xenith gradcheck # градиенты против численных
|
||||
./bin/Xenith bench models/my.xnh # скорость
|
||||
Xenith info models/my.xnh # конфигурация и словарь
|
||||
```
|
||||
### update - проверка обновления
|
||||
|
||||
`gradcheck` стоит запускать после любой правки в `model.cpp` / `backward.cpp` —
|
||||
он ловит ошибку в обратном проходе за секунды.
|
||||
|
||||
80
_xenith
Normal file
80
_xenith
Normal file
@ -0,0 +1,80 @@
|
||||
# _xenith - Zsh completion function for Xenith
|
||||
_xenith() {
|
||||
local curcontext="$curcontext" state line
|
||||
typeset -A opt_args
|
||||
local commands=(
|
||||
'new:создать новую модель из корпуса'
|
||||
'train:обучить существующую модель'
|
||||
'gen:сгенерировать текст'
|
||||
'info:показать конфигурацию модели'
|
||||
'gradcheck:проверить градиенты численно'
|
||||
'bench:замерить скорость форварда и шага обучения'
|
||||
'update:автоматическое обновление с гита'
|
||||
'tui:графическо-консольный режим (тест)'
|
||||
)
|
||||
_arguments -C \
|
||||
'1: :->command' \
|
||||
'*:: :->args' && return 0
|
||||
case $state in
|
||||
command)
|
||||
_describe -t commands 'команды Xenith' commands
|
||||
;;
|
||||
args)
|
||||
case $line[1] in
|
||||
new)
|
||||
_arguments \
|
||||
'--corpus[текст для построения словаря]:файл:_files' \
|
||||
'--out[файл модели]:файл:_files -g "*.xnh"' \
|
||||
'--vocab[размер словаря]' \
|
||||
'--embd[размер эмбеддинга]' \
|
||||
'--layers[число слоёв]' \
|
||||
'--heads[число голов внимания]' \
|
||||
'--ctx[максимальный контекст]' \
|
||||
'--ffn[размер FFN]' \
|
||||
'--rope[база RoPE]' \
|
||||
'--seed[зерно инициализации]' \
|
||||
'--untied[использовать отдельную матрицу выхода]' \
|
||||
'--view-vocab[вывести словарь нейросети]'
|
||||
;;
|
||||
train)
|
||||
_arguments \
|
||||
'--corpus[обучающий текст]:файл:_files' \
|
||||
'--out[куда сохранить]:файл:_files -g "*.xnh"' \
|
||||
'--steps[количество шагов обучения]' \
|
||||
'--batch[размер батча]' \
|
||||
'--block[длина окна контекста]' \
|
||||
'--lr[скорость обучения]' \
|
||||
'--wd[weight decay]' \
|
||||
'--clip[клиппинг нормы градиента]' \
|
||||
'--warmup[шаги прогрева]' \
|
||||
'--seed[зерно генератора]' \
|
||||
'--ckpt-every[сохранять чекпоинт каждые N шагов]' \
|
||||
'--val-every[валидация каждые N шагов]' \
|
||||
'--val-tokens[размер валидационной выборки]' \
|
||||
'--threads[количество потоков (0 = все ядра)]' \
|
||||
'--resume[продолжить обучение с последнего чекпоинта]' \
|
||||
'--fast[экспериментальный флаг ускорения forward pass]'
|
||||
;;
|
||||
gen)
|
||||
_arguments \
|
||||
'--prompt[стартовый текст]' \
|
||||
'--n[количество генерируемых токенов]' \
|
||||
'--temp[температура выборки (0 = жадный)]' \
|
||||
'--top-k[top-k sampling]' \
|
||||
'--top-p[top-p nucleus sampling]' \
|
||||
'--seed[зерно генератора]' \
|
||||
'--no-stream[не печатать токены по мере генерации]' \
|
||||
'--batch[несколько промптов через ;]' \
|
||||
'--show-tokens[показывать ID токенов]'
|
||||
;;
|
||||
info|gradcheck|bench|update|tui)
|
||||
_arguments \
|
||||
'--fast[экспериментальный флаг ускорения]'
|
||||
;;
|
||||
esac
|
||||
;;
|
||||
esac
|
||||
}
|
||||
if [[ "$funcstack[1]" == "_xenith" ]]; then
|
||||
_xenith "$@"
|
||||
fi
|
||||
85
activate.fish
Normal file
85
activate.fish
Normal file
@ -0,0 +1,85 @@
|
||||
set -gx VIRTUAL_ENV_PROMPT "Xenith"
|
||||
complete -c Xenith -erase
|
||||
complete -c ./bin/Xenith -erase
|
||||
complete -c Xenith -f
|
||||
complete -c ./bin/Xenith -f
|
||||
set -l main_commands "new:создать новую модель из корпуса" \
|
||||
"train:обучить существующую модель" \
|
||||
"gen:сгенерировать текст" \
|
||||
"info:показать конфигурацию модели" \
|
||||
"gradcheck:проверить градиенты численно" \
|
||||
"bench:замерить скорость форварда и шага обучения" \
|
||||
"update:Автомотическое обновление с гита"
|
||||
for cmd in $main_commands
|
||||
set -l parts (string split ":" $cmd)
|
||||
complete -c Xenith -n "__fish_use_subcommand" -a $parts[1] -d $parts[2]
|
||||
complete -c ./bin/Xenith -n "__fish_use_subcommand" -a $parts[1] -d $parts[2]
|
||||
end
|
||||
set -l new_cond "__fish_seen_subcommand_from new"
|
||||
complete -c Xenith -n $new_cond -l corpus -r -F -d "текст для построения словаря (обязательно)"
|
||||
complete -c Xenith -n $new_cond -l out -r -F -d "файл модели (по умолчанию model.xnh)"
|
||||
complete -c Xenith -n $new_cond -l vocab -r -d "размер словаря (по умолчанию 512)"
|
||||
complete -c Xenith -n $new_cond -l embd -r -d "размер эмбеддинга (по умолчанию 128)"
|
||||
complete -c Xenith -n $new_cond -l layers -r -d "число слоёв (по умолчанию 4)"
|
||||
complete -c Xenith -n $new_cond -l heads -r -d "число голов (по умолчанию 4)"
|
||||
complete -c Xenith -n $new_cond -l ctx -r -d "максимальный контекст (по умолчанию 128)"
|
||||
complete -c Xenith -n $new_cond -l ffn -r -d "размер FFN (по умолчанию 4*embd)"
|
||||
complete -c Xenith -n $new_cond -l rope -r -d "база RoPE (по умолчанию 10000)"
|
||||
complete -c Xenith -n $new_cond -l seed -r -d "зерно инициализации (по умолчанию 1337)"
|
||||
complete -c Xenith -n $new_cond -l untied -d "отдельная матрица выхода вместо привязанной к wte"
|
||||
complete -c Xenith -n $new_cond -l view-vocab -d "выводит словарь нейросети"
|
||||
set -l train_cond "__fish_seen_subcommand_from train"
|
||||
complete -c Xenith -n $train_cond -l corpus -r -F -d "обучающий текст (обязательно)"
|
||||
complete -c Xenith -n $train_cond -l out -r -F -d "куда сохранить (по умолчанию перезапись входного файла)"
|
||||
complete -c Xenith -n $train_cond -l steps -r -d "шагов (по умолчанию 1000)"
|
||||
complete -c Xenith -n $train_cond -l batch -r -d "батч (по умолчанию 8)"
|
||||
complete -c Xenith -n $train_cond -l block -r -d "длина окна (по умолчанию 64)"
|
||||
complete -c Xenith -n $train_cond -l lr -r -d "скорость (по умолчанию 3e-4)"
|
||||
complete -c Xenith -n $train_cond -l wd -r -d "weight decay (по умолчанию 0.01)"
|
||||
complete -c Xenith -n $train_cond -l clip -r -d "клип нормы градиента (по умолчанию 1.0)"
|
||||
complete -c Xenith -n $train_cond -l warmup -r -d "прогрев (по умолчанию 100)"
|
||||
complete -c Xenith -n $train_cond -l seed -r -d "зерно (по умолчанию 1337)"
|
||||
complete -c Xenith -n $train_cond -l ckpt-every -r -d "промежуточное сохранение каждые N шагов"
|
||||
complete -c Xenith -n $train_cond -l val-every -r -d "валидация каждые N шагов"
|
||||
complete -c Xenith -n $train_cond -l val-tokens -r -d "размер валидации (по умолчанию 20000)"
|
||||
complete -c Xenith -n $train_cond -l threads -r -d "потоков (0 = все ядра)"
|
||||
complete -c Xenith -n $train_cond -l resume -d "продолжить с сохранённого step (Adam-state из файла)"
|
||||
set -l gen_cond "__fish_seen_subcommand_from gen"
|
||||
complete -c Xenith -n $gen_cond -l prompt -r -d "стартовый текст"
|
||||
complete -c Xenith -n $gen_cond -l n -r -d "сколько токенов (по умолчанию 200)"
|
||||
complete -c Xenith -n $gen_cond -l temp -r -d "температура, 0 = жадный выбор (по умолчанию 0.8)"
|
||||
complete -c Xenith -n $gen_cond -l top-k -r -d "top-k (0 = выкл)"
|
||||
complete -c Xenith -n $gen_cond -l top-p -r -d "top-p (1.0 = выкл)"
|
||||
complete -c Xenith -n $gen_cond -l seed -r -d "зерно"
|
||||
complete -c Xenith -n $gen_cond -l no-stream -d "не печатать в процессе"
|
||||
complete -c Xenith -n $gen_cond -l batch -r -d "несколько промптов через ;"
|
||||
complete -c Xenith -n $gen_cond -l show-tokens -d "показать id токенов"
|
||||
complete -c ./bin/Xenith -n "__fish_seen_subcommand_from new" -a "(complete -C'Xenith new ')"
|
||||
complete -c ./bin/Xenith -n "__fish_seen_subcommand_from train" -a "(complete -C'Xenith train ')"
|
||||
complete -c ./bin/Xenith -n "__fish_seen_subcommand_from gen" -a "(complete -C'Xenith gen ')"
|
||||
function Xenith
|
||||
./bin/Xenith $argv
|
||||
end
|
||||
if not set -q _OLD_VIRTUAL_PS1
|
||||
functions -c fish_prompt _old_fish_prompt
|
||||
function fish_prompt
|
||||
echo -n "($VIRTUAL_ENV_PROMPT) "
|
||||
_old_fish_prompt
|
||||
end
|
||||
set -g _OLD_VIRTUAL_PS1 1
|
||||
end
|
||||
echo "Xenith activate"
|
||||
function deactivate
|
||||
functions -e Xenith
|
||||
if functions -q _old_fish_prompt
|
||||
functions -e fish_prompt
|
||||
functions -c _old_fish_prompt fish_prompt
|
||||
functions -e _old_fish_prompt
|
||||
end
|
||||
set -e VIRTUAL_ENV_PROMPT
|
||||
set -e _OLD_VIRTUAL_PS1
|
||||
complete -c Xenith -erase
|
||||
complete -c ./bin/Xenith -erase
|
||||
functions -e deactivate
|
||||
echo "Xenith deactivated"
|
||||
end
|
||||
21
activate.sh
Normal file
21
activate.sh
Normal file
@ -0,0 +1,21 @@
|
||||
export VIRTUAL_ENV_PROMPT="\033[31mXenith\033[39m"
|
||||
alias Xenith='./bin/Xenith'
|
||||
if [ -z "${_OLD_VIRTUAL_PS1+x}" ]; then
|
||||
_OLD_VIRTUAL_PS1="$PS1"
|
||||
fi
|
||||
PS1="($VIRTUAL_ENV_PROMPT) $_OLD_VIRTUAL_PS1"
|
||||
deactivate_custom() {
|
||||
unalias Xenith 2>/dev/null
|
||||
unalias deactivate 2>/dev/null
|
||||
if [ -n "${_OLD_VIRTUAL_PS1+x}" ]; then
|
||||
PS1="$_OLD_VIRTUAL_PS1"
|
||||
unset _OLD_VIRTUAL_PS1
|
||||
fi
|
||||
if (( $+functions[_old_virtual_deactivate] )); then
|
||||
_old_virtual_deactivate
|
||||
elif (( $+functions[deactivate] )); then
|
||||
builtin deactivate 2>/dev/null || true
|
||||
fi
|
||||
unset -f deactivate_custom
|
||||
}
|
||||
alias deactivate='deactivate_custom'
|
||||
27
activate.zsh
Normal file
27
activate.zsh
Normal file
@ -0,0 +1,27 @@
|
||||
# activate.zsh
|
||||
export VIRTUAL_ENV_PROMPT="%F{red}Xenith%f"
|
||||
alias Xenith='./bin/Xenith'
|
||||
if [ -z "${_OLD_VIRTUAL_PS1+x}" ]; then
|
||||
_OLD_VIRTUAL_PS1="$PS1"
|
||||
fi
|
||||
PS1="(${VIRTUAL_ENV_PROMPT}) ${_OLD_VIRTUAL_PS1}"
|
||||
deactivate_custom() {
|
||||
unalias Xenith 2>/dev/null
|
||||
unalias deactivate 2>/dev/null
|
||||
if [ -n "${_OLD_VIRTUAL_PS1+x}" ]; then
|
||||
PS1="$_OLD_VIRTUAL_PS1"
|
||||
unset _OLD_VIRTUAL_PS1
|
||||
fi
|
||||
if (( $+functions[_old_virtual_deactivate] )); then
|
||||
_old_virtual_deactivate
|
||||
elif (( $+functions[deactivate] )); then
|
||||
builtin deactivate 2>/dev/null || true
|
||||
fi
|
||||
unset -f deactivate_custom
|
||||
compdel Xenith 2>/dev/null
|
||||
compdel ./bin/Xenith 2>/dev/null
|
||||
}
|
||||
alias deactivate='deactivate_custom'
|
||||
source "${0:A:h}/_xenith"
|
||||
compdef _xenith Xenith
|
||||
compdef _xenith ./bin/Xenith
|
||||
BIN
bin/Xenith
BIN
bin/Xenith
Binary file not shown.
@ -1,2 +1,9 @@
|
||||
|
||||
models:
|
||||
BiPy: models/bipy.bif
|
||||
BiPy:
|
||||
file: "models/my.xnh"
|
||||
details:
|
||||
parent_model: ""
|
||||
family: "Bfr"
|
||||
families:
|
||||
- Bfr
|
||||
207
help.h
207
help.h
@ -1,71 +1,154 @@
|
||||
#ifndef XENITH_HELP_H
|
||||
#define XENITH_HELP_H
|
||||
|
||||
inline const std::string HELP_TEXT = R"(ИСПОЛЬЗОВАНИЕ
|
||||
xenith <команда> [опции]
|
||||
#define NEW
|
||||
#define TRAIN
|
||||
#define GEN
|
||||
#define INFO
|
||||
#define GRADCHECK
|
||||
#define BENCH
|
||||
//#define OLLAMA
|
||||
#define UPDATE
|
||||
//#define TUI
|
||||
|
||||
КОМАНДЫ
|
||||
new создать новую модель из корпуса
|
||||
train обучить существующую модель
|
||||
gen сгенерировать текст
|
||||
info показать конфигурацию модели
|
||||
gradcheck проверить градиенты численно
|
||||
bench замерить скорость форварда и шага обучения
|
||||
ollama запустьтить сервер с api ollama
|
||||
update Автомотическое обновление с гита
|
||||
#ifdef NEW
|
||||
#define CMD_NEW " new создать новую модель из корпуса\n"
|
||||
#else
|
||||
#define CMD_NEW ""
|
||||
#endif
|
||||
#ifdef TRAIN
|
||||
#define CMD_TRAIN " train обучить существующую модель\n"
|
||||
#else
|
||||
#define CMD_TRAIN ""
|
||||
#endif
|
||||
#ifdef GEN
|
||||
#define CMD_GEN " gen сгенерировать текст\n"
|
||||
#else
|
||||
#define CMD_GEN ""
|
||||
#endif
|
||||
#ifdef INFO
|
||||
#define CMD_INFO " info показать конфигурацию модели\n"
|
||||
#else
|
||||
#define CMD_INFO ""
|
||||
#endif
|
||||
#ifdef GRADCHECK
|
||||
#define CMD_GRADCHECK " gradcheck проверить градиенты численно\n"
|
||||
#else
|
||||
#define CMD_GRADCHECK ""
|
||||
#endif
|
||||
#ifdef BENCH
|
||||
#define CMD_BENCH " bench замерить скорость форварда и шага обучения\n"
|
||||
#else
|
||||
#define CMD_BENCH ""
|
||||
#endif
|
||||
#ifdef OLLAMA
|
||||
#define CMD_OLLAMA " ollama запустьтить сервер с api ollama\n"
|
||||
#else
|
||||
#define CMD_OLLAMA ""
|
||||
#endif
|
||||
#ifdef UPDATE
|
||||
#define CMD_UPDATE " update Автомотическое обновление с гита\n"
|
||||
#else
|
||||
#define CMD_UPDATE ""
|
||||
#endif
|
||||
#ifdef TUI
|
||||
#define CMD_TUI " tui Графическо-консольный режим (находится в тестировании)\n"
|
||||
#else
|
||||
#define CMD_TUI ""
|
||||
#endif
|
||||
|
||||
ПРИМЕРЫ
|
||||
xenith new model.xnh --corpus data.txt --vocab 512 --embd 64 --layers 4 --heads 4
|
||||
xenith train model.xnh --corpus data.txt --out model_trained.xnh --steps 3000
|
||||
xenith train model_trained.xnh --corpus data.txt --out model_trained.xnh --steps 6000
|
||||
xenith gen model_trained.xnh --prompt "Привет" --n 200 --temp 0.8
|
||||
|
||||
ОПЦИИ new
|
||||
--corpus PATH текст для построения словаря (обязательно)
|
||||
--out PATH файл модели (по умолчанию model.xnh)
|
||||
--vocab N размер словаря (по умолчанию 512)
|
||||
--embd N размер эмбеддинга (по умолчанию 128)
|
||||
--layers N число слоёв (по умолчанию 4)
|
||||
--heads N число голов (по умолчанию 4)
|
||||
--ctx N максимальный контекст (по умолчанию 128)
|
||||
--ffn N размер FFN (по умолчанию 4*embd)
|
||||
--rope N база RoPE (по умолчанию 10000)
|
||||
--seed N зерно инициализации (по умолчанию 1337)
|
||||
--untied отдельная матрица выхода вместо привязанной к wte
|
||||
--view-vocab выводит словарь нейросети
|
||||
#ifdef NEW
|
||||
#define OPT_NEW "\nОПЦИИ new\n" \
|
||||
" --corpus PATH текст для построения словаря (обязательно)\n" \
|
||||
" --out PATH файл модели (по умолчанию model.xnh)\n" \
|
||||
" --vocab N размер словаря (по умолчанию 512)\n" \
|
||||
" --embd N размер эмбеддинга (по умолчанию 128)\n" \
|
||||
" --layers N число слоёв (по умолчанию 4)\n" \
|
||||
" --heads N число голов (по умолчанию 4)\n" \
|
||||
" --ctx N максимальный контекст (по умолчанию 128)\n" \
|
||||
" --ffn N размер FFN (по умолчанию 4*embd)\n" \
|
||||
" --rope N база RoPE (по умолчанию 10000)\n" \
|
||||
" --seed N зерно инициализации (по умолчанию 1337)\n" \
|
||||
" --untied отдельная матрица выхода вместо привязанной к wte\n" \
|
||||
" --view-vocab выводит словарь нейросети\n"
|
||||
#else
|
||||
#define OPT_NEW ""
|
||||
#endif
|
||||
#ifdef TRAIN
|
||||
#define OPT_TRAIN "\nОПЦИИ train\n" \
|
||||
" --corpus PATH обучающий текст (обязательно)\n" \
|
||||
" --out PATH куда сохранить (по умолчанию перезапись входного файла)\n" \
|
||||
" --steps N шагов (по умолчанию 1000)\n" \
|
||||
" --batch N батч (по умолчанию 8)\n" \
|
||||
" --block N длина окна (по умолчанию 64)\n" \
|
||||
" --lr F скорость (по умолчанию 3e-4)\n" \
|
||||
" --wd F weight decay (по умолчанию 0.01)\n" \
|
||||
" --clip F клип нормы градиента (по умолчанию 1.0)\n" \
|
||||
" --warmup N прогрев (по умолчанию 100)\n" \
|
||||
" --seed N зерно (по умолчанию 1337)\n" \
|
||||
" --ckpt-every N промежуточное сохранение каждые N шагов (0 = нет)\n" \
|
||||
" --val-every N валидация каждые N шагов (0 = нет)\n" \
|
||||
" --val-tokens N размер валидации (по умолчанию 20000)\n" \
|
||||
" --threads N потоков (0 = все ядра)\n" \
|
||||
" --resume продолжить с сохранённого step (Adam-state из файла)\n" \
|
||||
" --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5\n"
|
||||
#else
|
||||
#define OPT_TRAIN ""
|
||||
#endif
|
||||
#ifdef GEN
|
||||
#define OPT_GEN "\nОПЦИИ gen\n" \
|
||||
" --prompt STR стартовый текст\n" \
|
||||
" --n N сколько токенов (по умолчанию 200)\n" \
|
||||
" --temp F температура, 0 = жадный выбор (по умолчанию 0.8)\n" \
|
||||
" --top-k N top-k (0 = выкл)\n" \
|
||||
" --top-p F top-p (1.0 = выкл)\n" \
|
||||
" --seed N зерно\n" \
|
||||
" --no-stream не печатать в процессе\n" \
|
||||
" --batch \"a;b;c\" несколько промптов OPTчерез ;\n" \
|
||||
" --show-tokens показать id токенов\n"
|
||||
#else
|
||||
#define OPT_GEN ""
|
||||
#endif
|
||||
#ifdef GRADCHECK
|
||||
#define OPT_GRADCHECK "\nОПЦИИ GRADCHECK\n" \
|
||||
" --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5\n"
|
||||
#else
|
||||
#define OPT_GRADCHECK ""
|
||||
#endif
|
||||
#ifdef BENCH
|
||||
#define OPT_BENCH "\nОПЦИИ BENCH\n" \
|
||||
" --fast эксперементаль флаг который ускоряет foreword примерно в x2-x5\n"
|
||||
#else
|
||||
#define OPT_BENCH ""
|
||||
#endif
|
||||
#ifdef OLLAMA
|
||||
#define OPT_OLLAMA "\nОПЦИИ ollama\n" \
|
||||
" --port PORT порт на котором будет открыт сервер\n" \
|
||||
" --conf NAME.conf конфиг с настройками и списком моделей\n"
|
||||
#else
|
||||
#define OPT_OLLAMA ""
|
||||
#endif
|
||||
|
||||
ОПЦИИ train
|
||||
--corpus PATH обучающий текст (обязательно)
|
||||
--out PATH куда сохранить (по умолчанию перезапись входного файла)
|
||||
--steps N шагов (по умолчанию 1000)
|
||||
--batch N батч (по умолчанию 8)
|
||||
--block N длина окна (по умолчанию 64)
|
||||
--lr F скорость (по умолчанию 3e-4)
|
||||
--wd F weight decay (по умолчанию 0.01)
|
||||
--clip F клип нормы градиента (по умолчанию 1.0)
|
||||
--warmup N прогрев (по умолчанию 100)
|
||||
--seed N зерно (по умолчанию 1337)
|
||||
--log-every N как часто печатать (по умолчанию 50)
|
||||
--ckpt-every N промежуточное сохранение каждые N шагов (0 = нет)
|
||||
--val-every N валидация каждые N шагов (0 = нет)
|
||||
--val-tokens N размер валидации (по умолчанию 20000)
|
||||
--threads N потоков (0 = все ядра)
|
||||
--resume продолжить с сохранённого step (Adam-state из файла)
|
||||
|
||||
ОПЦИИ gen
|
||||
--prompt STR стартовый текст
|
||||
--n N сколько токенов (по умолчанию 200)
|
||||
--temp F температура, 0 = жадный выбор (по умолчанию 0.8)
|
||||
--top-k N top-k (0 = выкл)
|
||||
--top-p F top-p (1.0 = выкл)
|
||||
--seed N зерно
|
||||
--no-stream не печатать в процессе
|
||||
--batch "a;b;c" несколько промптов через ;
|
||||
--show-tokens показать id токенов
|
||||
|
||||
ОПЦИИ ollama
|
||||
--port PORT порт на котором будет открыт сервер
|
||||
--conf NAME.conf конфиг с настройками и списком моделей
|
||||
)";
|
||||
#define HELP_TEXT \
|
||||
"ИСПОЛЬЗОВАНИЕ\n" \
|
||||
" ./bin/Xenith <команда> [опции]\n\n" \
|
||||
"КОМАНДЫ\n" \
|
||||
CMD_NEW \
|
||||
CMD_TRAIN \
|
||||
CMD_GEN \
|
||||
CMD_INFO \
|
||||
CMD_GRADCHECK \
|
||||
CMD_BENCH \
|
||||
CMD_OLLAMA \
|
||||
CMD_UPDATE \
|
||||
CMD_TUI \
|
||||
OPT_NEW \
|
||||
OPT_TRAIN \
|
||||
OPT_GEN \
|
||||
OPT_GRADCHECK \
|
||||
OPT_BENCH \
|
||||
OPT_OLLAMA \
|
||||
"\n"
|
||||
|
||||
#endif
|
||||
88
main.cpp
88
main.cpp
@ -1,7 +1,8 @@
|
||||
#include "xenith/core/checkpoint.h"
|
||||
#include "xenith/core/generate.h"
|
||||
#include "xenith/core/train.h"
|
||||
#include "xenith/core/gradcheck.h"
|
||||
#include "xenith/core/checkpoint.hpp"
|
||||
#include "xenith/core/generate.hpp"
|
||||
#include "xenith/core/train.hpp"
|
||||
#include "xenith/core/gradcheck.hpp"
|
||||
#include "xenith/ui/tui.hpp"
|
||||
#include <cstdio>
|
||||
#include "help.h"
|
||||
#include <cstdlib>
|
||||
@ -148,6 +149,7 @@ static int cmd_train(const Args& a) {
|
||||
tc.val_every = a.geti("val-every", 0);
|
||||
tc.val_tokens = a.geti("val-tokens", 20000);
|
||||
tc.threads = a.geti("threads", 0);
|
||||
tc.fast = a.has("fast");
|
||||
|
||||
if (tc.threads > 0) set_threads(tc.threads);
|
||||
if (tc.block > ck.params.cfg.block_size) tc.block = ck.params.cfg.block_size;
|
||||
@ -228,7 +230,7 @@ static int cmd_info(const Args& a) {
|
||||
}
|
||||
|
||||
static int cmd_gradcheck(const Args& a) {
|
||||
return run_gradcheck(a.has("verbose"));
|
||||
return run_gradcheck(a.has("verbose"), a.has("fast"));
|
||||
}
|
||||
|
||||
static int cmd_bench(const Args& a) {
|
||||
@ -236,7 +238,7 @@ static int cmd_bench(const Args& a) {
|
||||
Checkpoint ck;
|
||||
std::string err;
|
||||
if (!load_checkpoint(a.pos[0], ck, err)) { fprintf(stderr, "bench: %s\n", err.c_str()); return 1; }
|
||||
return run_bench(ck, a.geti("iters", 20), a.geti("block", 64));
|
||||
return run_bench(ck, a.geti("iters", 20), a.geti("block", 64), a.has("fast"));
|
||||
}
|
||||
|
||||
static int cmd_ollama(const Args& a) {
|
||||
@ -263,7 +265,6 @@ Version parse_version(const std::string& str) {
|
||||
return v;
|
||||
}
|
||||
|
||||
// Правильное сравнение версий
|
||||
bool is_newer(const Version& remote, const Version& current) {
|
||||
if (remote.major != current.major) return remote.major > current.major;
|
||||
if (remote.minor != current.minor) return remote.minor > current.minor;
|
||||
@ -301,12 +302,6 @@ static bool available_new_version() {
|
||||
|
||||
Version remote_ver = parse_version(remote_str);
|
||||
Version local_ver = parse_version(local_str);
|
||||
|
||||
std::cout << "Локальная версия: " << local_ver.major << "."
|
||||
<< local_ver.minor << "." << local_ver.patch << std::endl;
|
||||
std::cout << "Удалённая версия: " << remote_ver.major << "."
|
||||
<< remote_ver.minor << "." << remote_ver.patch << std::endl;
|
||||
|
||||
return is_newer(remote_ver, local_ver);
|
||||
}
|
||||
|
||||
@ -336,7 +331,17 @@ static int cmd_update(const Args& a) {
|
||||
const std::string REPO_URL = "https://git.bipfr.ru/BIPfR/Xenith.git";
|
||||
const std::string TMP_CLONE = "/tmp/xenith_update";
|
||||
|
||||
system(("rm -rf " + TMP_CLONE).c_str());
|
||||
int ret = system(("rm -rf " + TMP_CLONE).c_str());
|
||||
if (ret == -1)
|
||||
{
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
} else {
|
||||
int exit_code = WEXITSTATUS(ret);
|
||||
|
||||
if (exit_code != 0) {
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
}
|
||||
}
|
||||
|
||||
std::cout << "Скачивание новой версии..." << std::endl;
|
||||
int result = system(("git clone --depth 1 " + REPO_URL + " " + TMP_CLONE).c_str());
|
||||
@ -348,18 +353,62 @@ static int cmd_update(const Args& a) {
|
||||
result = system(("cp -ra " + TMP_CLONE + "/. .").c_str());
|
||||
if (result != 0) {
|
||||
std::cerr << "Ошибка копирования файлов" << std::endl;
|
||||
system(("rm -rf " + TMP_CLONE).c_str());
|
||||
ret = system(("rm -rf " + TMP_CLONE).c_str());
|
||||
if (ret == -1)
|
||||
{
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
} else {
|
||||
int exit_code = WEXITSTATUS(ret);
|
||||
|
||||
if (exit_code != 0) {
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
}
|
||||
}
|
||||
return -1;
|
||||
}
|
||||
system(("rm -rf " + TMP_CLONE).c_str());
|
||||
ret = system(("rm -rf " + TMP_CLONE).c_str());
|
||||
if (ret == -1)
|
||||
{
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
} else {
|
||||
int exit_code = WEXITSTATUS(ret);
|
||||
|
||||
if (exit_code != 0) {
|
||||
std::cout << "Ошибка выполнения команду терминала" << std::endl;
|
||||
}
|
||||
}
|
||||
std::cout << "Обновление завершено!" << std::endl;
|
||||
return 0;
|
||||
}
|
||||
|
||||
static int cmd_tui(const Args& a)
|
||||
{
|
||||
auto app = ui::Tui();
|
||||
while (app.is_running())
|
||||
{
|
||||
app.update();
|
||||
}
|
||||
std::cout << "\033[2J\033[H" << std::flush;
|
||||
return 0;
|
||||
}
|
||||
|
||||
int main(int argc, char** argv) {
|
||||
if (argc < 2) { printf(HELP_TEXT.c_str()); return 1; }
|
||||
if (available_new_version()) {
|
||||
const std::string TMP_FILE = "./tmp/remote_version.txt";
|
||||
std::ifstream tmp_file(TMP_FILE);
|
||||
std::string remote_str = "0.0.0";
|
||||
if (tmp_file.is_open() && std::getline(tmp_file, remote_str)) {
|
||||
tmp_file.close();
|
||||
} else {
|
||||
std::cerr << "Ошибка чтения файла новой версии (Пожалуйста напишите о этом баге)" << std::endl;
|
||||
return 1;
|
||||
}
|
||||
std::cout << "\033[32mДоступна новая версия \033[33m" << remote_str << "\033[32m\nВыполните \"Xenith update\" чтобы обновится до последней версии\033[39m" << std::endl;
|
||||
}
|
||||
|
||||
if (argc < 2) { printf(HELP_TEXT); return 1; }
|
||||
const std::string cmd = argv[1];
|
||||
if (cmd == "-h" || cmd == "--help" || cmd == "help") { printf(HELP_TEXT.c_str()); return 0; }
|
||||
if (cmd == "-h" || cmd == "--help" || cmd == "help") { printf(HELP_TEXT); return 0; }
|
||||
|
||||
Args a(argc, argv, 2);
|
||||
if (cmd == "new") return cmd_new(a);
|
||||
@ -370,8 +419,9 @@ int main(int argc, char** argv) {
|
||||
if (cmd == "bench") return cmd_bench(a);
|
||||
if (cmd == "ollama") return cmd_ollama(a);
|
||||
if (cmd == "update") return cmd_update(a);
|
||||
if (cmd == "tui") return cmd_tui(a);
|
||||
|
||||
fprintf(stderr, "неизвестная команда: %s\n\n", cmd.c_str());
|
||||
printf(HELP_TEXT.c_str());
|
||||
printf(HELP_TEXT);
|
||||
return 1;
|
||||
}
|
||||
15
settings.h
15
settings.h
@ -1,15 +0,0 @@
|
||||
#ifndef XENITH_SETTINGS_H
|
||||
#define XENITH_SETTINGS_H
|
||||
#pragma once
|
||||
#include "xenith/preprocessing/tokenizer.h"
|
||||
|
||||
struct llm_prm
|
||||
{
|
||||
bool use_mmap = true;
|
||||
bool use_mlock = false;
|
||||
int num_thread = 1;
|
||||
unsigned int num_gpu = 256;
|
||||
|
||||
};
|
||||
|
||||
#endif
|
||||
@ -1 +1 @@
|
||||
0.0.4
|
||||
0.0.7
|
||||
@ -1,4 +1,4 @@
|
||||
#include "converter.h"
|
||||
#include "converter.hpp"
|
||||
|
||||
#include <cstring>
|
||||
#include <iostream>
|
||||
@ -51,7 +51,14 @@ Model load_from_bfr(const std::string& filename)
|
||||
}
|
||||
|
||||
uint32_t map_size;
|
||||
fread(&map_size, sizeof(map_size), 1, mdl.fd);
|
||||
if (fread(&map_size, sizeof(map_size), 1, mdl.fd))
|
||||
{
|
||||
|
||||
} else
|
||||
{
|
||||
std::cerr << "Ошибка чтения файла" << std::endl;
|
||||
return mdl;
|
||||
}
|
||||
|
||||
if (map_size != mdl.config.vocab_size) {
|
||||
std::cerr << "Warning: vocab size mismatch! "
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
#include "train.h"
|
||||
#include "train.hpp"
|
||||
#include <cmath>
|
||||
#include <algorithm>
|
||||
|
||||
@ -17,7 +17,7 @@ static float softmax_ce(const float* logits, int V, int target, float* dlogits)
|
||||
}
|
||||
|
||||
float backward(Params& p, const ForwardCache& fc,
|
||||
const int* x, const int* y, int n_pairs) {
|
||||
const int* x, const int* y, int n_pairs, bool fast) {
|
||||
const Config& c = p.cfg;
|
||||
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
||||
const int V = c.vocab_out();
|
||||
@ -40,14 +40,14 @@ float backward(Params& p, const ForwardCache& fc,
|
||||
Tensor dxf(B, d);
|
||||
if (p.tied()) {
|
||||
gemm_nn(dlogits, p.wte, dxf);
|
||||
gemm_tn(dlogits, fc.xf, p.gwte, true);
|
||||
gemm_tn(dlogits, fc.xf, p.gwte, true, fast);
|
||||
} else {
|
||||
gemm_nn(dlogits, p.lm_head, dxf);
|
||||
gemm_tn(dlogits, fc.xf, p.glm_head, true);
|
||||
gemm_tn(dlogits, fc.xf, p.glm_head, true, fast);
|
||||
}
|
||||
|
||||
Tensor dx;
|
||||
rmsnorm_backward(fc.x, p.rms_final, fc.inv_rms, dxf, dx, p.grms_final);
|
||||
rmsnorm_backward(fc.x, p.rms_final, fc.inv_rms, dxf, dx, p.grms_final, fast);
|
||||
|
||||
Tensor tmp, dxb, dx2b, dha, dh1, dh3, dq, dk, dv, datt, dattout;
|
||||
|
||||
@ -56,7 +56,7 @@ float backward(Params& p, const ForwardCache& fc,
|
||||
|
||||
const Tensor dres = dx;
|
||||
|
||||
gemm_tn(lc.ha, dres, p.g2[L], true); // dW2
|
||||
gemm_tn(lc.ha, dres, p.g2[L], true, fast); // dW2
|
||||
gemm_nt(dres, p.w2[L], dha); // dha = dres @ w2^T
|
||||
|
||||
dh1.resize(B, f);
|
||||
@ -73,8 +73,8 @@ float backward(Params& p, const ForwardCache& fc,
|
||||
d3[j] = ga * silu(g);
|
||||
}
|
||||
}
|
||||
gemm_tn(lc.x2b, dh1, p.g1[L], true); // dW1
|
||||
gemm_tn(lc.x2b, dh3, p.g3[L], true); // dW3
|
||||
gemm_tn(lc.x2b, dh1, p.g1[L], true, fast); // dW1
|
||||
gemm_tn(lc.x2b, dh3, p.g3[L], true, fast); // dW3
|
||||
|
||||
dx2b.resize(B, d);
|
||||
gemm_nt(dh1, p.w1[L], dx2b);
|
||||
@ -83,11 +83,11 @@ float backward(Params& p, const ForwardCache& fc,
|
||||
|
||||
{
|
||||
Tensor g;
|
||||
rmsnorm_backward(lc.x2, p.rms_ffn[L], lc.inv_rms2, dx2b, g, p.grms_ffn[L]);
|
||||
rmsnorm_backward(lc.x2, p.rms_ffn[L], lc.inv_rms2, dx2b, g, p.grms_ffn[L], fast);
|
||||
for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i);
|
||||
}
|
||||
|
||||
gemm_tn(lc.attout, dx, p.go[L], true);
|
||||
gemm_tn(lc.attout, dx, p.go[L], true, fast);
|
||||
dattout.resize(B, d);
|
||||
gemm_nt(dx, p.wo[L], dattout);
|
||||
|
||||
@ -161,9 +161,9 @@ float backward(Params& p, const ForwardCache& fc,
|
||||
}
|
||||
}
|
||||
|
||||
gemm_tn(lc.xb, dq, p.gq[L], true);
|
||||
gemm_tn(lc.xb, dk, p.gk[L], true);
|
||||
gemm_tn(lc.xb, dv, p.gv[L], true);
|
||||
gemm_tn(lc.xb, dq, p.gq[L], true, fast);
|
||||
gemm_tn(lc.xb, dk, p.gk[L], true, fast);
|
||||
gemm_tn(lc.xb, dv, p.gv[L], true, fast);
|
||||
|
||||
dxb.resize(B, d);
|
||||
gemm_nt(dq, p.wq[L], dxb);
|
||||
@ -174,7 +174,7 @@ float backward(Params& p, const ForwardCache& fc,
|
||||
|
||||
{
|
||||
Tensor g;
|
||||
rmsnorm_backward(lc.x, p.rms_attn[L], lc.inv_rms, dxb, g, p.grms_attn[L]);
|
||||
rmsnorm_backward(lc.x, p.rms_attn[L], lc.inv_rms, dxb, g, p.grms_attn[L], fast);
|
||||
for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i);
|
||||
}
|
||||
}
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
#include "checkpoint.h"
|
||||
#include "checkpoint.hpp"
|
||||
#include <cstdio>
|
||||
#include <cstring>
|
||||
#include <fstream>
|
||||
@ -241,7 +241,7 @@ bool load_checkpoint(const std::string& path, Checkpoint& ck, std::string& err)
|
||||
if (magic != XNH_MAGIC) {
|
||||
char buf[64];
|
||||
std::snprintf(buf, sizeof buf,
|
||||
"плохая магия 0x%08X (ожидалась 0x%08X) — это не Xenith-модель",
|
||||
"Файл не является поддерживаемым",
|
||||
magic, XNH_MAGIC);
|
||||
err = buf;
|
||||
return false;
|
||||
|
||||
@ -1,5 +1,5 @@
|
||||
#pragma once
|
||||
#include "model.h"
|
||||
#include "model.hpp"
|
||||
#include <string>
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
@ -1,4 +1,4 @@
|
||||
#include "generate.h"
|
||||
#include "generate.hpp"
|
||||
#include <cstdio>
|
||||
#include <cmath>
|
||||
#include <algorithm>
|
||||
@ -44,7 +44,7 @@ void rope_apply(float* x, int hd, int pos, const std::vector<float>& cs,
|
||||
|
||||
void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||
const std::vector<float>& cs, const std::vector<float>& sn,
|
||||
Tensor& logits) {
|
||||
Tensor& logits, bool fast) {
|
||||
const Config& c = p.cfg;
|
||||
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
||||
const float scale = 1.0f / std::sqrt((float)hd);
|
||||
@ -58,7 +58,7 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||
Tensor ln, lnf;
|
||||
|
||||
for (int L = 0; L < c.n_layer; ++L) {
|
||||
rmsnorm_forward(x, p.rms_attn[L], c.rms_eps, xb, ln);
|
||||
rmsnorm_forward(x, p.rms_attn[L], c.rms_eps, xb, ln, fast);
|
||||
gemm_nn(xb, p.wq[L], q);
|
||||
gemm_nn(xb, p.wk[L], k);
|
||||
gemm_nn(xb, p.wv[L], v);
|
||||
@ -71,8 +71,8 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||
for (int j = 0; j < hd; ++j) { kp[j] = k.at(h * hd + j); vp[j] = v.at(h * hd + j); }
|
||||
}
|
||||
|
||||
// attention по всем позициям 0..pos
|
||||
attout.resize(1, d);
|
||||
if (fast) attout.resize_fast(1, d);
|
||||
else attout.resize(1, d);
|
||||
attout.zero();
|
||||
for (int h = 0; h < H; ++h) {
|
||||
const float* qr = q.row(0) + h * hd;
|
||||
@ -99,17 +99,19 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
||||
gemm_nn(attout, p.wo[L], proj);
|
||||
for (int j = 0; j < d; ++j) x.at(j) += proj.at(j);
|
||||
|
||||
rmsnorm_forward(x, p.rms_ffn[L], c.rms_eps, x2b, ln);
|
||||
rmsnorm_forward(x, p.rms_ffn[L], c.rms_eps, x2b, ln, fast);
|
||||
gemm_nn(x2b, p.w1[L], h1);
|
||||
gemm_nn(x2b, p.w3[L], h3);
|
||||
ha.resize(1, f);
|
||||
if (fast) attout.resize_fast(1, f);
|
||||
else attout.resize(1, f);
|
||||
for (size_t i = 0; i < ha.n(); ++i) ha.at(i) = silu(h1.at(i)) * h3.at(i);
|
||||
gemm_nn(ha, p.w2[L], fo);
|
||||
for (int j = 0; j < d; ++j) x.at(j) += fo.at(j);
|
||||
}
|
||||
|
||||
rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln);
|
||||
logits.resize(1, c.vocab_out());
|
||||
rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln, fast);
|
||||
if (fast) attout.resize_fast(1, c.vocab_out());
|
||||
else attout.resize(1, c.vocab_out());
|
||||
if (p.tied()) gemm_nt(lnf, p.wte, logits);
|
||||
else gemm_nn(lnf, p.lm_head, logits);
|
||||
}
|
||||
@ -184,7 +186,7 @@ int sample(const Tensor& logits, const GenConfig& gc, Rng& rng, const std::vecto
|
||||
|
||||
std::string generate(const Params& p, const Tokenizer& tok,
|
||||
const std::string& prompt, const GenConfig& gc,
|
||||
std::vector<int>* out_ids) {
|
||||
std::vector<int>* out_ids, bool fast) {
|
||||
const Config& c = p.cfg;
|
||||
KvCache kv;
|
||||
kv.init(c.n_layer, c.block_size, c.n_head, c.head_dim());
|
||||
@ -213,13 +215,13 @@ std::string generate(const Params& p, const Tokenizer& tok,
|
||||
ids.resize(keep);
|
||||
kv.reset();
|
||||
for (int t = 0; t < keep; ++t) {
|
||||
forward_token(p, kv, ids[t], kv.len, cs, sn, logits);
|
||||
forward_token(p, kv, ids[t], kv.len, cs, sn, logits, fast);
|
||||
kv.len++;
|
||||
}
|
||||
}
|
||||
|
||||
const int last = ids.back();
|
||||
forward_token(p, kv, last, kv.len, cs, sn, logits);
|
||||
forward_token(p, kv, last, kv.len, cs, sn, logits, fast);
|
||||
kv.len++;
|
||||
|
||||
int next = sample(logits, gc, rng, ids);
|
||||
|
||||
@ -1,6 +1,6 @@
|
||||
// generate.h — генерация текста с KV-кэшем
|
||||
#pragma once
|
||||
#include "checkpoint.h"
|
||||
#include "checkpoint.hpp"
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
@ -18,7 +18,7 @@ struct GenConfig {
|
||||
|
||||
std::string generate(const Params& p, const Tokenizer& tok,
|
||||
const std::string& prompt, const GenConfig& gc,
|
||||
std::vector<int>* out_ids = nullptr);
|
||||
std::vector<int>* out_ids = nullptr, bool fast = false);
|
||||
|
||||
std::vector<std::string> generate_batch(const Params& p, const Tokenizer& tok,
|
||||
const std::vector<std::string>& prompts,
|
||||
@ -1,8 +1,8 @@
|
||||
// gradcheck.cpp — проверяем, что обратный проход совпадает с численным
|
||||
// и заодно меряем скорость.
|
||||
#include "gradcheck.h"
|
||||
#include "generate.h"
|
||||
#include "train.h"
|
||||
#include "gradcheck.hpp"
|
||||
#include "generate.hpp"
|
||||
#include "train.hpp"
|
||||
#include <cstdio>
|
||||
#include <cmath>
|
||||
#include <chrono>
|
||||
@ -14,9 +14,9 @@ namespace xt {
|
||||
namespace {
|
||||
|
||||
// loss на фиксированном батче, без обратного прохода
|
||||
float eval_loss(const Params& p, const std::vector<int>& x, const std::vector<int>& y, int n) {
|
||||
float eval_loss(const Params& p, const std::vector<int>& x, const std::vector<int>& y, int n, bool fast) {
|
||||
ForwardCache fc;
|
||||
forward(p, x.data(), n, fc);
|
||||
forward(p, x.data(), n, fc, fast);
|
||||
Tensor loss;
|
||||
return softmax_eval(p, fc, x.data(), y.data(), n, loss);
|
||||
}
|
||||
@ -30,7 +30,7 @@ struct Probe {
|
||||
|
||||
}
|
||||
|
||||
int run_gradcheck(bool verbose) {
|
||||
int run_gradcheck(bool verbose, bool fast) {
|
||||
Config c;
|
||||
c.vocab_size = 24;
|
||||
c.n_layer = 2;
|
||||
@ -53,8 +53,8 @@ int run_gradcheck(bool verbose) {
|
||||
}
|
||||
|
||||
ForwardCache fc;
|
||||
forward(p, x.data(), n, n, fc);
|
||||
backward(p, fc, x.data(), y.data(), n);
|
||||
forward(p, x.data(), n, n, fc, fast);
|
||||
backward(p, fc, x.data(), y.data(), n, fast);
|
||||
|
||||
auto tensors = p.all_w();
|
||||
auto grads = p.all_grads();
|
||||
@ -80,9 +80,9 @@ int run_gradcheck(bool verbose) {
|
||||
const float orig = W.at(i);
|
||||
|
||||
W.at(i) = orig + eps;
|
||||
const float lp = eval_loss(p, x, y, n);
|
||||
const float lp = eval_loss(p, x, y, n, fast);
|
||||
W.at(i) = orig - eps;
|
||||
const float lm = eval_loss(p, x, y, n);
|
||||
const float lm = eval_loss(p, x, y, n, fast);
|
||||
W.at(i) = orig;
|
||||
|
||||
const float num = (lp - lm) / (2.0f * eps);
|
||||
@ -112,7 +112,7 @@ int run_gradcheck(bool verbose) {
|
||||
return 1;
|
||||
}
|
||||
|
||||
int run_bench(Checkpoint& ck, int iters, int block) {
|
||||
int run_bench(Checkpoint& ck, int iters, int block, bool fast) {
|
||||
const Config& c = ck.params.cfg;
|
||||
if (block > c.block_size) block = c.block_size;
|
||||
const int n = block;
|
||||
@ -129,19 +129,19 @@ int run_bench(Checkpoint& ck, int iters, int block) {
|
||||
|
||||
// прогрев
|
||||
ForwardCache fc;
|
||||
forward(ck.params, x.data(), n, n, fc);
|
||||
backward(ck.params, fc, x.data(), y.data(), n);
|
||||
forward(ck.params, x.data(), n, n, fc, fast);
|
||||
backward(ck.params, fc, x.data(), y.data(), n, fast);
|
||||
|
||||
auto t0 = std::chrono::steady_clock::now();
|
||||
for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc);
|
||||
for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc, fast);
|
||||
auto t1 = std::chrono::steady_clock::now();
|
||||
double d_fwd = std::chrono::duration<double>(t1 - t0).count();
|
||||
float fwd_tps = (double)n * iters / d_fwd;
|
||||
|
||||
t0 = std::chrono::steady_clock::now();
|
||||
for (int i = 0; i < iters; ++i) {
|
||||
forward(ck.params, x.data(), n, n, fc);
|
||||
backward(ck.params, fc, x.data(), y.data(), n);
|
||||
forward(ck.params, x.data(), n, n, fc, fast);
|
||||
backward(ck.params, fc, x.data(), y.data(), n, fast);
|
||||
}
|
||||
t1 = std::chrono::steady_clock::now();
|
||||
double d_all = std::chrono::duration<double>(t1 - t0).count();
|
||||
|
||||
@ -1,8 +0,0 @@
|
||||
// gradcheck.h — численная проверка градиентов и бенчмарк
|
||||
#pragma once
|
||||
#include "checkpoint.h"
|
||||
|
||||
namespace xt {
|
||||
int run_gradcheck(bool verbose);
|
||||
int run_bench(Checkpoint& ck, int iters, int block);
|
||||
}
|
||||
8
xenith/core/gradcheck.hpp
Normal file
8
xenith/core/gradcheck.hpp
Normal file
@ -0,0 +1,8 @@
|
||||
// gradcheck.h — численная проверка градиентов и бенчмарк
|
||||
#pragma once
|
||||
#include "checkpoint.hpp"
|
||||
|
||||
namespace xt {
|
||||
int run_gradcheck(bool verbose, bool fast);
|
||||
int run_bench(Checkpoint& ck, int iters, int block, bool fast);
|
||||
}
|
||||
@ -1,4 +1,4 @@
|
||||
#include "model.h"
|
||||
#include "model.hpp"
|
||||
#include <cstdio>
|
||||
#include <cmath>
|
||||
#include <algorithm>
|
||||
@ -204,7 +204,7 @@ void rope_tables(int block, int head_dim, int base, std::vector<float>& cs, std:
|
||||
}
|
||||
}
|
||||
|
||||
void forward(const Params& p, const int* tokens, int n_tok, int seq_len, ForwardCache& fc) {
|
||||
void forward(const Params& p, const int* tokens, int n_tok, int seq_len, ForwardCache& fc, bool fast) {
|
||||
const Config& c = p.cfg;
|
||||
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
||||
int B = n_tok;
|
||||
@ -226,7 +226,8 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
const float* CS = fc.rope_cos.data();
|
||||
const float* SN = fc.rope_sin.data();
|
||||
|
||||
fc.x.resize(B, d);
|
||||
if (fast) fc.x.resize_fast(B, d);
|
||||
else fc.x.resize(B, d);
|
||||
for (int t = 0; t < B; ++t) {
|
||||
float* xr = fc.x.row(t);
|
||||
const float* er = p.wte.row(tokens[t]);
|
||||
@ -238,10 +239,47 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
|
||||
for (int L = 0; L < c.n_layer; ++L) {
|
||||
LayerCache& lc = fc.layers[L];
|
||||
lc.x.resize(B, d);
|
||||
for (size_t i = 0; i < fc.x.n(); ++i) lc.x.at(i) = fc.x.at(i);
|
||||
|
||||
rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms);
|
||||
if (fast)
|
||||
{
|
||||
lc.x.resize_fast(B, d);
|
||||
lc.xb.resize_fast(B, d);
|
||||
lc.q.resize_fast(B, d);
|
||||
lc.k.resize_fast(B, d);
|
||||
lc.v.resize_fast(B, d);
|
||||
lc.attout.resize_fast(B, d);
|
||||
lc.proj.resize_fast(B, d);
|
||||
|
||||
lc.x2.resize_fast(B, d);
|
||||
lc.x2b.resize_fast(B, d);
|
||||
lc.h1.resize_fast(B, f);
|
||||
lc.h3.resize_fast(B, f);
|
||||
lc.ha.resize_fast(B, f);
|
||||
lc.fo.resize_fast(B, d);
|
||||
} else
|
||||
{
|
||||
lc.x.resize(B, d);
|
||||
lc.xb.resize(B, d);
|
||||
lc.q.resize(B, d);
|
||||
lc.k.resize(B, d);
|
||||
lc.v.resize(B, d);
|
||||
lc.attout.resize(B, d);
|
||||
lc.proj.resize(B, d);
|
||||
|
||||
lc.x2.resize(B, d);
|
||||
lc.x2b.resize(B, d);
|
||||
lc.h1.resize(B, f);
|
||||
lc.h3.resize(B, f);
|
||||
lc.ha.resize(B, f);
|
||||
lc.fo.resize(B, d);
|
||||
}
|
||||
|
||||
lc.att.resize3d(n_seq * H, S, S);
|
||||
|
||||
if (fast) fc.xf.resize_fast(B, d);
|
||||
else fc.xf.resize(B, d);
|
||||
|
||||
rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms, fast);
|
||||
gemm_nn(lc.xb, p.wq[L], lc.q);
|
||||
gemm_nn(lc.xb, p.wk[L], lc.k);
|
||||
gemm_nn(lc.xb, p.wv[L], lc.v);
|
||||
@ -268,7 +306,8 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
}
|
||||
|
||||
lc.att.resize3d(n_seq * H, S, S);
|
||||
lc.attout.resize(B, d);
|
||||
if (fast) fc.x.resize_fast(B, d);
|
||||
else fc.x.resize(B, d);
|
||||
lc.attout.zero();
|
||||
|
||||
for (int s = 0; s < n_seq; ++s) {
|
||||
@ -307,32 +346,35 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
||||
gemm_nn(lc.attout, p.wo[L], lc.proj);
|
||||
for (size_t i = 0; i < fc.x.n(); ++i) fc.x.at(i) += lc.proj.at(i);
|
||||
|
||||
lc.x2.resize(B, d);
|
||||
if (fast) fc.x.resize_fast(B, d);
|
||||
else fc.x.resize(B, d);
|
||||
for (size_t i = 0; i < fc.x.n(); ++i) lc.x2.at(i) = fc.x.at(i);
|
||||
|
||||
rmsnorm_forward(lc.x2, p.rms_ffn[L], c.rms_eps, lc.x2b, lc.inv_rms2);
|
||||
rmsnorm_forward(lc.x2, p.rms_ffn[L], c.rms_eps, lc.x2b, lc.inv_rms2, fast);
|
||||
gemm_nn(lc.x2b, p.w1[L], lc.h1);
|
||||
gemm_nn(lc.x2b, p.w3[L], lc.h3);
|
||||
lc.ha.resize(B, f);
|
||||
if (fast) fc.x.resize_fast(B, f);
|
||||
else fc.x.resize(B, f);
|
||||
for (size_t i = 0; i < lc.ha.n(); ++i)
|
||||
lc.ha.at(i) = silu(lc.h1.at(i)) * lc.h3.at(i);
|
||||
gemm_nn(lc.ha, p.w2[L], lc.fo);
|
||||
for (size_t i = 0; i < fc.x.n(); ++i) fc.x.at(i) += lc.fo.at(i);
|
||||
}
|
||||
|
||||
rmsnorm_forward(fc.x, p.rms_final, c.rms_eps, fc.xf, fc.inv_rms);
|
||||
fc.logits.resize(B, c.vocab_out());
|
||||
rmsnorm_forward(fc.x, p.rms_final, c.rms_eps, fc.xf, fc.inv_rms, fast);
|
||||
if (fast) fc.x.resize_fast(B, c.vocab_out());
|
||||
else fc.x.resize(B, c.vocab_out());
|
||||
if (p.tied()) gemm_nt(fc.xf, p.wte, fc.logits);
|
||||
else gemm_nn(fc.xf, p.lm_head, fc.logits);
|
||||
}
|
||||
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc) {
|
||||
forward(p, tokens, n_tok, n_tok, fc);
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc, bool fast) {
|
||||
forward(p, tokens, n_tok, n_tok, fc, fast);
|
||||
}
|
||||
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits) {
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits, bool fast) {
|
||||
ForwardCache fc;
|
||||
forward(p, tokens, n_tok, n_tok, fc);
|
||||
forward(p, tokens, n_tok, n_tok, fc, fast);
|
||||
const int V = p.cfg.vocab_out();
|
||||
logits.resize(1, V);
|
||||
for (int j = 0; j < V; ++j) logits.at(j) = fc.logits.at((n_tok - 1) * V + j);
|
||||
|
||||
@ -1,5 +1,5 @@
|
||||
#pragma once
|
||||
#include "tensor.h"
|
||||
#include "tensor.hpp"
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
@ -81,10 +81,10 @@ struct ForwardCache {
|
||||
|
||||
void rope_tables(int block, int head_dim, int base,
|
||||
std::vector<float>& cs, std::vector<float>& sn);
|
||||
void forward(const Params& p, const int* tokens, int n_tok, int seq_len, ForwardCache& fc);
|
||||
void forward(const Params& p, const int* tokens, int n_tok, int seq_len, ForwardCache& fc, bool fast);
|
||||
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc);
|
||||
void forward(const Params& p, const int* tokens, int n_tok, ForwardCache& fc, bool fast);
|
||||
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits);
|
||||
void forward_last(const Params& p, const int* tokens, int n_tok, Tensor& logits, bool fast);
|
||||
|
||||
}
|
||||
@ -71,6 +71,7 @@ struct Tensor {
|
||||
Tensor(int r, int c) : rows(r), cols(c), d((size_t)r * c, 0.0f) {}
|
||||
|
||||
void resize(int r, int c) { rows = r; cols = c; d.assign((size_t)r * c, 0.0f); }
|
||||
void resize_fast(int r, int c) { if (rows != r || cols != c) { rows = r; cols = c; d.assign((size_t)r * c, 0.0f);} else {std::fill(d.begin(), d.end(), 0.0f); }}
|
||||
void resize3d(int h, int r, int c) { rows = r; cols = c; d.assign((size_t)h * r * c, 0.0f); }
|
||||
void zero() { std::fill(d.begin(), d.end(), 0.0f); }
|
||||
int R() const { return rows; }
|
||||
@ -131,10 +132,11 @@ inline void gemm_nt(const Tensor& A, const Tensor& B, Tensor& C) {
|
||||
});
|
||||
}
|
||||
|
||||
inline void gemm_tn(const Tensor& A, const Tensor& B, Tensor& C, bool accumulate = false) {
|
||||
inline void gemm_tn(const Tensor& A, const Tensor& B, Tensor& C, bool accumulate, bool fast) {
|
||||
const int N = A.R(), K = A.C(), M = B.C();
|
||||
if (!accumulate) {
|
||||
C.resize(K, M);
|
||||
if (fast) C.resize_fast(K, M);
|
||||
else C.resize(K, M);
|
||||
} else {
|
||||
if (C.R() != K || C.C() != M) {
|
||||
std::fprintf(stderr,
|
||||
@ -163,11 +165,15 @@ inline float dsilu(float x) {
|
||||
return s * (1.0f + x * (1.0f - s));
|
||||
}
|
||||
|
||||
|
||||
|
||||
inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps,
|
||||
Tensor& y, Tensor& inv_rms) {
|
||||
Tensor& y, Tensor& inv_rms, bool fast) {
|
||||
const int N = x.R(), D = x.C();
|
||||
y.resize(N, D);
|
||||
inv_rms.resize(N, 1);
|
||||
if (fast) y.resize_fast(N, D);
|
||||
else y.resize(N, D);
|
||||
if (fast) inv_rms.resize_fast(N, 1);
|
||||
else inv_rms.resize(N, 1);
|
||||
for (int n = 0; n < N; ++n) {
|
||||
const float* xr = x.row(n);
|
||||
float ss = 0.0f;
|
||||
@ -180,10 +186,12 @@ inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps,
|
||||
}
|
||||
|
||||
inline void rmsnorm_backward(const Tensor& x, const Tensor& w, const Tensor& inv_rms,
|
||||
const Tensor& dy, Tensor& dx, Tensor& dw) {
|
||||
const Tensor& dy, Tensor& dx, Tensor& dw, bool fast) {
|
||||
const int N = x.R(), D = x.C();
|
||||
dx.resize(N, D);
|
||||
dw.resize(1, D);
|
||||
if (fast) dx.resize_fast(N, D);
|
||||
else dx.resize(N, D);
|
||||
if (fast) dw.resize_fast(1, D);
|
||||
else dw.resize(1, D);
|
||||
dw.zero();
|
||||
for (int n = 0; n < N; ++n) {
|
||||
const float* xr = x.row(n);
|
||||
@ -1,12 +1,12 @@
|
||||
#pragma once
|
||||
#include "model.h"
|
||||
#include "checkpoint.h"
|
||||
#include "model.hpp"
|
||||
#include "checkpoint.hpp"
|
||||
#include <vector>
|
||||
|
||||
namespace xt {
|
||||
|
||||
float backward(Params& p, const ForwardCache& fc,
|
||||
const int* x, const int* y, int n_pairs);
|
||||
const int* x, const int* y, int n_pairs, bool fast);
|
||||
|
||||
float softmax_eval(const Params& p, const ForwardCache& fc,
|
||||
const int* x, const int* y, int n_pairs, Tensor& per_token_loss);
|
||||
@ -29,6 +29,7 @@ struct TrainConfig {
|
||||
int ckpt_every = 0;
|
||||
int val_every = 0;
|
||||
int val_tokens = 20000;
|
||||
bool fast = false;
|
||||
};
|
||||
|
||||
struct Dataset {
|
||||
@ -1,5 +1,5 @@
|
||||
#include "train.h"
|
||||
#include "checkpoint.h"
|
||||
#include "train.hpp"
|
||||
#include "checkpoint.hpp"
|
||||
#include <cstdio>
|
||||
#include <cmath>
|
||||
#include <fstream>
|
||||
@ -62,7 +62,7 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu
|
||||
val_ids.assign(ids.end() - tc.val_tokens, ids.end());
|
||||
ids.resize(ids.size() - tc.val_tokens);
|
||||
has_val = true;
|
||||
fprintf(stderr, "валидация: %zu токенов\n", val_ids.size());
|
||||
fprintf(stderr, "валидация: %zu токенов\n\033[?25l", val_ids.size());
|
||||
}
|
||||
|
||||
std::vector<int> xb(B * block), yb(B * block);
|
||||
@ -79,8 +79,8 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu
|
||||
}
|
||||
}
|
||||
|
||||
forward(p, xb.data(), B * block, block, fc);
|
||||
float loss = backward(p, fc, xb.data(), yb.data(), B * block);
|
||||
forward(p, xb.data(), B * block, block, fc, tc.fast);
|
||||
float loss = backward(p, fc, xb.data(), yb.data(), B * block, tc.fast);
|
||||
p.adam_step(lr_at(tc, step), tc.beta1, tc.beta2, tc.eps,
|
||||
tc.weight_decay, step + 1, tc.clip);
|
||||
|
||||
@ -91,7 +91,7 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu
|
||||
auto now = std::chrono::steady_clock::now();
|
||||
double dt = std::chrono::duration<double>(now - t0).count();
|
||||
double tps = dt > 0 ? st.tokens_seen / dt : 0;
|
||||
fprintf(stderr, "шаг %5d/%d loss %.4f ppl %8.2f lr %.2e %6.0f ток/с\n",
|
||||
fprintf(stderr, "шаг %5d/%d loss %.4f ppl %8.2f lr %.2e %6.0f ток/с\r",
|
||||
step, tc.steps, loss, std::exp(loss), lr_at(tc, step), tps);
|
||||
|
||||
|
||||
@ -102,7 +102,7 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu
|
||||
size_t off = (val_ids.size() - block - 1) * w / nwin;
|
||||
std::vector<int> vx(block);
|
||||
for (int t = 0; t < block; ++t) vx[t] = val_ids[off + t];
|
||||
forward(p, vx.data(), block, block, fc);
|
||||
forward(p, vx.data(), block, block, fc, tc.fast);
|
||||
Tensor lg;
|
||||
softmax_eval(p, fc, vx.data(), val_ids.data() + off + 1, block, lg);
|
||||
for (int t = 0; t < block; ++t) vl += lg.at(t);
|
||||
|
||||
@ -1,5 +1,2 @@
|
||||
//
|
||||
// Created by koder on 30.09.2026.
|
||||
//
|
||||
#include "handler.hpp"
|
||||
|
||||
#include "handler.h"
|
||||
|
||||
@ -1,13 +0,0 @@
|
||||
#ifndef XENITH_HANDLER_H
|
||||
#define XENITH_HANDLER_H
|
||||
|
||||
namespace api
|
||||
{
|
||||
class handler
|
||||
{
|
||||
|
||||
};
|
||||
}
|
||||
|
||||
|
||||
#endif //XENITH_HANDLER_H
|
||||
66
xenith/ollama_api/handler.hpp
Normal file
66
xenith/ollama_api/handler.hpp
Normal file
@ -0,0 +1,66 @@
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
namespace api
|
||||
{
|
||||
enum class QuantLevel
|
||||
{
|
||||
BF16,
|
||||
FP16,
|
||||
Q8,
|
||||
Q6,
|
||||
Q5,
|
||||
Q4,
|
||||
Q2,
|
||||
Q1,
|
||||
};
|
||||
|
||||
class QuantizationLevel
|
||||
{
|
||||
private:
|
||||
std::string value;
|
||||
static std::string enum_to_string(QuantLevel quant)
|
||||
{
|
||||
switch (quant) {
|
||||
case QuantLevel::BF16: return "BF16";
|
||||
case QuantLevel::FP16: return "FP16";
|
||||
case QuantLevel::Q8: return "Q8_0";
|
||||
case QuantLevel::Q6: return "Q6_K";
|
||||
case QuantLevel::Q5: return "Q5_K_M";
|
||||
case QuantLevel::Q4: return "Q4_K_M";
|
||||
case QuantLevel::Q2: return "Q2_K";
|
||||
case QuantLevel::Q1: return "Q1";
|
||||
default: return "UNKNOWN";
|
||||
}
|
||||
}
|
||||
public:
|
||||
QuantizationLevel() : value("UNKNOWN") {}
|
||||
QuantizationLevel(QuantLevel quant) : value(enum_to_string(quant)) {}
|
||||
QuantizationLevel(const std::string& str) : value(str) {}
|
||||
operator std::string() const { return value; }
|
||||
};
|
||||
|
||||
struct models_details
|
||||
{
|
||||
std::string parent_model;
|
||||
std::string format;
|
||||
std::string family;
|
||||
std::pmr::vector<std::string> families;
|
||||
uint32_t parameter_size;
|
||||
QuantizationLevel quantization_level;
|
||||
std::string expires_at;
|
||||
std::string size_vram;
|
||||
};
|
||||
|
||||
struct models {
|
||||
std::string name;
|
||||
std::string model;
|
||||
std::string modified_at;
|
||||
uint64_t size;
|
||||
std::string digest;
|
||||
models_details details;
|
||||
};
|
||||
}
|
||||
9
xenith/ollama_api/server.cpp
Normal file
9
xenith/ollama_api/server.cpp
Normal file
@ -0,0 +1,9 @@
|
||||
#include "server.hpp"
|
||||
|
||||
namespace api
|
||||
{
|
||||
server::server()
|
||||
{
|
||||
|
||||
}
|
||||
}
|
||||
@ -1,6 +0,0 @@
|
||||
#ifndef XENITH_SERVER_H
|
||||
#define XENITH_SERVER_H
|
||||
|
||||
|
||||
|
||||
#endif //XENITH_SERVER_H
|
||||
13
xenith/ollama_api/server.hpp
Normal file
13
xenith/ollama_api/server.hpp
Normal file
@ -0,0 +1,13 @@
|
||||
#pragma once
|
||||
|
||||
#include "handler.hpp"
|
||||
|
||||
namespace api
|
||||
{
|
||||
class server
|
||||
{
|
||||
server();
|
||||
~server();
|
||||
};
|
||||
}
|
||||
|
||||
@ -3,8 +3,8 @@
|
||||
#include <thread>
|
||||
#include <vector>
|
||||
#include <string>
|
||||
#include "libs/cpp-httplib-0.57.1/httplib.h"
|
||||
#include "libs/json-3.12.0/single_include/nlohmann/json.hpp"
|
||||
#include "./xenith/libs/cpp-httplib-0.57.1/httplib.hpp"
|
||||
#include "./xenith/libs/json-3.12.0/nlohmann/json.hpp"
|
||||
|
||||
|
||||
using json = nlohmann::json;
|
||||
@ -29,7 +29,7 @@ void log_ollama_options(const std::string& request_body) {
|
||||
}
|
||||
}
|
||||
|
||||
int main() {
|
||||
int main_() {
|
||||
httplib::Server svr;
|
||||
|
||||
svr.set_logger([](const httplib::Request& req, const httplib::Response& res) {
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
#include "tokenizer.h"
|
||||
#include "tokenizer.hpp"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cmath>
|
||||
|
||||
@ -4,7 +4,7 @@
|
||||
#include <string>
|
||||
#include <iostream>
|
||||
#include <unordered_map>
|
||||
#include "./../converter.h"
|
||||
#include "./../converter.hpp"
|
||||
#include "models/test_model.h"
|
||||
|
||||
void proccesing_promt(const std::string& text, Model& mdl, uint32_t layer_idx = 0);
|
||||
147
xenith/ui/tui.hpp
Normal file
147
xenith/ui/tui.hpp
Normal file
@ -0,0 +1,147 @@
|
||||
#pragma once
|
||||
|
||||
#include <sys/ioctl.h>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
#include <termios.h>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
#include <string>
|
||||
|
||||
namespace ui
|
||||
{
|
||||
class Mouse
|
||||
{
|
||||
private:
|
||||
struct pos_t { int x; int y; };
|
||||
inline static bool btn0 = false;
|
||||
inline static bool btn1 = false;
|
||||
inline static bool btn2 = false;
|
||||
inline static int scroll = 0;
|
||||
inline static pos_t pos = {};
|
||||
public:
|
||||
static bool is_pressed(int index_btn = 0);
|
||||
static int get_scroll();
|
||||
static void reset_scroll();
|
||||
static void parse_ansi(char buf[]);
|
||||
static pos_t get_pos();
|
||||
};
|
||||
|
||||
class Tui
|
||||
{
|
||||
public:
|
||||
enum class Color : int {
|
||||
DefaultText = 39, DefaultBg = 49, Black = 30, Red = 31,
|
||||
Green = 32, Yellow = 33, Blue = 34, Magenta = 35,
|
||||
Cyan = 36, White = 37
|
||||
};
|
||||
Tui();
|
||||
~Tui();
|
||||
static void update();
|
||||
[[nodiscard]] bool is_running() const;
|
||||
private:
|
||||
inline static auto need_update_chat = false;
|
||||
inline static auto mouse = Mouse();
|
||||
inline static std::string chat_input;
|
||||
inline static bool cursor = false;
|
||||
inline static uint32_t frame_count = 0;
|
||||
inline static int background = 49;
|
||||
inline static int foreground = 39;
|
||||
inline static bool running = true;
|
||||
inline static int width = 80;
|
||||
inline static int height = 24;
|
||||
struct msg_block
|
||||
{
|
||||
std::optional<std::string> left = std::nullopt;
|
||||
std::optional<std::string> right = std::nullopt;
|
||||
};
|
||||
inline static std::vector<msg_block> chat_history;
|
||||
struct SplitDesc {
|
||||
enum Type { VERTICAL, HORIZONTAL };
|
||||
Type type;
|
||||
int position;
|
||||
std::vector<std::unique_ptr<SplitDesc>> children;
|
||||
SplitDesc(const Type t, const int pos) : type(t), position(pos) {}
|
||||
};
|
||||
struct border_charset {
|
||||
std::string_view top_left, top_right, bottom_left, bottom_right;
|
||||
std::string_view bottom, top, horizontal, vertical, bg;
|
||||
std::string_view cross, tee_up, tee_down, tee_left, tee_right;
|
||||
};
|
||||
struct scroll_chars
|
||||
{
|
||||
std::string_view top;
|
||||
std::string_view full;
|
||||
std::string_view bottom;
|
||||
};
|
||||
struct TuiCell {
|
||||
std::string ch;
|
||||
int width;
|
||||
TuiCell() : ch(" "), width(1) {}
|
||||
TuiCell(std::string c, const int w) : ch(std::move(c)), width(w) {}
|
||||
};
|
||||
std::string proces_chars = "⠙⠸⠴⠦⠇⠋";
|
||||
inline static struct termios orig_termios;
|
||||
inline static constexpr scroll_chars main_scroll_charset {
|
||||
.top = "▗",
|
||||
.full = "▐",
|
||||
.bottom = "▝",
|
||||
};
|
||||
inline static constexpr border_charset border_clip {
|
||||
.top_left="╭", .top_right="╮", .bottom_left="╰", .bottom_right="╯",
|
||||
.bottom="─", .top="─", .horizontal="─", .vertical="│", .bg=" ",
|
||||
.cross="┼", .tee_up="┴", .tee_down="┬", .tee_left="┤", .tee_right="├"
|
||||
};
|
||||
inline static constexpr border_charset border_ascii {
|
||||
.top_left="+", .top_right="+", .bottom_left="+", .bottom_right="+",
|
||||
.bottom="-", .top="-", .horizontal="-", .vertical="|", .bg=" ",
|
||||
.cross="+", .tee_up="+", .tee_down="+", .tee_left="+",
|
||||
.tee_right="+"
|
||||
};
|
||||
inline static constexpr border_charset border_bold {
|
||||
.top_left="┏", .top_right="┓", .bottom_left="┗", .bottom_right="┛",
|
||||
.bottom="━", .top="━", .horizontal="━", .vertical="┃",
|
||||
.bg=" ", .cross = "╋", .tee_up = "┻", .tee_down = "┳", .tee_left = "┨",
|
||||
.tee_right = "┣"
|
||||
};
|
||||
inline static constexpr border_charset border_solid {
|
||||
.top_left="▟", .top_right="▙", .bottom_left="▜", .bottom_right="▛",
|
||||
.bottom="█", .top="█", .horizontal = "█", .vertical="█",
|
||||
.bg=" ", .cross = "█", .tee_up = "█", .tee_down = "█", .tee_left = "█",
|
||||
.tee_right = "█",
|
||||
};
|
||||
inline static constexpr border_charset border_btn {
|
||||
.top_left="▟", .top_right="▙", .bottom_left="▜", .bottom_right="▛",
|
||||
.bottom="█", .top="█", .horizontal = "█", .vertical="█",
|
||||
.bg="█", .cross = "█", .tee_up = "█", .tee_down = "█", .tee_left = "█",
|
||||
.tee_right = "█"
|
||||
};
|
||||
static void init_draw();
|
||||
static winsize get_console_size();
|
||||
static bool wait_for_input_or_timeout(int timeout_ms);
|
||||
static void apply_color();
|
||||
static void disable_raw_mode();
|
||||
static void enable_raw_mode();
|
||||
static void move_cursor(int x, int y);
|
||||
static void clear();
|
||||
static void draw_border(int w, int h, border_charset charset = border_clip, int x = 1, int y = 1);
|
||||
static void hide_cursor();
|
||||
static void show_cursor();
|
||||
static void set_background(Color clr);
|
||||
static void set_foreground(Color clr);
|
||||
static void move_cursor_right(int n);
|
||||
static void move_cursor_x(int col);
|
||||
static void move_cursor_y(int row);
|
||||
static void invert_color();
|
||||
static void draw_button(int x, int y, int w, int h, std::string_view label, border_charset charset = border_solid);
|
||||
static void invert_color_off();
|
||||
static void draw_split_border(int w, int h, border_charset charset, int x, int y, const SplitDesc* root);
|
||||
static void apply_splits_recursive(std::vector<std::vector<TuiCell>>& grid, int offset_x, int offset_y, int w, int h, border_charset charset, const SplitDesc* split);
|
||||
static void apply_single_split(std::vector<std::vector<TuiCell>>& grid, int offset_x, int offset_y, int w, int h, border_charset charset, const SplitDesc* split);
|
||||
static void draw_border_to_buffer(std::vector<std::vector<TuiCell>>& grid, int w, int h, border_charset charset);
|
||||
static void push_msg_left(const std::string& msg);
|
||||
static void push_msg_right(const std::string& msg);
|
||||
static void draw_msg();
|
||||
};
|
||||
}
|
||||
68
xenith/ui/tui/message.cpp
Normal file
68
xenith/ui/tui/message.cpp
Normal file
@ -0,0 +1,68 @@
|
||||
#include <iostream>
|
||||
#include <ostream>
|
||||
|
||||
#include "./../tui.hpp"
|
||||
|
||||
namespace ui
|
||||
{
|
||||
void Tui::push_msg_left(const std::string& msg)
|
||||
{
|
||||
if (!chat_history.empty() && !chat_history.back().left.has_value())
|
||||
{
|
||||
chat_history.back().left = msg;
|
||||
}
|
||||
else
|
||||
{
|
||||
msg_block block = {
|
||||
.left = msg,
|
||||
.right = std::nullopt
|
||||
};
|
||||
chat_history.push_back(block);
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::push_msg_right(const std::string& msg)
|
||||
{
|
||||
if (!chat_history.empty() && !chat_history.back().right.has_value())
|
||||
{
|
||||
chat_history.back().right = msg;
|
||||
}
|
||||
else
|
||||
{
|
||||
msg_block block = {
|
||||
.left = std::nullopt,
|
||||
.right = msg,
|
||||
};
|
||||
chat_history.push_back(block);
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::draw_msg()
|
||||
{
|
||||
const auto root = std::make_unique<SplitDesc>(SplitDesc::HORIZONTAL, 2);
|
||||
|
||||
int total_hg = chat_history.size() * 5;
|
||||
|
||||
int offsetY = 0;
|
||||
for (int q = 0; q < chat_history.size(); q++)
|
||||
{
|
||||
int wh = std::max(3, static_cast<int>(chat_history[q].right.value_or("").length()));
|
||||
draw_split_border(wh + 2, 5, border_clip, width - wh - 2, 2 + offsetY, root.get());
|
||||
move_cursor(width - wh - 1, 3 + offsetY);
|
||||
std::cout << "You" << std::flush;
|
||||
move_cursor(width - wh - 1, 5 + offsetY);
|
||||
std::cout << chat_history[q].right.value_or("") << std::flush;
|
||||
offsetY += 5;
|
||||
|
||||
|
||||
wh = std::max(22, static_cast<int>(chat_history[q].left.value_or("").length()));
|
||||
draw_split_border(wh + 2, 5, border_clip, (width / 100) * 27 + 2, 2 + offsetY, root.get());
|
||||
move_cursor((width / 100) * 27 + 3, 3 + offsetY);
|
||||
std::cout << "Ai (Processing prompt█)" << std::flush;
|
||||
move_cursor((width / 100) * 27 + 3, 5 + offsetY);
|
||||
std::cout << chat_history[q].left.value_or("") << std::flush;
|
||||
offsetY += 5;
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
51
xenith/ui/tui/mouse.cpp
Normal file
51
xenith/ui/tui/mouse.cpp
Normal file
@ -0,0 +1,51 @@
|
||||
#include "../tui.hpp"
|
||||
|
||||
namespace ui
|
||||
{
|
||||
void Mouse::parse_ansi(char buf[])
|
||||
{
|
||||
int button = 0;
|
||||
char state = '\0';
|
||||
if (sscanf(buf + 3, "%d;%d;%d%c", &button, &pos.x, &pos.y, &state) == 4)
|
||||
{
|
||||
auto btn_index = static_cast<uint8_t>(button);
|
||||
uint8_t click_type = btn_index & 0xC3;
|
||||
if (click_type == 0b00000000)
|
||||
{
|
||||
btn0 = (state == 'M');
|
||||
}
|
||||
else if (click_type == 0b00000001)
|
||||
{
|
||||
btn1 = (state == 'M');
|
||||
}
|
||||
else if (click_type == 0b00000010)
|
||||
{
|
||||
btn2 = (state == 'M');
|
||||
}
|
||||
else if (click_type == 0b01000000)
|
||||
{
|
||||
scroll++;
|
||||
}
|
||||
else if (click_type == 0b01000001)
|
||||
{
|
||||
scroll--;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
bool Mouse::is_pressed(int index_btn)
|
||||
{
|
||||
switch (index_btn)
|
||||
{
|
||||
case 0: return btn0;
|
||||
case 1: return btn1;
|
||||
case 2: return btn2;
|
||||
default: return false;
|
||||
}
|
||||
}
|
||||
|
||||
void Mouse::reset_scroll() { scroll = 0; }
|
||||
int Mouse::get_scroll() { return scroll; }
|
||||
Mouse::pos_t Mouse::get_pos() { return pos; }
|
||||
}
|
||||
419
xenith/ui/tui/tui.cpp
Normal file
419
xenith/ui/tui/tui.cpp
Normal file
@ -0,0 +1,419 @@
|
||||
#include "../tui.hpp"
|
||||
|
||||
#include <chrono>
|
||||
#include <csignal>
|
||||
#include <cstring>
|
||||
#include <unistd.h>
|
||||
#include <sys/select.h>
|
||||
#include <iostream>
|
||||
|
||||
|
||||
namespace ui
|
||||
{
|
||||
Tui::Tui()
|
||||
{
|
||||
init_draw();
|
||||
|
||||
enable_raw_mode();
|
||||
}
|
||||
|
||||
void Tui::init_draw()
|
||||
{
|
||||
winsize wins = get_console_size();
|
||||
width = wins.ws_col;
|
||||
height = wins.ws_row;
|
||||
const int panel_w = (width / 100) * 27;
|
||||
clear();
|
||||
hide_cursor();
|
||||
set_foreground(Color::White);
|
||||
apply_color();
|
||||
|
||||
const auto root = std::make_unique<SplitDesc>(SplitDesc::VERTICAL, (width / 100) * 27);
|
||||
|
||||
root->children.push_back(nullptr);
|
||||
|
||||
auto chat_split = std::make_unique<SplitDesc>(SplitDesc::HORIZONTAL, height - 3);
|
||||
chat_split->children.push_back(nullptr);
|
||||
auto send_split = std::make_unique<SplitDesc>(SplitDesc::VERTICAL, (width - (width / 100) * 27) - 11);
|
||||
chat_split->children.push_back(std::move(send_split));
|
||||
root->children.push_back(std::move(chat_split));
|
||||
|
||||
draw_split_border(width, height, border_clip, 0, 0, root.get());
|
||||
|
||||
draw_button(2, 2, panel_w - 2, 3, "Chat", border_clip);
|
||||
draw_button(2, 2 + 3, panel_w - 2, 3, "Models", border_clip);
|
||||
draw_button(2, 2 + 6, panel_w - 2, 3, "Settings", border_clip);
|
||||
draw_button(2, 2 + 9, panel_w - 2, 3, "Server", border_clip);
|
||||
|
||||
move_cursor(width - 8, height - 1);
|
||||
std::cout << "[enter>" <<std::flush;
|
||||
}
|
||||
|
||||
Tui::~Tui()
|
||||
= default;
|
||||
|
||||
bool Tui::wait_for_input_or_timeout(int timeout_us) {
|
||||
fd_set fds;
|
||||
FD_ZERO(&fds);
|
||||
FD_SET(STDIN_FILENO, &fds);
|
||||
|
||||
struct timeval tv;
|
||||
tv.tv_sec = timeout_us / 1000000;
|
||||
tv.tv_usec = timeout_us % 1000000;
|
||||
|
||||
int retval = select(STDIN_FILENO + 1, &fds, NULL, NULL, &tv);
|
||||
return retval > 0;
|
||||
}
|
||||
|
||||
|
||||
void Tui::disable_raw_mode()
|
||||
{
|
||||
std::cout << "\033[?1006l\033[?1002l\033[?1000l" << std::flush;
|
||||
tcsetattr(STDIN_FILENO, TCSAFLUSH, &orig_termios);
|
||||
}
|
||||
|
||||
void Tui::enable_raw_mode() {
|
||||
tcgetattr(STDIN_FILENO, &orig_termios);
|
||||
struct termios raw = orig_termios;
|
||||
std::signal(SIGINT, [](int sig)
|
||||
{
|
||||
disable_raw_mode();
|
||||
show_cursor();
|
||||
running = false;
|
||||
});
|
||||
raw.c_lflag &= ~(ECHO | ICANON);
|
||||
tcsetattr(STDIN_FILENO, TCSAFLUSH, &raw);
|
||||
std::cout << "\033[?1000h\033[?1002h\033[?1006h" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::update()
|
||||
{
|
||||
bool input_changed = false;
|
||||
if (wait_for_input_or_timeout(10000)) {
|
||||
char buf[64];
|
||||
int bytes_read = read(STDIN_FILENO, buf, sizeof(buf) - 1);
|
||||
|
||||
if (bytes_read > 0) {
|
||||
buf[bytes_read] = '\0';
|
||||
int i = 0;
|
||||
|
||||
while (i < bytes_read) {
|
||||
unsigned char c = buf[i];
|
||||
|
||||
if (c == '\033') {
|
||||
if (i + 2 < bytes_read && strncmp(&buf[i], "\033[<", 3) == 0) {
|
||||
mouse.parse_ansi(&buf[i]);
|
||||
while (i < bytes_read && buf[i] != 'm' && buf[i] != 'M') {
|
||||
i++;
|
||||
}
|
||||
i++;
|
||||
continue;
|
||||
}
|
||||
i++;
|
||||
while (i < bytes_read && buf[i] >= 0x40 && buf[i] <= 0x7E) {
|
||||
i++;
|
||||
}
|
||||
}
|
||||
else if (c == 127 || c == 8) {
|
||||
if (!chat_input.empty()) {
|
||||
while (!chat_input.empty() && (static_cast<unsigned char>(chat_input.back()) & 0xC0) == 0x80) {
|
||||
chat_input.pop_back();
|
||||
}
|
||||
if (!chat_input.empty()) {
|
||||
chat_input.pop_back();
|
||||
}
|
||||
input_changed = true;
|
||||
}
|
||||
i++;
|
||||
}
|
||||
else if (c == 13 || c == 10) {
|
||||
if (!chat_input.empty()) {
|
||||
push_msg_right(chat_input);
|
||||
push_msg_left("Hi! How can I help you today?");
|
||||
chat_input.clear();
|
||||
input_changed = true;
|
||||
draw_msg();
|
||||
}
|
||||
i++;
|
||||
}
|
||||
else if (c >= 32) {
|
||||
int spaces_to_draw = std::max(0, (width - ((width / 100) * 27) - 15 - static_cast<int>(chat_input.length())));
|
||||
if (spaces_to_draw > 0)
|
||||
{
|
||||
chat_input += c;
|
||||
input_changed = true;
|
||||
}
|
||||
i++;
|
||||
}
|
||||
else {
|
||||
i++;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
static auto last_blink = std::chrono::steady_clock::now();
|
||||
auto now = std::chrono::steady_clock::now();
|
||||
|
||||
if (input_changed || std::chrono::duration_cast<std::chrono::milliseconds>(now - last_blink).count() >= 500) {
|
||||
if (input_changed) {
|
||||
last_blink = now;
|
||||
cursor = true;
|
||||
} else {
|
||||
cursor = !cursor;
|
||||
}
|
||||
|
||||
winsize wins = get_console_size();
|
||||
if (width != wins.ws_col || height != wins.ws_row)
|
||||
{
|
||||
init_draw();
|
||||
width = wins.ws_col;
|
||||
height = wins.ws_row;
|
||||
}
|
||||
|
||||
last_blink = now;
|
||||
|
||||
int input_x = ((width / 100) * 27) + 3;
|
||||
move_cursor(input_x, height - 1);
|
||||
int spaces_to_draw = std::max(0, (width - ((width / 100) * 27) - 15 - static_cast<int>(chat_input.length())));
|
||||
std::cout << chat_input << (cursor ? "_ " : " ") << std::string(spaces_to_draw, ' ') << std::flush;
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
bool Tui::is_running() const
|
||||
{
|
||||
return running;
|
||||
}
|
||||
|
||||
winsize Tui::get_console_size()
|
||||
{
|
||||
winsize w;
|
||||
ioctl(STDOUT_FILENO, TIOCGWINSZ, &w);
|
||||
return w;
|
||||
}
|
||||
|
||||
void Tui::apply_color()
|
||||
{
|
||||
std::cout << "\033["<< foreground << ";" << background <<"m" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::move_cursor(int x, int y)
|
||||
{
|
||||
std::cout << "\033[" << y << ";" << x << "H" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::clear()
|
||||
{
|
||||
std::cout << "\033c" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::draw_border(int w, int h, border_charset charset, int x, int y)
|
||||
{
|
||||
move_cursor(x, y);
|
||||
std::cout << charset.top_left;
|
||||
for (int q = 0; q < w - 2; q++) std::cout << charset.top;
|
||||
std::cout << charset.top_right <<"\n";
|
||||
move_cursor_x(x);
|
||||
for (int i = 0; i < h - 2; i++)
|
||||
{
|
||||
std::cout << charset.vertical;
|
||||
for (int q = 0; q < w - 2; q++) std::cout << charset.bg;
|
||||
std::cout << charset.vertical;
|
||||
std::cout << "\n";
|
||||
move_cursor_x(x);
|
||||
}
|
||||
std::cout << charset.bottom_left;
|
||||
for (int q = 0; q < w - 2; q++) std::cout << charset.bottom;
|
||||
std::cout << charset.bottom_right << std::flush;
|
||||
}
|
||||
|
||||
void Tui::hide_cursor()
|
||||
{
|
||||
std::cout << "\033[?25l" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::show_cursor()
|
||||
{
|
||||
std::cout << "\033[?25h" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::set_background(Color clr)
|
||||
{
|
||||
background = static_cast<int>(clr);
|
||||
}
|
||||
|
||||
void Tui::set_foreground(Color clr)
|
||||
{
|
||||
foreground = static_cast<int>(clr);
|
||||
}
|
||||
|
||||
void Tui::move_cursor_right(int n)
|
||||
{
|
||||
if (n > 0) {
|
||||
std::cout << "\033[" << n << "C" << std::flush;
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::move_cursor_x(int col)
|
||||
{
|
||||
if (col > 0) {
|
||||
std::cout << "\033[" << col << "G" << std::flush;
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::move_cursor_y(int row)
|
||||
{
|
||||
if (row > 0) {
|
||||
std::cout << "\033[" << row << "d" << std::flush;
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::invert_color()
|
||||
{
|
||||
std::cout << "\033[7m" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::draw_button(int x, int y, int w, int h, std::string_view label, border_charset charset)
|
||||
{
|
||||
draw_border(w, h, charset, x, y);
|
||||
int text_y = y + (h / 2);
|
||||
int text_x = x + ((w - static_cast<int>(label.length())) / 2);
|
||||
move_cursor(text_x, text_y);
|
||||
std::cout << label;
|
||||
std::cout << std::flush;
|
||||
}
|
||||
|
||||
void Tui::invert_color_off()
|
||||
{
|
||||
std::cout << "\033[27m" << std::flush;
|
||||
}
|
||||
|
||||
void Tui::draw_split_border(int w, int h, border_charset charset, int x, int y, const SplitDesc* root)
|
||||
{
|
||||
std::vector<std::vector<TuiCell>> grid(h, std::vector<TuiCell>(w, TuiCell(std::string(charset.bg), 1)));
|
||||
|
||||
draw_border_to_buffer(grid, w, h, charset);
|
||||
|
||||
if (root) {
|
||||
apply_splits_recursive(grid, 0, 0, w, h, charset, root);
|
||||
}
|
||||
|
||||
move_cursor(x, y);
|
||||
|
||||
// Посимвольный вывод без искажения UTF-8
|
||||
for (int i = 0; i < h; i++) {
|
||||
for (int j = 0; j < w; j++) {
|
||||
std::cout << grid[i][j].ch;
|
||||
}
|
||||
if (i < h - 1) {
|
||||
std::cout << "\033[1B\033[" << x << "G" << std::flush;
|
||||
}
|
||||
}
|
||||
|
||||
std::cout << std::flush;
|
||||
}
|
||||
|
||||
void Tui::apply_splits_recursive(std::vector<std::vector<TuiCell>>& grid,
|
||||
int offset_x, int offset_y,
|
||||
int w, int h,
|
||||
border_charset charset,
|
||||
const SplitDesc* split)
|
||||
{
|
||||
if (!split) return;
|
||||
|
||||
apply_single_split(grid, offset_x, offset_y, w, h, charset, split);
|
||||
|
||||
if (split->type == SplitDesc::VERTICAL) {
|
||||
if (split->children.size() >= 1) {
|
||||
int left_w = split->position;
|
||||
apply_splits_recursive(grid, offset_x, offset_y, left_w, h, charset, split->children[0].get());
|
||||
}
|
||||
if (split->children.size() >= 2) {
|
||||
int right_w = w - split->position;
|
||||
apply_splits_recursive(grid, offset_x + split->position, offset_y, right_w, h, charset, split->children[1].get());
|
||||
}
|
||||
} else { // HORIZONTAL
|
||||
if (split->children.size() >= 1) {
|
||||
int top_h = split->position;
|
||||
apply_splits_recursive(grid, offset_x, offset_y, w, top_h, charset, split->children[0].get());
|
||||
}
|
||||
if (split->children.size() >= 2) {
|
||||
int bottom_h = h - split->position;
|
||||
apply_splits_recursive(grid, offset_x, offset_y + split->position, w, bottom_h, charset, split->children[1].get());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::apply_single_split(std::vector<std::vector<TuiCell>>& grid, int offset_x, int offset_y, int w, int h, border_charset charset, const SplitDesc* split)
|
||||
{
|
||||
if (grid.empty()) {
|
||||
return;
|
||||
}
|
||||
if (split->type == SplitDesc::VERTICAL) {
|
||||
int x = offset_x + split->position;
|
||||
if (!grid[0].empty() && x >= 0 && x < static_cast<int>(grid[0].size())) {
|
||||
for (int i = 0; i < h; i++) {
|
||||
int y = offset_y + i;
|
||||
if (y >= 0 && y < static_cast<int>(grid.size())) {
|
||||
if (x < static_cast<int>(grid[y].size())) {
|
||||
TuiCell& cell = grid[y][x];
|
||||
if (i == 0) {
|
||||
if (cell.ch == charset.top) cell.ch = std::string(charset.tee_down);
|
||||
} else if (i == h - 1) {
|
||||
if (cell.ch == charset.bottom) cell.ch = std::string(charset.tee_up);
|
||||
} else {
|
||||
if (cell.ch == charset.bg) {
|
||||
cell.ch = std::string(charset.vertical);
|
||||
} else if (cell.ch == charset.horizontal || cell.ch == charset.top || cell.ch == charset.bottom) {
|
||||
cell.ch = std::string(charset.cross);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
int y = offset_y + split->position;
|
||||
if (y >= 0 && y < static_cast<int>(grid.size())) {
|
||||
for (int j = 0; j < w; j++) {
|
||||
int x = offset_x + j;
|
||||
if (x >= 0 && x < static_cast<int>(grid[y].size())) {
|
||||
TuiCell& cell = grid[y][x];
|
||||
if (j == 0) {
|
||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_right);
|
||||
} else if (j == w - 1) {
|
||||
if (cell.ch == charset.vertical) cell.ch = std::string(charset.tee_left);
|
||||
} else {
|
||||
if (cell.ch == charset.bg) {
|
||||
cell.ch = std::string(charset.horizontal);
|
||||
} else if (cell.ch == charset.vertical) {
|
||||
cell.ch = std::string(charset.cross);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void Tui::draw_border_to_buffer(std::vector<std::vector<TuiCell>>& grid, int w, int h, border_charset charset)
|
||||
{
|
||||
if (w <= 0 || h <= 0) return;
|
||||
|
||||
// Верхняя строка
|
||||
grid[0][0] = TuiCell(std::string(charset.top_left), 1);
|
||||
for (int j = 1; j < w - 1; j++) grid[0][j] = TuiCell(std::string(charset.top), 1);
|
||||
grid[0][w - 1] = TuiCell(std::string(charset.top_right), 1);
|
||||
|
||||
// Средние строки
|
||||
for (int i = 1; i < h - 1; i++) {
|
||||
grid[i][0] = TuiCell(std::string(charset.vertical), 1);
|
||||
grid[i][w - 1] = TuiCell(std::string(charset.vertical), 1);
|
||||
}
|
||||
|
||||
// Нижняя строка
|
||||
grid[h - 1][0] = TuiCell(std::string(charset.bottom_left), 1);
|
||||
for (int j = 1; j < w - 1; j++) grid[h - 1][j] = TuiCell(std::string(charset.bottom), 1);
|
||||
grid[h - 1][w - 1] = TuiCell(std::string(charset.bottom_right), 1);
|
||||
}
|
||||
}
|
||||
Loading…
x
Reference in New Issue
Block a user