pulsatrix
Loading...
Searching...
No Matches
transformer_block.hpp
Go to the documentation of this file.
1
7#pragma once
8
9#include <optional>
10#include "pulsatrix/module.hpp"
14
15namespace pulsatrix {
16
50class TransformerBlock : public Module {
51public:
65 TransformerBlock(int64_t d_model, int64_t num_heads, int64_t d_ff, DeviceBackend* backend, bool use_rope = true,
66 bool use_qk_norm = false);
67
78 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
79
81 [[nodiscard]] OpType op_type() const override { return OpType::Elementwise; }
82
94 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
95
97 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override;
98
100 void set_training(bool training) override;
101
102 [[nodiscard]] int64_t d_model() const { return mha_.d_model(); }
103
106 [[nodiscard]] RMSNormModule& norm1() { return norm1_; }
107 [[nodiscard]] MultiHeadAttentionModule& mha() { return mha_; }
108 [[nodiscard]] RMSNormModule& norm2() { return norm2_; }
109 [[nodiscard]] SwiGLUModule& swiglu() { return swiglu_; }
111
112
114 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
115
116protected:
123 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
124
125private:
126 int64_t d_model_;
127 DeviceBackend* backend_;
128
129 RMSNormModule norm1_;
131 RMSNormModule norm2_;
132 SwiGLUModule swiglu_;
133
134 // Forward caches -- the four operands the two residual splits need.
135 Shape last_input_shape_ = Shape({0});
136 Tensor last_x_;
137 Tensor last_attn_out_;
138 Tensor last_y1_;
139 Tensor last_ffn_out_;
140 bool has_forwarded_ = false;
141};
142
143} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
virtual DeviceType device() const noexcept=0
Which device this backend's buffers reside on.
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
softmax(Q @ K^T / sqrt(head_dim)) @ V, multi-head, with optional RoPE and optional QK-Norm....
Definition multihead_attention_module.hpp:52
int64_t d_model() const
Definition multihead_attention_module.hpp:137
y_{n,i} = gamma_i * x_{n,i} / rms(x_n), rms(x_n) = sqrt(mean_i(x_{n,i}^2) + eps), computed independen...
Definition rms_norm_module.hpp:32
An N-dimensional shape. A plain aggregate of dimensions with no invariant beyond "non-negative dimens...
Definition shape.hpp:24
down_proj(silu(gate_proj(x)) * up_proj(x)), the gated feedforward block used in place of a plain two-...
Definition swiglu_module.hpp:49
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
y1 = x + MHA(RMSNorm(x)), y2 = y1 + SwiGLU(RMSNorm(y1)). Shape (N, L, d_model) -> (N,...
Definition transformer_block.hpp:50
std::vector< NamedParamRef > named_parameters() override
norm1_'s, mha_'s, norm2_'s, and swiglu_'s parameters, flattened.
RMSNormModule & norm1()
Definition transformer_block.hpp:106
RMSNormModule & norm2()
Definition transformer_block.hpp:108
int64_t d_model() const
Definition transformer_block.hpp:102
MultiHeadAttentionModule & mha()
Definition transformer_block.hpp:107
OpType op_type() const override
Elementwise per this module's own op_type() note above.
Definition transformer_block.hpp:81
SwiGLUModule & swiglu()
Definition transformer_block.hpp:109
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition transformer_block.hpp:114
Tensor backward(const Tensor &grad_output) override
Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside norm1_/mha_/nor...
TransformerBlock(int64_t d_model, int64_t num_heads, int64_t d_ff, DeviceBackend *backend, bool use_rope=true, bool use_qk_norm=false)
Constructs a transformer block with zero-initialized sub-module parameters.
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
LRP relevance propagation: two residual epsilon/z-rule splits composed with norm1_'s/mha_'s/norm2_'s/...
Tensor forward_impl(const Tensor &input) override
Runs: norm1 -> attention -> residual add -> norm2 -> SwiGLU -> residual add.
void set_training(bool training) override
Cascades to every sub-module, the same way SequentialModule/MultiHeadAttentionModule do.
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
Multi-head scaled dot-product attention – this codebase's first Module composed out of other real Mod...
Definition acquisition_functions.hpp:16
OpType
The op-type tag a Node carries. Charter Part 2 §3: nodes are tagged by a small closed set of op types...
Definition op_type.hpp:19
RMS normalization layer (Zhang & Sennrich, 2019) – this codebase's first normalization Module,...
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57
SwiGLU gated feedforward block – second module composed from real LinearModule sub-objects,...