pulsatrix
Loading...
Searching...
No Matches
swiglu_module.hpp
Go to the documentation of this file.
1
6#pragma once
7
8#include <optional>
10#include "pulsatrix/module.hpp"
11
12namespace pulsatrix {
13
49class SwiGLUModule : public Module {
50public:
60 SwiGLUModule(int64_t d_model, int64_t d_ff, DeviceBackend* backend);
61
73 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
74
76 [[nodiscard]] OpType op_type() const override { return OpType::Elementwise; }
77
90 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
91
93 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override;
94
95 [[nodiscard]] int64_t d_model() const { return d_model_; }
96 [[nodiscard]] int64_t d_ff() const { return d_ff_; }
97
100 [[nodiscard]] LinearModule& gate_proj() { return gate_proj_; }
101 [[nodiscard]] LinearModule& up_proj() { return up_proj_; }
102 [[nodiscard]] LinearModule& down_proj() { return down_proj_; }
104
105
107 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
108
109protected:
116 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
117
118private:
119 int64_t d_model_;
120 int64_t d_ff_;
121 DeviceBackend* backend_;
122
123 LinearModule gate_proj_;
124 LinearModule up_proj_;
125 LinearModule down_proj_;
126
127 // Forward caches -- the operands the gate multiply's forward/backward/LRP rule needs.
128 // gate_proj_/up_proj_/down_proj_ also cache their own inputs internally; this module
129 // additionally caches these three for its own composed math, same choice
130 // MultiHeadAttentionModule made for its Q/K/V/scores caches.
131 Shape last_input_shape_ = Shape({0});
132 int64_t last_n_flat_ = 0;
133 Tensor last_gate_pre_;
134 Tensor last_gate_post_;
135 Tensor last_up_;
136 bool has_forwarded_ = false;
137};
138
139} // 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.
y = x @ W + b, batched (x is (N, in_features), y is (N, out_features)) – migrated from the original u...
Definition linear_module.hpp:37
Base class for every layer type (LinearModule, Conv2DModule, activations, ...).
Definition module.hpp:58
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
OpType op_type() const override
Elementwise per this module's own op_type() note above.
Definition swiglu_module.hpp:76
LinearModule & up_proj()
Definition swiglu_module.hpp:101
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition swiglu_module.hpp:107
SwiGLUModule(int64_t d_model, int64_t d_ff, DeviceBackend *backend)
Constructs a SwiGLU block with zero-initialized projections.
int64_t d_ff() const
Definition swiglu_module.hpp:96
Tensor forward_impl(const Tensor &input) override
Runs: project (gate, up) -> silu(gate) -> gate*up -> project (down).
Tensor backward(const Tensor &grad_output) override
Gradient w.r.t. this module's input; sub-module parameter gradients accumulate inside gate_proj_/up_p...
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
LRP relevance propagation: down_proj_'s epsilon rule, then the diagonal Eq. 15 split into gate/up sha...
int64_t d_model() const
Definition swiglu_module.hpp:95
std::vector< NamedParamRef > named_parameters() override
gate_proj_'s, up_proj_'s, and down_proj_'s parameters, flattened.
LinearModule & gate_proj()
Definition swiglu_module.hpp:100
LinearModule & down_proj()
Definition swiglu_module.hpp:102
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
Dense/fully-connected layer – the reference Module implementation.
Abstract base every layer subclasses – NVI forward(), pure-virtual LRP contract.
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
Configuration for LRP relevance propagation: which rule a module applies and its hyperparameters....
Definition lrp_rule_config.hpp:57