From ce70fa45d1e92d361ab73f96d9b75c8a9a8ad809 Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Thu, 11 Jun 2026 16:00:31 -0700 Subject: [PATCH 1/3] Add validation --- src/beam_search_scorer.cpp | 2 +- src/config.cpp | 8 ++++---- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/src/beam_search_scorer.cpp b/src/beam_search_scorer.cpp index 2b846636f6..ffb0d5b71f 100644 --- a/src/beam_search_scorer.cpp +++ b/src/beam_search_scorer.cpp @@ -66,7 +66,7 @@ BeamSearchScorer::BeamSearchScorer(const GeneratorParams& parameters) next_beam_indices_ = parameters.p_device->Allocate(batch_beam_size); // Space to store intermediate sequence - size_t const per_beam = (max_length_ * (max_length_ + 1)) / 2; + size_t const per_beam = (static_cast(max_length_) * (static_cast(max_length_) + 1)) / 2; hypothesis_buffer_ = device.Allocate(batch_beam_size * per_beam); memset(next_beam_scores_.Span().data(), 0, next_beam_scores_.Span().size_bytes()); diff --git a/src/config.cpp b/src/config.cpp index a32a986c19..a18a02a15e 100644 --- a/src/config.cpp +++ b/src/config.cpp @@ -1111,7 +1111,7 @@ struct Model_Element : JSON::Element { } else if (name == "vocab_size") { v_.vocab_size = static_cast(JSON::Get(value)); } else if (name == "context_length") { - v_.context_length = static_cast(JSON::Get(value)); + v_.context_length = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "pad_token_id") { v_.pad_token_id = static_cast(JSON::Get(value)); } else if (name == "eos_token_id") { @@ -1248,11 +1248,11 @@ struct Search_Element : JSON::Element { if (name == "min_length") { v_.min_length = static_cast(JSON::Get(value)); } else if (name == "max_length") { - v_.max_length = static_cast(JSON::Get(value)); + v_.max_length = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "batch_size") { - v_.batch_size = static_cast(JSON::Get(value)); + v_.batch_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_beams") { - v_.num_beams = static_cast(JSON::Get(value)); + v_.num_beams = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_return_sequences") { v_.num_return_sequences = static_cast(JSON::Get(value)); } else if (name == "top_k") { From 428e357f101ecb92e8d70620b467af1e80c33784 Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Thu, 11 Jun 2026 16:15:25 -0700 Subject: [PATCH 2/3] Address comments --- src/config.cpp | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/src/config.cpp b/src/config.cpp index a18a02a15e..68e4241bdc 100644 --- a/src/config.cpp +++ b/src/config.cpp @@ -1102,6 +1102,8 @@ struct Embedding_Element : JSON::Element { EmbeddingOutputs_Element outputs_{v_.outputs}; }; +int SafeDoubleToInt(double x, std::string_view name); + struct Model_Element : JSON::Element { explicit Model_Element(Config::Model& v) : v_{v} {} @@ -1112,6 +1114,8 @@ struct Model_Element : JSON::Element { v_.vocab_size = static_cast(JSON::Get(value)); } else if (name == "context_length") { v_.context_length = SafeDoubleToInt(JSON::Get(value), name); + if (v_.context_length <= 0) + throw std::out_of_range("context_length must be > 0, got " + std::to_string(v_.context_length)); } else if (name == "pad_token_id") { v_.pad_token_id = static_cast(JSON::Get(value)); } else if (name == "eos_token_id") { @@ -1236,8 +1240,14 @@ int SafeDoubleToInt(double x, std::string_view name) { throw std::overflow_error(ss.str()); } - // 3. Perform the cast. This truncates any fractional part (e.g., 3.9 becomes 3). - // If rounding is desired, use `return static_cast(std::round(x));` + // 3. Reject fractional values — these fields must be integral. + if (x != std::trunc(x)) { + std::stringstream ss; + ss << "Field '" << name << "' value " << x << " is not an integer"; + throw std::invalid_argument(ss.str()); + } + + // 4. Perform the cast. return static_cast(x); } From ecb0ef1bd5525fe7d990b24cea1563cf9e7b0e04 Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Fri, 12 Jun 2026 16:38:53 -0700 Subject: [PATCH 3/3] Address comments --- src/config.cpp | 114 ++++++++++++++++++++++++------------------------- src/config.h | 2 + 2 files changed, 58 insertions(+), 58 deletions(-) diff --git a/src/config.cpp b/src/config.cpp index 68e4241bdc..15eafdfed1 100644 --- a/src/config.cpp +++ b/src/config.cpp @@ -1,4 +1,4 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. // Modifications Copyright (C) 2026 Advanced Micro Devices, Inc. All rights reserved. // Portions of this file consist of AI generated content. @@ -77,7 +77,7 @@ struct Int_Array_Element : JSON::Element { explicit Int_Array_Element(std::vector& v) : v_{v} {} void OnValue(std::string_view name, JSON::Value value) override { - v_.emplace_back(static_cast(JSON::Get(value))); + v_.emplace_back(SafeDoubleToInt(JSON::Get(value), name)); } private: @@ -186,13 +186,13 @@ struct SessionOptions_Element : JSON::Element { } else if (name == "enable_profiling") { v_.enable_profiling = JSON::Get(value); } else if (name == "intra_op_num_threads") { - v_.intra_op_num_threads = static_cast(JSON::Get(value)); + v_.intra_op_num_threads = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "inter_op_num_threads") { - v_.inter_op_num_threads = static_cast(JSON::Get(value)); + v_.inter_op_num_threads = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "log_severity_level") { - v_.log_severity_level = static_cast(JSON::Get(value)); + v_.log_severity_level = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "log_verbosity_level") { - v_.log_verbosity_level = static_cast(JSON::Get(value)); + v_.log_verbosity_level = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "enable_cpu_mem_arena") { v_.enable_cpu_mem_arena = JSON::Get(value); } else if (name == "enable_mem_pattern") { @@ -408,7 +408,7 @@ struct IntArray_Element : JSON::Element { explicit IntArray_Element(std::vector& v) : v_{v} {} void OnValue(std::string_view name, JSON::Value value) override { - v_.push_back(static_cast(JSON::Get(value))); + v_.push_back(SafeDoubleToInt(JSON::Get(value), name)); } private: @@ -439,7 +439,7 @@ struct PipelineModel_Element : JSON::Element { } else if (name == "is_lm_head") { v_.is_lm_head = JSON::Get(value); } else if (name == "reset_session_idx") { - v_.reset_session_idx = static_cast(JSON::Get(value)); + v_.reset_session_idx = SafeDoubleToInt(JSON::Get(value), name); } else { throw JSON::unknown_value_error{}; } @@ -512,9 +512,9 @@ struct SlidingWindow_Element : JSON::Element { void OnValue(std::string_view name, JSON::Value value) override { if (name == "window_size") { - v_->window_size = static_cast(JSON::Get(value)); + v_->window_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "pad_value") { - v_->pad_value = static_cast(JSON::Get(value)); + v_->pad_value = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "alignment") { v_->alignment = JSON::Get(value); } else if (name == "slide_key_value_cache") { @@ -549,15 +549,15 @@ struct Encoder_Element : JSON::Element { if (name == "filename") { v_.filename = JSON::Get(value); } else if (name == "hidden_size") { - v_.hidden_size = static_cast(JSON::Get(value)); + v_.hidden_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_attention_heads") { - v_.num_attention_heads = static_cast(JSON::Get(value)); + v_.num_attention_heads = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_hidden_layers") { - v_.num_hidden_layers = static_cast(JSON::Get(value)); + v_.num_hidden_layers = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_key_value_heads") { - v_.num_key_value_heads = static_cast(JSON::Get(value)); + v_.num_key_value_heads = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "head_size") { - v_.head_size = static_cast(JSON::Get(value)); + v_.head_size = SafeDoubleToInt(JSON::Get(value), name); } else { throw JSON::unknown_value_error{}; } @@ -598,17 +598,17 @@ struct Decoder_Element : JSON::Element { if (name == "filename") { v_.filename = JSON::Get(value); } else if (name == "hidden_size") { - v_.hidden_size = static_cast(JSON::Get(value)); + v_.hidden_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_attention_heads") { - v_.num_attention_heads = static_cast(JSON::Get(value)); + v_.num_attention_heads = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_hidden_layers") { - v_.num_hidden_layers = static_cast(JSON::Get(value)); + v_.num_hidden_layers = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_key_value_heads") { - v_.num_key_value_heads = static_cast(JSON::Get(value)); + v_.num_key_value_heads = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "head_size") { - v_.head_size = static_cast(JSON::Get(value)); + v_.head_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "conv_cache_size") { - v_.conv_cache_size = static_cast(JSON::Get(value)); + v_.conv_cache_size = SafeDoubleToInt(JSON::Get(value), name); } else { throw JSON::unknown_value_error{}; } @@ -784,15 +784,15 @@ struct Vision_Element : JSON::Element { } else if (name == "adapter_filename") { v_.adapter_filename = JSON::Get(value); } else if (name == "spatial_merge_size") { - v_.spatial_merge_size = static_cast(JSON::Get(value)); + v_.spatial_merge_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "tokens_per_second") { v_.tokens_per_second = static_cast(JSON::Get(value)); } else if (name == "patch_size") { - v_.patch_size = static_cast(JSON::Get(value)); + v_.patch_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_visual_tokens") { - v_.num_visual_tokens = static_cast(JSON::Get(value)); + v_.num_visual_tokens = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "window_size") { - v_.window_size = static_cast(JSON::Get(value)); + v_.window_size = SafeDoubleToInt(JSON::Get(value), name); } else { throw JSON::unknown_value_error{}; } @@ -999,9 +999,9 @@ struct VAD_Element : JSON::Element { } else if (name == "threshold") { v_.threshold = static_cast(JSON::Get(value)); } else if (name == "silence_duration_ms") { - v_.silence_duration_ms = static_cast(JSON::Get(value)); + v_.silence_duration_ms = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "prefix_padding_ms") { - v_.prefix_padding_ms = static_cast(JSON::Get(value)); + v_.prefix_padding_ms = SafeDoubleToInt(JSON::Get(value), name); } else { throw JSON::unknown_value_error{}; } @@ -1102,8 +1102,6 @@ struct Embedding_Element : JSON::Element { EmbeddingOutputs_Element outputs_{v_.outputs}; }; -int SafeDoubleToInt(double x, std::string_view name); - struct Model_Element : JSON::Element { explicit Model_Element(Config::Model& v) : v_{v} {} @@ -1111,39 +1109,39 @@ struct Model_Element : JSON::Element { if (name == "type") { v_.type = JSON::Get(value); } else if (name == "vocab_size") { - v_.vocab_size = static_cast(JSON::Get(value)); + v_.vocab_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "context_length") { v_.context_length = SafeDoubleToInt(JSON::Get(value), name); if (v_.context_length <= 0) throw std::out_of_range("context_length must be > 0, got " + std::to_string(v_.context_length)); } else if (name == "pad_token_id") { - v_.pad_token_id = static_cast(JSON::Get(value)); + v_.pad_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "eos_token_id") { - v_.eos_token_id.assign(1, static_cast(JSON::Get(value))); + v_.eos_token_id.assign(1, SafeDoubleToInt(JSON::Get(value), name)); } else if (name == "bos_token_id") { - v_.bos_token_id = static_cast(JSON::Get(value)); + v_.bos_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "decoder_start_token_id") { - v_.decoder_start_token_id = static_cast(JSON::Get(value)); + v_.decoder_start_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "sep_token_id") { - v_.sep_token_id = static_cast(JSON::Get(value)); + v_.sep_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "image_token_id") { - v_.image_token_id = static_cast(JSON::Get(value)); + v_.image_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "audio_token_id") { - v_.audio_token_id = static_cast(JSON::Get(value)); + v_.audio_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "boa_token_id") { - v_.boa_token_id = static_cast(JSON::Get(value)); + v_.boa_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "video_token_id") { - v_.video_token_id = static_cast(JSON::Get(value)); + v_.video_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "vision_start_token_id") { - v_.vision_start_token_id = static_cast(JSON::Get(value)); + v_.vision_start_token_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_mels") { - v_.num_mels = static_cast(JSON::Get(value)); + v_.num_mels = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "fft_size") { - v_.fft_size = static_cast(JSON::Get(value)); + v_.fft_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "hop_length") { - v_.hop_length = static_cast(JSON::Get(value)); + v_.hop_length = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "win_length") { - v_.win_length = static_cast(JSON::Get(value)); + v_.win_length = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "preemph") { v_.preemph = static_cast(JSON::Get(value)); } else if (name == "log_eps") { @@ -1151,25 +1149,25 @@ struct Model_Element : JSON::Element { } else if (name == "norm_eps") { v_.norm_eps = static_cast(JSON::Get(value)); } else if (name == "subsampling_factor") { - v_.subsampling_factor = static_cast(JSON::Get(value)); + v_.subsampling_factor = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "left_context") { - v_.left_context = static_cast(JSON::Get(value)); + v_.left_context = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "conv_context") { - v_.conv_context = static_cast(JSON::Get(value)); + v_.conv_context = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "pre_encode_cache_size") { - v_.pre_encode_cache_size = static_cast(JSON::Get(value)); + v_.pre_encode_cache_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "sample_rate") { - v_.sample_rate = static_cast(JSON::Get(value)); + v_.sample_rate = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "chunk_samples") { - v_.chunk_samples = static_cast(JSON::Get(value)); + v_.chunk_samples = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "blank_id") { - v_.blank_id = static_cast(JSON::Get(value)); + v_.blank_id = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "max_symbols_per_step") { - v_.max_symbols_per_step = static_cast(JSON::Get(value)); + v_.max_symbols_per_step = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "left_context_samples") { - v_.left_context_samples = static_cast(JSON::Get(value)); + v_.left_context_samples = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "right_context_samples") { - v_.right_context_samples = static_cast(JSON::Get(value)); + v_.right_context_samples = SafeDoubleToInt(JSON::Get(value), name); } else { throw JSON::unknown_value_error{}; } @@ -1256,7 +1254,7 @@ struct Search_Element : JSON::Element { void OnValue(std::string_view name, JSON::Value value) override { if (name == "min_length") { - v_.min_length = static_cast(JSON::Get(value)); + v_.min_length = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "max_length") { v_.max_length = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "batch_size") { @@ -1264,9 +1262,9 @@ struct Search_Element : JSON::Element { } else if (name == "num_beams") { v_.num_beams = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "num_return_sequences") { - v_.num_return_sequences = static_cast(JSON::Get(value)); + v_.num_return_sequences = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "top_k") { - v_.top_k = static_cast(JSON::Get(value)); + v_.top_k = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "top_p") { v_.top_p = static_cast(JSON::Get(value)); } else if (name == "temperature") { @@ -1276,7 +1274,7 @@ struct Search_Element : JSON::Element { } else if (name == "length_penalty") { v_.length_penalty = static_cast(JSON::Get(value)); } else if (name == "no_repeat_ngram_size") { - v_.no_repeat_ngram_size = static_cast(JSON::Get(value)); + v_.no_repeat_ngram_size = SafeDoubleToInt(JSON::Get(value), name); } else if (name == "diversity_penalty") { v_.diversity_penalty = static_cast(JSON::Get(value)); } else if (name == "length_penalty") { diff --git a/src/config.h b/src/config.h index 8be9de3ecb..6f4573a216 100644 --- a/src/config.h +++ b/src/config.h @@ -457,6 +457,8 @@ void SetSearchBool(Config::Search& search, std::string_view name, bool value); void ClearProviders(Config& config); void SetProviderOption(Config& config, std::string_view provider_name, std::string_view option_name, std::string_view option_value); void OverlayConfig(Config& config, std::string_view json); +int SafeDoubleToInt(double x, std::string_view name); + bool IsGraphCaptureEnabled(const Config::SessionOptions& session_options); bool IsMultiProfileEnabled(const Config::SessionOptions& session_options);