|
Mila
Deep Neural Network Library
|
Device-templated fully connected (linear) component. More...
Public Types | |
| using | ComponentBase = Component<TDeviceType, TComputePrecision> |
| using | MR = typename DeviceTypeTraits<TDeviceType>::memory_resource |
| using | OpType = typename OperationTraits<OperationType::LinearOp, TDeviceType, TComputePrecision, TWeightQuant>::type |
| using | TensorType = Tensor<TComputePrecision, MR> |
| using | WeightScaleTensorType = Tensor<TWeightQuant::kScaleDtype, MR> |
| using | WeightTensorType = Tensor<kWeightDtype, MR> |
Public Member Functions | |
| Linear (const std::string &name, const LinearConfig &config, std::optional< DeviceId > device_id=std::nullopt) | |
| Construct a Linear component. | |
| TensorType & | backward (const TensorType &input, const TensorType &output_grad) |
| Perform backward pass. | |
| TensorType & | forward (const TensorType &input) |
| Perform forward pass: output = input * weight^T + bias. | |
| const LinearConfig & | getConfig () const noexcept |
| DeviceId | getDeviceId () const override |
| Get the compute device id associated with this component. | |
| std::vector< ITensor * > | getGradients () const override |
| Return non-owning pointers to parameter gradient tensors. | |
| MemoryStats | getMemoryStats () const override |
| Return the current memory allocation breakdown for this component. | |
| std::vector< std::string > | getParameterNames () const override |
| Canonical parameter names, in the order save_() and loadParameter() use. | |
| std::vector< ITensor * > | getParameters () const override |
| Return non-owning pointers to parameter tensors. | |
| MemoryStats | getRequiredMemory (const BuildContext &context) const override |
| What onBuilding() would allocate for this context, without allocating. | |
| const ComponentType | getType () const override |
| Get the component type identifier. | |
| bool | hasBias () const noexcept |
| void | installSharedOutput (std::shared_ptr< TensorType > output) |
| Install a shared output slot (activation pooling). | |
| void | installSharedWeight (std::shared_ptr< WeightTensorType > shared_weight) |
| Replace the owned weight with a shared tensor (e.g. | |
| void | installSharedWeight (std::shared_ptr< WeightTensorType > shared_weight, std::shared_ptr< WeightScaleTensorType > shared_scales) |
| Replace the owned weight and scales with shared tensors – the tied FP8 embedding/lm_head table (D4 Design B). | |
| void | loadParameter (const std::string &name, const ITensorBlob &blob) override |
| Load a named parameter from a serialized blob. | |
| dim_t | parameterCount () const override |
| Return number of trainable parameters. | |
| void | save_ (ModelArchive &archive, SerializationMode mode) const override |
| Save component state to a ModelArchive. | |
| void | saveFlatTensors (Serialization::SafeTensorsWriter &writer, const std::string &prefix, Serialization::TensorSavePass pass) const override |
| Drive this Linear's tensors through one pass of a flat safetensors save. | |
| void | synchronize () override |
| std::string | toString () const override |
| Produce a short, human-readable description of the component. | |
| void | zeroGradients () override |
| Clear all model-owned gradients for this component. | |
| Public Member Functions inherited from Mila::Dnn::Component< TDeviceType, TComputePrecision > | |
| 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). | |
| const std::string | getName () const |
| Get the component's name identifier. | |
| TrainingMode | getTrainingMode () const noexcept |
| The current runtime behavioral mode of this Component. | |
| 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 | requireSerializableParameters () const |
| Verify this component can serialize whatever parameters it owns. | |
| void | setTrainingMode (TrainingMode mode) |
| Set the runtime behavioral mode for this Component. | |
Static Public Attributes | |
| static constexpr bool | kIsQuantized = TWeightQuant::kIsQuantized |
| static constexpr TensorDataType | kWeightDtype |
Protected Member Functions | |
| void | onBuilding (const BuildContext &context) override |
| Hook invoked by build() to allocate component buffers. | |
| void | onExecutionContextSet () override |
| Lifecycle hook: Called immediately after ExecutionContext is set. | |
| void | onTrainingModeChanging (TrainingMode mode) override |
| Hook called before TrainingMode transitions. | |
| Protected Member Functions inherited from Mila::Dnn::Component< TDeviceType, TComputePrecision > | |
| IExecutionContext * | getExecutionContext () const |
| Get the shared execution context. | |
| bool | hasExecutionContext () const noexcept |
| Check if execution context has been set. | |
| 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. | |
| 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>". | |
| 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. | |
Additional Inherited Members | |
| Static Public Member Functions inherited from Mila::Dnn::Component< TDeviceType, TComputePrecision > | |
| 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 inherited from Mila::Dnn::Component< TDeviceType, TComputePrecision > | |
| using | HostStagingMemoryResource |
| Host memory a device-resident parameter stages through. | |
| Protected Attributes inherited from Mila::Dnn::Component< TDeviceType, TComputePrecision > | |
| BuildContext | build_context_ |
| The BuildContext stored at build time. | |
Device-templated fully connected (linear) component.
Delegates compute to a device-specific operation resolved at compile time via OperationTraits<LinearOp, TDeviceType, TComputePrecision, TWeightQuant>. TWeightQuant defaults to NoWeightQuant for unquantized paths.
When TWeightQuant::kIsQuantized is true, the weight tensor is allocated at the reduced-precision storage dtype (kWeightDtype = TWeightQuant::kStorageDtype) rather than TComputePrecision. Per-channel FP32 scale factors (weight_scales_) are allocated alongside the weight tensor and bound to the backend operation via setWeightScales() before the first forward pass. The backend operation receives both the quantized weight tensor and its scales and is responsible for dequantization during the GEMM.
Weight quantization is performed once at model load time (quantize-on-load) during loadParameter(). The source checkpoint blob is always at TComputePrecision.
| TDeviceType | Target device. |
| TComputePrecision | Activation and accumulation precision. |
| TWeightQuant | Weight quantization policy. Must satisfy WeightQuantPolicy. Defaults to NoWeightQuant (identity – no quantization). |
|
inlineexplicit |
Construct a Linear component.
Constructs with a name and configuration. If device_id is provided, the component creates and owns an ExecutionContext (standalone mode) and registers it with the base Component via setExecutionContext(). If device_id is not provided, the component expects a shared ExecutionContext to be provided later via setExecutionContext().
| name | Component name. |
| config | Layer configuration (validated on construction). |
| device_id | Optional device identifier. When present the component creates an owned ExecutionContext for the device. |
| std::invalid_argument | if config is invalid or device type mismatches. |
| std::runtime_error | if ExecutionContext creation fails. |
|
inline |
Perform backward pass.
Pre-zeros the component-owned input gradient buffer, then delegates to the backend operation. The backend accumulates weight and bias gradients into the buffers bound via setGradients() using += semantics; pre-zeroing ensures clean gradient state across calls.
Not supported on quantized paths (kIsQuantized == true) – the backend operation will throw std::logic_error if backward is attempted.
| input | Original forward-pass input tensor. |
| output_grad | Upstream gradient tensor (same shape as the forward output). |
| std::runtime_error | if the component has not been built. |
| std::runtime_error | if called while in inference (eval) mode. |
|
inline |
Perform forward pass: output = input * weight^T + bias.
Delegates to the backend operation using the component-owned output buffer allocated at build time. When the runtime input shape differs from the build-time shape (e.g. a shorter decode sequence vs. the prefill shape), a lightweight view over the output buffer is returned that reflects the true output shape without reallocating device memory.
| input | Input tensor (device-bound, rank >= 2). The last dimension must equal the configured input feature count. |
| std::runtime_error | if the component has not been built. |
| std::invalid_argument | if the input feature dimension does not match the config. |
|
inlineoverridevirtual |
Get the compute device id associated with this component.
Must return the device on which parameters and operations execute.
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
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. |
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
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.
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
Canonical parameter names, in the order save_() and loadParameter() use.
Reimplemented from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
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).
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
What onBuilding() would allocate for this context, without allocating.
Mirrors initializeParameters() and onBuilding() below. The two must agree; the drift gate in the test suite compares this against getMemoryStats() after a real build. See Specifications/MemoryFootprint.md section 7.
Reimplemented from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
Get the component type identifier.
Used for serialization and runtime type identification.
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inline |
Install a shared output slot (activation pooling).
Must be called before build(): onBuilding then skips output self-allocation after validating the slot's storage covers the build shape, and forward() always returns a shape-adjusted view so a wider slot never leaks its geometry to callers. Mirrors installSharedWeight; self-allocation remains the default. The slot is owned and memory-accounted by the installer.
|
inline |
Replace the owned weight with a shared tensor (e.g.
a tied lm_head sharing the token embedding table). See WeightTying.md.
May be called BEFORE build (onBuilding then skips weight self-allocation and wires the installed weight) or AFTER build (rebinds the live operation). The former avoids allocating a weight that tying would immediately free. Quantized instantiations must use the (weight, scales) overload – a quantized weight is meaningless without its dequantization scales.
| shared_weight | Shared device tensor; must match the configured shape. |
|
inline |
Replace the owned weight and scales with shared tensors – the tied FP8 embedding/lm_head table (D4 Design B).
Only per-channel policies are installable: the per-output-channel scale axis IS the vocabulary row the embedding gathers, so one scale tensor serves both consumers. Per-group scales sit on the input axis and do not transfer to a row gather – those instantiations throw.
| shared_weight | Shared quantized device tensor [out_features, in_features]. |
| shared_scales | Shared FP32 scale tensor [out_features]. |
|
inlineoverridevirtual |
Load a named parameter from a serialized blob.
Weight loading dispatches at compile time on kIsQuantized:
Bias is always stored and loaded at TComputePrecision regardless of TWeightQuant.
| name | Parameter name: "weight" or "bias". |
| blob | Serialized tensor blob from PretrainedModelReader. |
| std::invalid_argument | if the blob dtype does not match the expected source precision, if the blob shape does not match the config, or if name is neither "weight" nor "bias". |
Reimplemented from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverrideprotectedvirtual |
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 from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverrideprotectedvirtual |
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 from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverrideprotectedvirtual |
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 from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
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.
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
Save component state to a ModelArchive.
Writes a "meta.json" blob with component type and name, a "config.json" blob with input/output feature dimensions and bias flag, and raw tensor blobs for each name in getParameterNames() under "tensors/".
On CUDA devices each tensor is staged through a host buffer of the same dtype, so the blob carries the parameter's own storage bytes. Refuses outright on the quantized path – see the body.
| archive | ModelArchive to write to (scoped by caller). |
| mode | Serialization mode (currently unused; reserved for future use). |
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
Drive this Linear's tensors through one pass of a flat safetensors save.
This is the path the archive cannot serve. A quantized weight is packed storage plus a scale companion, and save_() refuses it outright because ModelArchive has no representation for that pairing; the flat artifact expresses it as two sibling tensors, which is what the ecosystem does and what the reader already handles.
Declare and write share one ordered body because the writer requires bodies in declaration order. Two separate walks could drift with no diagnostic until the file failed to read back.
The scales are emitted as "<prefix>.weight_scale" – an underscore, not a dot. parseParameterPath() splits a flat name on its LAST dot, so a dotted "weight.scales" would resolve to a component named "<prefix>.weight", which does not exist, and the artifact could be written but never read back. The underscore also matches the compressed-tensors spelling, so the name is conventional as well as loadable.
They are deliberately absent from getParameterNames(): that vector is the join between the archive's save_ and load_, and widening it would break the blob-count invariant those rest on.
| writer | Writer being driven. |
| prefix | Fully qualified component path, e.g. "tf_layer_0.qkv_proj". |
| pass | Declare reserves byte ranges; Write streams bytes. |
Reimplemented from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
@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.
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
Produce a short, human-readable description of the component.
Implementations should keep output concise and avoid throwing.
Implements Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
inlineoverridevirtual |
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 from Mila::Dnn::Component< TDeviceType, TComputePrecision >.
|
staticconstexpr |