Назад к блогу

PSSA: языковая модель без трансформеров на Rust

PSSA: языковая модель без трансформеров на Rust

Новая работа о PSSA — языковой модели на Rust, которая отказывается от трансформерной архитектуры в пользу непрерывной SSM-динамики и эпизодической памяти в гиперболическом пространстве. Автор разбирает, как из дифференциального уравнения получается пошаговое обновление состояния, как устроен обратный проход через такую рекуррентность и какие приёмы кэширования позволяют не пересчитывать преобразования на каждом шаге. Особый интерес представляет сравнение с параметр-матченным трансформером — оно показывает, чего можно добиться без механизма внимания.

PSSA — это языковая модель, собранная без трансформера: вместо внимания здесь непрерывная SSM и эпизодическая память в гиперболическом пространстве. Вопрос, который стоит разобрать: как из непрерывной динамики получается пошаговое обновление, как устроен обратный проход через такую рекуррентность, и что показывают замеры против параметр-матченного трансформера.

Непрерывный SSM-блок: от дифференциального уравнения к шагу

Непрерывная система дискретизируется. С шагом delta это даёт пошаговое обновление состояния. Формулы перехода:

Abar_ij = exp(delta_i * A_ij)
Bbar_ij = delta_i * B_j
h_ij <- Abar_ij * h_ij + Bbar_ij * x_i
y_i = sum_j C_j * h_ij

Здесь индекс i нумерует каналы (их d_m), а j — состояния внутри канала (их d_s); поэтому матрица A имеет размерность d_m x d_s, и A_ij — это скорость перехода для пары (канал i, состояние j). Матрица A строится как A = -softplus(A_raw); delta вычисляется как softplus(W_delta x) и имеет размерность d_m. Переход диагональный — по одной скорости на пару (канал, состояние), и она остаётся отрицательной по построению, чтобы рекуррентность не расходилась.

Одно вычисление преобразования на изменившуюся скорость

Пересчёт exp(delta*A) не делается на каждом шаге. Есть снимок «сырых» значений ssm_raw_snapshot: при инициализации он заполняется NaN, а кэши ssm_rates и ssm_rate_derivatives — нулями. Функция refresh_ssm_rates проходит по всем элементам a_mat.data и для каждого сравнивает текущее сырое значение с сохранённым в снимке; пересчёт происходит только при различии. Тогда вычисляются ssm_rates[i] = -softplus(raw) и ssm_rate_derivatives[i] = -sigmoid(raw), после чего снимок обновляется значением raw.

Принцип сформулирован прямо: «One transform evaluation per changed raw rate, not per timestep. Compare values rather than an optimizer version: callers may mutate public weights directly, including finite-difference probes and checkpoint decoding.» Сравнивать именно значения, а не версию оптимизатора, приходится потому, что веса могут менять напрямую — например, при численном дифференцировании или декодировании чекпоинта.

В последовательном скане refresh_ssm_rates вызывается один раз перед циклом по токенам, а внутри цикла exp(delta*A) берётся из кэша: (d_i * m.ssm_rates[idx]).exp().

Эпизодическая память: чтение через шар Пуанкаре

Запрос к памяти формируется как сумма двух линейных проекций: q = W_qx x + W_qh y, где x — это x_norm, а y — это y_ssm. Затем диффеоморфная проекция proj переводит запрос в шар Пуанкаре: qh = proj(q), при этом |qh| < 1.

Чтение — это softmax по отрицательному гиперболическому расстоянию, делённому на температуру tau_mem: w = softmax(-d_H(qh, k_s) / tau_mem) по четырём ближайшим слотам, где k_s — ключ слота, а значение собирается как m = sum_k w_k * v_k, где v_k — значение слота.

Почему именно четыре слота, а не весь банк: гиперболическое расстояние растёт к границе шара, поэтому слоты с общим контекстом и слоты с конкретным эпизодом остаются разделимыми без расширения чтения. Четыре слота — фиксированная стоимость на токен независимо от того, сколько всего хранит банк.

При этом обратный проход устроен иначе: там извлечение описано как полнобанковый softmax по -гиперболическое_расстояние/tau, перебираются все записи банка, а градиент по запросу накапливается с множителем, обратным tau_mem.

Запись: консолидация через ridge regression

Быстрые пластичные обновления сворачиваются обратно в базовую матрицу переходов по формуле A_base <- A_base + (H^T H + lambda I)^-1 H^T dH, названной closed-form ridge regression. Смысл шага — не дать банку слотов остаться единственным местом хранения долговременной структуры.

Гейт, адаптер и MLP: порядок операций

Сначала из x_norm вычисляется гейт g = sigmoid(W_gate x), и отдельно проецируется память m_proj = W_proj m; эти два результата перемножаются поэлементно в m_inj. Параллельно адаптер считает adapter_hidden = down_proj · x_norm, к нему поэлементно применяется SiLU (h*sigmoid(h)), и up-проекция даёт вклад adapter(x) = (up_proj + consolidated_up) · adapter_act.

Латентная агрегация собирается как z_raw = y_ssm*ssm_scale + m_inj + adapter(x), то есть s*y + g + adapter(x). Затем MLP-расширение: z_raw умножается на mlp_w1, применяется SiLU, результат умножается на mlp_w2, давая u; итог z_out = z_raw + u. В обратном проходе остаток прибавляется к градиенту z_raw.

Ouro-итерации и активационная лента

Ouro-итерации получает одну непрерывную ленту активаций. При loops > 1 число сохраняемых слотов равно loops + 1, и промежуточные массивы получают длину, пропорциональную числу слотов. Публичные logits, probs, losses при этом остаются формы max_l.

Адресация срезов идёт двумя функциями: loop_offset возвращает loop_index * self.max_l * width + token * width, а loop_state_offset — loop_index * (self.max_l + 1) * width + token * width, то есть состояние на токен включает дополнительную строку.

Перенос состояния между виртуальными проходами делает swap_loop_carry: для ненулевого loop_index он вычисляет start = self.loop_carry_start() + (loop_index - 1) * self.h_persistent.len() и попарно меняет местами элементы h_persistent и tape.h_states[start..]. Копирование полей ленты выполняют copy_loop_field_to_base и copy_base_field_to_loop: при loop_index == 0 они сразу возвращаются, иначе берут len = max_l * width, offset = loop_index * len и копируют base[..len] из/в rest[..len].

Что разделяется, а что своё

Между виртуальными проходами разделяются все веса, моменты оптимизатора и банк памяти. Своё у каждого прохода — только причинно-временной перенос, то есть состояние рекуррентности h_persistent, хранимое в runtime-хранилище ленты, чтобы не портить историческое состояние чекпоинта.

reset_recurrent_state обнуляет базовый перенос и все сохранённые слоты: он делает self.h_persistent.fill(0.0) и заполняет нулями tape.h_states начиная с loop_carry_start(). saved_loop_slot определяет, в какой слот ленты сохраняются активации данного прохода: для loop_index == 0 при нескольких проходах он возвращает loop_count(), иначе сам loop_index — так нулевой проход переживает последующие для reverse-mode.

Число отсоединённых скаляров переноса, сохраняемых между чанками, — depth * loops * d_latent * d_state. Runtime-значение Ouro намеренно отсутствует в чекпоинтах.

Обучение: прямой и обратный проход

forward_train_chunk — обёртка, вызывающая forward_train_chunk_loop с loop_index = 0. Внутри одна Ouro-итерация кладёт активации в отдельный слот ленты, вычисляя смещения по loop_index, копирует вход в слот и прогоняет цикл по токенам. В конце итерации перенос сохраняется в h_persistent и, если это нулевая итерация при нескольких циклах, базовый тейп копируется в слот.

backward_chunk — обёртка над backward_chunk_loop с loop_index = 0. backward_chunk_loop выполняет обратный проход по времени для одной Ouro-итерации, накапливая градиенты в общие параметры, при этом временной перенос отсекается на границе чанка, а межпроходные адъюнкты распространяет внешний residual-цикл. На уровне модели backward_chunk при loops > 1 идёт по итерациям в обратном порядке, масштабирует адъюнкты на loop_scale = 1.0 / loops as f32 и добавляет полученные residual_input_adjoints обратно в residual_block_adjoints.

Обратный проход через SSM

Adjoint запроса к банку памяти уже посчитан в bwd_stage_memory и не входит в рекуррентность SSM. Для каждого токена в обратном времени строится аффинное отображение p -> A_t*p + A_t*r_t, где r_t — прямой adjoint считывания. Исключающий обратный префикс даёт p_(t+1), а терминальный adjoint равен нулю, потому что TBPTT отсекает границы чанков.

Для этого в ssm_scan_a и ssm_scan_b записываются A_t и A_t*r_t, после чего вызывается affine_scan_in_place, снимающий обратную временную зависимость. Затем локальные производные пишут в непересекающиеся строки токенов, используя future_p из ssm_scan_b, и заполняют bwd_ssm_delta, bwd_ssm_b, bwd_ssm_c, bwd_ssm_a и bwd_g_xnorm.

Для retrieval полный softmax-VJP по банку выполняется в memory_query_adjoint: он читает q_poincare, q_euc, g_m_val, m_val, weights и tau, накапливает g_query_pnc по всем записям банка, а затем через projection_adjoint получает g_query_euc.

AdamW и глобальный клиппинг

apply_adamw_with_grad_clip сначала собирает все тензоры оптимизатора, затем одним проходом по их градиентам накапливает сумму квадратов в f64 и берёт sqrt — это глобальная L2-норма. Если норма не конечна, градиенты обнуляются и возвращается GradientClipOutcome::Skipped { norm }, при этом счётчик шагов не увеличивается.

Иначе вычисляется масштаб: scale = max_norm as f64 / norm, если norm > max_norm as f64, иначе ровно 1.0 — то есть min(1, max_norm/norm) в виде ветвления. Затем для каждого тензора вызывается tensor.step(&cfg, lr, step, scale), и внутри AdamTensor::step при scale < 1.0 градиент пересчитывается как (self.grad[i] as f64 * scale) as f32 и записывается обратно, а при scale == 1.0 градиент не трогается вовсе.

Клиппинг делается в f64 до f32-моментов, потому что моменты m и v обновляются уже по масштабированному g. Так огромные конечные градиенты (например 1e21) сначала приводятся к норме max_norm, и только потом попадают в f32-накопители второго момента, что предотвращает их переполнение.

Нефинитные градиенты

При нефинитной норме (бесконечность или NaN) вызывается zero_gradients() и возвращается Skipped, при этом step_counter не увеличивается. Атомарность обеспечивается тем, что до вызова tensor.step(...) дело не доходит: моменты не обновляются, веса не меняются, счётчик шагов не растёт.

Проверка нефинитности нормы выполняется до вычисления step = step_counter.checked_add(1) и до его использования, поэтому при step_counter = usize::MAX и NaN-градиенте шаг всё равно пропускается, счётчик остаётся usize::MAX, состояние оптимизатора не меняется, а градиенты обнуляются.

Dreaming: офлайн-репетиция

dream_replay_with_options_and_guard принимает режим, число повторов, длину сновидения, температуру, learning rate и число шагов репетиции, а также необязательную «свежую» задачу в виде пары input_ids/target_ids; из неё строится guard-последовательность с проверками непустоты, равенства длин, длины не больше chunk_len и всех id меньше d_vocab.

Память для репетиции формируется через remember_dream_sequence: она копирует последовательность в dream_sequences, а затем, если размер превышает mem_capacity.max(1), удаляет самые старые записи — кэш ограничен по количеству.

Проекция градиента в apply_dream_projected_sgd вычисляет скалярное произведение dot старого и свежего градиентов и квадрат нормы fresh_norm_sq по выбранным элементам; если dot < 0 и fresh_norm_sq > 0, projection = dot / fresh_norm_sq, иначе 0, после чего вес обновляется как weight -= lr * (old - projection * fresh).

Backtracking в rehearse_dream_sequence_guarded сохраняет снимок параметров, пробует шаг с step_lr, применяет проекцию и проверяет candidate = dream_loss(guard, recurrent); шаг принимается, если candidate конечен и candidate <= baseline + LOSS_TOLERANCE (1e-6), иначе step_lr умножается на 0.5 до 10 попыток, а при неудаче параметры восстанавливаются из снимка. Гарантия состоит в том, что принимается только шаг, не ухудшающий fresh-task loss более чем на 1e-6, а проекция служит лишь первым приближением — точная проверка идёт по фактическому loss.

Каждый memory seed стартует с одного и того же live carry: перед обработкой каждого значения рекуррентное состояние восстанавливается из сохранённой копии, что делает K последовательностей независимыми при детерминированном состоянии внутри каждой. Training replay использует CPU workspaces даже при выбранном GPU, потому что вызывающая сторона синхронизирует safeguarded weights перед входом, а сам метод временно подменяет устройство на CPU, после чего восстанавливает backend и инвалидирует его кэш.

Бэкенды: CPU, CUDA, WebGPU

Выбор устройства выполняет Device::try_gpu: при включённой фиче cuda он сначала пытается поднять CUDA и при успехе сразу возвращает Device::Cuda(ctx), иначе сохраняет ошибку и пробует WebGPU через WgpuContext::init_blocking(), возвращая Device::Gpu(ctx); если оба не удались, возвращается Err с причинами обоих бэкендов.

GpuDispatch. Флаг accelerates_backward показывает, что устройство умеет считать обратный проход рекуррентности и памяти нативно, а не на CPU. Он включён и для Wgpu, и для Cuda, то есть оба устройства имеют device-native recurrent и memory backward стадии. При этом у WebGPU плотные transpose-хелперы всё ещё используют детерминированную CPU-реализацию, но SSM reverse scan и memory VJPs диспетчеризуются на устройстве.

В sequence_batch.rs::backward при cuda_pending вызывается self.backward_cuda(m), и при ошибке CUDA-путь откатывается на CPU-реплей с rebuild_host_tapes и stages::bwd_stage_memory; иначе сразу выполняется stages::bwd_stage_memory(m, n). Для нерезидентного backward выбирается WebGPU-ветка, если устройство даёт GpuDispatch::Wgpu, и тогда для каждой lane вызывается lane.backward_wgpu(m, gpu), а при ошибке — CPU lane.backward(m).

Проекционные адъюнкты выполняются на устройстве только если m.device.gpu().filter(|g| g.accelerates_backward()) даёт Some, и тогда вызываются gpu.gemm_nn_into и gpu.gemm_tn_accumulate_into, иначе — CPU-функции stages::dense_input_adjoint и stages::dense_weight_adjoint. Финальный scatter по embedding и RMSNorm-градиенты в bwd_stage_ssm_cuda остаются на хосте: на устройстве остаются reverse affine maps, tiled scan и token-local derivatives, а на host-owned tape возвращаются только существующая dense projection boundary и финальный RMSNorm/embedding scatter.

WGSL-шейдеры

affine_scan_main — compute-шейдер с @workgroup_size(256,1,1), где каждый invocation владеет одной строкой. Он загружает a и b в workgroup-массивы по 256 элементов и выполняет workgroup-local Blelloch-подобный up-sweep: на каждом шаге offset удваивается, и при (local+1)%(2*offset)==0 выполняется аффинное слияние scan_shared_a[local] = old_a * scan_shared_a[left]; scan_shared_b[local] = old_a * scan_shared_b[left] + old_b. Затем invocation с local==0 записывает tile-сумму в summary-массивы, а последний элемент обнуляется нейтральным (1.0, 0.0) для последующего down-sweep в affine_scan_apply_main, который идёт от down=128 к 0 и применяет родительский префикс перед левым поддеревом.

Произвольные длины обрабатываются рекурсивно: launch_scan вычисляет (capacity, tiles) через scan_shape, запускает scan с grid_dim=(stride, tiles, 1) и block_dim=(capacity,1,1), а если tiles>1, рекурсивно вызывает launch_scan для summary-массивов, получая prefix_a/prefix_b, и затем запускает scan_apply, который для каждого элемента вычисляет tile = time / capacity и prefix = tile*stride + channel, применяя apply_a[flat] = a*pa; apply_b[flat] = a*pb + b.

ssm_materialize_main для каждого элемента вычисляет t и i, итерирует по состояниям, восстанавливает before = ssm_scan_a_in[index]*ssm_initial[state] + ssm_scan_b_in[index], затем h = ssm_local_a[index]*before + ssm_local_b[index]*ssm_x[...], пишет состояние и накапливает выход.

memory_distance — чистая функция гиперболического расстояния: z=(dx*dx+dy*dy)/denom, return 2.0*log(sqrt(z)+sqrt(1.0+z)) при denom=(1.0-qsq)*(1.0-ksq).

memory_forward_main для каждого токена сначала вычисляет q-вектор, затем делает scaled norm: находит max_abs, делит каждый q на max_abs, накапливает scaled_sq, вычисляет scaled_radius = sqrt(scaled_sq) и radius = min(max_abs*scaled_radius, 3.402823e+38). Это и есть защита от переполнения f32: деление на max_abs до возведения в квадрат не даёт произведению q*q выйти за пределы f32. Далее project_scale выбирается по условию radius > 2.0e+6, где max_radius = 1.0 - 8.0*1.1920929e-7. Затем для каждого слота вычисляется расстояние, находится минимум, веса пересчитываются как exp((min_dist - weight)/tau), нормируются делением на сумму и используются для накопления взвешенной суммы значений.

CUDA-путь

launch_scan сначала вычисляет форму через scan_shape, проверяет длины всех четырёх буферов и аллоцирует два device-only буфера summary. Затем запускает scan. Если tiles == 1, функция сразу возвращает Ok; иначе аллоцирует prefix-буферы и рекурсивно вызывает launch_scan на summary-буферах, после чего запускает scan_apply.

make_ssm_buffers проверяет длины входов и при несовпадении возвращает "CUDA SSM forward shape mismatch", затем копирует входы на устройство, аллоцирует нулями рабочие буферы, копирует initial в срез состояний, запускает prepare, вызывает launch_scan, затем materialize.

readback_ssm проверяет длины выходных срезов и при несовпадении возвращает "CUDA SSM output shape mismatch", после чего копирует обратно на хост bar_a, bar_b, состояния и выход.

run_memory_forward проверяет длины буферов, копирует веса и банк на устройство, выполняет gemm для запросов, складывает их, запускает memory_forward; при count == 0 делает memset_zeros, иначе вызывает BLAS gemm на активных значениях и весах, затем gemm для гейта, sigmoid, gemm для проекции и sigmoid_mul. readback_memory копирует обратно все результаты и вызывает stream.synchronize().

Ошибки JIT-лога драйвера обрабатываются в stage_kernels: при ошибке load_module формируется сообщение через stage_load_error, которое вызывает driver_error и jit_error_log; последний повторно загружает PTX через cuModuleLoadDataEx с опциями CU_JIT_ERROR_LOG_BUFFER и CU_JIT_ERROR_LOG_BUFFER_SIZE_BYTES, читает лог до первого нулевого байта и возвращает его, если результат не CUDA_SUCCESS и сообщение непустое; иначе возвращается "CUDA driver returned no JIT error log".

SequenceBatch: упаковка и параллельные линии

SequenceBatch хранит независимые линии как упакованные строки: каждая линия получает offset и len, а токены копируются в общий лента без padding, так что плотные стадии видят сумму активных длин как строки GEMM.

Перед любой мутацией carry/tape forward проверяет весь батч: уникальность lane, непустые и совпадающие по длине inputs/targets не длиннее chunk_len, и что все id меньше d_vocab. backward требует pending forward и конечный accumulation_scale, а при stacked/looped моделях вызывает backward_stacked.

publish_stacked_terminals копирует losses и терминальные q_poincare/z_final каждой активной линии в соответствующие позиции ленты основного и extra_blocks. copy_carry_to_model/copy_carry_from_model переносят состояние через m.copy_recurrent_state_from/to, а forward_stacked сохраняет initial_carry линии в replay и восстанавливает model_carry после прохода.

Поведение при сбое резидентного backward

При сбое резидентного CUDA-бэкварда (например, при отсутствии воркспейса) печатается предупреждение, выставляется cuda_carry_dirty = true, и выполняется ровно один CPU-реплей: rebuild_host_tapes и bwd_stage_memory, после чего resident_backward остаётся false. Правки carry сохраняются потому, что state_mut помечает carry как изменённый и возвращает ссылку на carry линии. После реплея carry совпадает у reference и recovered, последний шаг не использовал резидентный CUDA-путь, а следующий forward стартует от сохранённого carry, а не продвигает его второй раз.

Инференс и генерация

Инференс-пути не аллоцируют буферы: forward_inference использует заранее выделенные inf_features и inf_block_out, а forward_dream_input — те же буферы, но принимает input длиной d_latent вместо индекса токена.

Carry читается через swap_loop_carry: forward_inference_loop сначала вызывает swap_loop_carry(loop_index), затем forward_continuous_inference, затем снова swap_loop_carry(loop_index), чтобы подменить h_persistent на сохранённый carry нужного лупа.

forward_dream_input превращает эпизодический латент в генерацию, копируя input в inf_features и прогоняя те же блоки, после чего применяет unembed_w.matvec и масштабирует логиты на 1/sqrt(d).

Формат сохранения V7

V7 — это полный payload обучения/возобновления V6 плюс ограниченный по длине стандартный JSON токенизатора с префиксом длины. Такой BPE-чекпоинт самодостаточен и восстанавливает свой точный упорядоченный словарь без доступа к корпусу обучения или оценки. При загрузке V7 BPE восстанавливается именно встроенный токенизатор, и он никогда не пересобирается из выбранного набора данных. Команды generate и chat отклоняют --data для V7 BPE, поскольку переобучение токенизатора на внешних данных не подтвердило бы происхождение. Для V7 word-чекпоинтов и V6-чекпоинтов сохраняется устаревшее необязательное сравнение точного словаря через --data. Проверенные артефакты V5 остаются только для инференса и требуют --data, так как они никогда не содержали сведений о происхождении токенизатора.

Конфигурация и валидация

validate проверяет корректность конфигурации через assert: depth должен быть в диапазоне 1..=32, все размерности (d_vocab, d_latent, d_state, d_mem_key, mem_capacity, chunk_len) должны быть положительными, lr/eps/tau_mem — конечными и положительными, beta1/beta2 — в [0,1), weight_decay неотрицательным, ema_alpha в [0,1].

Значения по умолчанию: depth=1, d_vocab=10_000, d_latent=256, d_state=16, d_mem_key=32, mem_capacity=512, chunk_len=64, lr=1e-3, beta1=0.9, beta2=0.999, weight_decay=0.01, eps=1e-8, tau_mem=0.1, ema_alpha=0.01.

validate_loops_config проверяет, что loops в 1..=MAX_LOOPS, затем вызывает validate_model_config, вычисляет базовый размер и добавляет оценку памяти под дополнительные проходы:

per_pass = l * (15 * d + 2 * s + 2 * k + mem + 34 + 3 * d * s) + 2 * d * s
bytes = base + 4 * cfg.depth * (loops - 1) * per_pass

Если bytes превышает MAX_LOAD_ALLOCATION_BYTES — возвращается ошибка. resize_loop_tape пересоздаёт tape для основного и extra_blocks и создаёт ScanExecutor::with_loops(loops); set_loops проверяет диапазон, и если loops изменился, клонирует cfg с chunk_len = tape.max_l, вызывает validate_loops_config и resize_loop_tape.

Проверка webgpu limits без устройства выполняется функцией checked_wgpu_sizes, которая принимает wgpu::Limits (например, Limits::default()) и размеры, вызывает checked_gemm_sizes и shared_gemm_rows, проверяет, что размеры не превышают u32::MAX, что байты буферов не превышают max_buffer_size и max_storage_buffer_binding_size, что число workgroup-групп не превышает max_compute_workgroups_per_dimension, и что max_compute_workgroup_size_x/y >= 16, max_compute_invocations_per_workgroup >= 256, max_compute_workgroup_storage_size >= 2048.

Как сравнивали и что получилось

Стенд и методология

Сравнение PSSA и параметр-матченного трансформера проводилось на CPU-к-CPU стенде. Для честного сравнения обе модели обучались на одной и той же CPU-машине — 2-vCPU контейнер без GPU, на одном и том же срезе корпуса в 199,059 токенов, с seed 42 и одинаковым числом обновлений.

Размер моделей — 1.5M параметров на 12.7M токенов. Конфигурация baseline: 1,541,120 параметров, 1 слой, ширина 256, 4 головы, FFN 448, vocab 2,048. Конфигурация PSSA: latent 256, recurrent state 16, 512 memory slots, key width 32, vocab 2,048. Оптимизация у обоих одинаковая: 30,000-update cosine horizon, без warm-up restart, 512 supervised target tokens per update, seed 42.

Отдельно указано, что headline-числа throughput не согласованы по железу: PSSA обучалась на Kaggle T4 (~900 токенов/с), а baseline — только на CPU, потому что train-transformer не имеет GPU-пути, и держала 212 токенов/с.

Воспроизведение выполняется двумя командами: bash kaggle/kaggle_continue.sh для цепочки PSSA и bash kaggle/kaggle_transformer_baseline.sh для параметрически согласованного baseline. Оба скрипта читают из окружения переменные TOTAL, WINDOW и FRESH и пишут --loss-csv, чтобы кривая сохранялась при обрыве сессии. kaggle_continue.sh задаёт размер окна и число звеньев, идёт по корпусу со смещением и возобновляет каждое звено с предыдущего чекпоинта, а status затем показывает каждый чекпоинт в цепочке с его формой и числом шагов оптимизатора.

Пропускная способность

На согласованном CPU-стенде PSSA показала 1,716 токенов/с против 415 у baseline, то есть 4.1x на согласованном железе и согласованной работе.

Кривая обучения по звеньям

End-of-link training cross-entropy по каждому звену цепочки:

ЗвеноТокеновPSSATransformer
ck01200,0005.7336.461
ck051,000,0004.6175.467
ck102,000,0004.4475.082
ck153,000,0004.2924.858
ck204,000,0004.1854.704
ck255,000,0004.

Где смотреть в коде

Источники

Похожее