От gate перед Maxpool я избавился из-за отсутствия особой разницы, в связи с этим стал разумным шаг - перенести линейный слой с активацией на место gate, до Maxpool, т.е. новый класс стал совсем простым:
Код:
CausalMaxpool = Linear + ReLU + Maxpool
Хочу обратить внимание (может кто-то ещё не въехал в тему), что Maxpool здесь необычный -
каузальный, а именно:
1. Сравнение двух соседних токенов: kernel_size=2
2. Плотная структура, не разреженная: stride=1
3. Левосторонний паддинг: padding=(1, 0)
4. Возможность использования растущего dilation > 1 для более эффективного расширения рецептивного поля.
Именно такой Maxpool имеет каузальные свойства, то есть "не подглядывает в будущие токены".
Обновлённый класс (удалён gate, перенесёны Linear+ReLU):
Код:
import torch
import torch.nn as nn
import torch.nn.functional as F
class CausalMaxPool(nn.Module):
def __init__(
self,
d_model: int,
window_size: int = 2,
dilation: int = 1,
):
super().__init__()
if d_model < 1:
raise ValueError("d_model must be >= 1")
if window_size < 2:
raise ValueError("window_size must be >= 2")
if dilation < 1:
raise ValueError("dilation must be >= 1")
self.d_model = int(d_model)
self.window_size = int(window_size)
self.dilation = int(dilation)
self.proj = nn.Linear(d_model, d_model)
self.act = nn.ReLU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
x: (..., T, d_model)
returns: (..., T, d_model)
"""
if x.ndim < 2:
raise ValueError("Expected x with shape (..., T, d_model)")
if x.size(-1) != self.d_model:
raise ValueError(
f"Expected last dim {self.d_model}, got {x.size(-1)}"
)
x = self.act(self.proj(x))
original_shape = x.shape
T, D = original_shape[-2:]
batch_shape = original_shape[:-2]
# (..., T, D) -> (N, D, T)
x = x.reshape(-1, T, D).transpose(1, 2)
left_pad = (self.window_size - 1) * self.dilation
# Важно: ручной asymmetric padding.
# Аргумент padding у max_pool1d здесь не подходит:
# он симметричный и ограничен.
x = F.pad(x, (left_pad, 0), value=float("-inf"))
# (N, D, T + left_pad) -> (N, D, T)
pooled = F.max_pool1d(
x,
kernel_size=self.window_size,
stride=1,
padding=0,
dilation=self.dilation,
)
# (N, D, T) -> (..., T, D)
pooled = pooled.transpose(1, 2).reshape(*batch_shape, T, D)
return pooled
Итак, я провёл предварительные исследования, произвёл отбор кандидатов среди следующих:
1. Равномерное перцептивное поле (dilation = 1, 1, 1, ...)
2. Фиббоначчиево перцептивное поле (dilation = 1, 2, 3, 5, 8, 13, ...)
3. Экспоненциальное перцептивное поле (dilation = 1, 2, 4, 8, 16, 32, ...)
4. Комбинированную версию (dilation = 1, 1, 1, 2, 3, 5, 10, 20, ...)
Мне они все показались одинаково рабочими, поэтому я пока решаюсь остановиться на экспоненциальном варианте как на самом эффективном.
Провёл pretrained-этап создания GPT-модели, обучал на статьях русской википедии, далее чисто теоретически можно обучать в качестве чат-бота. Модель научилась закрывать скобки, довольно часто правильно ставит запятые и точки. Немного зацикливается, придумывает новые слова и словосочетания, короче, галлюцинирует.
Далее я испытал глубокую гибридную модель по принципу использования GatedDeltaNet в Qwen/Kimi:
1. Слои
CausalMaxpool (4 шт.)
2. Positional Encoding
3. Dropout
4. Слои трансформеров (2 шт.)
5. Слои
CausalMaxpool (4 шт.)
6. Positional Encoding
7. Dropout
8. Слои трансформеров (2 шт.)
Трансформеры сами по себе очень плохо обучаются, требуют пониженного learning rate и scheduling (learning rate изменяется по расписанию по определённому закону с "прогревом"). Гибрид из-за этого, в общем-то, тоже также следует обучать. Гибрид немного медленнее обучается, но предел сходимости (loss) гораздо ниже.
В общем CausalMaxpool следует рассматривать как "убийцу" GatedDeltaNet, а не самих трансформеров. Если GatedDeltaNet являются трансформерной структурой с линейным (точнее с субквадратичным) вниманием, то CausalMaxpool - это структура совершенно другой природы. Есть общее: а. вычислительная сложность также субквадратична, б. рецептивное поле также является локальным (не глобальным как у чистых трансформеров).
Далее предстоит сравнить CausalMaxpool с GatedDeltaNet в обучении.