|
Mila
Deep Neural Network Library
|
LLaMA-style transformer (decoder-only) for autoregressive token prediction. More...
Public Types | |
| using | ComponentPtr = typename NetworkBase::ComponentPtr |
| using | LinearType = Linear<TDeviceType, TPrecision, TWeightQuantization> |
| using | LmHeadLinearType = Linear<TDeviceType, TPrecision> |
| using | MR = typename DeviceTypeTraits<TDeviceType>::memory_resource |
| using | NetworkBase = LanguageNetwork<TDeviceType, TPrecision> |
| using | RmsNormType = RmsNorm<TDeviceType, TPrecision> |
| using | TensorType = Tensor<TPrecision, MR> |
| using | TokenEmbeddingType = TokenEmbedding<TDeviceType, dtype_t::INT32, TPrecision> |
| using | TokenIndexType = Tensor<dtype_t::INT32, MR> |
| using | TransformerBlockType = LlamaBlock<TDeviceType, TPrecision, TWeightQuantization, TKvCachePolicy> |
| Public Types inherited from Mila::Dnn::LanguageNetwork< TDeviceType, TPrecision > | |
| using | MR = typename DeviceTypeTraits<TDeviceType>::memory_resource |
| using | NetworkBase = Network<TDeviceType, TPrecision> |
| using | StageProbe = std::function<void( std::string_view stage, const TensorType& value )> |
| Observer called with each intermediate activation during prefill. | |
| using | TensorType = Tensor<TPrecision, MR> |
| using | TokenIndexType = Tensor<TensorDataType::INT32, MR> |
| Public Types inherited from Mila::Dnn::Network< TDeviceType, TPrecision > | |
| 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 | |
| LlamaTransformer (const std::string &name, const LlamaConfig &config, DeviceId device_id) | |
| TokenIndexType & | backward (const TokenIndexType &input, const TensorType &output_grad) override |
| TensorType & | decode (const TokenIndexType &input, dim_t position) override |
| TensorType & | forward (const TokenIndexType &input) override |
| IExecutionContext * | getExecutionContext () const |
| MemoryStats | getMemoryStats () const override |
| Return the current memory allocation breakdown for this component. | |
| ModelType | getModelType () const |
| MemoryStats | getRequiredMemory (const BuildContext &context) const override |
| What build( context ) would allocate for the whole model, without allocating. | |
| void | loadParameters (PretrainedModelReader &reader) |
| TensorType & | prefill (const TokenIndexType &input) override |
| std::string | toString () const override |
| Generate a human-readable description. | |
| void | zeroGradients () override |
| Clear all model-owned gradients for this component. | |
| Public Member Functions inherited from Mila::Dnn::LanguageNetwork< TDeviceType, TPrecision > | |
| LanguageNetwork (const std::string &name) | |
| virtual TokenIndexType & | backward (const TokenIndexType &input, const TensorType &output_grad)=0 |
| Full backward pass (training). | |
| virtual TensorType & | decode (const TokenIndexType &input, dim_t position)=0 |
| Inference decode – single-token autoregressive step. | |
| virtual TensorType & | forward (const TokenIndexType &input)=0 |
| Full-sequence forward pass. | |
| virtual TensorType & | prefill (const TokenIndexType &input)=0 |
| Inference prefill – process full prompt and populate the KV cache. | |
| virtual TensorType & | prefillFrom (const TokenIndexType &input, dim_t start_offset) |
| Chunked prefill starting at an absolute position (prompt-prefix reuse). | |
| virtual bool | rewindKvCache (dim_t position) |
| Rewind the KV caches to position for prompt-prefix reuse (PromptCaching.md). | |
| virtual void | setStageProbe (StageProbe probe) |
| Install a stage probe, or clear it by passing an empty function. | |
| Public Member Functions inherited from Mila::Dnn::Network< TDeviceType, TPrecision > | |
| 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. | |
| IExecutionContext * | getExecutionContext () 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. | |
| CompositeComponent & | addComponent (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. | |
| CompositeComponent & | operator= (CompositeComponent &&) noexcept=default |
| CompositeComponent & | operator= (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). | |
| 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. | |
| 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. | |
Protected Member Functions | |
| void | onBuilding (const BuildContext &context) override |
| Hook invoked by build() to allocate component buffers. | |
| void | onTrainingModeChanging (TrainingMode training_mode) override |
| Hook invoked when training mode is about to change. | |
| std::size_t | ropeCacheBytes (dim_t head_dim) const noexcept |
| Bytes one RoPE cos/sin cache occupies for a given head width. | |
| void | save_ (ModelArchive &archive, SerializationMode) const override |
| Hook for concrete classes to save type-specific state. | |
| Protected Member Functions inherited from Mila::Dnn::Network< TDeviceType, TPrecision > | |
| virtual void | load_ (ModelArchive &archive, SerializationMode mode) override |
| Hook for concrete classes to validate type-specific state on load. | |
| 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. | |
| 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 > | |
| IExecutionContext * | getExecutionContext () const |
| Get the shared execution context. | |
| bool | hasExecutionContext () const noexcept |
| Check if execution context has been set. | |
| template<TensorDataType TParameterPrecision, typename TMemoryResource> | |
| void | loadParameterFromBlob (const std::string ¶m_name, const Serialization::ITensorBlob &blob, Tensor< TParameterPrecision, TMemoryResource > &target, const shape_t &expected_shape) |
| Load a tensor blob into a parameter tensor with validation. | |
| template<TensorDataType TParameterPrecision, typename TMemoryResource> | |
| void | saveParameterToArchive (ModelArchive &archive, const std::string ¶meter_name, const Tensor< TParameterPrecision, TMemoryResource > ¶meter) const |
| Write one parameter tensor into the archive under "tensors/<name>". | |
| template<TensorDataType TParameterPrecision, typename TMemoryResource> | |
| void | saveParameterToWriter (Serialization::SafeTensorsWriter &writer, const std::string &flat_name, const Tensor< TParameterPrecision, TMemoryResource > ¶meter, Serialization::TensorSavePass pass) const |
| Drive one parameter through one pass of a flat safetensors save. | |
| void | setExecutionContext (IExecutionContext *context) |
| Set the execution context for this component. | |
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. | |
LLaMA-style transformer (decoder-only) for autoregressive token prediction.
Graph: TokenEmbedding -> RoPE -> LlamaBlock x N -> RmsNorm -> Linear (lm_head). RoPE is applied to the full embedding stream after the token lookup; each LlamaBlock receives rotary-encoded embeddings as input.
Template parameters:
|
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, TPrecision >.
|
inlineoverridevirtual |
What build( context ) would allocate for the whole model, without allocating.
Mirrors onBuilding(): resolve the prefill chunk first, recurse with the same per-child contexts, then add the shared GQA workspace this transformer owns.
Two corrections Gemma needs are absent here, and their absence is the finding rather than an omission. Llama does not pool per-block activations, so there is no installed-output adjustment; and it does not tie the embedding to the head, so the two largest tensors are counted separately and in full. See BACKLOG, Models.
Reimplemented from Mila::Dnn::Component< TDeviceType, TPrecision >.
|
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, TPrecision >.
|
inlineoverrideprotectedvirtual |
Hook invoked when training mode is about to change.
Propagates the new mode to all child components. The hook runs with the Component's training mutex held; it MUST NOT call setTrainingMode().
| training_mode | New training mode (Normal or Eval) |
Reimplemented from Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >.
|
inlineprotectednoexcept |
Bytes one RoPE cos/sin cache occupies for a given head width.
MUST match CudaRopeOp::getRequiredStateMemorySize – FP32 regardless of the model precision, half the head dimension, two caches. Duplicated here because the deduplication is the transformer's to apply and it needs the per-key size; the model-level comparison against getMemoryStats is what holds the two together.
|
inlineoverrideprotectedvirtual |
Hook for concrete classes to save type-specific state.
REQUIRED override for concrete networks. Must write:
This metadata enables the concrete class's Load() method to reconstruct the network.
Example implementation:
| archive | Archive to write to |
| mode | Serialization mode (passed from save()) |
Implements Mila::Dnn::Network< TDeviceType, TPrecision >.
|
inlineoverridevirtual |
Generate a human-readable description.
Reimplemented from Mila::Dnn::CompositeComponent< TDeviceType, TPrecision >.
|
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, TPrecision >.