diff --git a/nisps-core/include/nisps/mlp.hpp b/nisps-core/include/nisps/mlp.hpp index cfa6dc7..ae459f6 100644 --- a/nisps-core/include/nisps/mlp.hpp +++ b/nisps-core/include/nisps/mlp.hpp @@ -130,7 +130,8 @@ public: float learning_rate, int max_iterations = 5000, float min_error_cost = 0.001, - bool output_log = true); + bool output_log = true, + const std::vector* sample_weights = nullptr); /** * @brief Training with batch support diff --git a/nisps-core/include/nisps/mlp_impl.hpp b/nisps-core/include/nisps/mlp_impl.hpp index 6d17d3b..8f9acf9 100644 --- a/nisps-core/include/nisps/mlp_impl.hpp +++ b/nisps-core/include/nisps/mlp_impl.hpp @@ -516,39 +516,33 @@ T MLP::Train(const training_pair_t& training_sample_set_with_bias, float learning_rate, int max_iterations, float min_error_cost, - bool) { + bool, + const std::vector* sample_weights) { int i = 0; T current_iteration_cost_function = 0.f; - T sampleSizeReciprocal = 1.f / training_sample_set_with_bias.first.size(); + const size_t n_samples = training_sample_set_with_bias.first.size(); + T sampleSizeReciprocal = 1.f / n_samples; for (i = 0; i < max_iterations; i++) { current_iteration_cost_function = 0.f; auto training_features = training_sample_set_with_bias.first; auto training_labels = training_sample_set_with_bias.second; - auto t_feat = training_features.begin(); - auto t_label = training_labels.begin(); - while (t_feat != training_features.end() || t_label != training_labels.end()) { + for (size_t s = 0; s < n_samples; s++) { + T weight = sample_weights ? (*sample_weights)[s] : sampleSizeReciprocal; - // Payload current_iteration_cost_function += - _TrainOnExample(*t_feat, *t_label, learning_rate, sampleSizeReciprocal); - - // \Payload - if (t_feat != training_features.end()) - { - ++t_feat; - } - if (t_label != training_labels.end()) - { - ++t_label; - } + _TrainOnExample(training_features[s], training_labels[s], learning_rate, weight); } - current_iteration_cost_function *= sampleSizeReciprocal; + // When using custom weights (already normalized to sum to 1), loss is already scaled. + // With uniform weights, multiply by sampleSizeReciprocal for backward compat. + if (!sample_weights) { + current_iteration_cost_function *= sampleSizeReciprocal; + } ReportProgress(true, 100, i, current_iteration_cost_function); diff --git a/playground/wasm/nisps.wasm b/playground/wasm/nisps.wasm index e7bd20c..28e296c 100755 Binary files a/playground/wasm/nisps.wasm and b/playground/wasm/nisps.wasm differ diff --git a/playground/wasm/nisps_bindings.cpp b/playground/wasm/nisps_bindings.cpp index 2d011c8..bd54df6 100644 --- a/playground/wasm/nisps_bindings.cpp +++ b/playground/wasm/nisps_bindings.cpp @@ -90,6 +90,7 @@ EMSCRIPTEN_KEEPALIVE float nisps_mlp_train(void* ptr, float* features_flat, int n_samples, int feature_dim, float* labels_flat, int label_dim, + float* sample_weights, float learning_rate, int max_iterations, float min_error) { auto* mlp = static_cast*>(ptr); @@ -104,8 +105,16 @@ float nisps_mlp_train(void* ptr, labels_flat + (i + 1) * label_dim); } + // Build weights vector if provided (non-null pointer) + std::vector weights_vec; + const std::vector* weights_ptr = nullptr; + if (sample_weights) { + weights_vec.assign(sample_weights, sample_weights + n_samples); + weights_ptr = &weights_vec; + } + nisps::MLP::training_pair_t data(features, labels); - return mlp->Train(data, learning_rate, max_iterations, min_error, false); + return mlp->Train(data, learning_rate, max_iterations, min_error, false, weights_ptr); } // ---- Weight manipulation with spread ----