Mila
Deep Neural Network Library
Loading...
Searching...
No Matches
Mila::Dnn::Component< TDeviceType, TPrecision > Class Template Referenceabstractexport

Abstract base class for neural network components. More...

Inheritance diagram for Mila::Dnn::Component< TDeviceType, TPrecision >:
Mila::Dnn::Lpe< TDeviceType, dtype_t::INT32, TPrecision > Mila::Dnn::Swiglu< TDeviceType, TPrecision, TGate > Mila::Dnn::Swiglu< TDeviceType, TPrecision, ActivationType::Gelu > Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision, TableQuantizationPolicy > Mila::Dnn::TokenEmbedding< TDeviceType, dtype_t::INT32, TPrecision > Mila::Dnn::Activation< TDeviceType, TPrecision, TFn > Mila::Dnn::CompositeComponent< TDeviceType, TPrecision > Mila::Dnn::Gelu< TDeviceType, TPrecision > Mila::Dnn::LayerNorm< TDeviceType, TPrecision > Mila::Dnn::Loss< 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 >

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

IExecutionContextgetExecutionContext () 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 &param_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 &parameter_name, const Tensor< TParameterPrecision, TMemoryResource > &parameter) 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 > &parameter, 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.

Detailed Description

template<DeviceType TDeviceType, TensorDataType TPrecision>
requires PrecisionSupportedOnDevice<TPrecision, TDeviceType>
class Mila::Dnn::Component< TDeviceType, TPrecision >

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.

Ownership model

  • Components NEVER own ExecutionContext.
  • Parent (Network/CompositeComponent) owns and provides shared context.
  • Tests explicitly create context and pass raw pointer to components.

Component build lifecycle

Components progress through a well-defined lifecycle. Each stage has a single responsibility and a designated hook for subclass extension.

Stage 1 – Construction

The component is constructed with its name and component config. No device resources are allocated. No ExecutionContext is required.

auto linear = std::make_unique<Linear>( "fc", config, context );

Stage 2 – Build [ onBuilding() ]

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():

  • output_ forward output buffer sized by allocationSeqLen()
  • decode_output_ decode output buffer (decode-capable components)
  • kv_cache_ KV cache (decode-capable components)
  • operation_ buffers via operation_->build()

NOT allocated in onBuilding():

  • gradient buffers (deferred to first setEvaluation( false ))
  • backward state (deferred to first setEvaluation( false ))

The same BuildContext is cascaded unchanged through CompositeComponent and Network to all child components.

BuildContext config( shape_t{ batch_size, seq_length } );
config.withExecutionMode( ExecutionMode::Training );
model->build( config );
TensorShape shape_t
Row-major shape descriptor for tensor dimensional sizes.
Definition Tensor.Types.ixx:173
Build-time context for Component::build().
Definition Component.BuildContext.ixx:62

Stage 3 – Evaluation mode [ onEvaluationChanging() ]

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 ):

  • input_grad_ input gradient buffer
  • weight_grad_ weight gradient buffer
  • bias_grad_ bias gradient buffer
model->setEvaluation( true ); // suspend backward -- eval checkpoint
generateSample( model );
model->setEvaluation( false ); // resume training

Stage 4 – Forward / Decode / Backward

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_.

Lifecycle invariants

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

Base class provides

  • build() / lifecycle management with protected onBuilding() hook
  • execution mode query: isInferenceMode(), isTrainingMode()
  • evaluation mode transitions with serialized onEvaluationChanging() hook
  • parameter and gradient access
  • synchronization and serialization
  • short human-readable diagnostics via toString()
Template Parameters
TDeviceTypeCompile-time device identifier for this component.
TPrecisionTensor data precision for this component.
Trusted Collaborators
This class grants private access to:
  • CompositeComponent: Parent components that manage child execution contexts and aggregate parameters from child components.
  • Network: Top-level graph coordinator that collects parameters and gradients for optimization and serialization.

Member Typedef Documentation

◆ HostStagingMemoryResource

template<DeviceType TDeviceType, TensorDataType TPrecision>
using Mila::Dnn::Component< TDeviceType, TPrecision >::HostStagingMemoryResource
protected
Initial value:

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.

Constructor & Destructor Documentation

◆ Component()

template<DeviceType TDeviceType, TensorDataType TPrecision>
Mila::Dnn::Component< TDeviceType, TPrecision >::Component ( const std::string & name)
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.

Parameters
nameComponent name identifier (mandatory).
Exceptions
std::invalid_argumentif name is not a valid identifier.

Member Function Documentation

◆ build()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::build ( const BuildContext & context)
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.

Parameters
contextBuild-time configuration carrying the leading shape { B, T, ... }, ExecutionMode, and optional micro-batching settings.
Exceptions
std::runtime_errorif the component is already built.
std::runtime_errorif no ExecutionContext has been set.
std::invalid_argumentif config.validate() fails.
Anyexception from onBuilding().

◆ getDeviceId()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual DeviceId Mila::Dnn::Component< TDeviceType, TPrecision >::getDeviceId ( ) const
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 >.

◆ getExecutionContext()

template<DeviceType TDeviceType, TensorDataType TPrecision>
IExecutionContext * Mila::Dnn::Component< TDeviceType, TPrecision >::getExecutionContext ( ) const
inlineprotected

Get the shared execution context.

Provides access to the execution context for derived classes to:

  • Query device information
  • Create tensors on the correct device
  • Pass to backend operations
  • Synchronize device work
Returns
Non-owning pointer to execution context (guaranteed non-null).
Exceptions
std::runtime_errorif context has not been set.

◆ getGradients()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual std::vector< ITensor * > Mila::Dnn::Component< TDeviceType, TPrecision >::getGradients ( ) const
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.

Returns
Vector of gradient pointers; empty when built for inference.
Exceptions
std::runtime_errorif 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 >.

◆ getMemoryStats()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual MemoryStats Mila::Dnn::Component< TDeviceType, TPrecision >::getMemoryStats ( ) const
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.

Returns
MemoryStats reflecting current allocations.

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 >.

◆ getName()

template<DeviceType TDeviceType, TensorDataType TPrecision>
const std::string Mila::Dnn::Component< TDeviceType, TPrecision >::getName ( ) const
inline

Get the component's name identifier.

The name is used for logging, diagnostics, and serialization.

Returns
Component full hierarchical path.

◆ getParameterNames()

◆ getParameters()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual std::vector< ITensor * > Mila::Dnn::Component< TDeviceType, TPrecision >::getParameters ( ) const
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 >.

◆ getRequiredMemory()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual MemoryStats Mila::Dnn::Component< TDeviceType, TPrecision >::getRequiredMemory ( const BuildContext & context) const
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.

Parameters
contextThe context that would be passed to build().
Returns
MemoryStats the component would allocate.
Exceptions
std::logic_errorif 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 >.

◆ getTrainingMode()

template<DeviceType TDeviceType, TensorDataType TPrecision>
TrainingMode Mila::Dnn::Component< TDeviceType, TPrecision >::getTrainingMode ( ) const
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.

Returns
Current TrainingMode.

◆ getType()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual const ComponentType Mila::Dnn::Component< TDeviceType, TPrecision >::getType ( ) const
pure virtual

Get the component type identifier.

Used for serialization and runtime type identification.

Returns
Component type enum value.

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 >.

◆ hasExecutionContext()

template<DeviceType TDeviceType, TensorDataType TPrecision>
bool Mila::Dnn::Component< TDeviceType, TPrecision >::hasExecutionContext ( ) const
inlineprotectednoexcept

Check if execution context has been set.

Returns
true if context is set, false otherwise.

◆ load_()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::load_ ( ModelArchive & archive,
SerializationMode mode )
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.

Parameters
archiveArchive to read from, already scoped to this component.
modeSerialization mode (accepted for symmetry with save_).
Exceptions
std::runtime_errorif 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 >.

◆ loadParameter()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::loadParameter ( const std::string & ,
const Serialization::ITensorBlob &  )
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:

  • Precision conversion (blob dtype -> parameter dtype)
  • Device upload (CPU bytes -> target device)

Takes the parameter name used to locate the target tensor, and a blob holding serialized tensor metadata and raw bytes.

Exceptions
std::runtime_errorif component has no parameters to load.
std::runtime_errorif 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 >.

◆ loadParameterFromBlob()

template<DeviceType TDeviceType, TensorDataType TPrecision>
template<TensorDataType TParameterPrecision, typename TMemoryResource>
void Mila::Dnn::Component< TDeviceType, TPrecision >::loadParameterFromBlob ( const std::string & param_name,
const Serialization::ITensorBlob & blob,
Tensor< TParameterPrecision, TMemoryResource > & target,
const shape_t & expected_shape )
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.

Parameters
param_nameParameter name (used in error messages).
blobSource tensor blob from the model archive.
targetDestination tensor (must be initialized).
expected_shapeExpected tensor shape for validation.
Exceptions
std::invalid_argumentif dtype or shape mismatch.

◆ onBuilding()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::onBuilding ( const BuildContext & )
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.

// Example -- Linear component:
shape_t out_shape =
{
config.batchSize(),
config.allocationSeqLen(), // 1 for Inference, T for Training
config_.getOutputFeatures()
};
output_ = std::make_unique<TensorType>( device, out_shape,
this->getName() + ".output" );
const std::string getName() const
Get the component's name identifier.
Definition Component.ixx:533

The default implementation forwards to the legacy onBuilding( const shape_t& ) overload for backwards compatibility. New components should override this overload directly.

Note
Do not call build() or onBuilding() from within this hook.
Implementations should either succeed fully or leave no partial state, as a failed build() may be retried.

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 >.

◆ onExecutionContextSet()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::onExecutionContextSet ( )
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:

  • Composite components: Create and configure child components.
  • Device resource allocation: Query device capabilities.

Default implementation does nothing.

Exceptions
Anyexception 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 >.

◆ onTrainingModeChanging()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::onTrainingModeChanging ( TrainingMode )
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 >.

◆ parameterCount()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual dim_t Mila::Dnn::Component< TDeviceType, TPrecision >::parameterCount ( ) const
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 >.

◆ requireSerializableParameters()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::requireSerializableParameters ( ) const
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.

Exceptions
std::runtime_errorif the component has parameters but no names for them.

Reimplemented in Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >.

◆ save_()

◆ saveFlatTensors()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::saveFlatTensors ( Serialization::SafeTensorsWriter & writer,
const std::string & prefix,
Serialization::TensorSavePass pass ) const
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_().

Parameters
writerWriter being driven.
prefixFully qualified component path, e.g. "tf_layer_0.qkv_proj".
passDeclare 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 >.

◆ saveParameterToArchive()

template<DeviceType TDeviceType, TensorDataType TPrecision>
template<TensorDataType TParameterPrecision, typename TMemoryResource>
void Mila::Dnn::Component< TDeviceType, TPrecision >::saveParameterToArchive ( ModelArchive & archive,
const std::string & parameter_name,
const Tensor< TParameterPrecision, TMemoryResource > & parameter ) const
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.

Parameters
archiveArchive to write into, already scoped to this component.
parameter_nameCanonical name from getParameterNames().
parameterParameter tensor to serialize.

◆ saveParameterToWriter()

template<DeviceType TDeviceType, TensorDataType TPrecision>
template<TensorDataType TParameterPrecision, typename TMemoryResource>
void Mila::Dnn::Component< TDeviceType, TPrecision >::saveParameterToWriter ( Serialization::SafeTensorsWriter & writer,
const std::string & flat_name,
const Tensor< TParameterPrecision, TMemoryResource > & parameter,
Serialization::TensorSavePass pass ) const
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.

Parameters
writerWriter being driven.
flat_nameFully qualified tensor name.
parameterParameter tensor.
passDeclare reserves the byte range; Write streams the bytes.

◆ setExecutionContext()

template<DeviceType TDeviceType, TensorDataType TPrecision>
void Mila::Dnn::Component< TDeviceType, TPrecision >::setExecutionContext ( IExecutionContext * context)
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:

  • The component itself (standalone mode with owned context)
  • Parent composite when adding child (shared context mode)
  • ComponentFactory during deserialization

After setting the context, the onExecutionContextSet() hook is invoked to allow the component to perform context-dependent initialization.

Parameters
contextNon-owning pointer to execution context (must be non-null).
Exceptions
std::invalid_argumentif context is null.
std::runtime_errorif context has already been set.
std::invalid_argumentif context device type doesn't match TDeviceType.
std::runtime_errorif onExecutionContextSet() throws; context is restored to nullptr on failure.

◆ setTrainingMode()

template<DeviceType TDeviceType, TensorDataType TPrecision>
void Mila::Dnn::Component< TDeviceType, TPrecision >::setTrainingMode ( TrainingMode mode)
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.

State transitions

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.

Parameters
modeTrainingMode::Normal or TrainingMode::Eval.
Exceptions
std::runtime_errorif the component is not built.
std::runtime_errorif built with RuntimeMode::Inference.

◆ synchronize()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::synchronize ( )
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.
        &zwj;/

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();
    }

====================================================================

Synchronization

    /**
       @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 >.

◆ toString()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual std::string Mila::Dnn::Component< TDeviceType, TPrecision >::toString ( ) const
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 >.

◆ zeroGradients()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Component< TDeviceType, TPrecision >::zeroGradients ( )
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 >.

Member Data Documentation

◆ build_context_

template<DeviceType TDeviceType, TensorDataType TPrecision>
BuildContext Mila::Dnn::Component< TDeviceType, TPrecision >::build_context_ { shape_t{ 1 }, RuntimeMode::Training }
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:

  • build_config_.allocationSeqLen() – use when sizing output buffers in onBuilding(). Returns 1 for Inference, leading_shape[1] for Training.
  • build_config_.isInference() / isTrainingMode() – query the policy.
  • build_config_.batchSize() – the batch dimension.

Initialized to a placeholder before build() completes. Only valid after isBuilt() returns true.


The documentation for this class was generated from the following file: