|
|
| GemmaModel (const GemmaModel &)=delete |
|
| GemmaModel (GemmaModel &&)=default |
|
dim_t | contextLength () const noexcept |
| | Deployment context length: the KV-cache depth the network was built with.
|
| std::string | fingerprintPrefill (const std::vector< int32_t > &token_ids) |
| | Logits fingerprint for a fixed prompt, for comparing two loads of one model.
|
|
const GemmaModelConfig & | getModelConfig () const noexcept |
| | Deployment configuration (context length, weight-quant, kv-compression) this model was loaded with.
|
|
const GemmaConfig & | getNetworkConfig () const noexcept |
| | Architecture/network configuration (read from the checkpoint metadata).
|
|
GemmaModel & | operator= (const GemmaModel &)=delete |
|
GemmaModel & | operator= (GemmaModel &&)=default |
|
void | profilePrefill (const std::vector< int32_t > &token_ids) |
| std::string | toString () const override |
| | Human-readable summary of this model's configuration.
|
|
| 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.
|
|
| 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.
|
|
| dim_t | maxSequenceLength () const noexcept override |
| 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 + decode implementation hook.
|
| void | onTraining () override |
| | Training loop hook – derived class owns the implementation.
|
| dim_t | vocabSize () const noexcept override |
| | 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.
|
|
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.
|
| | Model (std::unique_ptr< NetworkType > network, RuntimeMode runtime_mode) |
| | Construct with a fully built network and runtime mode.
|
template<
DeviceType TDeviceType,
TensorDataType TPrecision>
requires PrecisionSupportedOnDevice<TPrecision, TDeviceType>
class Mila::Dnn::GemmaModel< TDeviceType, TPrecision >
Gemma 4 compatible inference model.
Owns a loaded, built GemmaTransformer and drives the prefill + KV-cache decode two-phase generation loop. Construction is only possible via fromPretrained(); the network is always built, weights-loaded, and in inference mode when generation runs.
Thread safety: not thread-safe; external synchronization required if shared.
| std::string Mila::Dnn::GemmaModel< TDeviceType, TPrecision >::fingerprintPrefill |
( |
const std::vector< int32_t > & | token_ids | ) |
|
|
inline |
Logits fingerprint for a fixed prompt, for comparing two loads of one model.
Diagnostic. Parameters and config can be proven byte-identical between a quantize-on-load and a pre-quantized load while the models still behave differently, and nothing upstream of inference can see that. This runs one prefill and reports what the model actually computed.
Raw token ids rather than text so no tokenizer is involved and two runs are comparable by construction.
- Returns
- Digest of the last-position logits, the argmax token, and its value.
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, without reading a weight, and therefore without needing the device to have room. Answers before a multi-gigabyte download and for hardware the caller does not own.
Returns measurements only. Whether a given headroom is too tight is a deployment policy and belongs to the adaptor, not here; on Windows in particular WDDM oversubscribes rather than failing, so "fits" is not a property the runtime can decide. See Specifications/MemoryFootprint.md.
- Exceptions
-
| std::invalid_argument | on device type mismatch or zero context length. |
| std::runtime_error | on an unreadable artifact or unsupported quantization. |
|
|
inlineoverrideprotectedvirtual |
Training loop hook – derived class owns the implementation.
Called by train() after precondition enforcement. The derived class has total control over data loading, optimizer construction, loss computation, backward pass, checkpointing, and sampling.
Pure virtual – a model declaring RuntimeMode::Training must provide a training loop.
Implements Mila::Dnn::Model< TDeviceType, TPrecision >.