WIP
This commit is contained in:
+15
-15
@@ -194,7 +194,7 @@ const char* Compiled::name() const {
|
||||
}
|
||||
|
||||
std::vector<Shape> Compiled::output_shapes(const std::vector<array>& 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<array> 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<std::vector<array>, ParentsMap> compile_dfs(
|
||||
std::unordered_set<std::uintptr_t> original_input_set;
|
||||
std::unordered_map<std::uintptr_t, std::vector<std::pair<array, int>>>
|
||||
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<std::vector<array>, 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<bool> 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<array> 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<array> compile_replace(
|
||||
const std::vector<array>& inputs,
|
||||
bool shapeless) {
|
||||
std::unordered_map<uintptr_t, array> 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<array> 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])});
|
||||
}
|
||||
}
|
||||
|
||||
+23
-23
@@ -190,8 +190,8 @@ std::tuple<std::vector<PathNode>, size_t, int> greedy_path(
|
||||
|
||||
// Start by iterating over all possible combinations
|
||||
std::vector<std::pair<int, int>> 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<std::vector<PathNode>, size_t, int> greedy_path(
|
||||
std::vector<Contraction> 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<std::vector<PathNode>, 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<char, int> 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<array> 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<array>& 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<int> 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<int> sum_axes;
|
||||
@@ -675,9 +675,9 @@ std::pair<std::vector<PathNode>, 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<std::vector<PathNode>, 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<std::vector<PathNode>, PathInfo> einsum_path_helper(
|
||||
|
||||
std::unordered_map<char, ShapeElem> dim_map;
|
||||
std::vector<Subscript> 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<std::vector<PathNode>, PathInfo> einsum_path_helper(
|
||||
// Check repeat subscripts are valid
|
||||
if (in_set.size() < in.size()) {
|
||||
std::unordered_map<char, ShapeElem> 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<std::vector<PathNode>, 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<int> a_contract;
|
||||
std::vector<int> a_batch;
|
||||
std::vector<int> 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()) {
|
||||
|
||||
+10
-9
@@ -138,7 +138,7 @@ T deserialize(Reader& is) {
|
||||
T v;
|
||||
auto size = deserialize<uint64_t>(is);
|
||||
v.reserve(size);
|
||||
for (int i = 0; i < size; ++i) {
|
||||
for (size_t i = 0; i < size; ++i) {
|
||||
v.push_back(deserialize<typename T::value_type>(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<std::pair<std::string, std::string>> 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<array> 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<uint64_t, array> array_map;
|
||||
auto trace_input_ids = deserialize<std::vector<uint64_t>>(is);
|
||||
auto trace_inputs = deserialize<std::vector<array>>(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<std::vector<uint64_t>>(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]);
|
||||
|
||||
+7
-7
@@ -13,11 +13,11 @@ std::vector<array> Custom::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>& outputs) {
|
||||
const std::vector<array>& /* outputs */) {
|
||||
auto [_, vjps] = mlx::core::vjp(fallback_, primals, cotangents);
|
||||
std::vector<array> 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<array> Custom::jvp(
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
std::vector<array> 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<array> RoPE::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>& outputs) {
|
||||
const std::vector<array>& /* 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;
|
||||
|
||||
+3
-3
@@ -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<size_t> valid_axes;
|
||||
std::vector<int> 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) {
|
||||
|
||||
+26
-26
@@ -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<array> meshgrid(
|
||||
"[meshgrid] Invalid indexing value. Valid values are 'xy' and 'ij'.");
|
||||
}
|
||||
|
||||
auto ndim = arrays.size();
|
||||
auto ndim = std::ssize(arrays);
|
||||
std::vector<array> outputs;
|
||||
for (int i = 0; i < ndim; ++i) {
|
||||
Shape shape(ndim, 1);
|
||||
@@ -1135,10 +1135,10 @@ array tile(
|
||||
std::vector<int> 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<int>& axes,
|
||||
const std::vector<int>& /* 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<array> 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<array> 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<array> broadcast_arrays(
|
||||
} else {
|
||||
// broadcasted array goes first followed by other stopgrad inputs
|
||||
std::vector<array> 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<array> 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<array> 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<int> 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<int> padding_lo(padding.size());
|
||||
std::vector<int> 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();
|
||||
|
||||
+111
-110
@@ -242,16 +242,16 @@ std::pair<std::vector<array>, std::vector<int>> Abs::vmap(
|
||||
}
|
||||
|
||||
std::vector<array> Add::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
return {
|
||||
tangents.size() > 1 ? add(tangents[0], tangents[1], stream())
|
||||
: tangents[0]};
|
||||
}
|
||||
|
||||
std::vector<array> Add::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
@@ -315,7 +315,7 @@ std::vector<array> AddMM::jvp(
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
std::vector<array> 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<array> ArgSort::jvp(
|
||||
std::vector<array> AsType::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>&) {
|
||||
if (cotangents[0].dtype() != dtype_) {
|
||||
throw std::invalid_argument(
|
||||
@@ -702,9 +702,9 @@ std::vector<array> AsType::vjp(
|
||||
}
|
||||
|
||||
std::vector<array> AsType::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
return {astype(tangents[0], dtype_, stream())};
|
||||
}
|
||||
|
||||
@@ -752,7 +752,7 @@ std::vector<array> AsStrided::vjp(
|
||||
std::vector<array> AsStrided::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
assert(primals.size() == 1);
|
||||
|
||||
return {as_strided(tangents[0], shape_, strides_, offset_, stream())};
|
||||
@@ -827,9 +827,9 @@ std::vector<array> Broadcast::vjp(
|
||||
}
|
||||
|
||||
std::vector<array> Broadcast::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* 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<array>& 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<array> BroadcastAxes::vjp(
|
||||
std::vector<array> BroadcastAxes::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
return {array(
|
||||
output_shape(primals, ignore_axes_),
|
||||
tangents[0].dtype(),
|
||||
@@ -895,8 +895,8 @@ std::vector<array> BroadcastAxes::jvp(
|
||||
}
|
||||
|
||||
std::pair<std::vector<array>, std::vector<int>> BroadcastAxes::vmap(
|
||||
const std::vector<array>& inputs,
|
||||
const std::vector<int>& axes) {
|
||||
const std::vector<array>& /* inputs */,
|
||||
const std::vector<int>& /* axes */) {
|
||||
throw std::invalid_argument("[BroadcastAxes] VMAP NYI");
|
||||
}
|
||||
|
||||
@@ -938,7 +938,7 @@ std::vector<array> Ceil::vjp(
|
||||
|
||||
std::vector<array> Ceil::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(primals.size() == 1);
|
||||
assert(argnums.size() == 1);
|
||||
@@ -1072,8 +1072,8 @@ std::vector<array> Concatenate::jvp(
|
||||
});
|
||||
|
||||
std::vector<array> 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<array>, std::vector<int>> 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<array>, std::vector<int>> Concatenate::vmap(
|
||||
std::vector<array> 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<Shape> Concatenate::output_shapes(
|
||||
const std::vector<array>& 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<array> Convolution::vjp(
|
||||
std::vector<int> padding_lo = padding_lo_;
|
||||
std::vector<int> 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<array> 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<array> 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<array> Depends::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>& outputs) {
|
||||
const std::vector<array>& /* outputs */) {
|
||||
std::vector<array> 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<array> Divide::vjp(
|
||||
|
||||
std::vector<array> DivMod::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<array>& /* cotangents */,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
std::vector<array> vjps;
|
||||
@@ -1749,8 +1750,8 @@ std::vector<array> DivMod::vjp(
|
||||
|
||||
std::vector<array> DivMod::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
return {zeros_like(primals[0], stream())};
|
||||
}
|
||||
|
||||
@@ -1848,7 +1849,7 @@ std::pair<std::vector<array>, std::vector<int>> Equal::vmap(
|
||||
|
||||
std::vector<array> Equal::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<array>& /* cotangents */,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
std::vector<array> vjps;
|
||||
@@ -1860,8 +1861,8 @@ std::vector<array> Equal::vjp(
|
||||
|
||||
std::vector<array> Equal::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape());
|
||||
return {zeros(shape, bool_, stream())};
|
||||
}
|
||||
@@ -1899,7 +1900,7 @@ std::pair<std::vector<array>, std::vector<int>> Erf::vmap(
|
||||
std::vector<array> ErfInv::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>& outputs) {
|
||||
auto dtype = primals[0].dtype();
|
||||
auto scale =
|
||||
@@ -1931,9 +1932,9 @@ std::pair<std::vector<array>, std::vector<int>> ErfInv::vmap(
|
||||
}
|
||||
|
||||
std::vector<array> Exp::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>& outputs) {
|
||||
return {multiply(cotangents[0], outputs[0], stream())};
|
||||
}
|
||||
@@ -1956,9 +1957,9 @@ std::pair<std::vector<array>, std::vector<int>> Exp::vmap(
|
||||
}
|
||||
|
||||
std::vector<array> Expm1::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>& outputs) {
|
||||
return {multiply(
|
||||
cotangents[0],
|
||||
@@ -2281,7 +2282,7 @@ std::vector<array> Floor::vjp(
|
||||
|
||||
std::vector<array> Floor::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(primals.size() == 1);
|
||||
assert(argnums.size() == 1);
|
||||
@@ -2346,7 +2347,7 @@ std::pair<std::vector<array>, std::vector<int>> 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<array>, std::vector<int>> Greater::vmap(
|
||||
|
||||
std::vector<array> Greater::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<array>& /* cotangents */,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
std::vector<array> vjps;
|
||||
@@ -2527,8 +2528,8 @@ std::vector<array> Greater::vjp(
|
||||
|
||||
std::vector<array> Greater::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape());
|
||||
return {zeros(shape, bool_, stream())};
|
||||
}
|
||||
@@ -2542,7 +2543,7 @@ std::pair<std::vector<array>, std::vector<int>> GreaterEqual::vmap(
|
||||
|
||||
std::vector<array> GreaterEqual::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<array>& /* cotangents */,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
std::vector<array> vjps;
|
||||
@@ -2554,8 +2555,8 @@ std::vector<array> GreaterEqual::vjp(
|
||||
|
||||
std::vector<array> GreaterEqual::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape());
|
||||
return {zeros(shape, bool_, stream())};
|
||||
}
|
||||
@@ -2599,7 +2600,7 @@ std::pair<std::vector<array>, std::vector<int>> Less::vmap(
|
||||
|
||||
std::vector<array> Less::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<array>& /* cotangents */,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
std::vector<array> vjps;
|
||||
@@ -2611,8 +2612,8 @@ std::vector<array> Less::vjp(
|
||||
|
||||
std::vector<array> Less::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape());
|
||||
return {zeros(shape, bool_, stream())};
|
||||
}
|
||||
@@ -2626,7 +2627,7 @@ std::pair<std::vector<array>, std::vector<int>> LessEqual::vmap(
|
||||
|
||||
std::vector<array> LessEqual::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<array>& /* cotangents */,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
std::vector<array> vjps;
|
||||
@@ -2638,8 +2639,8 @@ std::vector<array> LessEqual::vjp(
|
||||
|
||||
std::vector<array> LessEqual::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape());
|
||||
return {zeros(shape, bool_, stream())};
|
||||
}
|
||||
@@ -2748,7 +2749,7 @@ std::vector<array> LogicalAnd::vjp(
|
||||
|
||||
std::vector<array> LogicalAnd::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(primals.size() == 2);
|
||||
assert(argnums.size() <= 2);
|
||||
@@ -2780,7 +2781,7 @@ std::vector<array> LogicalOr::vjp(
|
||||
|
||||
std::vector<array> LogicalOr::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(primals.size() == 2);
|
||||
assert(argnums.size() <= 2);
|
||||
@@ -2859,7 +2860,7 @@ std::pair<std::vector<array>, std::vector<int>> LogSumExp::vmap(
|
||||
std::vector<array> LogSumExp::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>&) {
|
||||
assert(primals.size() == 1);
|
||||
assert(cotangents.size() == 1);
|
||||
@@ -2872,7 +2873,7 @@ std::vector<array> LogSumExp::vjp(
|
||||
std::vector<array> LogSumExp::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
assert(primals.size() == 1);
|
||||
assert(tangents.size() == 1);
|
||||
return {multiply(
|
||||
@@ -2920,7 +2921,7 @@ std::vector<array> Matmul::jvp(
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
std::vector<array> 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<array> 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<array>, std::vector<int>> NotEqual::vmap(
|
||||
|
||||
std::vector<array> NotEqual::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<array>& /* cotangents */,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
std::vector<array> vjps;
|
||||
@@ -3185,14 +3186,14 @@ std::vector<array> NotEqual::vjp(
|
||||
|
||||
std::vector<array> NotEqual::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
auto shape = broadcast_shapes(primals[0].shape(), primals[1].shape());
|
||||
return {zeros(shape, bool_, stream())};
|
||||
}
|
||||
|
||||
std::vector<array> Pad::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
@@ -3213,7 +3214,7 @@ std::vector<array> Pad::vjp(
|
||||
}
|
||||
|
||||
std::vector<array> Pad::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(argnums.size() == 1 && argnums[0] == 0);
|
||||
@@ -3229,8 +3230,8 @@ std::vector<array> Pad::jvp(
|
||||
}
|
||||
|
||||
std::pair<std::vector<array>, std::vector<int>> Pad::vmap(
|
||||
const std::vector<array>& inputs,
|
||||
const std::vector<int>& axes) {
|
||||
const std::vector<array>& /* inputs */,
|
||||
const std::vector<int>& /* axes */) {
|
||||
throw std::runtime_error("Pad vmap is NYI.");
|
||||
}
|
||||
|
||||
@@ -3244,7 +3245,7 @@ bool Pad::is_equivalent(const Primitive& other) const {
|
||||
std::vector<array> Partition::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>&) {
|
||||
auto sort_idx = argpartition(primals[0], kth_, axis_, stream());
|
||||
return {put_along_axis(
|
||||
@@ -3258,7 +3259,7 @@ std::vector<array> Partition::vjp(
|
||||
std::vector<array> Partition::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* 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<array>, std::vector<int>> QuantizedMatmul::vmap(
|
||||
const std::vector<array>& inputs,
|
||||
const std::vector<int>& axes) {
|
||||
const std::vector<array>& /* inputs */,
|
||||
const std::vector<int>& /* axes */) {
|
||||
throw std::runtime_error("[QuantizedMatmul::vmap] NYI");
|
||||
}
|
||||
|
||||
@@ -3450,8 +3451,8 @@ std::vector<Shape> QuantizedMatmul::output_shapes(
|
||||
}
|
||||
|
||||
std::pair<std::vector<array>, std::vector<int>> GatherQMM::vmap(
|
||||
const std::vector<array>& inputs,
|
||||
const std::vector<int>& axes) {
|
||||
const std::vector<array>& /* inputs */,
|
||||
const std::vector<int>& /* axes */) {
|
||||
throw std::runtime_error("GatherQMM::vmap NYI");
|
||||
}
|
||||
|
||||
@@ -3573,9 +3574,9 @@ std::vector<array> GatherQMM::vjp(
|
||||
}
|
||||
|
||||
std::vector<array> GatherQMM::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* 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<Shape> Reshape::output_shapes(const std::vector<array>& inputs) {
|
||||
std::vector<array> Reduce::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>& outputs) {
|
||||
auto in = primals[0];
|
||||
|
||||
@@ -3776,7 +3777,7 @@ std::vector<array> 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<array> 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<array> Round::vjp(
|
||||
|
||||
std::vector<array> Round::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(primals.size() == 1);
|
||||
assert(argnums.size() == 1);
|
||||
@@ -4021,7 +4022,7 @@ std::vector<array> Scan::vjp(
|
||||
}
|
||||
|
||||
std::vector<array> Scan::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(tangents.size() == 1);
|
||||
@@ -4099,7 +4100,7 @@ std::vector<array> 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<array> Scatter::vjp(
|
||||
}
|
||||
|
||||
std::vector<array> Scatter::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& /* argnums */) {
|
||||
throw std::runtime_error("[scatter] JVP not yet implemented");
|
||||
}
|
||||
|
||||
@@ -4167,7 +4168,7 @@ std::pair<std::vector<array>, std::vector<int>> 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<array>, std::vector<int>> 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<array> Sigmoid::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>& outputs) {
|
||||
auto& s = outputs[0];
|
||||
auto sprime =
|
||||
@@ -4369,7 +4370,7 @@ std::vector<array> Sign::vjp(
|
||||
|
||||
std::vector<array> Sign::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<array>& /* tangents */,
|
||||
const std::vector<int>& argnums) {
|
||||
assert(primals.size() == 1);
|
||||
assert(argnums.size() == 1);
|
||||
@@ -4453,7 +4454,7 @@ std::pair<std::vector<array>, std::vector<int>> Slice::vmap(
|
||||
std::vector<array> Slice::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>&) {
|
||||
// Check inputs
|
||||
assert(primals.size() == 1);
|
||||
@@ -4465,7 +4466,7 @@ std::vector<array> Slice::vjp(
|
||||
std::vector<array> Slice::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
// Check inputs
|
||||
assert(primals.size() == 1);
|
||||
return {slice(tangents[0], start_indices_, end_indices_, strides_, stream())};
|
||||
@@ -4562,7 +4563,7 @@ std::vector<array> SliceUpdate::vjp(
|
||||
std::vector<array> SliceUpdate::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
// Check inputs
|
||||
assert(primals.size() == 2);
|
||||
return {slice_update(
|
||||
@@ -4743,7 +4744,7 @@ std::pair<std::vector<array>, std::vector<int>> Softmax::vmap(
|
||||
std::vector<array> Softmax::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>& outputs) {
|
||||
assert(primals.size() == 1);
|
||||
assert(cotangents.size() == 1);
|
||||
@@ -4757,7 +4758,7 @@ std::vector<array> Softmax::vjp(
|
||||
std::vector<array> Softmax::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
assert(primals.size() == 1);
|
||||
assert(tangents.size() == 1);
|
||||
auto s = softmax(primals[0], std::vector<int>{-1}, precise_, stream());
|
||||
@@ -4793,7 +4794,7 @@ std::vector<array> Sort::vjp(
|
||||
std::vector<array> Sort::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* 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<array>, std::vector<int>> Split::vmap(
|
||||
}
|
||||
|
||||
std::vector<array> Split::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>&) {
|
||||
return {concatenate(cotangents, axis_, stream())};
|
||||
}
|
||||
|
||||
std::vector<array> Split::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
return split(tangents[0], indices_, axis_, stream());
|
||||
}
|
||||
|
||||
@@ -4846,7 +4847,7 @@ std::vector<array> Square::vjp(
|
||||
std::vector<array> Square::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* argnums */) {
|
||||
assert(primals.size() == 1);
|
||||
assert(tangents.size() == 1);
|
||||
return {multiply(
|
||||
@@ -4866,7 +4867,7 @@ std::pair<std::vector<array>, std::vector<int>> Square::vmap(
|
||||
std::vector<array> Sqrt::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<int>& /* argnums */,
|
||||
const std::vector<array>& outputs) {
|
||||
assert(primals.size() == 1);
|
||||
assert(cotangents.size() == 1);
|
||||
@@ -4919,7 +4920,7 @@ std::pair<std::vector<array>, std::vector<int>> StopGradient::vmap(
|
||||
}
|
||||
|
||||
std::vector<array> Subtract::vjp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& cotangents,
|
||||
const std::vector<int>& argnums,
|
||||
const std::vector<array>&) {
|
||||
@@ -4935,7 +4936,7 @@ std::vector<array> Subtract::vjp(
|
||||
}
|
||||
|
||||
std::vector<array> Subtract::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& /* primals */,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& 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<int>& 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<array> Transpose::vjp(
|
||||
assert(primals.size() == 1);
|
||||
assert(argnums.size() == 1);
|
||||
std::vector<int> 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<array> Transpose::vjp(
|
||||
std::vector<array> Transpose::jvp(
|
||||
const std::vector<array>& primals,
|
||||
const std::vector<array>& tangents,
|
||||
const std::vector<int>& argnums) {
|
||||
const std::vector<int>& /* 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<Shape> Transpose::output_shapes(const std::vector<array>& 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)};
|
||||
|
||||
+8
-7
@@ -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<size_t>& axes,
|
||||
const std::vector<int>& axes,
|
||||
bool inverse,
|
||||
bool real)
|
||||
: UnaryPrimitive(stream), axes_(axes), inverse_(inverse), real_(real) {}
|
||||
@@ -1089,7 +1089,7 @@ class FFT : public UnaryPrimitive {
|
||||
}
|
||||
|
||||
private:
|
||||
std::vector<size_t> axes_;
|
||||
std::vector<int> 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<Shape> output_shapes(const std::vector<array>& inputs) override {
|
||||
std::vector<Shape> output_shapes(
|
||||
const std::vector<array>& /* inputs */) override {
|
||||
return {{}};
|
||||
}
|
||||
std::tuple<std::vector<int>, bool, Dtype> state() const {
|
||||
|
||||
@@ -89,6 +89,7 @@ inline array uniform(
|
||||
const Shape& shape,
|
||||
const std::optional<array>& key = std::nullopt,
|
||||
StreamOrDevice s = {}) {
|
||||
(void)s;
|
||||
return uniform(shape, float32, key);
|
||||
}
|
||||
|
||||
|
||||
+2
-2
@@ -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<std::mutex> 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<std::mutex> lk(mtx);
|
||||
n_active_tasks_--;
|
||||
|
||||
@@ -24,9 +24,6 @@ struct _MLX_BFloat16 {
|
||||
// Default constructor
|
||||
_MLX_BFloat16() = default;
|
||||
|
||||
// Default copy constructor
|
||||
_MLX_BFloat16(_MLX_BFloat16 const&) = default;
|
||||
|
||||
// Appease std::vector<bool> for being special
|
||||
_MLX_BFloat16& operator=(std::vector<bool>::reference x) {
|
||||
bits_ = x;
|
||||
|
||||
+2
-2
@@ -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_;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user