pulsatrix
Loading...
Searching...
No Matches
embedding_module.hpp
Go to the documentation of this file.
1
5#pragma once
6
7#include <initializer_list>
8#include <vector>
9
10#include "pulsatrix/module.hpp"
11
12namespace pulsatrix {
13
39class EmbeddingModule : public Module {
40public:
50 EmbeddingModule(int64_t num_embeddings, int64_t embedding_dim, DeviceBackend* backend);
51
63 [[nodiscard]] Tensor backward(const Tensor& grad_output) override;
64
66 [[nodiscard]] OpType op_type() const override { return OpType::Embedding; }
67
69 void set_weight(std::initializer_list<float> values);
71 void set_weight(const std::vector<float>& values);
72
73 [[nodiscard]] const Tensor& weight() const { return weight_; }
74 [[nodiscard]] const Tensor& weight_grad() const { return weight_grad_; }
75
88 [[nodiscard]] Tensor propagate_relevance(const Tensor& relevance_out, const LRPRuleConfig& config) override;
89
90 [[nodiscard]] std::vector<NamedParamRef> named_parameters() override {
91 return {
92 {"weight", {&weight_, &weight_grad_}},
93 };
94 }
95
96protected:
103 [[nodiscard]] Tensor forward_impl(const Tensor& input) override;
104
105private:
106 int64_t num_embeddings_;
107 int64_t embedding_dim_;
108 DeviceBackend* backend_;
109 Tensor weight_; // shape (num_embeddings, embedding_dim)
110 Tensor weight_grad_;
111 Shape last_input_shape_ = Shape({0});
112 // Flat N*L validated indices as whole-number floats on the weights' device (exact below
113 // 2^24 -- far above any vocabulary this module is built for), consumed by gather_rows /
114 // scatter_add_rows (GPU-native-kernels Mission 2).
115 Tensor last_indices_ = Tensor(Shape({0}), backend_);
116 bool has_forwarded_ = false;
117};
118
119} // namespace pulsatrix
Vendor-agnostic compute/memory backend. CPUBackend, CUDABackend (Phase 1.5), and HIPBackend (Phase 1....
Definition device_backend.hpp:219
Embedding lookup table, rank-2 input (N, L) of float-encoded indices -> rank-3 output (N,...
Definition embedding_module.hpp:39
OpType op_type() const override
Embedding per charter's closed OpType set.
Definition embedding_module.hpp:66
void set_weight(const std::vector< float > &values)
Vector overload for runtime-sized sources – see Tensor's own vector ctor.
Tensor propagate_relevance(const Tensor &relevance_out, const LRPRuleConfig &config) override
Sum-over-embedding-dimension LRP relevance aggregation (Arras et al. 2017).
std::vector< NamedParamRef > named_parameters() override
This module's trainable parameters, each with its hierarchical name – the one place a module declares...
Definition embedding_module.hpp:90
const Tensor & weight() const
Definition embedding_module.hpp:73
Tensor backward(const Tensor &grad_output) override
Scatter-adds grad_output into the corresponding rows of weight_grad_.
void set_weight(std::initializer_list< float > values)
Overwrites the weight buffer – test/initialization use only.
const Tensor & weight_grad() const
Definition embedding_module.hpp:74
EmbeddingModule(int64_t num_embeddings, int64_t embedding_dim, DeviceBackend *backend)
Constructs an embedding table with a zero-initialized weight matrix.
Tensor forward_impl(const Tensor &input) override
The actual forward computation – per-position row copy from weight_.
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
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