|
Mila
Deep Neural Network Library
|
LLaMA 3 compatible inference model. More...
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 LlamaConfig & | getConfig () 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 ¶ms={}, 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 ¶ms, 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 ¶ms) |
| 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 ¶ms) |
| 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. | |
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.
| 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.
|
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:
| path | Path to the pretrained Llama model artifact. |
| model_config | Deployment configuration for this load. |
| device_id | Target device; must match TDeviceType. |
| std::invalid_argument | on device type mismatch or zero context length. |
| std::runtime_error | on load or parameter binding failure. |
| std::runtime_error | if model_config requests unsupported quantization (e.g. FP4). |
|
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.
| std::invalid_argument | on device type mismatch or zero context length. |
| std::runtime_error | on an unreadable artifact or unsupported quantization. |
|
inlineoverrideprotectedvirtualnoexcept |
Maximum sequence length from LLaMA config.
Implements Mila::Dnn::LanguageModel< TDeviceType, TPrecision >.
|
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.
| prompt_tokens | Input token ids; truncated from the start if they exceed the model's max sequence length. |
| on_token | Callback invoked once per generated token (not EOS). |
| params | Per-call generation parameters (loop bound + sampling). |
| stop | Stop token for cooperative cancellation. |
Implements Mila::Dnn::LanguageModel< TDeviceType, TPrecision >.
|
inlineoverrideprotectedvirtual |
Training loop – not yet implemented for LlamaModel.
| std::runtime_error | always. |
Implements Mila::Dnn::Model< TDeviceType, TPrecision >.
|
inlineoverridevirtual |
Human-readable summary of this model's configuration.
Implements Mila::Dnn::Model< TDeviceType, TPrecision >.
|
inlineoverrideprotectedvirtualnoexcept |
Vocabulary size from LLaMA config.
Implements Mila::Dnn::LanguageModel< TDeviceType, TPrecision >.