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 f003a101c7..bef34b9e74 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. @@ -88,7 +88,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: @@ -197,13 +197,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") { @@ -419,7 +419,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: @@ -450,7 +450,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{}; } @@ -523,9 +523,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") { @@ -560,15 +560,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{}; } @@ -609,17 +609,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{}; } @@ -795,15 +795,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{}; } @@ -1010,9 +1010,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{}; } @@ -1122,37 +1122,39 @@ struct Model_Element : JSON::Element { } else if (name == "tokenizer_dir") { v_.tokenizer_dir = 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 = static_cast(JSON::Get(value)); + 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") { @@ -1160,25 +1162,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{}; } @@ -1249,8 +1251,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); } @@ -1259,17 +1267,17 @@ 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 = 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)); + 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") { @@ -1279,7 +1287,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 145cc93c00..3a572fd858 100644 --- a/src/config.h +++ b/src/config.h @@ -465,6 +465,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);