From 953b2f5be2026cef23b483408ab29c27bccd162b Mon Sep 17 00:00:00 2001 From: Ronan Collobert Date: Wed, 29 Oct 2025 16:11:32 -0700 Subject: [PATCH] WIP --- mlx/compile.cpp | 30 +++--- mlx/einsum.cpp | 46 +++++----- mlx/export.cpp | 19 ++-- mlx/fast.cpp | 14 +-- mlx/fft.cpp | 6 +- mlx/ops.cpp | 52 +++++------ mlx/primitives.cpp | 221 +++++++++++++++++++++++---------------------- mlx/primitives.h | 15 +-- mlx/random.h | 1 + mlx/scheduler.h | 4 +- mlx/types/bf16.h | 3 - mlx/utils.h | 4 +- 12 files changed, 208 insertions(+), 207 deletions(-) diff --git a/mlx/compile.cpp b/mlx/compile.cpp index d762c8d1..b0909589 100644 --- a/mlx/compile.cpp +++ b/mlx/compile.cpp @@ -194,7 +194,7 @@ const char* Compiled::name() const { } std::vector Compiled::output_shapes(const std::vector& inputs) { - size_t nd = 0; + int nd = 0; for (auto& in : inputs) { nd = std::max(nd, in.ndim()); } @@ -256,7 +256,7 @@ void merge(array& dst, array& src, ParentsMap& parents_map) { auto sources = src.outputs(); auto dests = dst.outputs(); // For each src parent, point it to the corresponding dst - for (int i = 0; i < sources.size(); ++i) { + for (int i = 0; i < std::ssize(sources); ++i) { merge_one(dests[i], sources[i], parents_map); } } @@ -327,7 +327,7 @@ class CompilerCache { if (in1.size() != in2.size()) { return false; } - for (size_t i = 0; i < in1.size(); ++i) { + for (int i = 0; i < std::ssize(in1); ++i) { if (in1[i].ndim() != in2[i].ndim()) { return false; } @@ -399,7 +399,7 @@ compile_trace( // Run the function on placeholder inputs // to get compute graph std::vector tracer_inputs; - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { array in(inputs[i].shape(), inputs[i].dtype(), nullptr, {}); in.set_tracer(true); tracer_inputs.push_back(std::move(in)); @@ -420,7 +420,7 @@ std::pair, ParentsMap> compile_dfs( std::unordered_set original_input_set; std::unordered_map>> parents_map; - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { input_set.insert(inputs[i].id()); original_input_set.insert(original_inputs[i].id()); } @@ -436,7 +436,7 @@ std::pair, ParentsMap> compile_dfs( if (cache.find(id) != cache.end()) { return; } - for (int i = 0; i < a.inputs().size(); i++) { + for (int i = 0; i < std::ssize(a.inputs()); i++) { auto& in = a.inputs()[i]; parents_map[in.id()].push_back({a, i}); for (auto& s : a.siblings()) { @@ -534,7 +534,7 @@ void compile_simplify( return false; } - for (int i = 0; i < a.inputs().size(); i++) { + for (int i = 0; i < std::ssize(a.inputs()); i++) { if (a.inputs()[i].id() != b.inputs()[i].id()) { return false; } @@ -599,7 +599,7 @@ void compile_simplify( auto maybe_merge_parents = [&](auto& a) { auto parents = parents_map.find(a.id()); if (parents != parents_map.end()) { - auto N = parents->second.size(); + auto N = std::ssize(parents->second); std::vector mask(N, false); auto try_merge = [&](int dst_idx, int src_idx) { @@ -642,11 +642,11 @@ void compile_simplify( it->second.push_back(i); } for (auto& [_, group] : dst_map) { - for (int i = 0; i < group.size(); ++i) { + for (int i = 0; i < std::ssize(group); ++i) { if (mask[group[i]]) { continue; } - for (int j = i + 1; j < group.size(); ++j) { + for (int j = i + 1; j < std::ssize(group); ++j) { if (mask[group[j]]) { continue; } @@ -847,7 +847,7 @@ void compile_fuse( std::vector old_outputs; // Add to global cache and add any global outputs to outputs // of new primitive - for (int j = 0; j < fused_tape.size() - 1; ++j) { + for (int j = 0; j < std::ssize(fused_tape) - 1; ++j) { auto& f = fused_tape[j]; if (output_map.find(f.id()) != output_map.end()) { old_outputs.push_back(f); @@ -903,7 +903,7 @@ void compile_fuse( new_tape.push_back(compiled_outputs.back()); // Replace inputs old parents with compiled_outputs - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { auto& pairs = parents_map[inputs[i].id()]; pairs.erase( std::remove_if( @@ -918,7 +918,7 @@ void compile_fuse( // - Update outputs parents to point to compiled outputs // - Update any overall graph outputs to be compiled outputs - for (int o = 0; o < old_outputs.size(); ++o) { + for (int o = 0; o < std::ssize(old_outputs); ++o) { merge_one(compiled_outputs[o], old_outputs[o], parents_map); if (auto it = output_map.find(old_outputs[o].id()); it != output_map.end()) { @@ -943,7 +943,7 @@ std::vector compile_replace( const std::vector& inputs, bool shapeless) { std::unordered_map trace_to_real; - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { trace_to_real.insert({trace_inputs[i].id(), inputs[i]}); } @@ -989,7 +989,7 @@ std::vector compile_replace( } auto real_out = array::make_arrays( std::move(shapes), types, a.primitive_ptr(), real_inputs); - for (int i = 0; i < trace_out.size(); ++i) { + for (int i = 0; i < std::ssize(trace_out); ++i) { trace_to_real.insert({trace_out[i].id(), std::move(real_out[i])}); } } diff --git a/mlx/einsum.cpp b/mlx/einsum.cpp index 62908877..5dad17ba 100644 --- a/mlx/einsum.cpp +++ b/mlx/einsum.cpp @@ -190,8 +190,8 @@ std::tuple, size_t, int> greedy_path( // Start by iterating over all possible combinations std::vector> pos_pairs; - for (int i = 0; i < inputs.size(); ++i) { - for (int j = i + 1; j < inputs.size(); ++j) { + for (int i = 0; i < std::ssize(inputs); ++i) { + for (int j = i + 1; j < std::ssize(inputs); ++j) { pos_pairs.emplace_back(i, j); } } @@ -200,13 +200,13 @@ std::tuple, size_t, int> greedy_path( std::vector possible_contractions; size_t path_cost = 0; int path_scaling = 0; - auto num_in = inputs.size(); + auto num_in = std::ssize(inputs); for (int i = 0; i < num_in - 1; ++i) { auto add_contraction = [&](int p1, int p2) { CharSet new_term; CharSet contractions(inputs[p1].set.begin(), inputs[p1].set.end()); contractions.insert(inputs[p2].set.begin(), inputs[p2].set.end()); - for (int i = 0; i < inputs.size(); i++) { + for (int i = 0; i < std::ssize(inputs); i++) { if (i == p1 || i == p2) { continue; } @@ -321,7 +321,7 @@ std::tuple, size_t, int> greedy_path( } pos_pairs.clear(); - for (int i = 0; i < inputs.size() - 1; ++i) { + for (int i = 0; i < std::ssize(inputs) - 1; ++i) { pos_pairs.emplace_back(i, inputs.size() - 1); } path_cost += best.cost; @@ -360,7 +360,7 @@ array batch_tensordot( { auto a_shape = a.shape(); auto b_shape = b.shape(); - for (int i = 0; i < a_contract.size(); ++i) { + for (int i = 0; i < std::ssize(a_contract); ++i) { auto d = std::max(a.shape(a_contract[i]), b.shape(b_contract[i])); a_shape[a_contract[i]] = d; b_shape[b_contract[i]] = d; @@ -430,7 +430,7 @@ array collapse_repeats(array in, Subscript& subscript, StreamOrDevice s) { std::string repeat_str; std::string no_repeat_str; std::unordered_map counts; - for (int i = 0; i < str.size(); ++i) { + for (int i = 0; i < std::ssize(str); ++i) { auto [it, _] = counts.insert({str[i], 0}); it->second++; } @@ -455,7 +455,7 @@ array collapse_repeats(array in, Subscript& subscript, StreamOrDevice s) { std::vector indices; int n_expand = repeats.size(); for (auto [c, v] : repeats) { - for (int i = 0; i < str.size(); ++i) { + for (int i = 0; i < std::ssize(str); ++i) { if (str[i] == c) { slice_sizes[i] = 1; axes.push_back(i); @@ -494,7 +494,7 @@ void preprocess_einsum_inputs( std::vector& operands, StreamOrDevice s) { // Collapse repeat indices - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { auto& in = inputs[i]; if (in.set.size() < in.str.size()) { operands[positions[i]] = collapse_repeats(operands[positions[i]], in, s); @@ -514,10 +514,10 @@ void preprocess_einsum_inputs( auto inserted = counts.insert({c, 0}); inserted.first->second++; } - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { auto& in = inputs[i]; std::vector sum_axes; - for (int ax = 0; ax < in.str.size(); ++ax) { + for (int ax = 0; ax < std::ssize(in.str); ++ax) { if (counts[in.str[ax]] == 1) { sum_axes.push_back(ax); } @@ -549,12 +549,12 @@ array einsum_naive( } // Expand and transpose inputs as needed - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { int pos = positions[i]; auto& op = operands[pos]; // Add missing dimensions at the end - if (op.ndim() != char_to_ax.size()) { + if (op.ndim() != std::ssize(char_to_ax)) { auto shape = op.shape(); shape.insert(shape.end(), char_to_ax.size() - shape.size(), 1); op = reshape(op, std::move(shape), s); @@ -597,7 +597,7 @@ array einsum_naive( // Multiply and sum auto out = operands[positions[0]]; - for (int i = 1; i < positions.size(); ++i) { + for (int i = 1; i < std::ssize(positions); ++i) { out = multiply(out, operands[positions[i]], s); } std::vector sum_axes; @@ -675,9 +675,9 @@ std::pair, PathInfo> einsum_path_helper( int operand_idx) { bool have_ellipsis = false; int cnt_before = 0, cnt_after = 0; - for (int i = 0; i < subscript.size(); i++) { + for (int i = 0; i < std::ssize(subscript); i++) { if (!isalpha(subscript[i])) { - if (i + 2 >= subscript.size() || subscript[i] != '.' || + if (i + 2 >= std::ssize(subscript) || subscript[i] != '.' || subscript[i + 1] != '.' || subscript[i + 2] != '.') { std::ostringstream msg; msg << "[" << fn_name << "] Subscripts must be letters, but got '" @@ -732,7 +732,7 @@ std::pair, PathInfo> einsum_path_helper( } }; - for (int i = 0; i < operands.size(); i++) { + for (int i = 0; i < std::ssize(operands); i++) { check_letters_and_expand_ellipsis(in_subscripts[i], &operands[i], i); } check_letters_and_expand_ellipsis(out_subscript, nullptr, -1); @@ -747,12 +747,12 @@ std::pair, PathInfo> einsum_path_helper( std::unordered_map dim_map; std::vector inputs; - for (int i = 0; i < in_subscripts.size(); ++i) { + for (int i = 0; i < std::ssize(in_subscripts); ++i) { auto& in = in_subscripts[i]; CharSet in_set(in.begin(), in.end()); inputs.emplace_back(in, in_set); - if (in.size() != operands[i].ndim()) { + if (std::ssize(in) != operands[i].ndim()) { std::ostringstream msg; msg << "[" << fn_name << "] Invalid number of subscripts " << in.size() << " for input " << i << " with " << operands[i].ndim() @@ -763,7 +763,7 @@ std::pair, PathInfo> einsum_path_helper( // Check repeat subscripts are valid if (in_set.size() < in.size()) { std::unordered_map local_dims; - for (int j = 0; j < in.size(); ++j) { + for (int j = 0; j < std::ssize(in); ++j) { auto dim = operands[i].shape(j); auto inserted = local_dims.insert({in[j], dim}); if (!inserted.second) { @@ -778,7 +778,7 @@ std::pair, PathInfo> einsum_path_helper( } } - for (int j = 0; j < in.size(); j++) { + for (int j = 0; j < std::ssize(in); j++) { auto c = in[j]; auto dim = operands[i].shape(j); auto inserted = dim_map.insert({c, dim}); @@ -864,7 +864,7 @@ array einsum( std::vector a_contract; std::vector a_batch; std::vector a_concat; - for (int i = 0; i < in_a.str.size(); ++i) { + for (int i = 0; i < std::ssize(in_a.str); ++i) { auto c = in_a.str[i]; if (out.set.find(c) == out.set.end()) { // Not in the output, contraction @@ -887,7 +887,7 @@ array einsum( for (auto a_i : a_batch) { b_batch.push_back(in_b.str.find(in_a.str[a_i])); } - for (int i = 0; i < in_b.str.size(); ++i) { + for (int i = 0; i < std::ssize(in_b.str); ++i) { auto c = in_b.str[i]; if (out.set.find(c) != out.set.end() && in_a.set.find(c) == in_a.set.end()) { diff --git a/mlx/export.cpp b/mlx/export.cpp index 3448178e..ce26141c 100644 --- a/mlx/export.cpp +++ b/mlx/export.cpp @@ -138,7 +138,7 @@ T deserialize(Reader& is) { T v; auto size = deserialize(is); v.reserve(size); - for (int i = 0; i < size; ++i) { + for (size_t i = 0; i < size; ++i) { v.push_back(deserialize(is)); } return v; @@ -487,11 +487,11 @@ struct FunctionTable { int n = 1; for (auto& [_, vec] : table) { for (auto& fun : vec) { - auto npos = fun.inputs.size() - fun.kwarg_keys.size(); + auto npos = std::ssize(fun.inputs) - std::ssize(fun.kwarg_keys); os << " " << n++ << ". Function with " << npos - << " positional inputs and " << fun.kwarg_keys.size() + << " positional inputs and " << std::ssize(fun.kwarg_keys) << " keyword inputs:\n"; - for (int j = 0; j < fun.inputs.size(); ++j) { + for (int j = 0; j < std::ssize(fun.inputs); ++j) { auto& in = fun.inputs[j]; if (j < npos) { os << " " << j + 1 << ": "; @@ -536,7 +536,7 @@ bool FunctionTable::match( }; int i = 0; - for (; i < args.size(); ++i) { + for (; i < std::ssize(args); ++i) { if (!match_inputs(args[i], fun.inputs[i])) { return false; } @@ -627,7 +627,8 @@ void FunctionExporter::export_with_callback( // Callback on the inputs callback({{"type", "inputs"}, {"inputs", to_vector_data(inputs)}}); std::vector> keyword_inputs; - for (int i = inputs.size() - kwarg_keys.size(), j = 0; i < inputs.size(); + for (int i = std::ssize(inputs) - std::ssize(kwarg_keys), j = 0; + i < std::ssize(inputs); ++i, ++j) { keyword_inputs.emplace_back(kwarg_keys[j], namer.get_name(inputs[i])); } @@ -928,7 +929,7 @@ std::vector ImportedFunction::operator()( ftable->print_functions(msg); msg << "\nCalled with " << args.size() << " positional inputs and " << kwargs.size() << " keyword inputs:\n"; - for (int i = 0; i < args.size(); ++i) { + for (int i = 0; i < std::ssize(args); ++i) { auto& in = args[i]; msg << " " << i + 1 << ": " << in.shape() << " " << in.dtype() << "\n"; } @@ -970,7 +971,7 @@ ImportedFunction::ImportedFunction(const std::string& file) std::unordered_map array_map; auto trace_input_ids = deserialize>(is); auto trace_inputs = deserialize>(is); - for (int i = 0; i < trace_inputs.size(); ++i) { + for (int i = 0; i < std::ssize(trace_inputs); ++i) { array_map.emplace(trace_input_ids[i], trace_inputs[i]); } auto trace_output_ids = deserialize>(is); @@ -1006,7 +1007,7 @@ ImportedFunction::ImportedFunction(const std::string& file) std::move(types), std::move(prim), std::move(inputs)); - for (int i = 0; i < arrays.size(); ++i) { + for (int i = 0; i < std::ssize(arrays); ++i) { auto sid = ids[i]; if (sid == id) { tape.push_back(arrays[i]); diff --git a/mlx/fast.cpp b/mlx/fast.cpp index 0f34aec9..e88527a8 100644 --- a/mlx/fast.cpp +++ b/mlx/fast.cpp @@ -13,11 +13,11 @@ std::vector Custom::vjp( const std::vector& primals, const std::vector& cotangents, const std::vector& argnums, - const std::vector& outputs) { + const std::vector& /* outputs */) { auto [_, vjps] = mlx::core::vjp(fallback_, primals, cotangents); std::vector vjp_outs; - for (int i = 0, j = 0; i < vjps.size(); ++i) { - if (j < argnums.size() && i == argnums[j]) { + for (int i = 0, j = 0; i < std::ssize(vjps); ++i) { + if (j < std::ssize(argnums) && i == argnums[j]) { vjp_outs.push_back(vjps[i]); j++; } @@ -30,8 +30,8 @@ std::vector Custom::jvp( const std::vector& tangents, const std::vector& argnums) { std::vector all_tangents; - for (int i = 0, j = 0; i < primals.size(); i++) { - if (j < argnums.size() && i == argnums[j]) { + for (int i = 0, j = 0; i < std::ssize(primals); i++) { + if (j < std::ssize(argnums) && i == argnums[j]) { all_tangents.emplace_back(tangents[j++]); } else { all_tangents.emplace_back(zeros_like(primals[i])); @@ -536,7 +536,7 @@ std::vector RoPE::vjp( const std::vector& primals, const std::vector& cotangents, const std::vector& argnums, - const std::vector& outputs) { + const std::vector& /* outputs */) { auto s = stream(); auto fallback = [dims = dims_, traditional = traditional_, @@ -635,7 +635,7 @@ array scaled_dot_product_attention( throw std::invalid_argument(msg.str()); } - const size_t batch_dim = queries.shape(0); + const int batch_dim = queries.shape(0); for (const auto& tensor : {keys, values}) { if (tensor.shape(0) != batch_dim) { std::ostringstream msg; diff --git a/mlx/fft.cpp b/mlx/fft.cpp index 6510faec..33f5e763 100644 --- a/mlx/fft.cpp +++ b/mlx/fft.cpp @@ -20,14 +20,14 @@ array fft_impl( throw std::invalid_argument( "[fftn] Requires array with at least one dimension."); } - if (n.size() != axes.size()) { + if (n.size() != std::ssize(axes)) { throw std::invalid_argument("[fftn] Shape and axes have different sizes."); } if (axes.empty()) { return a; } - std::vector valid_axes; + std::vector valid_axes; for (int ax : axes) { valid_axes.push_back(ax < 0 ? ax + a.ndim() : ax); } @@ -59,7 +59,7 @@ array fft_impl( } auto in_shape = a.shape(); - for (int i = 0; i < valid_axes.size(); ++i) { + for (int i = 0; i < std::ssize(valid_axes); ++i) { in_shape[valid_axes[i]] = n[i]; } if (real && inverse) { diff --git a/mlx/ops.cpp b/mlx/ops.cpp index 30e934f8..e1f26457 100644 --- a/mlx/ops.cpp +++ b/mlx/ops.cpp @@ -390,7 +390,7 @@ array unflatten( throw std::invalid_argument(msg.str()); } - size_t size = 1; + int64_t size = 1; int infer_idx = -1; for (int i = 0; i < shape.size(); ++i) { if (shape[i] == -1) { @@ -687,10 +687,10 @@ void normalize_dynamic_slice_inputs( << "."; throw std::invalid_argument(msg.str()); } - if (start.size() != axes.size()) { + if (start.size() != std::ssize(axes)) { std::ostringstream msg; msg << prefix << " Number of starting indices " << start.size() - << " does not match number of axes " << axes.size() << "."; + << " does not match number of axes " << std::ssize(axes) << "."; throw std::invalid_argument(msg.str()); } if (!issubdtype(start.dtype(), integer)) { @@ -847,7 +847,7 @@ array slice_update( // Broadcast update with unspecified axes auto up_shape = update.shape(); - auto dim_diff = std::max(src.ndim() - update.ndim(), size_t(0)); + auto dim_diff = std::max(src.ndim() - update.ndim(), 0); up_shape.insert( up_shape.begin(), src.shape().begin(), src.shape().begin() + dim_diff); for (int d = dim_diff; d < src.ndim(); ++d) { @@ -957,7 +957,7 @@ std::vector meshgrid( "[meshgrid] Invalid indexing value. Valid values are 'xy' and 'ij'."); } - auto ndim = arrays.size(); + auto ndim = std::ssize(arrays); std::vector outputs; for (int i = 0; i < ndim; ++i) { Shape shape(ndim, 1); @@ -1135,10 +1135,10 @@ array tile( std::vector reps, StreamOrDevice s /* = {} */) { auto shape = arr.shape(); - if (reps.size() < shape.size()) { + if (std::ssize(reps) < shape.size()) { reps.insert(reps.begin(), shape.size() - reps.size(), 1); } - if (reps.size() > shape.size()) { + if (std::ssize(reps) > shape.size()) { shape.insert(shape.begin(), reps.size() - shape.size(), 1); } @@ -1162,7 +1162,7 @@ array tile( array edge_pad( const array& a, - const std::vector& axes, + const std::vector& /* axes */, const Shape& low_pad_size, const Shape& high_pad_size, const Shape& out_shape, @@ -1214,17 +1214,17 @@ array pad( const array& pad_value /*= array(0)*/, const std::string& mode /*= "constant"*/, StreamOrDevice s /* = {}*/) { - if (axes.size() != low_pad_size.size() || - axes.size() != high_pad_size.size()) { + if (std::ssize(axes) != low_pad_size.size() || + std::ssize(axes) != high_pad_size.size()) { std::ostringstream msg; msg << "Invalid number of padding sizes passed to pad " - << "with axes of size " << axes.size(); + << "with axes of size " << std::ssize(axes); throw std::invalid_argument(msg.str()); } auto out_shape = a.shape(); - for (int i = 0; i < axes.size(); i++) { + for (int i = 0; i < std::ssize(axes); i++) { if (low_pad_size[i] < 0) { std::ostringstream msg; msg << "Invalid low padding size (" << low_pad_size[i] @@ -1365,7 +1365,7 @@ array transpose( for (auto& ax : axes) { ax = ax < 0 ? ax + a.ndim() : ax; } - if (axes.size() != a.ndim()) { + if (std::ssize(axes) != a.ndim()) { std::ostringstream msg; msg << "[transpose] Recived " << axes.size() << " axes for array with " << a.ndim() << " dimensions."; @@ -1387,7 +1387,7 @@ array transpose( shape[ax] = 1; } - for (int i = 0; i < axes.size(); ++i) { + for (int i = 0; i < std::ssize(axes); ++i) { shape[i] = a.shape()[axes[i]]; } return array( @@ -1444,7 +1444,7 @@ std::vector broadcast_arrays( auto shape = BroadcastAxes::output_shape(inputs, ignore_axes); auto check_and_get_shape = [&shape, &ignore_axes](const array& in) { auto out_shape = shape; - for (int i = 0; i < ignore_axes.size(); ++i) { + for (int i = 0; i < std::ssize(ignore_axes); ++i) { auto ax = ignore_axes[i]; auto pos_ax = in.ndim() + ax; if (pos_ax < 0 || pos_ax > in.ndim() || @@ -1478,7 +1478,7 @@ std::vector broadcast_arrays( stop_grad_inputs.push_back(stop_gradient(in, s)); } - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { auto& in = inputs[i]; auto out_shape = check_and_get_shape(in); if (in.shape() == out_shape) { @@ -1486,7 +1486,7 @@ std::vector broadcast_arrays( } else { // broadcasted array goes first followed by other stopgrad inputs std::vector p_inputs = {in}; - for (int j = 0; j < inputs.size(); ++j) { + for (int j = 0; j < std::ssize(inputs); ++j) { if (j == i) { continue; } @@ -1530,14 +1530,14 @@ std::vector broadcast_arrays( for (auto& in : inputs) { stop_grad_inputs.push_back(stop_gradient(in, s)); } - for (int i = 0; i < inputs.size(); ++i) { + for (int i = 0; i < std::ssize(inputs); ++i) { auto& in = inputs[i]; if (in.shape() == shape) { outputs.push_back(in); } else { // broadcasted array goes first followed by other stopgrad inputs std::vector p_inputs = {in}; - for (int j = 0; j < inputs.size(); ++j) { + for (int j = 0; j < std::ssize(inputs); ++j) { if (j == i) { continue; } @@ -1961,7 +1961,7 @@ array median( auto dtype = at_least_float(a.dtype()); std::vector transpose_axes; for (int i = 0, j = 0; i < a.ndim(); ++i) { - if (j < sorted_axes.size() && i == sorted_axes[j]) { + if (j < std::ssize(sorted_axes) && i == sorted_axes[j]) { j++; continue; } @@ -3010,7 +3010,7 @@ array gather( const Shape& slice_sizes, StreamOrDevice s /* = {} */) { // Checks that indices, dimensions, and slice_sizes are all valid - if (indices.size() > a.ndim()) { + if (std::ssize(indices) > a.ndim()) { std::ostringstream msg; msg << "[gather] Too many index arrays. Got " << indices.size() << " index arrays for input with " << a.ndim() << " dimensions."; @@ -3312,7 +3312,7 @@ array scatter( Scatter::ReduceType mode, StreamOrDevice s) { // Checks that indices, dimensions, and slice_sizes are all valid - if (indices.size() > a.ndim()) { + if (std::ssize(indices) > a.ndim()) { std::ostringstream msg; msg << "[scatter] Too many index arrays. Got " << indices.size() << " index arrays for input with " << a.ndim() << " dimensions."; @@ -3820,7 +3820,7 @@ array conv_transpose_general( StreamOrDevice s) { std::vector padding_lo(padding.size()); std::vector padding_hi(padding.size()); - for (int i = 0; i < padding.size(); ++i) { + for (int i = 0; i < std::ssize(padding); ++i) { int wt_size = 1 + dilation[i] * (weight.shape(1 + i) - 1); padding_lo[i] = wt_size - padding[i] - 1; @@ -4632,7 +4632,7 @@ array tensordot( int csize = 1; auto x = a; auto y = b; - for (int i = 0; i < axes_a.size(); i++) { + for (int i = 0; i < std::ssize(axes_a); i++) { if (x.shape(axes_a.at(i)) == y.shape(axes_b.at(i))) { csize *= x.shape(axes_a.at(i)); } else { @@ -5560,7 +5560,7 @@ array roll( return a; } - if (shift.size() < axes.size()) { + if (shift.size() < std::ssize(axes)) { std::ostringstream msg; msg << "[roll] At least one shift value per axis is required, " << shift.size() << " provided for " << axes.size() << " axes."; @@ -5568,7 +5568,7 @@ array roll( } array result = a; - for (int i = 0; i < axes.size(); i++) { + for (int i = 0; i < std::ssize(axes); i++) { int ax = axes[i]; if (ax < 0) { ax += a.ndim(); diff --git a/mlx/primitives.cpp b/mlx/primitives.cpp index 0b335e76..152487de 100644 --- a/mlx/primitives.cpp +++ b/mlx/primitives.cpp @@ -242,16 +242,16 @@ std::pair, std::vector> Abs::vmap( } std::vector Add::jvp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { return { tangents.size() > 1 ? add(tangents[0], tangents[1], stream()) : tangents[0]}; } std::vector Add::vjp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& cotangents, const std::vector& argnums, const std::vector&) { @@ -315,7 +315,7 @@ std::vector AddMM::jvp( const std::vector& tangents, const std::vector& argnums) { std::vector jvp; - for (int i = 0; i < argnums.size(); ++i) { + for (int i = 0; i < std::ssize(argnums); ++i) { auto arg = argnums[i]; if (arg == 0) { if (jvp.empty()) { @@ -692,7 +692,7 @@ std::vector ArgSort::jvp( std::vector AsType::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector&) { if (cotangents[0].dtype() != dtype_) { throw std::invalid_argument( @@ -702,9 +702,9 @@ std::vector AsType::vjp( } std::vector AsType::jvp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { return {astype(tangents[0], dtype_, stream())}; } @@ -752,7 +752,7 @@ std::vector AsStrided::vjp( std::vector AsStrided::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { assert(primals.size() == 1); return {as_strided(tangents[0], shape_, strides_, offset_, stream())}; @@ -827,9 +827,9 @@ std::vector Broadcast::vjp( } std::vector Broadcast::jvp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { return {array( shape_, tangents[0].dtype(), @@ -858,7 +858,7 @@ bool Broadcast::is_equivalent(const Primitive& other) const { Shape Broadcast::output_shape(const std::vector& inputs) { auto shape = inputs[0].shape(); - for (int i = 1; i < inputs.size(); ++i) { + for (int i = 1; i < std::ssize(inputs); ++i) { shape = broadcast_shapes(shape, inputs[i].shape()); } return shape; @@ -886,7 +886,7 @@ std::vector BroadcastAxes::vjp( std::vector BroadcastAxes::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { return {array( output_shape(primals, ignore_axes_), tangents[0].dtype(), @@ -895,8 +895,8 @@ std::vector BroadcastAxes::jvp( } std::pair, std::vector> BroadcastAxes::vmap( - const std::vector& inputs, - const std::vector& axes) { + const std::vector& /* inputs */, + const std::vector& /* axes */) { throw std::invalid_argument("[BroadcastAxes] VMAP NYI"); } @@ -938,7 +938,7 @@ std::vector Ceil::vjp( std::vector Ceil::jvp( const std::vector& primals, - const std::vector& tangents, + const std::vector& /* tangents */, const std::vector& argnums) { assert(primals.size() == 1); assert(argnums.size() == 1); @@ -1072,8 +1072,8 @@ std::vector Concatenate::jvp( }); std::vector vals; - for (int i = 0, j = 0; i < primals.size(); ++i) { - if (j < argnums.size() && argnums[argidx[j]] == i) { + for (int i = 0, j = 0; i < std::ssize(primals); ++i) { + if (j < std::ssize(argnums) && argnums[argidx[j]] == i) { vals.push_back(tangents[argidx[j++]]); } else { vals.push_back(zeros_like(primals[i], stream())); @@ -1089,7 +1089,7 @@ std::pair, std::vector> Concatenate::vmap( int first_vmap = -1; // Find the first vmapped input - for (int i = 0; i < axes.size(); i++) { + for (int i = 0; i < std::ssize(axes); i++) { if (axes[i] >= 0) { out_ax = axes[i]; first_vmap = i; @@ -1107,7 +1107,7 @@ std::pair, std::vector> Concatenate::vmap( std::vector t_inputs; int axis = axis_ + (axis_ >= out_ax); auto cat_shape = inputs[first_vmap].shape(); - for (int i = 0; i < axes.size(); i++) { + for (int i = 0; i < std::ssize(axes); i++) { if (axes[i] >= 0) { if (out_ax != axes[i]) { t_inputs.push_back(moveaxis(inputs[i], axes[i], out_ax, stream())); @@ -1132,7 +1132,7 @@ bool Concatenate::is_equivalent(const Primitive& other) const { std::vector Concatenate::output_shapes( const std::vector& inputs) { auto shape = inputs[0].shape(); - for (int i = 1; i < inputs.size(); ++i) { + for (int i = 1; i < std::ssize(inputs); ++i) { shape[axis_] += inputs[i].shape(axis_); } return {std::move(shape)}; @@ -1272,28 +1272,29 @@ Shape Convolution::conv_out_shape( int spatial_dims = in_shape.size() - 2; - if (strides.size() != spatial_dims) { + if (std::ssize(strides) != spatial_dims) { std::ostringstream msg; msg << "[conv] Invalid strides " << strides << " for " << spatial_dims << "D convolution."; throw std::invalid_argument(msg.str()); } - if (pads_lo.size() != spatial_dims || pads_hi.size() != spatial_dims) { + if (std::ssize(pads_lo) != spatial_dims || + std::ssize(pads_hi) != spatial_dims) { std::ostringstream msg; msg << "[conv] Invalid padding " << pads_lo << " | " << pads_hi << " for " << spatial_dims << "D convolution."; throw std::invalid_argument(msg.str()); } - if (kernel_dilation.size() != spatial_dims) { + if (std::ssize(kernel_dilation) != spatial_dims) { std::ostringstream msg; msg << "[conv] Invalid kernel dilation " << kernel_dilation << " for " << spatial_dims << "D convolution."; throw std::invalid_argument(msg.str()); } - if (input_dilation.size() != spatial_dims) { + if (std::ssize(input_dilation) != spatial_dims) { std::ostringstream msg; msg << "[conv] Invalid input dilation " << input_dilation << " for " << spatial_dims << "D convolution."; @@ -1386,7 +1387,7 @@ std::vector Convolution::vjp( std::vector padding_lo = padding_lo_; std::vector padding_hi = padding_hi_; - for (int i = 0; i < padding_lo.size(); ++i) { + for (int i = 0; i < std::ssize(padding_lo); ++i) { int wt_size = 1 + kernel_dilation_[i] * (wt.shape(1 + i) - 1); padding_lo[i] = wt_size - padding_lo_[i] - 1; @@ -1440,7 +1441,7 @@ std::vector Convolution::vjp( else if (a == 1) { bool no_dilation = true; - for (int i = 0; i < input_dilation_.size(); i++) { + for (int i = 0; i < std::ssize(input_dilation_); i++) { no_dilation &= (input_dilation_[i] == 1) && (kernel_dilation_[i] == 1); } @@ -1451,7 +1452,7 @@ std::vector Convolution::vjp( } else { auto padding_hi = padding_lo_; - for (int i = 0; i < padding_hi.size(); ++i) { + for (int i = 0; i < std::ssize(padding_hi); ++i) { int in_size = 1 + input_dilation_[i] * (in.shape(1 + i) - 1); int out_size = 1 + kernel_strides_[i] * (cotan.shape(1 + i) - 1); int wt_size = 1 + kernel_dilation_[i] * (wt.shape(1 + i) - 1); @@ -1699,11 +1700,11 @@ std::vector Depends::vjp( const std::vector& primals, const std::vector& cotangents, const std::vector& argnums, - const std::vector& outputs) { + const std::vector& /* outputs */) { std::vector vjps; for (auto arg : argnums) { - if (arg < cotangents.size()) { + if (arg < std::ssize(cotangents)) { vjps.push_back(cotangents[arg]); } else { vjps.push_back(zeros_like(primals[arg])); @@ -1737,7 +1738,7 @@ std::vector Divide::vjp( std::vector DivMod::vjp( const std::vector& primals, - const std::vector& cotangents, + const std::vector& /* cotangents */, const std::vector& argnums, const std::vector&) { std::vector vjps; @@ -1749,8 +1750,8 @@ std::vector DivMod::vjp( std::vector DivMod::jvp( const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* tangents */, + const std::vector& /* argnums */) { return {zeros_like(primals[0], stream())}; } @@ -1848,7 +1849,7 @@ std::pair, std::vector> Equal::vmap( std::vector Equal::vjp( const std::vector& primals, - const std::vector& cotangents, + const std::vector& /* cotangents */, const std::vector& argnums, const std::vector&) { std::vector vjps; @@ -1860,8 +1861,8 @@ std::vector Equal::vjp( std::vector Equal::jvp( const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* tangents */, + const std::vector& /* argnums */) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); return {zeros(shape, bool_, stream())}; } @@ -1899,7 +1900,7 @@ std::pair, std::vector> Erf::vmap( std::vector ErfInv::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector& outputs) { auto dtype = primals[0].dtype(); auto scale = @@ -1931,9 +1932,9 @@ std::pair, std::vector> ErfInv::vmap( } std::vector Exp::vjp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector& outputs) { return {multiply(cotangents[0], outputs[0], stream())}; } @@ -1956,9 +1957,9 @@ std::pair, std::vector> Exp::vmap( } std::vector Expm1::vjp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector& outputs) { return {multiply( cotangents[0], @@ -2281,7 +2282,7 @@ std::vector Floor::vjp( std::vector Floor::jvp( const std::vector& primals, - const std::vector& tangents, + const std::vector& /* tangents */, const std::vector& argnums) { assert(primals.size() == 1); assert(argnums.size() == 1); @@ -2346,7 +2347,7 @@ std::pair, std::vector> Gather::vmap( // Reorder all the index arrays so the vmap axis is in the same spot. if (indices_vmapped) { - for (int i = 1; i < axes.size(); ++i) { + for (int i = 1; i < std::ssize(axes); ++i) { if (out_ax != axes[i] && axes[i] >= 0) { indices[i - 1] = moveaxis(indices[i - 1], axes[i], out_ax, stream()); } else if (axes[i] < 0) { @@ -2515,7 +2516,7 @@ std::pair, std::vector> Greater::vmap( std::vector Greater::vjp( const std::vector& primals, - const std::vector& cotangents, + const std::vector& /* cotangents */, const std::vector& argnums, const std::vector&) { std::vector vjps; @@ -2527,8 +2528,8 @@ std::vector Greater::vjp( std::vector Greater::jvp( const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* tangents */, + const std::vector& /* argnums */) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); return {zeros(shape, bool_, stream())}; } @@ -2542,7 +2543,7 @@ std::pair, std::vector> GreaterEqual::vmap( std::vector GreaterEqual::vjp( const std::vector& primals, - const std::vector& cotangents, + const std::vector& /* cotangents */, const std::vector& argnums, const std::vector&) { std::vector vjps; @@ -2554,8 +2555,8 @@ std::vector GreaterEqual::vjp( std::vector GreaterEqual::jvp( const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* tangents */, + const std::vector& /* argnums */) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); return {zeros(shape, bool_, stream())}; } @@ -2599,7 +2600,7 @@ std::pair, std::vector> Less::vmap( std::vector Less::vjp( const std::vector& primals, - const std::vector& cotangents, + const std::vector& /* cotangents */, const std::vector& argnums, const std::vector&) { std::vector vjps; @@ -2611,8 +2612,8 @@ std::vector Less::vjp( std::vector Less::jvp( const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* tangents */, + const std::vector& /* argnums */) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); return {zeros(shape, bool_, stream())}; } @@ -2626,7 +2627,7 @@ std::pair, std::vector> LessEqual::vmap( std::vector LessEqual::vjp( const std::vector& primals, - const std::vector& cotangents, + const std::vector& /* cotangents */, const std::vector& argnums, const std::vector&) { std::vector vjps; @@ -2638,8 +2639,8 @@ std::vector LessEqual::vjp( std::vector LessEqual::jvp( const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* tangents */, + const std::vector& /* argnums */) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); return {zeros(shape, bool_, stream())}; } @@ -2748,7 +2749,7 @@ std::vector LogicalAnd::vjp( std::vector LogicalAnd::jvp( const std::vector& primals, - const std::vector& tangents, + const std::vector& /* tangents */, const std::vector& argnums) { assert(primals.size() == 2); assert(argnums.size() <= 2); @@ -2780,7 +2781,7 @@ std::vector LogicalOr::vjp( std::vector LogicalOr::jvp( const std::vector& primals, - const std::vector& tangents, + const std::vector& /* tangents */, const std::vector& argnums) { assert(primals.size() == 2); assert(argnums.size() <= 2); @@ -2859,7 +2860,7 @@ std::pair, std::vector> LogSumExp::vmap( std::vector LogSumExp::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector&) { assert(primals.size() == 1); assert(cotangents.size() == 1); @@ -2872,7 +2873,7 @@ std::vector LogSumExp::vjp( std::vector LogSumExp::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { assert(primals.size() == 1); assert(tangents.size() == 1); return {multiply( @@ -2920,7 +2921,7 @@ std::vector Matmul::jvp( const std::vector& tangents, const std::vector& argnums) { std::vector jvp; - for (int i = 0; i < argnums.size(); ++i) { + for (int i = 0; i < std::ssize(argnums); ++i) { auto arg = argnums[i]; if (arg == 0 && i == 0) { jvp.push_back(matmul(tangents[0], primals[1], stream())); @@ -3096,7 +3097,7 @@ std::vector Select::jvp( }; array jvp = jvp_fun(argnums[0]); - for (int i = 1; i < argnums.size(); i++) { + for (int i = 1; i < std::ssize(argnums); i++) { jvp = add(jvp, jvp_fun(argnums[i])); } return {jvp}; @@ -3173,7 +3174,7 @@ std::pair, std::vector> NotEqual::vmap( std::vector NotEqual::vjp( const std::vector& primals, - const std::vector& cotangents, + const std::vector& /* cotangents */, const std::vector& argnums, const std::vector&) { std::vector vjps; @@ -3185,14 +3186,14 @@ std::vector NotEqual::vjp( std::vector NotEqual::jvp( const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* tangents */, + const std::vector& /* argnums */) { auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape()); return {zeros(shape, bool_, stream())}; } std::vector Pad::vjp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& cotangents, const std::vector& argnums, const std::vector&) { @@ -3213,7 +3214,7 @@ std::vector Pad::vjp( } std::vector Pad::jvp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& tangents, const std::vector& argnums) { assert(argnums.size() == 1 && argnums[0] == 0); @@ -3229,8 +3230,8 @@ std::vector Pad::jvp( } std::pair, std::vector> Pad::vmap( - const std::vector& inputs, - const std::vector& axes) { + const std::vector& /* inputs */, + const std::vector& /* axes */) { throw std::runtime_error("Pad vmap is NYI."); } @@ -3244,7 +3245,7 @@ bool Pad::is_equivalent(const Primitive& other) const { std::vector Partition::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector&) { auto sort_idx = argpartition(primals[0], kth_, axis_, stream()); return {put_along_axis( @@ -3258,7 +3259,7 @@ std::vector Partition::vjp( std::vector Partition::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { assert(primals.size() == 1); assert(tangents.size() == 1); auto sort_idx = argpartition(primals[0], kth_, axis_, stream()); @@ -3344,8 +3345,8 @@ QuantizationMode string_to_quantization_mode(const std::string& mode) { } std::pair, std::vector> QuantizedMatmul::vmap( - const std::vector& inputs, - const std::vector& axes) { + const std::vector& /* inputs */, + const std::vector& /* axes */) { throw std::runtime_error("[QuantizedMatmul::vmap] NYI"); } @@ -3450,8 +3451,8 @@ std::vector QuantizedMatmul::output_shapes( } std::pair, std::vector> GatherQMM::vmap( - const std::vector& inputs, - const std::vector& axes) { + const std::vector& /* inputs */, + const std::vector& /* axes */) { throw std::runtime_error("GatherQMM::vmap NYI"); } @@ -3573,9 +3574,9 @@ std::vector GatherQMM::vjp( } std::vector GatherQMM::jvp( - const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* primals */, + const std::vector& /* tangents */, + const std::vector& /* argnums */) { throw std::runtime_error("GatherQMM::jvp NYI"); } @@ -3706,7 +3707,7 @@ bool Reshape::is_equivalent(const Primitive& other) const { } Shape Reshape::output_shape(const array& input, Shape shape) { - size_t size = 1; + int64_t size = 1; int infer_idx = -1; for (int i = 0; i < shape.size(); ++i) { if (shape[i] == -1) { @@ -3730,7 +3731,7 @@ Shape Reshape::output_shape(const array& input, Shape shape) { } // Check that the reshaping is valid - if (input.size() != size) { + if (std::ssize(input) != size) { std::ostringstream msg; msg << "[reshape] Cannot reshape array of size " << input.size() << " into shape " << shape << "."; @@ -3746,7 +3747,7 @@ std::vector Reshape::output_shapes(const std::vector& inputs) { std::vector Reduce::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector& outputs) { auto in = primals[0]; @@ -3776,7 +3777,7 @@ std::vector Reduce::vjp( // except the reduced over axes. int j = 0; for (int i = 0; i < in.ndim(); i++) { - if (j < axes_.size() && axes_[j] == i) { + if (j < std::ssize(axes_) && axes_[j] == i) { j++; } else { transpose_to.push_back(i); @@ -3788,7 +3789,7 @@ std::vector Reduce::vjp( } shape_flat.push_back(-1); transpose_back.resize(transpose_to.size()); - for (int i = 0; i < transpose_to.size(); i++) { + for (int i = 0; i < std::ssize(transpose_to); i++) { transpose_back[transpose_to[i]] = i; } } @@ -3886,7 +3887,7 @@ std::vector Round::vjp( std::vector Round::jvp( const std::vector& primals, - const std::vector& tangents, + const std::vector& /* tangents */, const std::vector& argnums) { assert(primals.size() == 1); assert(argnums.size() == 1); @@ -4021,7 +4022,7 @@ std::vector Scan::vjp( } std::vector Scan::jvp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& tangents, const std::vector& argnums) { assert(tangents.size() == 1); @@ -4099,7 +4100,7 @@ std::vector Scatter::vjp( // Should never reach here throw std::invalid_argument(""); } - } else if (num == primals.size() - 1) { + } else if (num == std::ssize(primals) - 1) { switch (reduce_type_) { case Scatter::None: case Scatter::Sum: { @@ -4140,9 +4141,9 @@ std::vector Scatter::vjp( } std::vector Scatter::jvp( - const std::vector& primals, - const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* primals */, + const std::vector& /* tangents */, + const std::vector& /* argnums */) { throw std::runtime_error("[scatter] JVP not yet implemented"); } @@ -4167,7 +4168,7 @@ std::pair, std::vector> Scatter::vmap( inputs[0] = repeat(expand_dims(inputs[0], 0, stream()), vmap_size, 0, stream()); } - for (int i = 1; i < vmap_axes.size() - 1; ++i) { + for (int i = 1; i < std::ssize(vmap_axes) - 1; ++i) { // vmap axis for indices goes to 0 if (vmap_axes[i] >= 0) { inputs[i] = moveaxis(inputs[i], vmap_axes[i], 0, stream()); @@ -4303,7 +4304,7 @@ std::pair, std::vector> ScatterAxis::vmap( } auto v_in = inputs; - for (int i = 0; i < axes.size(); ++i) { + for (int i = 0; i < std::ssize(axes); ++i) { if (axes[i] >= 0) { // if out_ax >= 0 move axis o/w set out_ax if (out_ax != axes[i]) { @@ -4329,9 +4330,9 @@ bool ScatterAxis::is_equivalent(const Primitive& other) const { } std::vector Sigmoid::vjp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector& outputs) { auto& s = outputs[0]; auto sprime = @@ -4369,7 +4370,7 @@ std::vector Sign::vjp( std::vector Sign::jvp( const std::vector& primals, - const std::vector& tangents, + const std::vector& /* tangents */, const std::vector& argnums) { assert(primals.size() == 1); assert(argnums.size() == 1); @@ -4453,7 +4454,7 @@ std::pair, std::vector> Slice::vmap( std::vector Slice::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector&) { // Check inputs assert(primals.size() == 1); @@ -4465,7 +4466,7 @@ std::vector Slice::vjp( std::vector Slice::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { // Check inputs assert(primals.size() == 1); return {slice(tangents[0], start_indices_, end_indices_, strides_, stream())}; @@ -4562,7 +4563,7 @@ std::vector SliceUpdate::vjp( std::vector SliceUpdate::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { // Check inputs assert(primals.size() == 2); return {slice_update( @@ -4743,7 +4744,7 @@ std::pair, std::vector> Softmax::vmap( std::vector Softmax::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector& outputs) { assert(primals.size() == 1); assert(cotangents.size() == 1); @@ -4757,7 +4758,7 @@ std::vector Softmax::vjp( std::vector Softmax::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { assert(primals.size() == 1); assert(tangents.size() == 1); auto s = softmax(primals[0], std::vector{-1}, precise_, stream()); @@ -4793,7 +4794,7 @@ std::vector Sort::vjp( std::vector Sort::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { assert(primals.size() == 1); assert(tangents.size() == 1); auto sort_idx = argsort(primals[0], axis_, stream()); @@ -4816,17 +4817,17 @@ std::pair, std::vector> Split::vmap( } std::vector Split::vjp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector&) { return {concatenate(cotangents, axis_, stream())}; } std::vector Split::jvp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { return split(tangents[0], indices_, axis_, stream()); } @@ -4846,7 +4847,7 @@ std::vector Square::vjp( std::vector Square::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { assert(primals.size() == 1); assert(tangents.size() == 1); return {multiply( @@ -4866,7 +4867,7 @@ std::pair, std::vector> Square::vmap( std::vector Sqrt::vjp( const std::vector& primals, const std::vector& cotangents, - const std::vector& argnums, + const std::vector& /* argnums */, const std::vector& outputs) { assert(primals.size() == 1); assert(cotangents.size() == 1); @@ -4919,7 +4920,7 @@ std::pair, std::vector> StopGradient::vmap( } std::vector Subtract::vjp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& cotangents, const std::vector& argnums, const std::vector&) { @@ -4935,7 +4936,7 @@ std::vector Subtract::vjp( } std::vector Subtract::jvp( - const std::vector& primals, + const std::vector& /* primals */, const std::vector& tangents, const std::vector& argnums) { auto jvp_fun = [&](int i) { @@ -4994,7 +4995,7 @@ bool Squeeze::is_equivalent(const Primitive& other) const { Shape Squeeze::output_shape(const array& input, const std::vector& axes) { Shape shape; for (int i = 0, j = 0; i < input.ndim(); ++i) { - if (j < axes.size() && i == axes[j]) { + if (j < std::ssize(axes) && i == axes[j]) { j++; } else { shape.push_back(input.shape(i)); @@ -5409,7 +5410,7 @@ std::vector Transpose::vjp( assert(primals.size() == 1); assert(argnums.size() == 1); std::vector iaxes(axes_.size()); - for (int i = 0; i < axes_.size(); ++i) { + for (int i = 0; i < std::ssize(axes_); ++i) { iaxes[axes_[i]] = i; } return {transpose(cotangents[0], iaxes, stream())}; @@ -5418,7 +5419,7 @@ std::vector Transpose::vjp( std::vector Transpose::jvp( const std::vector& primals, const std::vector& tangents, - const std::vector& argnums) { + const std::vector& /* argnums */) { assert(primals.size() == 1); assert(tangents.size() == 1); return {transpose(tangents[0], axes_, stream())}; @@ -5449,7 +5450,7 @@ bool Transpose::is_equivalent(const Primitive& other) const { std::vector Transpose::output_shapes(const std::vector& inputs) { auto& in = inputs[0]; Shape shape(in.ndim(), 0); - for (int i = 0; i < axes_.size(); ++i) { + for (int i = 0; i < std::ssize(axes_); ++i) { shape[i] = in.shape()[axes_[i]]; } return {std::move(shape)}; diff --git a/mlx/primitives.h b/mlx/primitives.h index 2a843a0e..a37124db 100644 --- a/mlx/primitives.h +++ b/mlx/primitives.h @@ -31,9 +31,9 @@ return #PRIMITIVE; \ } -#define DEFINE_DEFAULT_IS_EQUIVALENT() \ - bool is_equivalent(const Primitive& other) const override { \ - return true; \ +#define DEFINE_DEFAULT_IS_EQUIVALENT() \ + bool is_equivalent(const Primitive& /* other */) const override { \ + return true; \ } #define DEFINE_INPUT_OUTPUT_SHAPE() \ @@ -104,7 +104,7 @@ class Primitive { virtual const char* name() const = 0; /** Equivalence check defaults to false unless overridden by the primitive */ - virtual bool is_equivalent(const Primitive& other) const { + virtual bool is_equivalent(const Primitive& /* other */) const { return false; } @@ -1071,7 +1071,7 @@ class FFT : public UnaryPrimitive { public: explicit FFT( Stream stream, - const std::vector& axes, + const std::vector& axes, bool inverse, bool real) : UnaryPrimitive(stream), axes_(axes), inverse_(inverse), real_(real) {} @@ -1089,7 +1089,7 @@ class FFT : public UnaryPrimitive { } private: - std::vector axes_; + std::vector axes_; bool inverse_; bool real_; }; @@ -1526,7 +1526,8 @@ class NumberOfElements : public UnaryPrimitive { DEFINE_VMAP() DEFINE_NAME(NumberOfElements) bool is_equivalent(const Primitive& other) const override; - std::vector output_shapes(const std::vector& inputs) override { + std::vector output_shapes( + const std::vector& /* inputs */) override { return {{}}; } std::tuple, bool, Dtype> state() const { diff --git a/mlx/random.h b/mlx/random.h index 0dfdab7a..c707889b 100644 --- a/mlx/random.h +++ b/mlx/random.h @@ -89,6 +89,7 @@ inline array uniform( const Shape& shape, const std::optional& key = std::nullopt, StreamOrDevice s = {}) { + (void)s; return uniform(shape, float32, key); } diff --git a/mlx/scheduler.h b/mlx/scheduler.h index 877fdd5f..65286bd6 100644 --- a/mlx/scheduler.h +++ b/mlx/scheduler.h @@ -103,7 +103,7 @@ class Scheduler { default_streams_.at(s.device.type) = s; } - void notify_new_task(const Stream& stream) { + void notify_new_task(const Stream& /* stream */) { { std::lock_guard lk(mtx); n_active_tasks_++; @@ -111,7 +111,7 @@ class Scheduler { completion_cv.notify_all(); } - void notify_task_completion(const Stream& stream) { + void notify_task_completion(const Stream& /* stream */) { { std::lock_guard lk(mtx); n_active_tasks_--; diff --git a/mlx/types/bf16.h b/mlx/types/bf16.h index 59519417..3e1f9d9a 100644 --- a/mlx/types/bf16.h +++ b/mlx/types/bf16.h @@ -24,9 +24,6 @@ struct _MLX_BFloat16 { // Default constructor _MLX_BFloat16() = default; - // Default copy constructor - _MLX_BFloat16(_MLX_BFloat16 const&) = default; - // Appease std::vector for being special _MLX_BFloat16& operator=(std::vector::reference x) { bits_ = x; diff --git a/mlx/utils.h b/mlx/utils.h index dbf79a71..a71e92a3 100644 --- a/mlx/utils.h +++ b/mlx/utils.h @@ -138,8 +138,8 @@ namespace env { int get_var(const char* name, int default_value); -inline int bfs_max_width() { - static int bfs_max_width_ = get_var("MLX_BFS_MAX_WIDTH", 20); +inline unsigned int bfs_max_width() { + static unsigned int bfs_max_width_ = get_var("MLX_BFS_MAX_WIDTH", 20); return bfs_max_width_; }