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

LLaMA 3 compatible inference model. More...

Inheritance diagram for Mila::Dnn::LlamaModel< TDeviceType, TPrecision >:
Mila::Dnn::LanguageModel< TDeviceType, TPrecision > Mila::Dnn::Model< TDeviceType, TPrecision >

Public Types

using LlamaKvPolicy = Quant::KvCache::NoKvCompression
 KV policy for the Llama chassis.
using ModelBase = LanguageModel<TDeviceType, TPrecision>
using MR = typename DeviceTypeTraits<TDeviceType>::memory_resource
using StagingMR = CpuMemoryResource
using TensorType = Tensor<TPrecision, MR>
using TokenIndexType = Tensor<dtype_t::INT32, MR>
Public Types inherited from Mila::Dnn::LanguageModel< TDeviceType, TPrecision >
using Base = Model<TDeviceType, TPrecision>
Public Types inherited from Mila::Dnn::Model< TDeviceType, TPrecision >
using NetworkType = Network<TDeviceType, TPrecision>

Public Member Functions

 LlamaModel (const LlamaModel &)=delete
 LlamaModel (LlamaModel &&)=default
const LlamaConfiggetConfig () const noexcept
LlamaModel & operator= (const LlamaModel &)=delete
LlamaModel & operator= (LlamaModel &&)=default
void profilePrefill (const std::vector< int32_t > &token_ids)
std::string toString () const override
 Human-readable summary of this model's configuration.
Public Member Functions inherited from Mila::Dnn::LanguageModel< TDeviceType, TPrecision >
 LanguageModel (const LanguageModel &)=delete
 LanguageModel (LanguageModel &&)=default
GenerateStatus generate (std::span< const int32_t > prompt_tokens, const std::function< void(int32_t)> &on_token, const GenerateParams &params={}, std::stop_token stop={})
 Generate tokens from a prompt, streaming each through on_token.
LanguageModel & operator= (const LanguageModel &)=delete
LanguageModel & operator= (LanguageModel &&)=default
void savePretrained (const std::filesystem::path &path) const
 Write this model's live weights as a safetensors artifact.
void seedSampler (uint64_t seed)
 Seed the sampler's RNG for reproducible generation.
Public Member Functions inherited from Mila::Dnn::Model< TDeviceType, TPrecision >
 Model (const Model &)=delete
 Model (Model &&)=default
DeviceId getDeviceId () const noexcept
 The device this model runs on.
MemoryStats getMemoryStats () const
 Current memory allocation breakdown for this model.
RuntimeMode getRuntimeMode () const noexcept
 The runtime mode this model was constructed for.
std::size_t getScratchHighWaterBytes () const
 Context-owned scratch device memory, in bytes, at its high-water mark.
bool isEval () const noexcept
 True if this model is currently in eval sub-state.
bool isInferenceMode () const noexcept
 True if this model was constructed for inference.
bool isTrainingMode () const noexcept
 True if this model was constructed for training.
Model & operator= (const Model &)=delete
Model & operator= (Model &&)=default
void setEval (bool eval)
 Toggle eval sub-state for this model.
void train ()
 Run the training loop for this model.

Static Public Member Functions

static std::unique_ptr< LlamaModel< TDeviceType, TPrecision > > fromPretrained (const std::filesystem::path &path, const LlamaModelConfig &model_config, DeviceId device_id=DeviceId{ TDeviceType, 0 })
 Load from third-party pretrained weights.
static MemoryStats getRequiredMemory (const std::filesystem::path &path, const LlamaModelConfig &model_config, DeviceId device_id=DeviceId{ TDeviceType, 0 })
 What loading this checkpoint at this context length would cost in VRAM.

Protected Member Functions

dim_t maxSequenceLength () const noexcept override
 Maximum sequence length from LLaMA config.
GenerateStatus onGenerating (std::span< const int32_t > prompt_tokens, const std::function< void(int32_t)> &on_token, const GenerateParams &params, std::stop_token stop) override
 Prefill + KV-cache decode loop with per-token streaming.
void onTraining () override
 Training loop – not yet implemented for LlamaModel.
dim_t vocabSize () const noexcept override
 Vocabulary size from LLaMA config.
Protected Member Functions inherited from Mila::Dnn::LanguageModel< TDeviceType, TPrecision >
 LanguageModel (std::unique_ptr< LanguageNetwork< TDeviceType, TPrecision > > network, RuntimeMode runtime_mode, Serialization::PretrainedMetadata source_metadata={}, WeightQuantization weight_quantization=WeightQuantization::None)
int32_t awaitSampledToken ()
 Block until the last enqueueSampleNext()'s token id is host-visible.
void enqueueSampleNext (const TensorType &logits, TokenTensor &token_out, const SamplingParams &params)
 Enqueue a sampling step without waiting for the host readback.
virtual float finalLogitSoftcap () const noexcept
 Optional final-logit softcap the sampler applies (0 disables).
const LanguageNetwork< TDeviceType, TPrecision > & getLanguageNetwork () const noexcept
LanguageNetwork< TDeviceType, TPrecision > & getLanguageNetwork () noexcept
int32_t sampleNext (const TensorType &logits, TokenTensor &token_out, const SamplingParams &params)
 Sample the next token from a logits row on the device.
Protected Member Functions inherited from Mila::Dnn::Model< TDeviceType, TPrecision >
 Model (std::unique_ptr< NetworkType > network, RuntimeMode runtime_mode)
 Construct with a fully built network and runtime mode.

Additional Inherited Members

Protected Types inherited from Mila::Dnn::LanguageModel< TDeviceType, TPrecision >
using MR = typename DeviceTypeTraits<TDeviceType>::memory_resource
using TensorType = Tensor<TPrecision, MR>
using TokenTensor = Tensor<TensorDataType::INT32, MR>
Protected Attributes inherited from Mila::Dnn::LanguageModel< TDeviceType, TPrecision >
Serialization::PretrainedMetadata source_metadata_
 The loaded artifact's metadata, carried so savePretrained can write it back verbatim.
WeightQuantization weight_quantization_ { WeightQuantization::None }
 What the live weights are, not what the source file was.
Protected Attributes inherited from Mila::Dnn::Model< TDeviceType, TPrecision >
std::unique_ptr< NetworkType > network_
 The owned Network instance.

Detailed Description

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

LLaMA 3 compatible inference model.

Owns a loaded, built LlamaTransformer and exposes generate() for autoregressive text generation. Supports the prefill + KV-cache decode two-phase generation loop.

Construction is only possible via fromPretrained(). The network is always in a built, weights-loaded, inference-mode state when generation is called.

Thread safety: not thread-safe; external synchronization required if shared.

Member Typedef Documentation

◆ LlamaKvPolicy

template<DeviceType TDeviceType, TensorDataType TPrecision>
using Mila::Dnn::LlamaModel< TDeviceType, TPrecision >::LlamaKvPolicy = Quant::KvCache::NoKvCompression

KV policy for the Llama chassis.

Every Llama layer is full-attention, so there is no sliding window to bound a ring against and the cache spans the whole context. Class scope rather than per-function so the load and footprint paths cannot be pointed at different policies – that would make a model report a figure for a cache it does not build.

Member Function Documentation

◆ fromPretrained()

template<DeviceType TDeviceType, TensorDataType TPrecision>
std::unique_ptr< LlamaModel< TDeviceType, TPrecision > > Mila::Dnn::LlamaModel< TDeviceType, TPrecision >::fromPretrained ( const std::filesystem::path & path,
const LlamaModelConfig & model_config,
DeviceId device_id = DeviceId{ TDeviceType, 0 } )
inlinestatic

Load from third-party pretrained weights.

Reads a Mila-compatible pretrained artifact (e.g. converted from a HuggingFace LLaMA checkpoint) via PretrainedModelReader. The network is built at the context length specified in model_config so RoPE embeddings and KV cache buffers cover the full range.

The model_config carries all deployment decisions:

  • context_length – maximum sequence length to build for
  • weight_quantization – compile-time dispatch to quantized or BF16 path
  • kv_cache_compression – compile-time dispatch to KV cache policy
Parameters
pathPath to the pretrained Llama model artifact.
model_configDeployment configuration for this load.
device_idTarget device; must match TDeviceType.
Returns
Inference-ready LlamaModel.
Exceptions
std::invalid_argumenton device type mismatch or zero context length.
std::runtime_erroron load or parameter binding failure.
std::runtime_errorif model_config requests unsupported quantization (e.g. FP4).

◆ getRequiredMemory()

template<DeviceType TDeviceType, TensorDataType TPrecision>
MemoryStats Mila::Dnn::LlamaModel< TDeviceType, TPrecision >::getRequiredMemory ( const std::filesystem::path & path,
const LlamaModelConfig & model_config,
DeviceId device_id = DeviceId{ TDeviceType, 0 } )
inlinestatic

What loading this checkpoint at this context length would cost in VRAM.

Reads the artifact header for geometry and constructs the graph, then reports what build() would allocate – without building it and without reading a weight. Returns measurements only; the fits/does-not verdict is adaptor policy. See Specifications/MemoryFootprint.md.

Exceptions
std::invalid_argumenton device type mismatch or zero context length.
std::runtime_erroron an unreadable artifact or unsupported quantization.

◆ maxSequenceLength()

template<DeviceType TDeviceType, TensorDataType TPrecision>
dim_t Mila::Dnn::LlamaModel< TDeviceType, TPrecision >::maxSequenceLength ( ) const
inlineoverrideprotectedvirtualnoexcept

Maximum sequence length from LLaMA config.

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

◆ onGenerating()

template<DeviceType TDeviceType, TensorDataType TPrecision>
GenerateStatus Mila::Dnn::LlamaModel< TDeviceType, TPrecision >::onGenerating ( std::span< const int32_t > prompt_tokens,
const std::function< void(int32_t)> & on_token,
const GenerateParams & params,
std::stop_token stop )
inlineoverrideprotectedvirtual

Prefill + KV-cache decode loop with per-token streaming.

Phase 1 (prefill): runs the full prompt through prefill() to populate the KV cache and samples the first new token from the last position. Phase 2 (decode): iterates one token at a time until max_new_tokens is reached, EOS is emitted, or stop is requested.

on_token is called for every generated token except EOS.

Parameters
prompt_tokensInput token ids; truncated from the start if they exceed the model's max sequence length.
on_tokenCallback invoked once per generated token (not EOS).
paramsPer-call generation parameters (loop bound + sampling).
stopStop token for cooperative cancellation.
Returns
Why generation stopped.

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

◆ onTraining()

template<DeviceType TDeviceType, TensorDataType TPrecision>
void Mila::Dnn::LlamaModel< TDeviceType, TPrecision >::onTraining ( )
inlineoverrideprotectedvirtual

Training loop – not yet implemented for LlamaModel.

Exceptions
std::runtime_erroralways.

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

◆ toString()

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

Human-readable summary of this model's configuration.

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

◆ vocabSize()

template<DeviceType TDeviceType, TensorDataType TPrecision>
dim_t Mila::Dnn::LlamaModel< TDeviceType, TPrecision >::vocabSize ( ) const
inlineoverrideprotectedvirtualnoexcept

Vocabulary size from LLaMA config.

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


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