|
Mila
Deep Neural Network Library
|
Abstract base class for neural network components. More...
Public Member Functions | |
| Component (const std::string &name) | |
| Construct component with required name identifier. | |
| virtual void | build (const BuildContext &context) final |
| Build the component with the provided BuildContext (canonical overload). | |
| virtual DeviceId | getDeviceId () const =0 |
| Get the compute device id associated with this component. | |
| virtual std::vector< ITensor * > | getGradients () const =0 |
| Return non-owning pointers to parameter gradient tensors. | |
| virtual MemoryStats | getMemoryStats () const =0 |
| Return the current memory allocation breakdown for this component. | |
| const std::string | getName () const |
| Get the component's name identifier. | |
| virtual std::vector< std::string > | getParameterNames () const |
| List all available parameter names for this component. | |
| virtual std::vector< ITensor * > | getParameters () const =0 |
| Return non-owning pointers to parameter tensors. | |
| virtual MemoryStats | getRequiredMemory (const BuildContext &context) const |
| Report what build( context ) would allocate, without allocating it. | |
| TrainingMode | getTrainingMode () const noexcept |
| The current runtime behavioral mode of this Component. | |
| virtual const ComponentType | getType () const =0 |
| Get the component type identifier. | |
| virtual bool | isBuilt () const final |
| Returns true if build() has completed successfully. | |
| virtual void | load_ (ModelArchive &archive, SerializationMode mode) |
| Restore this component's parameters from its archive scope. | |
| virtual void | loadParameter (const std::string &, const Serialization::ITensorBlob &) |
| Load a parameter from serialized tensor data. | |
| virtual dim_t | parameterCount () const =0 |
| Return number of trainable parameters. | |
| virtual void | requireSerializableParameters () const |
| Verify this component can serialize whatever parameters it owns. | |
| virtual void | save_ (ModelArchive &archive, SerializationMode mode) const =0 |
| virtual void | saveFlatTensors (Serialization::SafeTensorsWriter &writer, const std::string &prefix, Serialization::TensorSavePass pass) const |
| Drive this component's tensors through one pass of a flat safetensors save. | |
| void | setTrainingMode (TrainingMode mode) |
| Set the runtime behavioral mode for this Component. | |
| virtual void | synchronize ()=0 |
| virtual std::string | toString () const =0 |
| Produce a short, human-readable description of the component. | |
| virtual void | zeroGradients () |
| Clear all model-owned gradients for this component. | |
Static Public Member Functions | |
| static constexpr DeviceType | getDeviceType () |
| Compile-time device type for this component instance. | |
| static constexpr TensorDataType | getPrecision () noexcept |
| Compile-time tensor precision for this component instance. | |
Protected Types | |
| using | HostStagingMemoryResource |
| Host memory a device-resident parameter stages through. | |
Protected Member Functions | |
| IExecutionContext * | getExecutionContext () const |
| Get the shared execution context. | |
| bool | hasExecutionContext () const noexcept |
| Check if execution context has been set. | |
| template<TensorDataType TParameterPrecision, typename TMemoryResource> | |
| void | loadParameterFromBlob (const std::string ¶m_name, const Serialization::ITensorBlob &blob, Tensor< TParameterPrecision, TMemoryResource > &target, const shape_t &expected_shape) |
| Load a tensor blob into a parameter tensor with validation. | |
| virtual void | onBuilding (const BuildContext &) |
| Hook invoked by build() to allocate component buffers. | |
| virtual void | onExecutionContextSet () |
| Lifecycle hook: Called immediately after ExecutionContext is set. | |
| virtual void | onTrainingModeChanging (TrainingMode) |
| Hook called before TrainingMode transitions. | |
| template<TensorDataType TParameterPrecision, typename TMemoryResource> | |
| void | saveParameterToArchive (ModelArchive &archive, const std::string ¶meter_name, const Tensor< TParameterPrecision, TMemoryResource > ¶meter) const |
| Write one parameter tensor into the archive under "tensors/<name>". | |
| template<TensorDataType TParameterPrecision, typename TMemoryResource> | |
| void | saveParameterToWriter (Serialization::SafeTensorsWriter &writer, const std::string &flat_name, const Tensor< TParameterPrecision, TMemoryResource > ¶meter, Serialization::TensorSavePass pass) const |
| Drive one parameter through one pass of a flat safetensors save. | |
| void | setExecutionContext (IExecutionContext *context) |
| Set the execution context for this component. | |
Protected Attributes | |
| BuildContext | build_context_ { shape_t{ 1 }, RuntimeMode::Training } |
| The BuildContext stored at build time. | |
Friends | |
| template<DeviceType, TensorDataType> | |
| class | CompositeComponent |
| template<DeviceType, TensorDataType> | |
| class | Network |
| std::ostream & | operator<< (std::ostream &os, const Component &component) |
| Stream output uses toString() to provide a human-readable description of the component. | |
Abstract base class for neural network components.
Component enforces a single ownership model: all components receive a non-owning pointer to an IExecutionContext that is owned by the parent (Network or test fixture). This ensures consistent resource sharing and eliminates dual constructor patterns.
Components progress through a well-defined lifecycle. Each stage has a single responsibility and a designated hook for subclass extension.
The component is constructed with its name and component config. No device resources are allocated. No ExecutionContext is required.
build() is called with a BuildContext carrying the leading shape { B, T, ... } and the ExecutionMode that governs buffer allocation.
| Mode | allocationSeqLen() | Gradient buffers |
|---|---|---|
| Inference | 1 | never allocated |
| Training | leading_shape[1] | allocated on demand |
Allocated in onBuilding():
NOT allocated in onBuilding():
The same BuildContext is cascaded unchanged through CompositeComponent and Network to all child components.
Only valid for Training-built components. setEvaluation( false ) triggers gradient buffer allocation on the first call. Subsequent calls zero existing buffers without reallocating. setEvaluation( true ) zeros gradient buffers and disables the backward path.
Allocated on first setEvaluation( false ):
Runtime dimensions are read from the input tensor shape on each call. No shape information is cached from build time beyond what is in build_config_.
build() requires ExecutionContext to be set setEvaluation() requires build() to have completed setEvaluation() requires ExecutionMode::Training forward() requires build() to have completed backward() requires isTrainingMode() == true decode() requires build() to have completed
| TDeviceType | Compile-time device identifier for this component. |
| TPrecision | Tensor data precision for this component. |
|
protected |
Host memory a device-resident parameter stages through.
Every serialization path here has to get device bytes somewhere the host can read them. Naming the CUDA pinned resource directly would put a CUDA import in an always-compiled core module, which a CPU-only build cannot satisfy – the backend axis belongs in DeviceTypeTraits, which is itself backend-gated.
|
inlineexplicit |
Construct component with required name identifier.
The name is used for identification, logging, and serialization. Names must be valid identifiers: start with a letter, contain only letters, digits, '.', '_', '-', and be 1-128 characters long.
| name | Component name identifier (mandatory). |
| std::invalid_argument | if name is not a valid identifier. |
|
inlinefinalvirtual |
Build the component with the provided BuildContext (canonical overload).
Validates the config, stores it as build_config_, then invokes the onBuilding() hook for component-specific buffer allocation and initialization.
After onBuilding() returns without throwing, the component is marked built and isBuilt() returns true. If onBuilding() throws, built_ remains false and build() may be retried – but only if the onBuilding() implementation leaves component state coherent on failure.
The stored BuildContext is accessible to derived classes via the protected build_config_ member throughout the component lifetime.
| context | Build-time configuration carrying the leading shape { B, T, ... }, ExecutionMode, and optional micro-batching settings. |
| std::runtime_error | if the component is already built. |
| std::runtime_error | if no ExecutionContext has been set. |
| std::invalid_argument | if config.validate() fails. |
| Any | exception from onBuilding(). |
|
pure virtual |
Get the compute device id associated with this component.
Must return the device on which parameters and operations execute.
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Network< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlineprotected |
Get the shared execution context.
Provides access to the execution context for derived classes to:
| std::runtime_error | if context has not been set. |
|
pure virtual |
Return non-owning pointers to parameter gradient tensors.
Gradient buffers are allocated only when the component is built in training mode, so a component built for inference returns an empty vector. Stateless components return empty in either mode. This is the accessor counterpart to getParameters() and does not throw on mode.
| std::runtime_error | if called before the component has been built. |
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
pure virtual |
Return the current memory allocation breakdown for this component.
Reflects allocations at the moment of the call. The returned stats naturally track the component lifecycle:
After construction – nothing; construction allocates none After build( Inference ) – parameters + T=1 state buffers After build( Training ) – parameters + T=full state buffers After setTrainingMode( Train ) – parameters + state + gradients
For CompositeComponent and Network, the returned stats are the recursive aggregate of all child components.
May be called at any time – no lifecycle preconditions.
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::GatedMLP< TDeviceType, TPrecision, TGate >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GptBlock< TDeviceType, TPrecision >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuant, TKvPolicy >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MLP< TDeviceType, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inline |
Get the component's name identifier.
The name is used for logging, diagnostics, and serialization.
|
inlinevirtual |
List all available parameter names for this component.
Returns an empty vector by default. Leaf components with parameters should override to return their canonical parameter name list in the same stable order used by save_() and loadParameter().
Reimplemented in Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
pure virtual |
Return non-owning pointers to parameter tensors.
The returned tensor pointers remain valid for the lifetime of the component. Order should be canonical (weights before biases).
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlinevirtual |
Report what build( context ) would allocate, without allocating it.
The predictive counterpart to getMemoryStats(): what this component would need against what it currently has. Called on a constructed but unbuilt component, which is possible because construction commits no device memory – see Specifications/MemoryFootprint.md.
Implementations mirror their own onBuilding(), which is why the two take the same BuildContext: an implementation cannot answer for a different shape, runtime mode, or prefill size than the one build() would receive.
Composites mirror whatever their getMemoryStats() twin does – children, plus directly-owned buffers, minus any pooling, output-sharing or weight- tying correction. A plain sum over children overcounts wherever a buffer is installed rather than self-allocated.
| context | The context that would be passed to build(). |
| std::logic_error | if this component has not implemented the contract. |
Reimplemented in Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuant, TKvPolicy >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlinenoexcept |
The current runtime behavioral mode of this Component.
Returns the current TrainingMode for Components built with RuntimeMode::Training. For Components built with RuntimeMode::Inference the return value is always TrainingMode::Eval – inference components never compute gradients.
|
pure virtual |
Get the component type identifier.
Used for serialization and runtime type identification.
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::GatedMLP< TDeviceType, TPrecision, TGate >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::GptBlock< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuant, TKvPolicy >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MLP< TDeviceType, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Network< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlineprotectednoexcept |
Check if execution context has been set.
|
inlinevirtual |
Restore this component's parameters from its archive scope.
The inverse of save_(), and unlike save_() it is NOT pure virtual: the default below is the whole implementation for every leaf that owns parameters. It walks getParameterNames() – the same vector save_() walked – reads each blob back, and hands it to loadParameter(), which already validates dtype and shape and performs any conversion or device upload. The two directions cannot drift because they iterate the same names.
Restores into an ALREADY-CONSTRUCTED, ALREADY-BUILT graph. This is weight restoration, not model reconstruction: the caller builds the same topology it saved (from a config it supplies or reads separately) and then calls this.
Composites override to recurse; a component with no parameters inherits an empty loop and needs no override at all.
| archive | Archive to read from, already scoped to this component. |
| mode | Serialization mode (accepted for symmetry with save_). |
| std::runtime_error | if a named parameter has no blob in the archive. |
Reimplemented in Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, and Mila::Dnn::Network< TDeviceType, TPrecision >.
|
inlinevirtual |
Load a parameter from serialized tensor data.
Loads raw tensor bytes directly into an existing parameter tensor, handling precision conversion and device upload as needed.
The component validates that the blob's shape matches the parameter's expected shape, then delegates to the backend to perform:
Takes the parameter name used to locate the target tensor, and a blob holding serialized tensor metadata and raw bytes.
| std::runtime_error | if component has no parameters to load. |
| std::runtime_error | if blob shape doesn't match parameter shape. |
Reimplemented in Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlineprotected |
Load a tensor blob into a parameter tensor with validation.
Validates dtype and shape match then copies blob data into the tensor. Intended for use in loadParameter() overrides.
| param_name | Parameter name (used in error messages). |
| blob | Source tensor blob from the model archive. |
| target | Destination tensor (must be initialized). |
| expected_shape | Expected tensor shape for validation. |
| std::invalid_argument | if dtype or shape mismatch. |
|
inlineprotectedvirtual |
Hook invoked by build() to allocate component buffers.
Receives the stored BuildContext. Implementations must use config.allocationSeqLen() when sizing output buffers – this is the single call that makes Inference and Training allocate the correct buffer sizes automatically without per-component logic.
The default implementation forwards to the legacy onBuilding( const shape_t& ) overload for backwards compatibility. New components should override this overload directly.
Takes the build-time configuration; use its allocationSeqLen() to obtain the correct output buffer sequence dimension.
Reimplemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::GatedMLP< TDeviceType, TPrecision, TGate >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GptBlock< TDeviceType, TPrecision >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuant, TKvPolicy >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MLP< TDeviceType, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlineprotectedvirtual |
Lifecycle hook: Called immediately after ExecutionContext is set.
Override this to perform initialization that requires a valid ExecutionContext. At the time this is called, getExecutionContext() is guaranteed to return a valid context.
Common uses:
Default implementation does nothing.
| Any | exception thrown will cause setExecutionContext() to fail and restore the component to a "context not set" state. |
Reimplemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlineprotectedvirtual |
Hook called before TrainingMode transitions.
Called by setTrainingMode() after validation and lock acquisition, before the internal state is updated. Derived classes override to respond to the transition – e.g. zeroing gradient buffers on transition to Eval, or re-enabling dropout on transition to Training.
The default implementation is a no-op.
Takes the incoming TrainingMode.
Reimplemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::GatedMLP< TDeviceType, TPrecision, TGate >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GptBlock< TDeviceType, TPrecision >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuant, TKvPolicy >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MLP< TDeviceType, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
pure virtual |
Return number of trainable parameters.
For leaf components this is the element count of owned parameter tensors. CompositeComponent and Network implementations should return the recursive aggregate across all children.
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlinevirtual |
Verify this component can serialize whatever parameters it owns.
getParameterNames() is the vocabulary shared by save_() and loadParameter(). A component that owns parameters but names none of them has no way to round-trip them, and a save_() that silently writes nothing is indistinguishable from a successful save – which is how an archive missing most of a model's weights gets reported as written. The save traversal calls this before each component so the failure names the component instead of surfacing as a short archive.
Virtual because a composite reports its children's parameters through parameterCount() but names none of them itself – the recursion checks each child in turn, so CompositeComponent overrides this to a no-op.
| std::runtime_error | if the component has parameters but no names for them. |
Reimplemented in Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >.
|
pure virtual |
Implemented in Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Network< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, and Mila::Dnn::Softmax< TDeviceType, TPrecision >.
|
inlinevirtual |
Drive this component's tensors through one pass of a flat safetensors save.
This is the path save_() cannot serve. A quantized weight is packed storage plus a scale companion and ModelArchive has no representation for that pairing, so save_() refuses it; the flat container expresses it as sibling tensors and this writes them.
The vocabulary is the flat format's: a fully qualified dotted component path plus a parameter name, exactly what loadParameters() parses on the way back in. Composites contribute nothing themselves – they recurse and extend the prefix – so the emitted set matches what the reader will route back through loadParameter().
Concrete parameter types, not ITensor, on purpose: a type-erased walk would have to ask at runtime whether a tensor's memory is host-accessible, and MemoryResource::is_host_accessible is a compile-time constant. Keeping concrete types here is what avoids widening an exported core type for a serialization concern.
The default refuses rather than writing nothing. A component that owns parameters and does not implement this would otherwise contribute silently to an artifact that loads, runs, and produces garbage – the same failure mode Phase 0 removed from save_().
| writer | Writer being driven. |
| prefix | Fully qualified component path, e.g. "tf_layer_0.qkv_proj". |
| pass | Declare reserves byte ranges; Write streams bytes in declaration order. |
Reimplemented in Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, kGlobal, TWeightQuant, TKvPolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, false, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GemmaBlock< TDeviceType, TPrecision, true, TWeightQuantization, NoKvCompression >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlineprotected |
Write one parameter tensor into the archive under "tensors/<name>".
The save counterpart to loadParameterFromBlob(). Serialization moves the parameter's bytes as stored – the archive records the tensor's own dtype, and any precision conversion happens on the way back in through loadParameter().
A device-resident parameter is staged through a host tensor of the SAME dtype. Widening to FP32 here would pair a byte count derived from the device dtype with a buffer holding a wider one, writing a fraction of the staged bytes under the wrong type label. Same-dtype staging constrains the staging memory – see the comment on the device branch.
| archive | Archive to write into, already scoped to this component. |
| parameter_name | Canonical name from getParameterNames(). |
| parameter | Parameter tensor to serialize. |
|
inlineprotected |
Drive one parameter through one pass of a flat safetensors save.
The flat artifact is keyed by dotted tensor name, not by archive scope, so the caller supplies the fully qualified name rather than a prefix plus a convention.
Device parameters stage through PINNED host memory of the same dtype for the same reason saveParameterToArchive does: every reduced precision Mila serves in is is_device_only, so Tensor<TParameterPrecision, CpuMemoryResource> is not a valid template-id. The staging tensor is scoped to the write so only one parameter is resident at a time, which is what makes a 22 GB model writable.
| writer | Writer being driven. |
| flat_name | Fully qualified tensor name. |
| parameter | Parameter tensor. |
| pass | Declare reserves the byte range; Write streams the bytes. |
|
inlineprotected |
Set the execution context for this component.
Establishes the device and execution environment. Can only be called once – the execution context is immutable after setting.
Called by:
After setting the context, the onExecutionContextSet() hook is invoked to allow the component to perform context-dependent initialization.
| context | Non-owning pointer to execution context (must be non-null). |
| std::invalid_argument | if context is null. |
| std::runtime_error | if context has already been set. |
| std::invalid_argument | if context device type doesn't match TDeviceType. |
| std::runtime_error | if onExecutionContextSet() throws; context is restored to nullptr on failure. |
|
inline |
Set the runtime behavioral mode for this Component.
Toggles between Training and Eval behavioral states at runtime. Only valid on Components built with RuntimeMode::Training – throws if called on a Component built with RuntimeMode::Inference.
| From | To | Effect |
|---|---|---|
| Training | Eval | Gradients off, dropout off, running stats |
| Eval | Training | Gradients on, dropout on, batch stats |
Derived classes respond to the transition via the onTrainingModeChanging() hook, called before the state is updated.
| mode | TrainingMode::Normal or TrainingMode::Eval. |
| std::runtime_error | if the component is not built. |
| std::runtime_error | if built with RuntimeMode::Inference. |
|
pure virtual |
@brief Convenience accessor -- true if currently in Eval mode.
Equivalent to getTrainingMode() == TrainingMode::Eval.
Valid for both RuntimeMode::Inference and RuntimeMode::Training
built components.
@return true if in Eval mode.
‍/
bool isEvalMode() const noexcept { return getTrainingMode() == TrainingMode::Eval; }
RuntimeMode getRuntimeMode() const noexcept
{
return build_context_.getRuntimeMode();
}
bool isInferenceMode() const noexcept
{
return build_context_.isInferenceMode();
}
bool isTrainingMode() const noexcept
{
return build_context_.isTrainingMode();
}
====================================================================
/**
@brief Wait for outstanding device work submitted by this component.
On CPU this may be a no-op. Use to ensure results are visible to
the host or to measure synchronous timings.
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Network< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
pure virtual |
Produce a short, human-readable description of the component.
Implementations should keep output concise and avoid throwing.
Implemented in Mila::Dnn::Activation< TDeviceType, TPrecision, TFn >, Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >, Mila::Dnn::GatedMLP< TDeviceType, TPrecision, TGate >, Mila::Dnn::Gelu< TDeviceType, TPrecision >, Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GptBlock< TDeviceType, TPrecision >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TComputePrecision, TKvPolicy >, Mila::Dnn::GroupedQueryAttention< TDeviceType, TPrecision, TKvPolicy >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MLP< TDeviceType, TPrecision >, Mila::Dnn::MultiHeadAttention< TDeviceType, TPrecision >, Mila::Dnn::Network< TDeviceType, TPrecision >, Mila::Dnn::Residual< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::Softmax< TDeviceType, TPrecision >, Mila::Dnn::SoftmaxCrossEntropy< TDeviceType, TPrecision >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu >, Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
inlinevirtual |
Clear all model-owned gradients for this component.
Default implementation is a no-op. Composite components should override to recurse to children. Leaf components should override to zero their parameter and activation gradients using device-aware helpers.
Reimplemented in Mila::Dnn::GatedMLP< TDeviceType, TPrecision, TGate >, Mila::Dnn::GptBlock< TDeviceType, TPrecision >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, Mila::Dnn::LayerNorm< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TComputePrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision >, Mila::Dnn::Linear< TDeviceType, TPrecision, TableQuantizationPolicy >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuant >, Mila::Dnn::Linear< TDeviceType, TPrecision, TWeightQuantization >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuant, TKvPolicy >, Mila::Dnn::LlamaBlock< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::Lpe< TDeviceType, TIndex, TPrecision >, Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision >, Mila::Dnn::MLP< TDeviceType, TPrecision >, Mila::Dnn::RmsNorm< TDeviceType, TPrecision >, Mila::Dnn::Rope< TDeviceType, TPrecision >, Mila::Dnn::TokenEmbedding< TDeviceType, TIndex, TPrecision, TTableQuantization >, Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision >, and Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy >.
|
protected |
The BuildContext stored at build time.
Available to derived classes throughout the component lifetime – in onBuilding(), onEvaluationChanging(), forward(), backward(), and any other method that needs build-time configuration.
Key uses:
Initialized to a placeholder before build() completes. Only valid after isBuilt() returns true.