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

Root composite network container. More...

Inheritance diagram for Mila::Dnn::Network< TDeviceType, TPrecision >:
Mila::Dnn::CompositeComponent< TDeviceType, TPrecision > Mila::Dnn::Component< TDeviceType, TPrecision > Mila::Dnn::LanguageNetwork< TDeviceType, TPrecision > Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy > Mila::Dnn::GptTransformer< TDeviceType, TPrecision > Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >

Public Types

using ComponentPtr = typename CompositeBase::ComponentPtr
using CompositeBase = CompositeComponent<TDeviceType, TPrecision>
using MR = typename DeviceTypeTraits<TDeviceType>::memory_resource
Public Types inherited from Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >
using ComponentBase = Component<TDeviceType, TPrecision>
using ComponentPtr = std::shared_ptr<Component<TDeviceType, TPrecision>>

Public Member Functions

 Network (const std::string &name)
 Construct network (context managed by derived class).
template<typename TOptimizer, typename TConfig>
std::shared_ptr< TOptimizer > createOptimizer (const TConfig &config)
 Create and configure an optimizer for this network's parameters.
DeviceId getDeviceId () const noexcept override
 Get the compute device for this composite.
IExecutionContextgetExecutionContext () const
 Public access to the network's shared execution context.
const ComponentType getType () const override
 Get the component type identifier.
void load (ModelArchive &archive, SerializationMode mode)
 Restore this network's parameters from an archive.
void save (ModelArchive &archive, SerializationMode mode) const
 Save network to archive.
void synchronize () override
 Synchronize all child components.
std::string toString () const override
 Generate a human-readable description.
Public Member Functions inherited from Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >
 CompositeComponent (CompositeComponent &&) noexcept=default
 CompositeComponent (const CompositeComponent &)=delete
 CompositeComponent (const std::string &name)
 Construct composite component with name.
CompositeComponentaddComponent (ComponentPtr component)
 Add a pre-constructed child component (chainable).
size_t childCount () const noexcept
 Get the number of direct children.
void clearComponents ()
 Clear all child components.
ComponentPtr findComponent (const std::string &path) const
 Resolve a dot-separated component path within this composite.
ComponentPtr getComponent (const std::string &name) const
 Retrieve a direct child component by name.
const std::vector< ComponentPtr > & getComponents () const
 Get all child components in insertion order.
std::vector< ITensor * > getGradients () const override
 Get all parameter gradients from all children.
std::vector< ITensor * > getParameters () const override
 Get all parameters from all children.
bool hasChildren () const noexcept
 Check if this composite has any children.
bool hasComponent (const std::string &name) const
 Check if a named child component exists.
CompositeComponentoperator= (CompositeComponent &&) noexcept=default
CompositeComponentoperator= (const CompositeComponent &)=delete
dim_t parameterCount () const override
 Count parameters across all children.
bool removeComponent (const std::string &name)
void saveFlatTensors (Serialization::SafeTensorsWriter &writer, const std::string &prefix, Serialization::TensorSavePass pass) const override
 Recurse into children, extending the flat dotted prefix.
ComponentPtr tryFindComponent (const std::string &path) const
 Try to resolve a dot-separated component path within this composite.
Public Member Functions inherited from Mila::Dnn::Component< TDeviceType, TPrecision >
 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 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 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 bool isBuilt () const final
 Returns true if build() has completed successfully.
virtual void loadParameter (const std::string &, const Serialization::ITensorBlob &)
 Load a parameter from serialized tensor data.
void setTrainingMode (TrainingMode mode)
 Set the runtime behavioral mode for this Component.
virtual void zeroGradients ()
 Clear all model-owned gradients for this component.

Protected Member Functions

virtual void load_ (ModelArchive &archive, SerializationMode mode) override
 Hook for concrete classes to validate type-specific state on load.
virtual void save_ (ModelArchive &archive, SerializationMode mode) const override=0
 Hook for concrete classes to save type-specific state.
void verifyArchitectureCompatibility (const PretrainedMetadata &metadata)
 Verify that imported model is compatible with network architecture.
Protected Member Functions inherited from Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >
template<typename TComponent>
std::shared_ptr< TComponent > getComponentAs (const std::string &name) const
 Retrieve a typed child component by name.
void onExecutionContextSet () override
 Hook invoked after ExecutionContext is set.
void onTrainingModeChanging (TrainingMode training_mode) override
 Hook invoked when training mode is about to change.
virtual void optimize ()
 Virtual hook for graph optimization after construction.
void requireSerializableParameters () const override
 No-op override: a composite names no parameters of its own.
Protected Member Functions inherited from Mila::Dnn::Component< TDeviceType, TPrecision >
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.
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.

Additional Inherited Members

Static Public Member Functions inherited from Mila::Dnn::Component< TDeviceType, TPrecision >
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, TPrecision >
using HostStagingMemoryResource
 Host memory a device-resident parameter stages through.
Protected Attributes inherited from Mila::Dnn::Component< TDeviceType, TPrecision >
BuildContext build_context_ { shape_t{ 1 }, RuntimeMode::Training }
 The BuildContext stored at build time.

Detailed Description

template<DeviceType TDeviceType, TensorDataType TPrecision>
class Mila::Dnn::Network< TDeviceType, TPrecision >

Root composite network container.

Network is a specialized CompositeComponent that represents a complete neural network model and serves as the top-level entry point. It provides high-level serialization semantics while delegating lifecycle management to concrete subclasses.

Ownership Model:

Construction Pattern (Concrete Subclass):

class MnistClassifier : public Network<DeviceType::Cpu, TensorDataType::FP32>
{
public:
explicit MnistClassifier(const std::string& name, int64_t batch_size, DeviceId device_id)
: Network(name),
owned_context_(createExecutionContext(device_id)),
batch_size_(batch_size)
{
// 1. Create component graph (context-independent)
createGraph();
// 2. Propagate context to self and children
this->setExecutionContext(owned_context_.get());
}
private:
std::unique_ptr<IExecutionContext> owned_context_; // Concrete class owns context
};
void setExecutionContext(IExecutionContext *context)
Set the execution context for this component.
Definition Component.ixx:782
Network(const std::string &name)
Construct network (context managed by derived class).
Definition Network.ixx:137
Lightweight identifier for a compute device.
Definition DeviceId.ixx:38

Serialization Contract:

  • Base class (Network): Saves component graph topology and generic metadata
  • Concrete class: MUST override save_() to write type identifier and configuration
  • Concrete class: MUST provide static Load() factory method for deserialization

Deserialization Pattern (Concrete Subclass):

// REQUIRED: Static factory method for type-safe deserialization
static std::unique_ptr<MnistClassifier> Load(ModelArchive& archive, DeviceId device_id)
{
// 1. Read concrete-specific metadata
json meta = archive.readJson("network/classifier_meta.json");
std::string name = meta.at("name");
int64_t batch_size = meta.at("batch_size");
// 2. Construct via normal constructor path
auto classifier = std::make_unique<MnistClassifier>(name, batch_size, device_id);
// 3. Build with saved input shape
shape_t input_shape = meta.at("input_shape");
classifier->build(input_shape);
// 4. Load component weights
// (Base class handles graph traversal; weights loaded into already-built components)
return classifier;
}
TensorShape shape_t
Row-major shape descriptor for tensor dimensional sizes.
Definition Tensor.Types.ixx:173
ModelArchive provides high-level helpers for component serialization.
Definition ModelArchive.ixx:47

Design Rationale:

  • Concrete classes control infrastructure (context) lifecycle
  • Network base class focuses on container semantics and serialization
  • Clear initialization order: create context ? pass to base ? build graph
  • Enables future flexibility (custom contexts, multi-device, etc.)
  • Type-safe deserialization via concrete class Load() methods

Constructor & Destructor Documentation

◆ Network()

template<DeviceType TDeviceType, TensorDataType TPrecision>
Mila::Dnn::Network< TDeviceType, TPrecision >::Network ( const std::string & name)
inlineexplicit

Construct network (context managed by derived class).

Base constructor for concrete network classes. Derived classes are responsible for:

  1. Creating and owning ExecutionContext
  2. Building the component graph via createGraph()
  3. Calling setExecutionContext() to propagate context to children
Parameters
nameNetwork name for identification and serialization
Exceptions
std::invalid_argumentif name is not a valid identifier

Member Function Documentation

◆ createOptimizer()

template<DeviceType TDeviceType, TensorDataType TPrecision>
template<typename TOptimizer, typename TConfig>
std::shared_ptr< TOptimizer > Mila::Dnn::Network< TDeviceType, TPrecision >::createOptimizer ( const TConfig & config)
inline

Create and configure an optimizer for this network's parameters.

Factory method that creates an optimizer, enables training mode on the network, and registers all network parameters and gradients in a single atomic operation.

Lifecycle:

  1. Enables training mode (allocates gradients for all components)
  2. Creates optimizer using network's ExecutionContext
  3. Registers all network parameters and gradients
  4. Returns ready-to-use optimizer

Usage Pattern:

// Build network
mnist_net->build(input_shape);
// Create optimizer in one step
auto optimizer = mnist_net->createOptimizer<AdamWOptimizer<DeviceType::Cuda, TensorDataType::FP32>>(
AdamWConfig()
.withLearningRate(0.001f)
.withWeightDecay(0.01f)
);
// Optimizer is ready to use
optimizer->step();
Template Parameters
TOptimizerOptimizer type (e.g., AdamWOptimizer, SGD)
TConfigOptimizer configuration type
Parameters
configOptimizer configuration
Returns
Shared pointer to configured and ready-to-use optimizer
Exceptions
std::runtime_errorif network is not built
std::runtime_errorif parameter/gradient count mismatch
Note
This method automatically calls setTraining(true), so explicit training mode activation is not required.

◆ getDeviceId()

template<DeviceType TDeviceType, TensorDataType TPrecision>
DeviceId Mila::Dnn::Network< TDeviceType, TPrecision >::getDeviceId ( ) const
inlineoverridevirtualnoexcept

Get the compute device for this composite.

Returns the device from the shared execution context.

Returns
DeviceId for this composite and its children

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

◆ getExecutionContext()

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

Public access to the network's shared execution context.

Exposes the Component-level context so model-level orchestrator tools (e.g. TokenSampler) can be constructed on the network's context. Wraps the protected Component accessor.

◆ getType()

template<DeviceType TDeviceType, TensorDataType TPrecision>
const ComponentType Mila::Dnn::Network< TDeviceType, TPrecision >::getType ( ) const
inlineoverridevirtual

Get the component type identifier.

Used for serialization and runtime type identification.

Returns
Component type enum value.

Implements Mila::Dnn::Component< TDeviceType, TPrecision >.

◆ load()

template<DeviceType TDeviceType, TensorDataType TPrecision>
void Mila::Dnn::Network< TDeviceType, TPrecision >::load ( ModelArchive & archive,
SerializationMode mode )
inline

Restore this network's parameters from an archive.

The inverse of save(). Restores into the LIVE graph: the network must already be constructed with the same topology and built, which is why this is a member and not a factory. Concrete models reconstruct themselves through their own fromCheckpoint() – read the config, construct, build, then call this.

The child manifest is not replayed. saveComponentGraph records names into per-child descriptor files rather than a list, so the live children are the authoritative enumeration; network/architecture.json is read only to check the count, which catches the common "loaded into a different model" mistake before any tensor is touched.

Parameters
archiveArchive to read from.
modeSerialization mode.
Exceptions
std::runtime_errorif the archive's component count disagrees with this network, or if any component's blobs are missing.

◆ load_()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Network< TDeviceType, TPrecision >::load_ ( ModelArchive & archive,
SerializationMode mode )
inlineoverrideprotectedvirtual

Hook for concrete classes to validate type-specific state on load.

Deliberately NOT pure, where save_() is. A concrete network must write its config, because nothing else can; it rarely needs to read it back here, because the config is consumed by the static factory before construction – by the time this runs the network already exists with that geometry. Override only to assert the archive matches what was built.

The empty body also keeps this from being a breaking change for every existing Network subclass.

Parameters
archiveArchive to read from.
modeSerialization mode.

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

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

◆ save()

template<DeviceType TDeviceType, TensorDataType TPrecision>
void Mila::Dnn::Network< TDeviceType, TPrecision >::save ( ModelArchive & archive,
SerializationMode mode ) const
inline

Save network to archive.

Saves component graph structure and delegates to concrete class via save_() hook for type-specific configuration.

Archive structure produced:

  • network/meta.json: Base metadata (name, version, num_components, timestamp)
  • network/architecture.json: Component topology (names, paths, ordering)
  • components/<name>/...: Child component state (recursive)
  • Concrete class writes additional files via save_() override
Parameters
archiveArchive to write to
modeSerialization mode (Checkpoint, WeightsOnly, Architecture)
Exceptions
std::runtime_errorif save_() is not overridden by concrete class

◆ save_()

template<DeviceType TDeviceType, TensorDataType TPrecision>
virtual void Mila::Dnn::Network< TDeviceType, TPrecision >::save_ ( ModelArchive & archive,
SerializationMode mode ) const
overrideprotectedpure virtual

Hook for concrete classes to save type-specific state.

REQUIRED override for concrete networks. Must write:

  • Type identifier (e.g., "type": "MnistClassifier")
  • Configuration parameters (batch_size, architecture constants)
  • Shape metadata (for validation during Load())

This metadata enables the concrete class's Load() method to reconstruct the network.

Example implementation:

void save_(ModelArchive& archive, SerializationMode mode) const override
{
json meta;
meta["type"] = "MnistClassifier"; // Type identifier for runtime dispatch
meta["batch_size"] = batch_size_;
meta["input_shape"] = leading_shape_;
// ... other configuration
archive.writeJson("network/classifier_meta.json", meta);
}
virtual void save_(ModelArchive &archive, SerializationMode mode) const override=0
Hook for concrete classes to save type-specific state.
Parameters
archiveArchive to write to
modeSerialization mode (passed from save())

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

Implemented in Mila::Dnn::GemmaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >, Mila::Dnn::GptTransformer< TDeviceType, TPrecision >, and Mila::Dnn::LlamaTransformer< TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy >.

◆ synchronize()

template<DeviceType TDeviceType, TensorDataType TPrecision>
void Mila::Dnn::Network< TDeviceType, TPrecision >::synchronize ( )
inlineoverridevirtual

Synchronize all child components.

Waits for outstanding device operations on all children.

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

◆ toString()

template<DeviceType TDeviceType, TensorDataType TPrecision>
std::string Mila::Dnn::Network< TDeviceType, TPrecision >::toString ( ) const
inlineoverridevirtual

Generate a human-readable description.

Returns
String representation showing network name and children

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


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