pulsatrix
Loading...
Searching...
No Matches
rope_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <optional>
9
10namespace pulsatrix {
11
35class RoPEModule : public Module {
36public:
49 RoPEModule(int64_t head_dim, DeviceBackend* backend, float base = 10000.0f);
50
67 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
68
84 [[nodiscard]] OpType op_type() const override { return OpType::Elementwise; }
85
117 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
118
119
121 [[nodiscard]] std::optional<DeviceType> compute_device() const override { return backend_->device(); }
122
123protected:
131 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
132
133private:
134 int64_t head_dim_;
135 float base_;
136 DeviceBackend* backend_;
137 Tensor last_input_;
138 Tensor last_output_;
139 bool has_forwarded_ = false;
140
148 void ensure_tables(int64_t seq_len);
149 Tensor cos_table_;
150 Tensor sin_table_;
151 int64_t table_seq_len_ = 0;
152};
153
154} // 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
Rotary Position Embedding (RoPE, Su et al. 2021): a fixed, non-learnable, position-dependent rotation...
Definition rope_module.hpp:35
Tensor forward_impl(const Tensor &input) override
Applies the per-position pair rotation.
std::optional< DeviceType > compute_device() const override
Where this layer computes, so forward() rejects an input on another device (FND-8).
Definition rope_module.hpp:121
RoPEModule(int64_t head_dim, DeviceBackend *backend, float base=10000.0f)
Constructs a RoPE module.
OpType op_type() const override
Elementwise per the charter's closed OpType set.
Definition rope_module.hpp:84
Tensor backward(const Tensor &grad_output) override
Computes the gradient w.r.t. this module's input.
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Standard weighted-connection epsilon-rule LRP for this fixed linear map.
N-dimensional tensor. Owns its data buffer exclusively; a DeviceBackend* is injected (not owned) – th...
Definition tensor.hpp:29
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