fix
This commit is contained in:
parent
515a434928
commit
0c065a0931
@ -138,13 +138,12 @@ source activate.fish
|
|||||||
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
--fast эксперементаль флаг который ускоряет foreword примерно в x2-x5
|
||||||
```
|
```
|
||||||
|
|
||||||
### update - проверка обновления
|
|
||||||
|
|
||||||
### info
|
### info
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
Xenith info models/my.xnh # конфигурация и словарь
|
Xenith info models/my.xnh # конфигурация и словарь
|
||||||
```
|
```
|
||||||
|
### update - проверка обновления
|
||||||
|
|
||||||
`gradcheck` стоит запускать после любой правки в `model.cpp` / `backward.cpp` —
|
`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
|
||||||
@ -1,13 +1,13 @@
|
|||||||
|
# activate.zsh
|
||||||
export VIRTUAL_ENV_PROMPT="%F{red}Xenith%f"
|
export VIRTUAL_ENV_PROMPT="%F{red}Xenith%f"
|
||||||
alias Xenith='./bin/Xenith'
|
alias Xenith='./bin/Xenith'
|
||||||
if [ -z "${_OLD_VIRTUAL_PS1+x}" ]; then
|
if [ -z "${_OLD_VIRTUAL_PS1+x}" ]; then
|
||||||
_OLD_VIRTUAL_PS1="$PS1"
|
_OLD_VIRTUAL_PS1="$PS1"
|
||||||
fi
|
fi
|
||||||
PS1="%{($VIRTUAL_ENV_PROMPT)%} $_OLD_VIRTUAL_PS1"
|
PS1="(${VIRTUAL_ENV_PROMPT}) ${_OLD_VIRTUAL_PS1}"
|
||||||
deactivate_custom() {
|
deactivate_custom() {
|
||||||
unalias Xenith 2>/dev/null
|
unalias Xenith 2>/dev/null
|
||||||
unalias deactivate 2>/dev/null
|
unalias deactivate 2>/dev/null
|
||||||
|
|
||||||
if [ -n "${_OLD_VIRTUAL_PS1+x}" ]; then
|
if [ -n "${_OLD_VIRTUAL_PS1+x}" ]; then
|
||||||
PS1="$_OLD_VIRTUAL_PS1"
|
PS1="$_OLD_VIRTUAL_PS1"
|
||||||
unset _OLD_VIRTUAL_PS1
|
unset _OLD_VIRTUAL_PS1
|
||||||
@ -18,5 +18,10 @@ deactivate_custom() {
|
|||||||
builtin deactivate 2>/dev/null || true
|
builtin deactivate 2>/dev/null || true
|
||||||
fi
|
fi
|
||||||
unset -f deactivate_custom
|
unset -f deactivate_custom
|
||||||
|
compdel Xenith 2>/dev/null
|
||||||
|
compdel ./bin/Xenith 2>/dev/null
|
||||||
}
|
}
|
||||||
alias deactivate='deactivate_custom'
|
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.
11
main.cpp
11
main.cpp
@ -302,15 +302,6 @@ static bool available_new_version() {
|
|||||||
|
|
||||||
Version remote_ver = parse_version(remote_str);
|
Version remote_ver = parse_version(remote_str);
|
||||||
Version local_ver = parse_version(local_str);
|
Version local_ver = parse_version(local_str);
|
||||||
|
|
||||||
if (is_newer(remote_ver, local_ver))
|
|
||||||
{
|
|
||||||
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);
|
return is_newer(remote_ver, local_ver);
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -412,7 +403,7 @@ int main(int argc, char** argv) {
|
|||||||
std::cerr << "Ошибка чтения файла новой версии (Пожалуйста напишите о этом баге)" << std::endl;
|
std::cerr << "Ошибка чтения файла новой версии (Пожалуйста напишите о этом баге)" << std::endl;
|
||||||
return 1;
|
return 1;
|
||||||
}
|
}
|
||||||
std::cout << "\033[32mДоступна новая версия \033[33m" << remote_str << "\033[32m\nВыполните \"Xenith update\" чтобы обновится до последней версии\033[39m\n" << std::endl;
|
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; }
|
if (argc < 2) { printf(HELP_TEXT); return 1; }
|
||||||
|
|||||||
@ -17,7 +17,7 @@ static float softmax_ce(const float* logits, int V, int target, float* dlogits)
|
|||||||
}
|
}
|
||||||
|
|
||||||
float backward(Params& p, const ForwardCache& fc,
|
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 Config& c = p.cfg;
|
||||||
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
const int d = c.n_embd, H = c.n_head, hd = c.head_dim(), f = c.ffn();
|
||||||
const int V = c.vocab_out();
|
const int V = c.vocab_out();
|
||||||
@ -40,14 +40,14 @@ float backward(Params& p, const ForwardCache& fc,
|
|||||||
Tensor dxf(B, d);
|
Tensor dxf(B, d);
|
||||||
if (p.tied()) {
|
if (p.tied()) {
|
||||||
gemm_nn(dlogits, p.wte, dxf);
|
gemm_nn(dlogits, p.wte, dxf);
|
||||||
gemm_tn(dlogits, fc.xf, p.gwte, true);
|
gemm_tn(dlogits, fc.xf, p.gwte, true, fast);
|
||||||
} else {
|
} else {
|
||||||
gemm_nn(dlogits, p.lm_head, dxf);
|
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;
|
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;
|
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;
|
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
|
gemm_nt(dres, p.w2[L], dha); // dha = dres @ w2^T
|
||||||
|
|
||||||
dh1.resize(B, f);
|
dh1.resize(B, f);
|
||||||
@ -73,8 +73,8 @@ float backward(Params& p, const ForwardCache& fc,
|
|||||||
d3[j] = ga * silu(g);
|
d3[j] = ga * silu(g);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
gemm_tn(lc.x2b, dh1, p.g1[L], true); // dW1
|
gemm_tn(lc.x2b, dh1, p.g1[L], true, fast); // dW1
|
||||||
gemm_tn(lc.x2b, dh3, p.g3[L], true); // dW3
|
gemm_tn(lc.x2b, dh3, p.g3[L], true, fast); // dW3
|
||||||
|
|
||||||
dx2b.resize(B, d);
|
dx2b.resize(B, d);
|
||||||
gemm_nt(dh1, p.w1[L], dx2b);
|
gemm_nt(dh1, p.w1[L], dx2b);
|
||||||
@ -83,11 +83,11 @@ float backward(Params& p, const ForwardCache& fc,
|
|||||||
|
|
||||||
{
|
{
|
||||||
Tensor g;
|
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);
|
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);
|
dattout.resize(B, d);
|
||||||
gemm_nt(dx, p.wo[L], dattout);
|
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, dq, p.gq[L], true, fast);
|
||||||
gemm_tn(lc.xb, dk, p.gk[L], true);
|
gemm_tn(lc.xb, dk, p.gk[L], true, fast);
|
||||||
gemm_tn(lc.xb, dv, p.gv[L], true);
|
gemm_tn(lc.xb, dv, p.gv[L], true, fast);
|
||||||
|
|
||||||
dxb.resize(B, d);
|
dxb.resize(B, d);
|
||||||
gemm_nt(dq, p.wq[L], dxb);
|
gemm_nt(dq, p.wq[L], dxb);
|
||||||
@ -174,7 +174,7 @@ float backward(Params& p, const ForwardCache& fc,
|
|||||||
|
|
||||||
{
|
{
|
||||||
Tensor g;
|
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);
|
for (size_t i = 0; i < dx.n(); ++i) dx.at(i) += g.at(i);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@ -58,7 +58,7 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
|||||||
Tensor ln, lnf;
|
Tensor ln, lnf;
|
||||||
|
|
||||||
for (int L = 0; L < c.n_layer; ++L) {
|
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.wq[L], q);
|
||||||
gemm_nn(xb, p.wk[L], k);
|
gemm_nn(xb, p.wk[L], k);
|
||||||
gemm_nn(xb, p.wv[L], v);
|
gemm_nn(xb, p.wv[L], v);
|
||||||
@ -99,7 +99,7 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
|||||||
gemm_nn(attout, p.wo[L], proj);
|
gemm_nn(attout, p.wo[L], proj);
|
||||||
for (int j = 0; j < d; ++j) x.at(j) += proj.at(j);
|
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.w1[L], h1);
|
||||||
gemm_nn(x2b, p.w3[L], h3);
|
gemm_nn(x2b, p.w3[L], h3);
|
||||||
if (fast) attout.resize_fast(1, f);
|
if (fast) attout.resize_fast(1, f);
|
||||||
@ -109,7 +109,7 @@ void forward_token(const Params& p, KvCache& kv, int token, int pos,
|
|||||||
for (int j = 0; j < d; ++j) x.at(j) += fo.at(j);
|
for (int j = 0; j < d; ++j) x.at(j) += fo.at(j);
|
||||||
}
|
}
|
||||||
|
|
||||||
rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln);
|
rmsnorm_forward(x, p.rms_final, c.rms_eps, lnf, ln, fast);
|
||||||
if (fast) attout.resize_fast(1, c.vocab_out());
|
if (fast) attout.resize_fast(1, c.vocab_out());
|
||||||
else attout.resize(1, c.vocab_out());
|
else attout.resize(1, c.vocab_out());
|
||||||
if (p.tied()) gemm_nt(lnf, p.wte, logits);
|
if (p.tied()) gemm_nt(lnf, p.wte, logits);
|
||||||
|
|||||||
@ -54,7 +54,7 @@ int run_gradcheck(bool verbose, bool fast) {
|
|||||||
|
|
||||||
ForwardCache fc;
|
ForwardCache fc;
|
||||||
forward(p, x.data(), n, n, fc, fast);
|
forward(p, x.data(), n, n, fc, fast);
|
||||||
backward(p, fc, x.data(), y.data(), n);
|
backward(p, fc, x.data(), y.data(), n, fast);
|
||||||
|
|
||||||
auto tensors = p.all_w();
|
auto tensors = p.all_w();
|
||||||
auto grads = p.all_grads();
|
auto grads = p.all_grads();
|
||||||
@ -130,7 +130,7 @@ int run_bench(Checkpoint& ck, int iters, int block, bool fast) {
|
|||||||
// прогрев
|
// прогрев
|
||||||
ForwardCache fc;
|
ForwardCache fc;
|
||||||
forward(ck.params, x.data(), n, n, fc, fast);
|
forward(ck.params, x.data(), n, n, fc, fast);
|
||||||
backward(ck.params, fc, x.data(), y.data(), n);
|
backward(ck.params, fc, x.data(), y.data(), n, fast);
|
||||||
|
|
||||||
auto t0 = std::chrono::steady_clock::now();
|
auto t0 = std::chrono::steady_clock::now();
|
||||||
for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc, fast);
|
for (int i = 0; i < iters; ++i) forward(ck.params, x.data(), n, n, fc, fast);
|
||||||
@ -141,7 +141,7 @@ int run_bench(Checkpoint& ck, int iters, int block, bool fast) {
|
|||||||
t0 = std::chrono::steady_clock::now();
|
t0 = std::chrono::steady_clock::now();
|
||||||
for (int i = 0; i < iters; ++i) {
|
for (int i = 0; i < iters; ++i) {
|
||||||
forward(ck.params, x.data(), n, n, fc, fast);
|
forward(ck.params, x.data(), n, n, fc, fast);
|
||||||
backward(ck.params, fc, x.data(), y.data(), n);
|
backward(ck.params, fc, x.data(), y.data(), n, fast);
|
||||||
}
|
}
|
||||||
t1 = std::chrono::steady_clock::now();
|
t1 = std::chrono::steady_clock::now();
|
||||||
double d_all = std::chrono::duration<double>(t1 - t0).count();
|
double d_all = std::chrono::duration<double>(t1 - t0).count();
|
||||||
|
|||||||
@ -240,12 +240,46 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
|||||||
for (int L = 0; L < c.n_layer; ++L) {
|
for (int L = 0; L < c.n_layer; ++L) {
|
||||||
LayerCache& lc = fc.layers[L];
|
LayerCache& lc = fc.layers[L];
|
||||||
|
|
||||||
if (fast) fc.x.resize_fast(B, d);
|
if (fast)
|
||||||
else fc.x.resize(B, d);
|
{
|
||||||
|
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);
|
||||||
|
|
||||||
for (size_t i = 0; i < fc.x.n(); ++i) lc.x.at(i) = fc.x.at(i);
|
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);
|
||||||
|
|
||||||
rmsnorm_forward(lc.x, p.rms_attn[L], c.rms_eps, lc.xb, lc.inv_rms);
|
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.wq[L], lc.q);
|
||||||
gemm_nn(lc.xb, p.wk[L], lc.k);
|
gemm_nn(lc.xb, p.wk[L], lc.k);
|
||||||
gemm_nn(lc.xb, p.wv[L], lc.v);
|
gemm_nn(lc.xb, p.wv[L], lc.v);
|
||||||
@ -316,7 +350,7 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
|||||||
else fc.x.resize(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);
|
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.w1[L], lc.h1);
|
||||||
gemm_nn(lc.x2b, p.w3[L], lc.h3);
|
gemm_nn(lc.x2b, p.w3[L], lc.h3);
|
||||||
if (fast) fc.x.resize_fast(B, f);
|
if (fast) fc.x.resize_fast(B, f);
|
||||||
@ -327,7 +361,7 @@ void forward(const Params& p, const int* tokens, int n_tok, int seq_len, Forward
|
|||||||
for (size_t i = 0; i < fc.x.n(); ++i) fc.x.at(i) += lc.fo.at(i);
|
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);
|
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());
|
if (fast) fc.x.resize_fast(B, c.vocab_out());
|
||||||
else fc.x.resize(B, c.vocab_out());
|
else fc.x.resize(B, c.vocab_out());
|
||||||
if (p.tied()) gemm_nt(fc.xf, p.wte, fc.logits);
|
if (p.tied()) gemm_nt(fc.xf, p.wte, fc.logits);
|
||||||
|
|||||||
@ -132,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();
|
const int N = A.R(), K = A.C(), M = B.C();
|
||||||
if (!accumulate) {
|
if (!accumulate) {
|
||||||
C.resize(K, M);
|
if (fast) C.resize_fast(K, M);
|
||||||
|
else C.resize(K, M);
|
||||||
} else {
|
} else {
|
||||||
if (C.R() != K || C.C() != M) {
|
if (C.R() != K || C.C() != M) {
|
||||||
std::fprintf(stderr,
|
std::fprintf(stderr,
|
||||||
@ -167,10 +168,12 @@ inline float dsilu(float x) {
|
|||||||
|
|
||||||
|
|
||||||
inline void rmsnorm_forward(const Tensor& x, const Tensor& w, float eps,
|
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();
|
const int N = x.R(), D = x.C();
|
||||||
y.resize(N, D);
|
if (fast) y.resize_fast(N, D);
|
||||||
inv_rms.resize(N, 1);
|
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) {
|
for (int n = 0; n < N; ++n) {
|
||||||
const float* xr = x.row(n);
|
const float* xr = x.row(n);
|
||||||
float ss = 0.0f;
|
float ss = 0.0f;
|
||||||
@ -183,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,
|
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();
|
const int N = x.R(), D = x.C();
|
||||||
dx.resize(N, D);
|
if (fast) dx.resize_fast(N, D);
|
||||||
dw.resize(1, D);
|
else dx.resize(N, D);
|
||||||
|
if (fast) dw.resize_fast(1, D);
|
||||||
|
else dw.resize(1, D);
|
||||||
dw.zero();
|
dw.zero();
|
||||||
for (int n = 0; n < N; ++n) {
|
for (int n = 0; n < N; ++n) {
|
||||||
const float* xr = x.row(n);
|
const float* xr = x.row(n);
|
||||||
|
|||||||
@ -6,7 +6,7 @@
|
|||||||
namespace xt {
|
namespace xt {
|
||||||
|
|
||||||
float backward(Params& p, const ForwardCache& fc,
|
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,
|
float softmax_eval(const Params& p, const ForwardCache& fc,
|
||||||
const int* x, const int* y, int n_pairs, Tensor& per_token_loss);
|
const int* x, const int* y, int n_pairs, Tensor& per_token_loss);
|
||||||
|
|||||||
@ -80,7 +80,7 @@ TrainStats train_model(Params& p, const Tokenizer& tok, const std::string& corpu
|
|||||||
}
|
}
|
||||||
|
|
||||||
forward(p, xb.data(), B * block, block, fc, tc.fast);
|
forward(p, xb.data(), B * block, block, fc, tc.fast);
|
||||||
float loss = backward(p, fc, xb.data(), yb.data(), B * block);
|
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,
|
p.adam_step(lr_at(tc, step), tc.beta1, tc.beta2, tc.eps,
|
||||||
tc.weight_decay, step + 1, tc.clip);
|
tc.weight_decay, step + 1, tc.clip);
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user