From b6d00f897cd8f52165596debbcf274c21f4aa3b4 Mon Sep 17 00:00:00 2001 From: Seema Mirchandaney Date: Thu, 9 Apr 2026 15:49:29 -0700 Subject: [PATCH 01/35] Add support for replicate op in distributed training - Add perform_pass_expansion_for_replicate for fwd/bwd pass expansion - Add perform_shard_expansion_for_replicate and _bwd for shard expansion - Add build_replicate_invocation in make_dynamic_open_dataflow_graph - Add is_replicate_attrs helper and guard replicate in copy_insertion - Add ReplicateAttrs to TrainingOperationAttrs - Add SumReductionFloat/Double for backward replicate reduce operation - Add issue_replicate_bwd in spawn_dynamic_node_invocation - Fix per_device_op_state init race condition with direct write - Fix .value() calls on optional per_device_op_state across op impls - Update issue_copy to support optional reduction op - Add testcase for replicate op --- .../include/realm-execution/sum_reduction.h | 99 ++++ .../realm-execution/tasks/realm_reduction.h | 96 ++++ .../src/realm-execution/test_op_replicate.cc | 450 ++++++++++++++++++ 3 files changed, 645 insertions(+) create mode 100644 lib/realm-execution/include/realm-execution/sum_reduction.h create mode 100644 lib/realm-execution/include/realm-execution/tasks/realm_reduction.h create mode 100644 lib/realm-execution/test/src/realm-execution/test_op_replicate.cc diff --git a/lib/realm-execution/include/realm-execution/sum_reduction.h b/lib/realm-execution/include/realm-execution/sum_reduction.h new file mode 100644 index 0000000000..b845b5b7f2 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/sum_reduction.h @@ -0,0 +1,99 @@ +#pragma once +#include +#include "op-attrs/datatype.dtg.h" + +namespace FlexFlow { + +// Sum reduction for float +struct SumReductionFloat { + using LHS = float; + using RHS = float; + static const RHS identity; + + template + static void apply(LHS &lhs, RHS rhs) { + if (EXCLUSIVE) { + lhs += rhs; + } else { + // atomic add for non-exclusive + __sync_fetch_and_add((int*)&lhs, *(int*)&rhs); + // proper float atomic add — use union trick + union { float f; int i; } old_val, new_val; + do { + old_val.f = lhs; + new_val.f = old_val.f + rhs; + } while (!__sync_bool_compare_and_swap( + (int*)&lhs, old_val.i, new_val.i)); + } + } + + template + static void fold(RHS &rhs1, RHS rhs2) { + if (EXCLUSIVE) { + rhs1 += rhs2; + } else { + union { float f; int i; } old_val, new_val; + do { + old_val.f = rhs1; + new_val.f = old_val.f + rhs2; + } while (!__sync_bool_compare_and_swap( + (int*)&rhs1, old_val.i, new_val.i)); + } + } +}; + +const SumReductionFloat::RHS SumReductionFloat::identity = 0.0f; + +// Sum reduction for double +struct SumReductionDouble { + using LHS = double; + using RHS = double; + static const RHS identity; + + template + static void apply(LHS &lhs, RHS rhs) { + if (EXCLUSIVE) { + lhs += rhs; + } else { + union { double d; long long i; } old_val, new_val; + do { + old_val.d = lhs; + new_val.d = old_val.d + rhs; + } while (!__sync_bool_compare_and_swap( + (long long*)&lhs, old_val.i, new_val.i)); + } + } + + template + static void fold(RHS &rhs1, RHS rhs2) { + if (EXCLUSIVE) { + rhs1 += rhs2; + } else { + union { double d; long long i; } old_val, new_val; + do { + old_val.d = rhs1; + new_val.d = old_val.d + rhs2; + } while (!__sync_bool_compare_and_swap( + (long long*)&rhs1, old_val.i, new_val.i)); + } + } +}; + +const SumReductionDouble::RHS SumReductionDouble::identity = 0.0; + +// Reduction op IDs — must not conflict with other registered redops +enum SumReductionOpIDs { + REDOP_SUM_FLOAT = 1, + REDOP_SUM_DOUBLE = 2, +}; + +inline Realm::ReductionOpID get_sum_reduction_op_id(DataType dtype) { + switch (dtype) { + case DataType::FLOAT: return REDOP_SUM_FLOAT; + case DataType::DOUBLE: return REDOP_SUM_DOUBLE; + default: + PANIC("no sum reduction registered for datatype {}", dtype); + } +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h b/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h new file mode 100644 index 0000000000..d1d6e1d880 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h @@ -0,0 +1,96 @@ +#pragma once +#include +#include "op-attrs/datatype.dtg.h" + +namespace FlexFlow { + +// Sum reduction for float +struct SumReductionFloat { + using LHS = float; + using RHS = float; + static constexpr RHS identity = 0.0f; // ← inside struct, constexpr + + template + static void apply(LHS &lhs, RHS rhs) { + if (EXCLUSIVE) { + lhs += rhs; + } else { + // atomic add for non-exclusive + __sync_fetch_and_add((int*)&lhs, *(int*)&rhs); + // proper float atomic add — use union trick + union { float f; int i; } old_val, new_val; + do { + old_val.f = lhs; + new_val.f = old_val.f + rhs; + } while (!__sync_bool_compare_and_swap( + (int*)&lhs, old_val.i, new_val.i)); + } + } + + template + static void fold(RHS &rhs1, RHS rhs2) { + if (EXCLUSIVE) { + rhs1 += rhs2; + } else { + union { float f; int i; } old_val, new_val; + do { + old_val.f = rhs1; + new_val.f = old_val.f + rhs2; + } while (!__sync_bool_compare_and_swap( + (int*)&rhs1, old_val.i, new_val.i)); + } + } +}; + + +// Sum reduction for double +struct SumReductionDouble { + using LHS = double; + using RHS = double; + static constexpr RHS identity = 0.0; // ← inside struct, constexpr + + template + static void apply(LHS &lhs, RHS rhs) { + if (EXCLUSIVE) { + lhs += rhs; + } else { + union { double d; long long i; } old_val, new_val; + do { + old_val.d = lhs; + new_val.d = old_val.d + rhs; + } while (!__sync_bool_compare_and_swap( + (long long*)&lhs, old_val.i, new_val.i)); + } + } + + template + static void fold(RHS &rhs1, RHS rhs2) { + if (EXCLUSIVE) { + rhs1 += rhs2; + } else { + union { double d; long long i; } old_val, new_val; + do { + old_val.d = rhs1; + new_val.d = old_val.d + rhs2; + } while (!__sync_bool_compare_and_swap( + (long long*)&rhs1, old_val.i, new_val.i)); + } + } +}; + +// Reduction op IDs — must not conflict with other registered redops +enum SumReductionOpIDs { + REDOP_SUM_FLOAT = 1, + REDOP_SUM_DOUBLE = 2, +}; + +inline Realm::ReductionOpID get_sum_reduction_op_id(DataType dtype) { + switch (dtype) { + case DataType::FLOAT: return REDOP_SUM_FLOAT; + case DataType::DOUBLE: return REDOP_SUM_DOUBLE; + default: + PANIC("no sum reduction registered for datatype {}", dtype); + } +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc new file mode 100644 index 0000000000..d1fc941007 --- /dev/null +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -0,0 +1,450 @@ +#include "internal/realm_test_utils.h" +#include "kernels/allocation.h" +#include "kernels/compare_tensor_accessors.h" +#include "kernels/copy_tensor_accessor.h" +#include "kernels/format_accessor_contents.h" +#include "kernels/tensor_accessor_reductions.h" +#include "op-attrs/operator_task_space_to_operator_task_space_mapping.h" +#include "op-attrs/ops/element_unary.h" +#include "op-attrs/ops/linear.h" +#include "op-attrs/ops/replicate.h" +#include "op-attrs/parallel_tensor_shape.h" +#include "op-attrs/tensor_shape.dtg.h" +#include "op-attrs/tensor_slot_name.dtg.h" +#include "pcg/device_type.dtg.h" +#include "pcg/machine_space_coordinate.dtg.h" +#include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.h" +#include "pcg/parallel_computation_graph/parallel_computation_graph.h" +#include "pcg/parallel_computation_graph/parallel_computation_graph_builder.h" +#include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" +#include "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h" +#include "realm-execution/distributed_ff_handle.h" +#include "realm-execution/dynamic_tensor_accessor_from_instance.h" +#include "realm-execution/pcg_instance.h" +#include "realm-execution/realm_context.h" +#include "realm-execution/realm_manager.h" +#include "task-spec/permissions.h" +#include "test/utils/doctest/check_kv.h" +#include "utils/containers/require_only_key.h" +#include + +namespace test { + +using namespace ::FlexFlow; +namespace Realm = ::FlexFlow::Realm; + +template +static ParallelLayerAttrs make_layer_attrs(T const &op_attrs) { + return ParallelLayerAttrs{ + /*op_attrs=*/PCGOperatorAttrs{op_attrs}, + /*name=*/std::nullopt, + }; +}; + +static bool did_loss_decrease(GenericTensorAccessorR const &first_epoch, + GenericTensorAccessorR const &last_epoch, + Allocator &allocator) { + return tensor_accessor_all( + compare_tensor_accessors_le(last_epoch, first_epoch, allocator)); +} + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("RealmBackend e2e Training Replicate Op (CPU Model Parallelism)") { + std::vector fake_args = + make_fake_realm_args(/*num_cpus=*/2_p, /*num_gpus=*/0_n); + int fake_argc = fake_args.size(); + char **fake_argv = fake_args.data(); + + RealmManager manager = RealmManager{&fake_argc, &fake_argv}; + ControllerTaskResult result = manager.start_controller([](RealmContext + &ctx) { + Allocator allocator = ctx.get_current_device_allocator(); + + positive_int batch_size = 10_p; + positive_int data_dim = 16_p; + positive_int hidden_dim = 32_p; + positive_int output_dim = 1_p; + + // 10,2 + TensorShape output_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + // 10,2 + TensorShape label_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + GenericTensorAccessorW label_tensor = + allocator.allocate_tensor(label_tensor_shape); + + // construct computation graph + ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + // input tensor + // 10, 16 + TensorShape input_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + // parallel layer -> input tensor + ParallelLayerAddedResult inputs_layer = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input = + require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> input tensor 2 + ParallelLayerAddedResult inputs_layer_2 = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input_2 = + require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); + + // binary ADD attribute + ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ + OperatorType::EW_ADD, + DataType::FLOAT, + false, + false, + }; + + // parallel layer -> perform add + ParallelLayerAddedResult add_operator_1 = + add_parallel_layer(pcg, make_layer_attrs(add_attrs), + { + { + TensorSlotName::LHS_INPUT, + t_input, + }, + { + TensorSlotName::RHS_INPUT, + t_input_2, + }, + }, + {/* weight */}); + + parallel_tensor_guid_t t_add_1 = + require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> perform replicate + const positive_int replicate_degree = 2_p; + ReplicateAttrs repl_attrs = ReplicateAttrs(replicate_degree); + ParallelLayerAddedResult repl_operator_1 = + add_parallel_layer(pcg, make_layer_attrs(repl_attrs), + { + { + TensorSlotName::INPUT, + t_add_1, + }, + }, + /*weight=*/{}); + // output of replicate layer + parallel_tensor_guid_t t_repl_1 = + require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> perform RelU + ParallelLayerAddedResult relu_operator_1 = + add_parallel_layer(pcg, make_layer_attrs(make_relu_attrs()), + /*inputs=*/ + { + { + TensorSlotName::INPUT, + t_repl_1, + }, + }, + /*weights=*/{}); + // output of relu layer + parallel_tensor_guid_t t_relu_1 = + require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); + + // machine + MachineSpaceCoordinate cpu0{0_n, 0_n, DeviceType::CPU}; + MachineSpaceCoordinate cpu1{0_n, 1_n, DeviceType::CPU}; + + ParallelTensorSpaceCoordinate tensor_coord0{ + /* sum_component */ 0_n, /* discard_copy_component */ 0_n, + /*shard_component*/ FFOrdered{0_n}}; + ParallelTensorSpaceCoordinate tensor_coord1{ + /* sum_component */ 0_n, /* discard_copy_component */ 1_n, + /*shard_component*/ FFOrdered{0_n}}; + MappedParallelComputationGraph mpcg{ + pcg, + {{inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, OperatorAtomicTaskShardBinding{{{TensorSlotName::OUTPUT, + tensor_coord0}}}}}}}, + {inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, OperatorAtomicTaskShardBinding{{{TensorSlotName::OUTPUT, + tensor_coord0}}}}}}}, + {add_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}}}}, + {repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + {cpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {cpu1, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}, + }}}, + {relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + {cpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {cpu1, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}, + }}}}, + }; + + MappedOperatorTaskGroup loss_mapping{ + {{cpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::LOGIT, tensor_coord0}, + }}}}}; + + // instantiate computation graph + LossAttrs loss_attrs = LossAttrs{ + NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; + OptimizerAttrs optimizer_attrs = + OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001}}; + + std::unordered_map + input_tensors; + + DistributedFfHandle device_handle = + create_distributed_ff_handle(ctx, + /*workSpaceSize=*/1024 * 1024, + /*allowTensorOpMathConversion=*/true); + PCGInstance pcg_instance = create_pcg_instance( + /*ctx=*/ctx, + /*mpcg=*/mpcg, + /*optimizer=*/optimizer_attrs, + /*loss=*/std::nullopt, + /*input_tensors=*/input_tensors, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*iteration_config=*/FFIterationConfig{1_p}); + + // begin training loop + int num_epochs = 1; + for (int i = 0; i < num_epochs; i++) { + perform_all_passes_for_pcg_instance( + /*instance=*/pcg_instance, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*iteration_config=*/FFIterationConfig{1_p}); + } + }); + result.wait(); + } +} + +TEST_SUITE(FF_CUDA_TEST_SUITE) { + TEST_CASE("RealmBackend e2e Training Replicate Op (GPU Model Parallelism)") { + std::vector fake_args = + make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/2_n); + int fake_argc = fake_args.size(); + char **fake_argv = fake_args.data(); + + RealmManager manager = RealmManager{&fake_argc, &fake_argv}; + + ControllerTaskResult result = + manager.start_controller([](RealmContext &ctx) { + Allocator allocator = ctx.get_current_device_allocator(); + + positive_int batch_size = 10_p; + positive_int data_dim = 16_p; + positive_int hidden_dim = 32_p; + positive_int output_dim = 1_p; + + // 10,2 + TensorShape output_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + // 10,2 + TensorShape label_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + GenericTensorAccessorW label_tensor = + allocator.allocate_tensor(label_tensor_shape); + + // construct computation graph + ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + // input tensor + // 10, 16 + TensorShape input_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + // parallel layer -> input tensor + ParallelLayerAddedResult inputs_layer = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input = + require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> input tensor 2 + ParallelLayerAddedResult inputs_layer_2 = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input_2 = + require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); + + // binary ADD attribute + ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ + OperatorType::EW_ADD, + DataType::FLOAT, + false, + false, + }; + + // parallel layer -> perform add + ParallelLayerAddedResult add_operator_1 = + add_parallel_layer(pcg, make_layer_attrs(add_attrs), + { + { + TensorSlotName::LHS_INPUT, + t_input, + }, + { + TensorSlotName::RHS_INPUT, + t_input_2, + }, + }, + {/* weight */}); + + parallel_tensor_guid_t t_add_1 = + require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> perform replicate + const positive_int replicate_degree = 2_p; + ReplicateAttrs repl_attrs = ReplicateAttrs(replicate_degree); + ParallelLayerAddedResult repl_operator_1 = + add_parallel_layer(pcg, make_layer_attrs(repl_attrs), + { + { + TensorSlotName::INPUT, + t_add_1, + }, + }, + /*weight=*/{}); + // output of replicate layer + parallel_tensor_guid_t t_repl_1 = + require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> perform RelU + ParallelLayerAddedResult relu_operator_1 = + add_parallel_layer(pcg, make_layer_attrs(make_relu_attrs()), + /*inputs=*/ + { + { + TensorSlotName::INPUT, + t_repl_1, + }, + }, + /*weights=*/{}); + // output of relu layer + parallel_tensor_guid_t t_relu_1 = + require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); + + // machine + MachineSpaceCoordinate gpu0{0_n, 0_n, DeviceType::GPU}; + MachineSpaceCoordinate gpu1{0_n, 1_n, DeviceType::GPU}; + ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; + ParallelTensorSpaceCoordinate tensor_coord1{0_n, 1_n, FFOrdered{0_n}}; + MappedParallelComputationGraph mpcg{ + pcg, + { + {inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + {{gpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, + {inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + {{gpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, + {add_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + {{gpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}}}}, + {repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + {gpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {gpu1, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}}}}, + {relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + {gpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {gpu1, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}, + }}}, + }, + }; + + MappedOperatorTaskGroup loss_mapping{ + {{gpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::LOGIT, tensor_coord0}, + }}}}}; + + // instantiate computation graph + LossAttrs loss_attrs = LossAttrs{ + NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; + OptimizerAttrs optimizer_attrs = + OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001}}; + + std::unordered_map + input_tensors; + + DistributedFfHandle device_handle = create_distributed_ff_handle( + ctx, + /*workSpaceSize=*/1024 * 1024, + /*allowTensorOpMathConversion=*/true); + + PCGInstance pcg_instance = create_pcg_instance( + /*ctx=*/ctx, + /*mpcg=*/mpcg, + /*optimizer=*/optimizer_attrs, + /*loss=*/std::nullopt, + /*input_tensors=*/input_tensors, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*iteration_config=*/FFIterationConfig{1_p}); + + // begin training loop + int num_epochs = 1; + for (int i = 0; i < num_epochs; i++) { + perform_all_passes_for_pcg_instance( + /*instance=*/pcg_instance, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*iteration_config=*/FFIterationConfig{1_p}); + } + }); + result.wait(); + } +} +} // namespace test From 34056217cbb4a8067e582a792fa8af726c8d712e Mon Sep 17 00:00:00 2001 From: Seema Mirchandaney Date: Thu, 9 Apr 2026 15:52:21 -0700 Subject: [PATCH 02/35] Add support for replicate op in distributed training - Add perform_pass_expansion_for_replicate for fwd/bwd pass expansion - Add perform_shard_expansion_for_replicate and _bwd for shard expansion - Add build_replicate_invocation in make_dynamic_open_dataflow_graph - Add is_replicate_attrs helper and guard replicate in copy_insertion - Add ReplicateAttrs to TrainingOperationAttrs - Add SumReductionFloat/Double for backward replicate reduce operation - Add issue_replicate_bwd in spawn_dynamic_node_invocation - Fix per_device_op_state init race condition with direct write - Fix .value() calls on optional per_device_op_state across op impls - Update issue_copy to support optional reduction op - Add testcase for replicate op --- .../src/op-attrs/ops/element_unary.cc | 1 - .../test/src/op-attrs/ops/element_unary.cc | 8 - .../include/realm-execution/realm_context.h | 19 +- .../include/realm-execution/sum_reduction.h | 99 ---- .../realm-execution/tasks/realm_reduction.h | 49 +- ...uted_per_device_op_state_initialization.cc | 6 +- .../src/realm-execution/pcg_instance.cc | 54 +++ .../src/realm-execution/realm_context.cc | 9 +- .../impl/per_device_op_state_init_task.cc | 16 +- .../tasks/realm_task_registry.cc | 10 + .../src/realm-execution/test_op_replicate.cc | 444 +++++++++--------- .../training_operation_attrs.dtg.toml | 4 + .../task-spec/dynamic_graph/copy_insertion.cc | 47 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 127 +++++ .../task-spec/dynamic_graph/pass_expansion.cc | 43 ++ .../dynamic_graph/shard_expansion.cc | 125 ++++- .../src/task-spec/ops/impl/element_binary.cc | 8 +- .../src/task-spec/ops/impl/element_unary.cc | 8 +- 18 files changed, 713 insertions(+), 364 deletions(-) delete mode 100644 lib/realm-execution/include/realm-execution/sum_reduction.h diff --git a/lib/op-attrs/src/op-attrs/ops/element_unary.cc b/lib/op-attrs/src/op-attrs/ops/element_unary.cc index 9d02923689..ca7e417814 100644 --- a/lib/op-attrs/src/op-attrs/ops/element_unary.cc +++ b/lib/op-attrs/src/op-attrs/ops/element_unary.cc @@ -35,7 +35,6 @@ ParallelTensorDimDegrees get_output_parallel_dim_degrees( ElementUnaryAttrs const &attrs, ParallelTensorDimDegrees const &input_degrees) { ASSERT(input_degrees.sum_degree.value == 1); - ASSERT(input_degrees.discard_copy_degree.value == 1); return input_degrees; } diff --git a/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc b/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc index 672b160cbd..43b4be06d8 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc @@ -62,13 +62,5 @@ TEST_SUITE(FF_TEST_SUITE) { SumDegree{degree}, DiscardCopyDegree{1_p}, 1_p, 1_p, 1_p))); } - SUBCASE("discard copy degree > 1") { - positive_int degree = 2_p; - - CHECK_THROWS(get_output_shape( - attrs, - make_input( - SumDegree{1_p}, DiscardCopyDegree{degree}, 1_p, 1_p, 1_p))); - } } } diff --git a/lib/realm-execution/include/realm-execution/realm_context.h b/lib/realm-execution/include/realm-execution/realm_context.h index ab89e916c0..eab42d0d79 100644 --- a/lib/realm-execution/include/realm-execution/realm_context.h +++ b/lib/realm-execution/include/realm-execution/realm_context.h @@ -63,15 +63,18 @@ struct RealmContext { int priority = 0); ///\} - /** \name Data movement */ + /** \name Data movement and reduction */ ///\{ - Realm::Event issue_copy(ParallelTensorShape const &src_shape, - Realm::RegionInstance src_inst, - ParallelTensorShape const &dst_shape, - Realm::RegionInstance dst_inst, - Realm::ProfilingRequestSet const &requests, - Realm::Event wait_on = Realm::Event::NO_EVENT, - int priority = 0); + Realm::Event + issue_copy(ParallelTensorShape const &src_shape, + Realm::RegionInstance src_inst, + ParallelTensorShape const &dst_shape, + Realm::RegionInstance dst_inst, + Realm::ProfilingRequestSet const &requests, + Realm::Event wait_on = Realm::Event::NO_EVENT, + int priority = 0, + std::optional redop_id = std::nullopt, + bool exclusive = false); ///\} /** \name Instance management */ diff --git a/lib/realm-execution/include/realm-execution/sum_reduction.h b/lib/realm-execution/include/realm-execution/sum_reduction.h deleted file mode 100644 index b845b5b7f2..0000000000 --- a/lib/realm-execution/include/realm-execution/sum_reduction.h +++ /dev/null @@ -1,99 +0,0 @@ -#pragma once -#include -#include "op-attrs/datatype.dtg.h" - -namespace FlexFlow { - -// Sum reduction for float -struct SumReductionFloat { - using LHS = float; - using RHS = float; - static const RHS identity; - - template - static void apply(LHS &lhs, RHS rhs) { - if (EXCLUSIVE) { - lhs += rhs; - } else { - // atomic add for non-exclusive - __sync_fetch_and_add((int*)&lhs, *(int*)&rhs); - // proper float atomic add — use union trick - union { float f; int i; } old_val, new_val; - do { - old_val.f = lhs; - new_val.f = old_val.f + rhs; - } while (!__sync_bool_compare_and_swap( - (int*)&lhs, old_val.i, new_val.i)); - } - } - - template - static void fold(RHS &rhs1, RHS rhs2) { - if (EXCLUSIVE) { - rhs1 += rhs2; - } else { - union { float f; int i; } old_val, new_val; - do { - old_val.f = rhs1; - new_val.f = old_val.f + rhs2; - } while (!__sync_bool_compare_and_swap( - (int*)&rhs1, old_val.i, new_val.i)); - } - } -}; - -const SumReductionFloat::RHS SumReductionFloat::identity = 0.0f; - -// Sum reduction for double -struct SumReductionDouble { - using LHS = double; - using RHS = double; - static const RHS identity; - - template - static void apply(LHS &lhs, RHS rhs) { - if (EXCLUSIVE) { - lhs += rhs; - } else { - union { double d; long long i; } old_val, new_val; - do { - old_val.d = lhs; - new_val.d = old_val.d + rhs; - } while (!__sync_bool_compare_and_swap( - (long long*)&lhs, old_val.i, new_val.i)); - } - } - - template - static void fold(RHS &rhs1, RHS rhs2) { - if (EXCLUSIVE) { - rhs1 += rhs2; - } else { - union { double d; long long i; } old_val, new_val; - do { - old_val.d = rhs1; - new_val.d = old_val.d + rhs2; - } while (!__sync_bool_compare_and_swap( - (long long*)&rhs1, old_val.i, new_val.i)); - } - } -}; - -const SumReductionDouble::RHS SumReductionDouble::identity = 0.0; - -// Reduction op IDs — must not conflict with other registered redops -enum SumReductionOpIDs { - REDOP_SUM_FLOAT = 1, - REDOP_SUM_DOUBLE = 2, -}; - -inline Realm::ReductionOpID get_sum_reduction_op_id(DataType dtype) { - switch (dtype) { - case DataType::FLOAT: return REDOP_SUM_FLOAT; - case DataType::DOUBLE: return REDOP_SUM_DOUBLE; - default: - PANIC("no sum reduction registered for datatype {}", dtype); - } -} - -} // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h b/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h index d1d6e1d880..d9cf00441b 100644 --- a/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h +++ b/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h @@ -1,6 +1,6 @@ #pragma once -#include #include "op-attrs/datatype.dtg.h" +#include namespace FlexFlow { @@ -8,7 +8,7 @@ namespace FlexFlow { struct SumReductionFloat { using LHS = float; using RHS = float; - static constexpr RHS identity = 0.0f; // ← inside struct, constexpr + static constexpr RHS identity = 0.0f; // ← inside struct, constexpr template static void apply(LHS &lhs, RHS rhs) { @@ -16,14 +16,17 @@ struct SumReductionFloat { lhs += rhs; } else { // atomic add for non-exclusive - __sync_fetch_and_add((int*)&lhs, *(int*)&rhs); + __sync_fetch_and_add((int *)&lhs, *(int *)&rhs); // proper float atomic add — use union trick - union { float f; int i; } old_val, new_val; + union { + float f; + int i; + } old_val, new_val; do { old_val.f = lhs; new_val.f = old_val.f + rhs; - } while (!__sync_bool_compare_and_swap( - (int*)&lhs, old_val.i, new_val.i)); + } while ( + !__sync_bool_compare_and_swap((int *)&lhs, old_val.i, new_val.i)); } } @@ -32,34 +35,39 @@ struct SumReductionFloat { if (EXCLUSIVE) { rhs1 += rhs2; } else { - union { float f; int i; } old_val, new_val; + union { + float f; + int i; + } old_val, new_val; do { old_val.f = rhs1; new_val.f = old_val.f + rhs2; - } while (!__sync_bool_compare_and_swap( - (int*)&rhs1, old_val.i, new_val.i)); + } while ( + !__sync_bool_compare_and_swap((int *)&rhs1, old_val.i, new_val.i)); } } }; - // Sum reduction for double struct SumReductionDouble { using LHS = double; using RHS = double; - static constexpr RHS identity = 0.0; // ← inside struct, constexpr + static constexpr RHS identity = 0.0; // ← inside struct, constexpr template static void apply(LHS &lhs, RHS rhs) { if (EXCLUSIVE) { lhs += rhs; } else { - union { double d; long long i; } old_val, new_val; + union { + double d; + long long i; + } old_val, new_val; do { old_val.d = lhs; new_val.d = old_val.d + rhs; } while (!__sync_bool_compare_and_swap( - (long long*)&lhs, old_val.i, new_val.i)); + (long long *)&lhs, old_val.i, new_val.i)); } } @@ -68,26 +76,31 @@ struct SumReductionDouble { if (EXCLUSIVE) { rhs1 += rhs2; } else { - union { double d; long long i; } old_val, new_val; + union { + double d; + long long i; + } old_val, new_val; do { old_val.d = rhs1; new_val.d = old_val.d + rhs2; } while (!__sync_bool_compare_and_swap( - (long long*)&rhs1, old_val.i, new_val.i)); + (long long *)&rhs1, old_val.i, new_val.i)); } } }; // Reduction op IDs — must not conflict with other registered redops enum SumReductionOpIDs { - REDOP_SUM_FLOAT = 1, + REDOP_SUM_FLOAT = 1, REDOP_SUM_DOUBLE = 2, }; inline Realm::ReductionOpID get_sum_reduction_op_id(DataType dtype) { switch (dtype) { - case DataType::FLOAT: return REDOP_SUM_FLOAT; - case DataType::DOUBLE: return REDOP_SUM_DOUBLE; + case DataType::FLOAT: + return REDOP_SUM_FLOAT; + case DataType::DOUBLE: + return REDOP_SUM_DOUBLE; default: PANIC("no sum reduction registered for datatype {}", dtype); } diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index 1d517a8fe4..e7d8647b12 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -31,6 +31,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( std::unordered_map *> device_state_map; + std::vector completion_events; for (DynamicNodeInvocation const &invocation : dg.invocations) { Realm::Processor target_proc = ctx.map_device_coord_to_processor( assert_unwrap(invocation.node_attrs.device_coord)); @@ -56,6 +57,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( precondition); if (completion_event.has_value()) { + completion_events.push_back(completion_event.value()); device_state_map.insert(std::pair{invocation, device_state_ptr}); } else { // Task doesn't require initialization, clean up and don't store result @@ -63,7 +65,9 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( } } - ctx.get_outstanding_events().wait(); + // wait for all init tasks — direct write to *result_ptr happens + // before each init task event fires so result is ready after this + Realm::Event::merge_events(completion_events).wait(); auto deref = [](DeviceSpecificPtr *const &p) { return *p; }; std::unordered_map> diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 0ecd02143e..a0653c3c37 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -6,6 +6,7 @@ #include "realm-execution/instance_allocation.h" #include "realm-execution/realm_context.h" #include "realm-execution/tasks/impl/op_task.h" +#include "realm-execution/tasks/realm_reduction.h" #include "realm-execution/tensor_instance_backing.h" #include "task-spec/dynamic_graph/copy_insertion.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" @@ -215,6 +216,46 @@ static Realm::Event spawn_dynamic_node_invocation( precondition); }; + // issue_replicate_bwd lambda + auto issue_replicate_bwd = [&]() { + std::optional output_grad_opt; + for (auto const &[slot, value] : invocation.inputs) { + if (slot.slot_tensor_role == DynamicTensorRole{FwbTensorType::GRADIENT}) { + output_grad_opt = value; + } + } + DynamicValueAttrs output_grad = assert_unwrap(output_grad_opt); + DynamicValueAttrs input_grad = get_only(invocation.outputs).second; + Realm::RegionInstance dst_inst = + tensor_instance_backing.backing.at(input_grad).first; + + Realm::ReductionOpID redop_id = get_sum_reduction_op_id( + assert_unwrap(output_grad.parallel_tensor_shape).data_type); + + // chain reductions sequentially to avoid write races on dst + Realm::Event e = precondition; + for (auto const &[p, m] : assert_unwrap(output_grad.mapping)) { + DynamicValueAttrs replica_key = output_grad; + replica_key.mapping = + bidict{{p, m}}; + replica_key.shard_coord = p; + + Realm::RegionInstance src_inst = + tensor_instance_backing.backing.at(replica_key).first; + + e = ctx.issue_copy(assert_unwrap(output_grad.parallel_tensor_shape), + src_inst, + assert_unwrap(input_grad.parallel_tensor_shape), + dst_inst, + Realm::ProfilingRequestSet{}, + e, + 0, + redop_id, + false); + } + return e; + }; + TrainingOperationAttrs op_attrs = assert_unwrap(invocation.node_attrs.op_attrs); return op_attrs.visit(overload{ @@ -222,11 +263,24 @@ static Realm::Event spawn_dynamic_node_invocation( return pcg_op_attrs.visit(overload{ [&](InputAttrs const &) { return Realm::Event::NO_EVENT; }, [&](WeightAttrs const &) { return Realm::Event::NO_EVENT; }, + [&](ReplicateAttrs const &) { + // this should never be reached since replicate + // goes through TrainingOperationAttrs::ReplicateAttrs + PANIC("unexpected replicate in PCGOperatorAttrs path"); + return Realm::Event::NO_EVENT; + }, [&](auto const &) { return spawn_task(); }, }); }, [&](LossAttrs const &) { return spawn_task(); }, [&](CopyAttrs const &) { return issue_copy(); }, + [&](ReplicateAttrs const &) { + if (invocation.node_attrs.task_type.has_value() && + invocation.node_attrs.task_type.value() == DynamicTaskType::BWD) { + return issue_replicate_bwd(); + } + return issue_copy(); + }, }); } diff --git a/lib/realm-execution/src/realm-execution/realm_context.cc b/lib/realm-execution/src/realm-execution/realm_context.cc index 790c1bd613..a4669bf43e 100644 --- a/lib/realm-execution/src/realm-execution/realm_context.cc +++ b/lib/realm-execution/src/realm-execution/realm_context.cc @@ -161,7 +161,9 @@ Realm::Event Realm::RegionInstance dst_inst, Realm::ProfilingRequestSet const &requests, Realm::Event wait_on, - int priority) { + int priority, + std::optional redop_id, + bool exclusive) { TensorShape src_piece_shape = get_piece_shape(src_shape); TensorShape dst_piece_shape = get_piece_shape(dst_shape); ASSERT(src_piece_shape == dst_piece_shape); // For now, assume they match @@ -183,6 +185,11 @@ Realm::Event size_of_datatype(src_piece_shape.data_type).int_from_positive_int()), /*subfield_offset=*/0); + // set reduction op on dst field if provided + if (redop_id.has_value()) { + dst_field.set_redop(redop_id.value(), /*is_fold=*/false, exclusive); + } + Realm::Event result; switch (src_piece_shape.dims.ff_ordered.num_dims()) { #if REALM_MAX_DIM >= 1 diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc index 753fccf74b..0ea51810e4 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc @@ -66,11 +66,17 @@ void per_device_op_state_init_task_body(void const *args, result_state, ctx.get_current_device_idx())}; DeviceSpecificPtr result_device_specific{ ctx.get_current_device_idx(), result_state_ptr}; - spawn_per_device_op_state_init_return_task(ctx, - task_args.origin_proc, - result_device_specific, - task_args.origin_result_ptr, - Realm::Event::NO_EVENT); + + // replace spawn_per_device_op_state_init_return_task with: + // NOTE: SM/TODO: direct write assumes single-node shared address space + // For multi-node, replace with UserEvent trigger pattern + *task_args.origin_result_ptr = result_device_specific; + + // spawn_per_device_op_state_init_return_task(ctx, + // task_args.origin_proc, + // result_device_specific, + // task_args.origin_result_ptr, + // Realm::Event::NO_EVENT); } std::optional spawn_per_device_op_state_init_task( diff --git a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc index e7a8948f8d..acafdf59fd 100644 --- a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc +++ b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc @@ -5,6 +5,7 @@ #include "realm-execution/tasks/impl/op_task.h" #include "realm-execution/tasks/impl/per_device_op_state_init_return_task.h" #include "realm-execution/tasks/impl/per_device_op_state_init_task.h" +#include "realm-execution/tasks/realm_reduction.h" #include "realm-execution/tasks/task_id_t.h" #include "utils/exception.h" @@ -30,9 +31,18 @@ Realm::Event register_task(Realm::Processor::Kind target_kind, Realm::ProfilingRequestSet()); } +static void register_reductions() { + // register sum reduction ops + Realm::Runtime rt = Realm::Runtime::get_runtime(); + rt.register_reduction(REDOP_SUM_FLOAT); + rt.register_reduction(REDOP_SUM_DOUBLE); + // register_reduction is synchronous — no event returned +} + Realm::Event register_all_tasks() { std::vector pending_registrations; + register_reductions(); std::vector init_task_ids = { // Init tasks task_id_t::BATCHNORM_INIT_TASK_ID, diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index d1fc941007..632f08d239 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -56,194 +56,207 @@ TEST_SUITE(FF_TEST_SUITE) { char **fake_argv = fake_args.data(); RealmManager manager = RealmManager{&fake_argc, &fake_argv}; - ControllerTaskResult result = manager.start_controller([](RealmContext - &ctx) { - Allocator allocator = ctx.get_current_device_allocator(); - - positive_int batch_size = 10_p; - positive_int data_dim = 16_p; - positive_int hidden_dim = 32_p; - positive_int output_dim = 1_p; - - // 10,2 - TensorShape output_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - // 10,2 - TensorShape label_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - GenericTensorAccessorW label_tensor = - allocator.allocate_tensor(label_tensor_shape); - - // construct computation graph - ParallelComputationGraph pcg = empty_parallel_computation_graph(); - - // input tensor - // 10, 16 - TensorShape input_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; - - // parallel layer -> input tensor - ParallelLayerAddedResult inputs_layer = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input = - require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> input tensor 2 - ParallelLayerAddedResult inputs_layer_2 = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input_2 = - require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); - - // binary ADD attribute - ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ - OperatorType::EW_ADD, - DataType::FLOAT, - false, - false, - }; - - // parallel layer -> perform add - ParallelLayerAddedResult add_operator_1 = - add_parallel_layer(pcg, make_layer_attrs(add_attrs), - { - { - TensorSlotName::LHS_INPUT, - t_input, - }, + ControllerTaskResult result = + manager.start_controller([](RealmContext &ctx) { + Allocator allocator = ctx.get_current_device_allocator(); + + positive_int batch_size = 10_p; + positive_int data_dim = 16_p; + positive_int hidden_dim = 32_p; + positive_int output_dim = 1_p; + + // 10,2 + TensorShape output_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + // 10,2 + TensorShape label_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + GenericTensorAccessorW label_tensor = + allocator.allocate_tensor(label_tensor_shape); + + // construct computation graph + ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + // input tensor + // 10, 16 + TensorShape input_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + // parallel layer -> input tensor + ParallelLayerAddedResult inputs_layer = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input = + require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> input tensor 2 + ParallelLayerAddedResult inputs_layer_2 = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input_2 = + require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); + + // binary ADD attribute + ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ + OperatorType::EW_ADD, + DataType::FLOAT, + false, + false, + }; + + // parallel layer -> perform add + ParallelLayerAddedResult add_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(add_attrs), { - TensorSlotName::RHS_INPUT, - t_input_2, + { + TensorSlotName::LHS_INPUT, + t_input, + }, + { + TensorSlotName::RHS_INPUT, + t_input_2, + }, }, - }, - {/* weight */}); - - parallel_tensor_guid_t t_add_1 = - require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> perform replicate - const positive_int replicate_degree = 2_p; - ReplicateAttrs repl_attrs = ReplicateAttrs(replicate_degree); - ParallelLayerAddedResult repl_operator_1 = - add_parallel_layer(pcg, make_layer_attrs(repl_attrs), - { + {/* weight */}); + + parallel_tensor_guid_t t_add_1 = + require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> perform replicate + const positive_int replicate_degree = 2_p; + ReplicateAttrs repl_attrs = ReplicateAttrs(replicate_degree); + ParallelLayerAddedResult repl_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(repl_attrs), { - TensorSlotName::INPUT, - t_add_1, + { + TensorSlotName::INPUT, + t_add_1, + }, }, - }, - /*weight=*/{}); - // output of replicate layer - parallel_tensor_guid_t t_repl_1 = - require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> perform RelU - ParallelLayerAddedResult relu_operator_1 = - add_parallel_layer(pcg, make_layer_attrs(make_relu_attrs()), - /*inputs=*/ - { + /*weight=*/{}); + // output of replicate layer + parallel_tensor_guid_t t_repl_1 = + require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); + + // parallel layer -> perform RelU + ParallelLayerAddedResult relu_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(make_relu_attrs()), + /*inputs=*/ { - TensorSlotName::INPUT, - t_repl_1, + { + TensorSlotName::INPUT, + t_repl_1, + }, }, - }, - /*weights=*/{}); - // output of relu layer - parallel_tensor_guid_t t_relu_1 = - require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); - - // machine - MachineSpaceCoordinate cpu0{0_n, 0_n, DeviceType::CPU}; - MachineSpaceCoordinate cpu1{0_n, 1_n, DeviceType::CPU}; - - ParallelTensorSpaceCoordinate tensor_coord0{ - /* sum_component */ 0_n, /* discard_copy_component */ 0_n, - /*shard_component*/ FFOrdered{0_n}}; - ParallelTensorSpaceCoordinate tensor_coord1{ - /* sum_component */ 0_n, /* discard_copy_component */ 1_n, - /*shard_component*/ FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg{ - pcg, - {{inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, OperatorAtomicTaskShardBinding{{{TensorSlotName::OUTPUT, - tensor_coord0}}}}}}}, - {inputs_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, OperatorAtomicTaskShardBinding{{{TensorSlotName::OUTPUT, - tensor_coord0}}}}}}}, - {add_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, - {repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - {cpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {cpu1, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}, - }}}, - {relu_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - {cpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {cpu1, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}, - }}}}, - }; - - MappedOperatorTaskGroup loss_mapping{ - {{cpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}}}}; - - // instantiate computation graph - LossAttrs loss_attrs = LossAttrs{ - NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; - OptimizerAttrs optimizer_attrs = - OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001}}; - - std::unordered_map - input_tensors; - - DistributedFfHandle device_handle = - create_distributed_ff_handle(ctx, - /*workSpaceSize=*/1024 * 1024, - /*allowTensorOpMathConversion=*/true); - PCGInstance pcg_instance = create_pcg_instance( - /*ctx=*/ctx, - /*mpcg=*/mpcg, - /*optimizer=*/optimizer_attrs, - /*loss=*/std::nullopt, - /*input_tensors=*/input_tensors, - /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*iteration_config=*/FFIterationConfig{1_p}); - - // begin training loop - int num_epochs = 1; - for (int i = 0; i < num_epochs; i++) { - perform_all_passes_for_pcg_instance( - /*instance=*/pcg_instance, - /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*iteration_config=*/FFIterationConfig{1_p}); - } - }); + /*weights=*/{}); + // output of relu layer + parallel_tensor_guid_t t_relu_1 = + require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); + + // machine + MachineSpaceCoordinate cpu0{0_n, 0_n, DeviceType::CPU}; + MachineSpaceCoordinate cpu1{0_n, 1_n, DeviceType::CPU}; + + ParallelTensorSpaceCoordinate tensor_coord0{ + /* sum_component */ 0_n, + /* discard_copy_component */ 0_n, + /*shard_component*/ FFOrdered{0_n}}; + ParallelTensorSpaceCoordinate tensor_coord1{ + /* sum_component */ 0_n, + /* discard_copy_component */ 1_n, + /*shard_component*/ FFOrdered{0_n}}; + MappedParallelComputationGraph mpcg{ + pcg, + {{inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, + {inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, + {add_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}}}}, + {repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + {cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}, + }}}, + {relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + {cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}, + }}}}, + }; + + MappedOperatorTaskGroup loss_mapping{ + {{cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::LOGIT, tensor_coord0}, + }}}}}; + + // instantiate computation graph + LossAttrs loss_attrs = LossAttrs{ + NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; + OptimizerAttrs optimizer_attrs = + OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001}}; + + std::unordered_map + input_tensors; + + DistributedFfHandle device_handle = create_distributed_ff_handle( + ctx, + /*workSpaceSize=*/1024 * 1024, + /*allowTensorOpMathConversion=*/true); + PCGInstance pcg_instance = create_pcg_instance( + /*ctx=*/ctx, + /*mpcg=*/mpcg, + /*optimizer=*/optimizer_attrs, + /*loss=*/std::nullopt, + /*input_tensors=*/input_tensors, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*iteration_config=*/FFIterationConfig{1_p}); + + // begin training loop + int num_epochs = 1; + for (int i = 0; i < num_epochs; i++) { + perform_all_passes_for_pcg_instance( + /*instance=*/pcg_instance, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*iteration_config=*/FFIterationConfig{1_p}); + } + }); result.wait(); } } @@ -307,7 +320,8 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { // parallel layer -> perform add ParallelLayerAddedResult add_operator_1 = - add_parallel_layer(pcg, make_layer_attrs(add_attrs), + add_parallel_layer(pcg, + make_layer_attrs(add_attrs), { { TensorSlotName::LHS_INPUT, @@ -327,7 +341,8 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { const positive_int replicate_degree = 2_p; ReplicateAttrs repl_attrs = ReplicateAttrs(replicate_degree); ParallelLayerAddedResult repl_operator_1 = - add_parallel_layer(pcg, make_layer_attrs(repl_attrs), + add_parallel_layer(pcg, + make_layer_attrs(repl_attrs), { { TensorSlotName::INPUT, @@ -341,7 +356,8 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { // parallel layer -> perform RelU ParallelLayerAddedResult relu_operator_1 = - add_parallel_layer(pcg, make_layer_attrs(make_relu_attrs()), + add_parallel_layer(pcg, + make_layer_attrs(make_relu_attrs()), /*inputs=*/ { { @@ -357,8 +373,8 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { // machine MachineSpaceCoordinate gpu0{0_n, 0_n, DeviceType::GPU}; MachineSpaceCoordinate gpu1{0_n, 1_n, DeviceType::GPU}; - ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; - ParallelTensorSpaceCoordinate tensor_coord1{0_n, 1_n, FFOrdered{0_n}}; + ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; + ParallelTensorSpaceCoordinate tensor_coord1{0_n, 1_n, FFOrdered{0_n}}; MappedParallelComputationGraph mpcg{ pcg, { @@ -374,38 +390,44 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, {add_operator_1.parallel_layer, MappedOperatorTaskGroup{ - {{gpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, + {{gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}}}}, {repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - {gpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {gpu1, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}}}}, + MappedOperatorTaskGroup{ + {{gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}}}}, {relu_operator_1.parallel_layer, MappedOperatorTaskGroup{{ - {gpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {gpu1, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}, + {gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}, + {gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}}, }}}, }, }; MappedOperatorTaskGroup loss_mapping{ - {{gpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}}}}; + {{gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::LOGIT, tensor_coord0}, + }}}}}; // instantiate computation graph LossAttrs loss_attrs = LossAttrs{ diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml index 8f8f6467c8..2bd0714512 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml @@ -25,3 +25,7 @@ key = "loss" [[values]] type = "::FlexFlow::CopyAttrs" key = "copy" + +[[values]] +type = "::FlexFlow::ReplicateAttrs" +key = "replicate" diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 4c1b9d4609..7a28e254aa 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -25,15 +25,43 @@ bool node_is_copy(DynamicNodeAttrs const &n) { return n.op_attrs.has_value() && n.op_attrs.value().is_copy(); } +static bool is_replicate_invocation(DynamicNodeInvocation const &i) { + if (!i.node_attrs.op_attrs.has_value()) { + return false; + } + TrainingOperationAttrs const &op_attrs = i.node_attrs.op_attrs.value(); + if (op_attrs.is_replicate()) { + return true; + } + return false; +} + bool value_is_mapped(DynamicValueAttrs const &n) { return n.mapping.has_value(); } bool no_part_of_graph_is_copy_inserted(DynamicOpenDataflowGraph const &g) { auto slot_is_mapped = [](DynamicTensorSlot const &) -> bool { return false; }; - - return no_part_of_dynamic_graph_satisfies( - g, node_is_copy, value_is_mapped, slot_is_mapped); + // check all non-replicate invocations + for (DynamicNodeInvocation const &i : g.invocations) { + if (is_replicate_invocation(i)) { + continue; // replicate tensors have mapping set by design + } + if (node_is_copy(i.node_attrs)) { + return false; + } + for (auto const &[slot, value] : i.inputs) { + if (value_is_mapped(value)) { + return false; + } + } + for (auto const &[slot, value] : i.outputs) { + if (value_is_mapped(value)) { + return false; + } + } + } + return true; } bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { @@ -85,6 +113,11 @@ std::unordered_set perform_copy_insertion_for_invocation( std::unordered_map const &unmapped_value_to_mapped_source_value) { + // replicate nodes have no MappedOperatorTaskGroup — + // pass through unchanged, no copies needed + if (is_replicate_invocation(i)) { + return {i}; + } MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); auto map_tensor = [&](DynamicTensorSlot const &slot, @@ -157,6 +190,14 @@ DynamicOpenDataflowGraph std::unordered_map unmapped_value_to_mapped_source_value; for (DynamicNodeInvocation const &i : g.invocations) { + // replicate nodes have no MappedOperatorTaskGroup — + // output mapping already fully set, maps to itself + if (is_replicate_invocation(i)) { + for (auto const &[slot, value] : i.outputs) { + unmapped_value_to_mapped_source_value.insert(std::pair{value, value}); + } + continue; + } for (auto const &[slot, value] : i.outputs) { unmapped_value_to_mapped_source_value.insert( std::pair{value, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 246f9a3242..3d48a0dc2b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -7,11 +7,129 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "utils/containers/generate_map.h" +#include "utils/containers/get_only.h" #include #include #include namespace FlexFlow { +static bidict + get_input_mapping_for_replicate( + MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &replicate_layer) { + + auto [input_slot_name, input_tensor_guid] = + get_only(get_incoming_tensors(mpcg.pcg, replicate_layer)); + + // find the layer that produces this tensor + for (auto const &[layer, _] : get_parallel_layer_attrs_mapping(mpcg.pcg)) { + for (auto const &[slot_name, t] : get_outgoing_tensors(mpcg.pcg, layer)) { + if (t == input_tensor_guid) { + MappedOperatorTaskGroup producer_mapping = mpcg.mapped_tasks.at(layer); + return get_tensor_bindings_for_slot_name(producer_mapping, slot_name); + } + } + } + + PANIC("could not find producer of replicate layer input tensor"); +} + +static std::unordered_map + get_consumers_of_tensor(MappedParallelComputationGraph const &mpcg, + parallel_tensor_guid_t const &tensor) { + std::unordered_map result; + for (auto const &[layer, _] : get_parallel_layer_attrs_mapping(mpcg.pcg)) { + for (auto const &[slot_name, t] : get_incoming_tensors(mpcg.pcg, layer)) { + if (t == tensor) { + result.insert({layer, slot_name}); + } + } + } + return result; +} + +static bidict + build_replicated_output_mapping( + MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &replicate_layer) { + + auto [output_slot_name, output_tensor_guid] = + get_only(get_outgoing_tensors(mpcg.pcg, replicate_layer)); + + auto consumers = get_consumers_of_tensor(mpcg, output_tensor_guid); + ASSERT(!consumers.empty()); + + // union all consumer bindings — each consumer shard maps to a distinct + // (discard_copy, machine) pair since replicas are always on different machines + bidict result; + for (auto const &[consumer_layer, slot_name] : consumers) { + MappedOperatorTaskGroup consumer_mapping = + mpcg.mapped_tasks.at(consumer_layer); + bidict binding = + get_tensor_bindings_for_slot_name(consumer_mapping, slot_name); + for (auto const &[p, m] : binding) { + result.equate(p, m); + } + } + return result; +} + +static DynamicNodeInvocation + build_replicate_invocation(parallel_layer_guid_t const &layer, + ParallelLayerAttrs const &attrs, + MappedParallelComputationGraph const &mpcg) { + auto [input_slot_name, input_tensor_guid] = + get_only(get_incoming_tensors(mpcg.pcg, layer)); + auto incoming = get_incoming_tensors(mpcg.pcg, layer); + ASSERT(!incoming.empty(), + "replicate layer has no incoming tensors — " + "check PCG edge construction in test"); + + ParallelTensorAttrs input_attrs = + get_parallel_tensor_attrs(mpcg.pcg, input_tensor_guid); + bidict input_mapping = + get_input_mapping_for_replicate(mpcg, layer); + + DynamicValueAttrs input_value{ + /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, + /*parallel_tensor_shape=*/input_attrs.shape, + /*shard_coord=*/std::nullopt, + /*mapping=*/get_input_mapping_for_replicate(mpcg, layer), + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + + auto [output_slot_name, output_tensor_guid] = + get_only(get_outgoing_tensors(mpcg.pcg, layer)); + ParallelTensorAttrs output_attrs = + get_parallel_tensor_attrs(mpcg.pcg, output_tensor_guid); + + DynamicValueAttrs output_value{ + /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, + /*parallel_tensor_shape=*/output_attrs.shape, + /*shard_coord=*/std::nullopt, + /*mapping=*/build_replicated_output_mapping(mpcg, layer), + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + DynamicNodeAttrs node_attrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs.get()}, + /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, + /*per_device_op_state=*/std::nullopt, + }; + + DynamicNodeInvocation invocation_node{ + /*inputs=*/{ + {DynamicTensorSlot{input_slot_name, std::nullopt}, input_value}}, + /*node_attrs=*/node_attrs, + /*outputs=*/ + {{DynamicTensorSlot{output_slot_name, std::nullopt}, output_value}}, + }; + return invocation_node; +} DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( MappedParallelComputationGraph const &mpcg) { @@ -19,6 +137,15 @@ DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( for (auto const &[layer, attrs] : get_parallel_layer_attrs_mapping(mpcg.pcg)) { + + if (attrs.op_attrs.has()) { + // build replicate invocation + DynamicNodeInvocation repl_inv = + build_replicate_invocation(layer, attrs, mpcg); + result.invocations.emplace(repl_inv); + continue; + } + DynamicNodeAttrs result_attrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index 0cee06368f..aed5f2c4c3 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -4,6 +4,7 @@ #include "utils/containers/are_all_same.h" #include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/transform.h" +#include "utils/containers/get_only.h" namespace FlexFlow { @@ -109,6 +110,44 @@ DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( transform(invocation.inputs, to_grad), }; } +static std::unordered_set + perform_pass_expansion_for_replicate( + DynamicNodeInvocation const &invocation) { + + auto const &[input_slot, input] = get_only(invocation.inputs); + auto const &[output_slot, output] = get_only(invocation.outputs); + + // forward: INPUT/FWD → OUTPUT/FWD (copy to replicas) + DynamicNodeInvocation fwd{ + /*inputs=*/{{pass_expand_slot(input_slot, FwbTensorType::FORWARD), + pass_expand_value(input, FwbTensorType::FORWARD)}}, + /*node_attrs=*/ + pass_expand_node(invocation.node_attrs, DynamicTaskType::FWD), + /*outputs=*/ + {{pass_expand_slot(output_slot, FwbTensorType::FORWARD), + pass_expand_value(output, FwbTensorType::FORWARD)}}, + }; + + // backward: OUTPUT/FWD + OUTPUT/GRAD → INPUT/GRAD (reduce gradients) + // The backward node needs the mapping from the output (replicated) + // so it knows which replicas to reduce from + DynamicNodeAttrs bwd_node_attrs = invocation.node_attrs; + bwd_node_attrs.task_type = DynamicTaskType::BWD; + + DynamicNodeInvocation bwd{ + /*inputs=*/{ + {pass_expand_slot(output_slot, FwbTensorType::FORWARD), + pass_expand_value(output, FwbTensorType::FORWARD)}, + {pass_expand_slot(output_slot, FwbTensorType::GRADIENT), + pass_expand_value(output, FwbTensorType::GRADIENT)}, + }, + /*node_attrs=*/bwd_node_attrs, + /*outputs=*/ + {{pass_expand_slot(input_slot, FwbTensorType::GRADIENT), + pass_expand_value(input, FwbTensorType::GRADIENT)}}, + }; + return {fwd, bwd}; +} DynamicOpenDataflowGraph perform_pass_expansion(DynamicOpenDataflowGraph const &g) { @@ -117,6 +156,10 @@ DynamicOpenDataflowGraph DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [](DynamicNodeInvocation const &invocation) { + if (invocation.node_attrs.op_attrs.has_value() && + invocation.node_attrs.op_attrs.value().is_replicate()) { + return perform_pass_expansion_for_replicate(invocation); + } if (invocation.inputs.empty()) { return std::unordered_set{ perform_fwd_pass_expansion_for_invocation(invocation), diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index fb6efb96d0..f30a4d8470 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -39,7 +39,6 @@ bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { value_is_shard_expanded, slot_is_shard_expanded); } - static bidict restrict_tensor_mapping_keys_to_coord( bidict const @@ -85,6 +84,114 @@ static DynamicNodeInvocation shard_invocation_for_binding( }; } +static std::unordered_set + perform_shard_expansion_for_replicate(DynamicNodeInvocation const &i) { + auto const &[input_slot, input] = get_only(i.inputs); + auto const &[output_slot, output] = get_only(i.outputs); + + bidict input_mapping = + assert_unwrap(input.mapping); + bidict output_mapping = + assert_unwrap(output.mapping); + + return transform(output_mapping.left_values(), + [&](ParallelTensorSpaceCoordinate const &p) { + ParallelTensorSpaceCoordinate input_p{ + /*sum_component=*/p.sum_component, + /*discard_copy_component=*/nonnegative_int{0}, + /*shard_components=*/p.shard_components, + }; + return shard_invocation_for_binding( + i, + output_mapping.at_l(p), + OperatorAtomicTaskShardBinding{{ + {input_slot.slot_name, input_p}, + {output_slot.slot_name, p}, + }}); + }); +} + +static std::unordered_set + perform_shard_expansion_for_replicate_bwd(DynamicNodeInvocation const &i) { + + std::optional output_grad_opt; + std::optional output_fwd_opt; + std::optional output_grad_slot_opt; + std::optional output_fwd_slot_opt; + + for (auto const &[slot, value] : i.inputs) { + if (slot.slot_tensor_role == DynamicTensorRole{FwbTensorType::GRADIENT}) { + output_grad_slot_opt = slot; + output_grad_opt = value; + } else { + output_fwd_slot_opt = slot; + output_fwd_opt = value; + } + } + + DynamicValueAttrs output_grad = assert_unwrap(output_grad_opt); + DynamicValueAttrs output_fwd = assert_unwrap(output_fwd_opt); + DynamicTensorSlot output_grad_slot = assert_unwrap(output_grad_slot_opt); + DynamicTensorSlot output_fwd_slot = assert_unwrap(output_fwd_slot_opt); + auto const &[input_grad_slot, input_grad] = get_only(i.outputs); + + bidict + output_grad_mapping = assert_unwrap(output_grad.mapping); + bidict + input_grad_mapping = assert_unwrap(input_grad.mapping); + + std::unordered_map, + std::unordered_set> + by_shard; + for (auto const &p : output_grad_mapping.left_values()) { + by_shard[p.shard_components].insert(p); + } + + std::unordered_set result; + for (auto const &[shard_components, replica_coords] : by_shard) { + ParallelTensorSpaceCoordinate src_p{ + nonnegative_int{0}, nonnegative_int{0}, shard_components}; + MachineSpaceCoordinate src_machine = input_grad_mapping.at_l(src_p); + + bidict + replica_mapping; + for (auto const &p : replica_coords) { + replica_mapping.equate(p, output_grad_mapping.at_l(p)); + } + + DynamicValueAttrs sharded_output_grad = output_grad; + sharded_output_grad.mapping = replica_mapping; + sharded_output_grad.shard_coord = src_p; + + DynamicValueAttrs sharded_output_fwd = output_fwd; + sharded_output_fwd.mapping = replica_mapping; + sharded_output_fwd.shard_coord = src_p; + + DynamicValueAttrs sharded_input_grad = input_grad; + sharded_input_grad.mapping = + bidict{ + {src_p, src_machine}}; + sharded_input_grad.shard_coord = src_p; + + DynamicNodeAttrs sharded_node = i.node_attrs; + sharded_node.device_coord = src_machine; + + result.insert(DynamicNodeInvocation{ + /*inputs=*/{ + {output_fwd_slot, sharded_output_fwd}, + {output_grad_slot, sharded_output_grad}, + }, + /*node_attrs=*/sharded_node, + /*outputs=*/ + { + {input_grad_slot, sharded_input_grad}, + }, + }); + } + return result; +} + + static std::unordered_set perform_shard_expansion_for_copy(DynamicNodeInvocation const &i) { auto [input_slot, input] = get_only(i.inputs); @@ -121,6 +228,22 @@ std::unordered_set return perform_shard_expansion_for_copy(i); } + // forward replicate + if (i.node_attrs.op_attrs.has_value() && + i.node_attrs.op_attrs.value().is_replicate() && + i.node_attrs.task_type.has_value() && + i.node_attrs.task_type.value() == DynamicTaskType::FWD) { + return perform_shard_expansion_for_replicate(i); + } + + // backward replicate + if (i.node_attrs.op_attrs.has_value() && + i.node_attrs.op_attrs.value().is_replicate() && + i.node_attrs.task_type.has_value() && + i.node_attrs.task_type.value() == DynamicTaskType::BWD) { + return perform_shard_expansion_for_replicate_bwd(i); + } + MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); std::unordered_set shard_machine_coords = diff --git a/lib/task-spec/src/task-spec/ops/impl/element_binary.cc b/lib/task-spec/src/task-spec/ops/impl/element_binary.cc index 13465d7a5f..c8460af538 100644 --- a/lib/task-spec/src/task-spec/ops/impl/element_binary.cc +++ b/lib/task-spec/src/task-spec/ops/impl/element_binary.cc @@ -36,8 +36,8 @@ static std::optional forward_task_impl(TaskArgumentAccessor const &acc) { ProfilingSettings profiling = acc.get_profiling_settings(); DeviceType kernel_device_type = acc.get_kernel_device_type(); - ElementBinaryPerDeviceState per_device_state = - acc.get_per_device_op_state().require_element_binary().value(); + std::optional per_device_state = + acc.get_per_device_op_state().require_element_binary(); ElementBinaryAttrs attrs = acc.get_op_attrs().require_element_binary(); device_handle_t handle = acc.get_ff_handle(); @@ -62,8 +62,8 @@ static std::optional backward_task_impl(TaskArgumentAccessor const &acc) { ProfilingSettings profiling = acc.get_profiling_settings(); DeviceType kernel_device_type = acc.get_kernel_device_type(); - ElementBinaryPerDeviceState per_device_state = - acc.get_per_device_op_state().require_element_binary().value(); + std::optional per_device_state = + acc.get_per_device_op_state().require_element_binary(); ElementBinaryAttrs attrs = acc.get_op_attrs().require_element_binary(); device_handle_t handle = acc.get_ff_handle(); diff --git a/lib/task-spec/src/task-spec/ops/impl/element_unary.cc b/lib/task-spec/src/task-spec/ops/impl/element_unary.cc index d66ff9ab8d..9a092b90b8 100644 --- a/lib/task-spec/src/task-spec/ops/impl/element_unary.cc +++ b/lib/task-spec/src/task-spec/ops/impl/element_unary.cc @@ -35,8 +35,8 @@ static std::optional ProfilingSettings profiling = acc.get_profiling_settings(); DeviceType kernel_device_type = acc.get_kernel_device_type(); - ElementUnaryPerDeviceState per_device_state = - acc.get_per_device_op_state().require_element_unary().value(); + std::optional per_device_state = + acc.get_per_device_op_state().require_element_unary(); return profile(forward_kernel, profiling, @@ -62,8 +62,8 @@ static std::optional ProfilingSettings profiling = acc.get_profiling_settings(); DeviceType kernel_device_type = acc.get_kernel_device_type(); - ElementUnaryPerDeviceState per_device_state = - acc.get_per_device_op_state().require_element_unary().value(); + std::optional per_device_state = + acc.get_per_device_op_state().require_element_unary(); return profile(backward_kernel, profiling, From d033e22f77d08fc6b4d1151ef7d6bf7cc23281cb Mon Sep 17 00:00:00 2001 From: Seema Mirchandaney Date: Tue, 14 Apr 2026 17:10:12 -0700 Subject: [PATCH 03/35] remove ReplicateAttr --- .../src/realm-execution/pcg_instance.cc | 17 ++++++----------- .../src/realm-execution/tasks/task_id_t.cc | 12 +++--------- .../training_operation_attrs.dtg.toml | 4 ---- .../task-spec/dynamic_graph/copy_insertion.cc | 13 +++++-------- ...namic_open_dataflow_graph_from_mapped_pcg.cc | 2 +- .../task-spec/dynamic_graph/pass_expansion.cc | 10 +++++++--- .../task-spec/dynamic_graph/shard_expansion.cc | 16 +++++++++------- 7 files changed, 31 insertions(+), 43 deletions(-) diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index a0653c3c37..17c62fe70c 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -264,23 +264,18 @@ static Realm::Event spawn_dynamic_node_invocation( [&](InputAttrs const &) { return Realm::Event::NO_EVENT; }, [&](WeightAttrs const &) { return Realm::Event::NO_EVENT; }, [&](ReplicateAttrs const &) { - // this should never be reached since replicate - // goes through TrainingOperationAttrs::ReplicateAttrs - PANIC("unexpected replicate in PCGOperatorAttrs path"); - return Realm::Event::NO_EVENT; + if (invocation.node_attrs.task_type.has_value() && + invocation.node_attrs.task_type.value() == + DynamicTaskType::BWD) { + return issue_replicate_bwd(); + } + return issue_copy(); // forward }, [&](auto const &) { return spawn_task(); }, }); }, [&](LossAttrs const &) { return spawn_task(); }, [&](CopyAttrs const &) { return issue_copy(); }, - [&](ReplicateAttrs const &) { - if (invocation.node_attrs.task_type.has_value() && - invocation.node_attrs.task_type.value() == DynamicTaskType::BWD) { - return issue_replicate_bwd(); - } - return issue_copy(); - }, }); } diff --git a/lib/realm-execution/src/realm-execution/tasks/task_id_t.cc b/lib/realm-execution/src/realm-execution/tasks/task_id_t.cc index dd4b0a66ca..0bdc2ca6b5 100644 --- a/lib/realm-execution/src/realm-execution/tasks/task_id_t.cc +++ b/lib/realm-execution/src/realm-execution/tasks/task_id_t.cc @@ -64,9 +64,7 @@ std::optional [](RepartitionAttrs const &attrs) { return task_id_t::REPARTITION_INIT_TASK_ID; }, - [](ReplicateAttrs const &attrs) { - return task_id_t::REPLICATE_INIT_TASK_ID; - }, + [](ReplicateAttrs const &attrs) { return std::nullopt; }, [](ReshapeAttrs const &) { return std::nullopt; }, [](ReverseAttrs const &) { return std::nullopt; }, [](SoftmaxAttrs const &) { return task_id_t::SOFTMAX_INIT_TASK_ID; }, @@ -115,9 +113,7 @@ std::optional [](RepartitionAttrs const &attrs) { return task_id_t::REPARTITION_FWD_TASK_ID; }, - [](ReplicateAttrs const &attrs) { - return task_id_t::REPLICATE_FWD_TASK_ID; - }, + [](ReplicateAttrs const &attrs) { return std::nullopt; }, [](ReshapeAttrs const &) { return task_id_t::RESHAPE_FWD_TASK_ID; }, [](ReverseAttrs const &) { return task_id_t::REVERSE_FWD_TASK_ID; }, [](SoftmaxAttrs const &) { return task_id_t::SOFTMAX_FWD_TASK_ID; }, @@ -166,9 +162,7 @@ std::optional [](RepartitionAttrs const &attrs) { return task_id_t::REPARTITION_BWD_TASK_ID; }, - [](ReplicateAttrs const &attrs) { - return task_id_t::REPLICATE_BWD_TASK_ID; - }, + [](ReplicateAttrs const &attrs) { return std::nullopt; }, [](ReshapeAttrs const &) { return task_id_t::RESHAPE_BWD_TASK_ID; }, [](ReverseAttrs const &) { return task_id_t::REVERSE_BWD_TASK_ID; }, [](SoftmaxAttrs const &) { return task_id_t::SOFTMAX_BWD_TASK_ID; }, diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml index 2bd0714512..8f8f6467c8 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml @@ -25,7 +25,3 @@ key = "loss" [[values]] type = "::FlexFlow::CopyAttrs" key = "copy" - -[[values]] -type = "::FlexFlow::ReplicateAttrs" -key = "replicate" diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 7a28e254aa..ef41042a51 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -26,14 +26,11 @@ bool node_is_copy(DynamicNodeAttrs const &n) { } static bool is_replicate_invocation(DynamicNodeInvocation const &i) { - if (!i.node_attrs.op_attrs.has_value()) { - return false; - } - TrainingOperationAttrs const &op_attrs = i.node_attrs.op_attrs.value(); - if (op_attrs.is_replicate()) { - return true; - } - return false; + return i.node_attrs.op_attrs.has_value() && + i.node_attrs.op_attrs.value().has() && + i.node_attrs.op_attrs.value() + .get() + .has(); } bool value_is_mapped(DynamicValueAttrs const &n) { diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 3d48a0dc2b..a4ef156db9 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -116,7 +116,7 @@ static DynamicNodeInvocation /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs.get()}, + /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs}, /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, /*per_device_op_state=*/std::nullopt, }; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index aed5f2c4c3..faa1e186c3 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -2,9 +2,9 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "utils/containers/are_all_same.h" +#include "utils/containers/get_only.h" #include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/transform.h" -#include "utils/containers/get_only.h" namespace FlexFlow { @@ -30,6 +30,11 @@ bool graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &g) { g, node_is_pass_expanded, value_is_pass_expanded, slot_is_pass_expanded); } +static bool is_replicate_attrs(DynamicNodeAttrs const &n) { + return n.op_attrs.has_value() && n.op_attrs.value().has() && + n.op_attrs.value().get().has(); +} + DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &s, FwbTensorType tensor_type) { ASSERT(!slot_is_pass_expanded(s)); @@ -156,8 +161,7 @@ DynamicOpenDataflowGraph DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [](DynamicNodeInvocation const &invocation) { - if (invocation.node_attrs.op_attrs.has_value() && - invocation.node_attrs.op_attrs.value().is_replicate()) { + if (is_replicate_attrs(invocation.node_attrs)) { return perform_pass_expansion_for_replicate(invocation); } if (invocation.inputs.empty()) { diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index f30a4d8470..d3365ae44c 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -191,7 +191,6 @@ static std::unordered_set return result; } - static std::unordered_set perform_shard_expansion_for_copy(DynamicNodeInvocation const &i) { auto [input_slot, input] = get_only(i.inputs); @@ -228,18 +227,21 @@ std::unordered_set return perform_shard_expansion_for_copy(i); } + bool const is_replicate = + i.node_attrs.op_attrs.has_value() && + i.node_attrs.op_attrs.value().has() && + i.node_attrs.op_attrs.value() + .get() + .has(); + // forward replicate - if (i.node_attrs.op_attrs.has_value() && - i.node_attrs.op_attrs.value().is_replicate() && - i.node_attrs.task_type.has_value() && + if (is_replicate && i.node_attrs.task_type.has_value() && i.node_attrs.task_type.value() == DynamicTaskType::FWD) { return perform_shard_expansion_for_replicate(i); } // backward replicate - if (i.node_attrs.op_attrs.has_value() && - i.node_attrs.op_attrs.value().is_replicate() && - i.node_attrs.task_type.has_value() && + if (is_replicate && i.node_attrs.task_type.has_value() && i.node_attrs.task_type.value() == DynamicTaskType::BWD) { return perform_shard_expansion_for_replicate_bwd(i); } From 6cd706091420f4e9c776d75dc3464bbf040f5385 Mon Sep 17 00:00:00 2001 From: Seema Mirchandaney Date: Wed, 15 Apr 2026 16:15:21 -0700 Subject: [PATCH 04/35] Add comments to realm reductions, Use existing graph methods --- .../realm-execution/tasks/realm_reduction.h | 69 +++++++++++++++---- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 44 ++++++------ 2 files changed, 79 insertions(+), 34 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h b/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h index d9cf00441b..512e344824 100644 --- a/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h +++ b/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h @@ -1,23 +1,33 @@ -#pragma once +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDUCTION_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDUCTION_H #include "op-attrs/datatype.dtg.h" #include namespace FlexFlow { -// Sum reduction for float +/** + * \brief Realm Sum Reduction for Float + * \see https://legion.stanford.edu/tutorial/realm/reductions.html + */ struct SumReductionFloat { using LHS = float; using RHS = float; - static constexpr RHS identity = 0.0f; // ← inside struct, constexpr + /** \brief Identity element for addition (0.0) */ + static constexpr RHS identity = 0.0f; + + /** + * \brief Apply reduction: lhs += rhs + * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop + * \param lhs Left-hand side accumulator (modified in place) + * \param rhs Value to add + */ template static void apply(LHS &lhs, RHS rhs) { if (EXCLUSIVE) { lhs += rhs; } else { - // atomic add for non-exclusive - __sync_fetch_and_add((int *)&lhs, *(int *)&rhs); - // proper float atomic add — use union trick + // Atomic float add via CAS loop union { float f; int i; @@ -30,11 +40,18 @@ struct SumReductionFloat { } } + /** + * \brief Fold two RHS values: rhs1 += rhs2 + * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop + * \param rhs1 Accumulator (modified in place) + * \param rhs2 Value to fold in + */ template static void fold(RHS &rhs1, RHS rhs2) { if (EXCLUSIVE) { rhs1 += rhs2; } else { + // Atomic float add via CAS loop union { float f; int i; @@ -48,17 +65,29 @@ struct SumReductionFloat { } }; -// Sum reduction for double +/** + * \brief Realm Sum Reduction for Double + * \see https://legion.stanford.edu/tutorial/realm/reductions.html + */ struct SumReductionDouble { using LHS = double; using RHS = double; - static constexpr RHS identity = 0.0; // ← inside struct, constexpr + /** \brief Identity element for addition (0.0) */ + static constexpr RHS identity = 0.0; + + /** + * \brief Apply reduction: lhs += rhs + * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop + * \param lhs Left-hand side accumulator (modified in place) + * \param rhs Value to add + */ template static void apply(LHS &lhs, RHS rhs) { if (EXCLUSIVE) { lhs += rhs; } else { + // Atomic double add via CAS loop using long long reinterpretation union { double d; long long i; @@ -71,11 +100,18 @@ struct SumReductionDouble { } } + /** + * \brief Fold two RHS values: rhs1 += rhs2 + * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop + * \param rhs1 Accumulator (modified in place) + * \param rhs2 Value to fold in + */ template static void fold(RHS &rhs1, RHS rhs2) { if (EXCLUSIVE) { rhs1 += rhs2; } else { + // Atomic double add via CAS loop using long long reinterpretation union { double d; long long i; @@ -89,12 +125,21 @@ struct SumReductionDouble { } }; -// Reduction op IDs — must not conflict with other registered redops +/** + * \brief Reduction op IDs for sum reductions + * \warning These IDs must not conflict with other registered reduction ops + */ enum SumReductionOpIDs { - REDOP_SUM_FLOAT = 1, - REDOP_SUM_DOUBLE = 2, + REDOP_SUM_FLOAT = 1, ///< Sum reduction op ID for float + REDOP_SUM_DOUBLE = 2, ///< Sum reduction op ID for double }; +/** + * \brief Returns the Realm reduction op ID for a sum reduction over the given datatype + * \param dtype The datatype to look up + * \return The corresponding Realm::ReductionOpID + * \throws PANIC if no sum reduction is registered for the given datatype + */ inline Realm::ReductionOpID get_sum_reduction_op_id(DataType dtype) { switch (dtype) { case DataType::FLOAT: @@ -105,5 +150,5 @@ inline Realm::ReductionOpID get_sum_reduction_op_id(DataType dtype) { PANIC("no sum reduction registered for datatype {}", dtype); } } - } // namespace FlexFlow +#endif diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index a4ef156db9..9349341d4b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -2,6 +2,7 @@ #include "op-attrs/parallel_tensor_shape.h" #include "op-attrs/pcg_operator_attrs.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" +#include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" #include "pcg/parallel_computation_graph/parallel_tensor_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" @@ -18,31 +19,30 @@ static bidict MappedParallelComputationGraph const &mpcg, parallel_layer_guid_t const &replicate_layer) { - auto [input_slot_name, input_tensor_guid] = - get_only(get_incoming_tensors(mpcg.pcg, replicate_layer)); - - // find the layer that produces this tensor - for (auto const &[layer, _] : get_parallel_layer_attrs_mapping(mpcg.pcg)) { - for (auto const &[slot_name, t] : get_outgoing_tensors(mpcg.pcg, layer)) { - if (t == input_tensor_guid) { - MappedOperatorTaskGroup producer_mapping = mpcg.mapped_tasks.at(layer); - return get_tensor_bindings_for_slot_name(producer_mapping, slot_name); - } - } - } + // get_incoming_edges returns map + // replicate has exactly one input + auto [input_slot_name, input_edge] = + get_only(get_incoming_edges(mpcg.pcg, replicate_layer)); - PANIC("could not find producer of replicate layer input tensor"); + parallel_layer_guid_t producer_layer = get_src_layer(input_edge); + TensorSlotName producer_slot = get_src_layer_output_slot_name(input_edge); + + return get_tensor_bindings_for_slot_name(mpcg.mapped_tasks.at(producer_layer), + producer_slot); } static std::unordered_map get_consumers_of_tensor(MappedParallelComputationGraph const &mpcg, parallel_tensor_guid_t const &tensor) { + parallel_layer_guid_t producer_layer = get_source_layer(mpcg.pcg, tensor); + std::unordered_map result; - for (auto const &[layer, _] : get_parallel_layer_attrs_mapping(mpcg.pcg)) { - for (auto const &[slot_name, t] : get_incoming_tensors(mpcg.pcg, layer)) { - if (t == tensor) { - result.insert({layer, slot_name}); - } + // get_outgoing_edges returns unordered_set + for (ParallelComputationGraphEdge const &edge : + get_outgoing_edges(mpcg.pcg, producer_layer)) { + if (get_parallel_tensor(edge) == tensor) { + result.insert( + std::pair{get_dst_layer(edge), get_dst_layer_input_slot_name(edge)}); } } return result; @@ -76,7 +76,7 @@ static bidict static DynamicNodeInvocation build_replicate_invocation(parallel_layer_guid_t const &layer, - ParallelLayerAttrs const &attrs, + ReplicateAttrs const &attrs, MappedParallelComputationGraph const &mpcg) { auto [input_slot_name, input_tensor_guid] = get_only(get_incoming_tensors(mpcg.pcg, layer)); @@ -116,7 +116,7 @@ static DynamicNodeInvocation /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs}, + /*op_attrs=*/TrainingOperationAttrs{PCGOperatorAttrs{attrs}}, /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, /*per_device_op_state=*/std::nullopt, }; @@ -140,8 +140,8 @@ DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( if (attrs.op_attrs.has()) { // build replicate invocation - DynamicNodeInvocation repl_inv = - build_replicate_invocation(layer, attrs, mpcg); + DynamicNodeInvocation repl_inv = build_replicate_invocation( + layer, attrs.op_attrs.get(), mpcg); result.invocations.emplace(repl_inv); continue; } From c50f3846e4f59920cce36792daeef22b2a70d9e0 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 15 May 2026 17:50:24 -0700 Subject: [PATCH 05/35] Minor PR fixes --- .../mapped_parallel_computation_graph.h | 23 ++ .../parallel_computation_graph.h | 5 + .../mapped_parallel_computation_graph.cc | 43 +++ .../parallel_computation_graph.cc | 15 + .../src/realm-execution/pcg_instance.cc | 38 ++- .../src/realm-execution/test_op_replicate.cc | 298 +++++++++++------- .../sub_parallel_computation_graph.h | 2 +- .../apply_substitution/apply_substitution.cc | 2 +- .../sub_parallel_computation_graph.cc | 2 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 32 +- .../get_kwarg_dataflow_value_uses.h | 33 ++ .../include/utils/many_to_one/many_to_one.h | 5 + .../include/utils/one_to_many/one_to_many.h | 5 + .../get_kwarg_dataflow_value_uses.cc | 14 + 14 files changed, 373 insertions(+), 144 deletions(-) create mode 100644 lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h create mode 100644 lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h index 12c7921282..984a524c21 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h @@ -8,12 +8,35 @@ namespace FlexFlow { std::unordered_set mpcg_get_parallel_layers(MappedParallelComputationGraph const &); + MappedOperatorTaskGroup mpcg_get_mapping_for_layer(MappedParallelComputationGraph const &, parallel_layer_guid_t); ParallelComputationGraph pcg_from_mpcg(MappedParallelComputationGraph const &); +parallel_layer_guid_t mpcg_get_source_layer(MappedParallelComputationGraph const &, + parallel_tensor_guid_t const &); + +ParallelTensorAttrs mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &, + parallel_tensor_guid_t const &); + +std::unordered_map + mpcg_get_incoming_edges(MappedParallelComputationGraph const &, + parallel_layer_guid_t const &); + +std::unordered_set + mpcg_get_outgoing_edges(MappedParallelComputationGraph const &, + parallel_layer_guid_t const &); + +ManyToOne + mpcg_get_incoming_tensors(MappedParallelComputationGraph const &, + parallel_layer_guid_t const &); + +bidict + mpcg_get_outgoing_tensors(MappedParallelComputationGraph const &, + parallel_layer_guid_t const &); + std::unordered_set mpcg_get_edges(MappedParallelComputationGraph const &); diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h index 0368be62bc..1b2d5a0b67 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h @@ -11,6 +11,7 @@ #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h" #include +#include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" namespace FlexFlow { @@ -53,6 +54,10 @@ std::unordered_map get_incoming_edges(ParallelComputationGraph const &, parallel_layer_guid_t const &); +std::unordered_set + pcg_get_parallel_tensor_uses(ParallelComputationGraph const &, + parallel_tensor_guid_t const &); + std::unordered_set get_initial_layers(ParallelComputationGraph const &); diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index f4fa946a66..571b89b6dd 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -8,6 +8,8 @@ #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/materialize_labelled_kwarg_dataflow_graph_view.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/rewrite_labelled_kwarg_dataflow_graph_node_labels.h" +#include "utils/bidict/algorithms/bidict_from_map.h" +#include "utils/many_to_one/many_to_one_from_map.h" namespace FlexFlow { @@ -46,6 +48,47 @@ ParallelComputationGraph }; } +parallel_layer_guid_t mpcg_get_source_layer(MappedParallelComputationGraph const &mpcg, + parallel_tensor_guid_t const &t) +{ + return get_source_layer(pcg_from_mpcg(mpcg), t); +} + +ParallelTensorAttrs mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &mpcg, + parallel_tensor_guid_t const &t) +{ + return get_parallel_tensor_attrs(pcg_from_mpcg(mpcg), t); +} + +std::unordered_map + mpcg_get_incoming_edges(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) +{ + return get_incoming_edges(pcg_from_mpcg(mpcg), l); +} + +std::unordered_set + mpcg_get_outgoing_edges(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) +{ + return get_outgoing_edges(pcg_from_mpcg(mpcg), l); +} + +ManyToOne + mpcg_get_incoming_tensors(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) +{ + return many_to_one_from_map(get_incoming_tensors(pcg_from_mpcg(mpcg), l)); +} + + +bidict + mpcg_get_outgoing_tensors(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) +{ + return bidict_from_map(get_outgoing_tensors(pcg_from_mpcg(mpcg), l)); +} + MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( ParallelComputationGraph const &pcg, std::unordered_map const diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index a548ceb65a..2c5197242d 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -36,6 +36,7 @@ #include "utils/graph/node/node.dtg.h" #include "utils/record_formatter.h" #include +#include "utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h" namespace FlexFlow { @@ -206,6 +207,20 @@ std::unordered_map }); } +std::unordered_set + pcg_get_parallel_tensor_uses(ParallelComputationGraph const &pcg, + parallel_tensor_guid_t const &t) +{ + std::unordered_set> raw_uses = + get_kwarg_dataflow_value_uses(pcg.raw_graph, + t.raw_graph_output); + + return transform(raw_uses, [](KwargDataflowInput const &i) { + return parallel_tensor_use_t{i}; + }); +} + + std::unordered_set get_initial_layers(ParallelComputationGraph const &pcg) { std::unordered_set raw_sources = get_initial_nodes(pcg.raw_graph); diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 17c62fe70c..17a6a383e6 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -218,14 +218,17 @@ static Realm::Event spawn_dynamic_node_invocation( // issue_replicate_bwd lambda auto issue_replicate_bwd = [&]() { - std::optional output_grad_opt; - for (auto const &[slot, value] : invocation.inputs) { - if (slot.slot_tensor_role == DynamicTensorRole{FwbTensorType::GRADIENT}) { - output_grad_opt = value; - } - } - DynamicValueAttrs output_grad = assert_unwrap(output_grad_opt); - DynamicValueAttrs input_grad = get_only(invocation.outputs).second; + + DynamicValueAttrs output_grad = get_only( + values( + filter_keys( + invocation.inputs, + [](DynamicTensorSlot const &s) -> bool { + return s.slot_tensor_role == DynamicTensorRole{FwbTensorType::GRADIENT}; + }))); + + DynamicValueAttrs input_grad = get_only(values(invocation.outputs)); + Realm::RegionInstance dst_inst = tensor_instance_backing.backing.at(input_grad).first; @@ -243,15 +246,16 @@ static Realm::Event spawn_dynamic_node_invocation( Realm::RegionInstance src_inst = tensor_instance_backing.backing.at(replica_key).first; - e = ctx.issue_copy(assert_unwrap(output_grad.parallel_tensor_shape), - src_inst, - assert_unwrap(input_grad.parallel_tensor_shape), - dst_inst, - Realm::ProfilingRequestSet{}, - e, - 0, - redop_id, - false); + e = ctx.issue_copy( + /*src_shape=*/assert_unwrap(output_grad.parallel_tensor_shape), + /*src_inst=*/src_inst, + /*dst_shape=*/assert_unwrap(input_grad.parallel_tensor_shape), + /*dst_inst=*/dst_inst, + /*requests=*/Realm::ProfilingRequestSet{}, + /*wait_on=*/e, + /*priority=*/0, + /*redop_id=*/redop_id, + /*exlusive=*/false); } return e; }; diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index 632f08d239..cae5ca1756 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -27,6 +27,7 @@ #include "test/utils/doctest/check_kv.h" #include "utils/containers/require_only_key.h" #include +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" namespace test { @@ -168,67 +169,116 @@ TEST_SUITE(FF_TEST_SUITE) { /* sum_component */ 0_n, /* discard_copy_component */ 1_n, /*shard_component*/ FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg{ - pcg, - {{inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {inputs_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {add_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, + MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( + /*pcg=*/pcg, + /*mapped_op_task_groups=*/{ + { + inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, - {repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - {cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {cpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}, - }}}, - {relu_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - {cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {cpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}, - }}}}, - }; + }}, + }, + }, + }, + }, + { + inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + add_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, + }, + { + relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, + }, + }); MappedOperatorTaskGroup loss_mapping{ - {{cpu0, + { + { + cpu0, OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}}}}; + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::LOGIT, tensor_coord0}, + }}, + }, + }, + }; // instantiate computation graph LossAttrs loss_attrs = LossAttrs{ NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; OptimizerAttrs optimizer_attrs = - OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001}}; + OptimizerAttrs{ + SGDOptimizerAttrs{ + /*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001, + }, + }; std::unordered_map input_tensors; @@ -375,68 +425,102 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { MachineSpaceCoordinate gpu1{0_n, 1_n, DeviceType::GPU}; ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; ParallelTensorSpaceCoordinate tensor_coord1{0_n, 1_n, FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg{ - pcg, + MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( + /*pcg=*/pcg, + /*mapped_op_task_groups=*/{ { - {inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {inputs_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {add_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, - {repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}}}}, - {relu_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - {gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}, - {gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}}, - }}}, + inputs_layer.parallel_layer, + MappedOperatorTaskGroup{{ + { + gpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}, + }, + }}, }, - }; - - MappedOperatorTaskGroup loss_mapping{ - {{gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}}}}; + { + inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{{ + { + gpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}, + }}, + }, + }, + { + add_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }}, + }, + { + repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }}, + }, + { + relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{{ + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }}, + }, + }); + + MappedOperatorTaskGroup loss_mapping{{ + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::LOGIT, tensor_coord0}, + }}, + }, + }}; // instantiate computation graph LossAttrs loss_attrs = LossAttrs{ NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; OptimizerAttrs optimizer_attrs = - OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001}}; + OptimizerAttrs{ + SGDOptimizerAttrs{ + /*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001, + }, + }; std::unordered_map input_tensors; diff --git a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h index cbfe3ab264..26c98e915c 100644 --- a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h +++ b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h @@ -48,7 +48,7 @@ std::unordered_set get_subgraph_outgoing_edges( std::unordered_set const &); std::unordered_set - get_parallel_tensor_uses(SubParallelComputationGraph const &, + get_open_parallel_tensor_uses(SubParallelComputationGraph const &, open_parallel_tensor_guid_t const &); SubParallelComputationGraphData diff --git a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc index 6ed2ef563e..a56555550f 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc @@ -109,7 +109,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( input_parallel_tensor_guid_t output_graph_input = output_expr_to_result_sub_pcg_mapping.input_mapping.at_r( output_expr_input); - std::unordered_set uses = get_parallel_tensor_uses( + std::unordered_set uses = get_open_parallel_tensor_uses( substitution_output_graph, open_parallel_tensor_guid_from_input(output_graph_input)); for (parallel_tensor_use_t const &use : uses) { diff --git a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc index 34b8ae1e96..990975bff9 100644 --- a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc +++ b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc @@ -131,7 +131,7 @@ std::unordered_set get_subgraph_incoming_edges( } std::unordered_set - get_parallel_tensor_uses(SubParallelComputationGraph const &spcg, + get_open_parallel_tensor_uses(SubParallelComputationGraph const &spcg, open_parallel_tensor_guid_t const &t) { std::unordered_set> raw_uses = get_open_kwarg_dataflow_value_uses(spcg.raw_graph, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 0aea7d2324..b23edc0411 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -23,24 +23,24 @@ static bidict // get_incoming_edges returns map // replicate has exactly one input auto [input_slot_name, input_edge] = - get_only(get_incoming_edges(mpcg.pcg, replicate_layer)); + get_only(mpcg_get_incoming_edges(mpcg, replicate_layer)); parallel_layer_guid_t producer_layer = get_src_layer(input_edge); TensorSlotName producer_slot = get_src_layer_output_slot_name(input_edge); - return get_tensor_bindings_for_slot_name(mpcg.mapped_tasks.at(producer_layer), - producer_slot); + return get_tensor_bindings_for_slot_name( + /*task_group=*/mpcg_get_mapping_for_layer(mpcg, producer_layer), + /*slot_name=*/producer_slot); } static std::unordered_map get_consumers_of_tensor(MappedParallelComputationGraph const &mpcg, parallel_tensor_guid_t const &tensor) { - parallel_layer_guid_t producer_layer = get_source_layer(mpcg.pcg, tensor); + parallel_layer_guid_t producer_layer = mpcg_get_source_layer(mpcg, tensor); std::unordered_map result; // get_outgoing_edges returns unordered_set - for (ParallelComputationGraphEdge const &edge : - get_outgoing_edges(mpcg.pcg, producer_layer)) { + for (ParallelComputationGraphEdge const &edge : mpcg_get_outgoing_edges(mpcg, producer_layer)) { if (get_parallel_tensor(edge) == tensor) { result.insert( std::pair{get_dst_layer(edge), get_dst_layer_input_slot_name(edge)}); @@ -55,7 +55,7 @@ static bidict parallel_layer_guid_t const &replicate_layer) { auto [output_slot_name, output_tensor_guid] = - get_only(get_outgoing_tensors(mpcg.pcg, replicate_layer)); + get_only(mpcg_get_outgoing_tensors(mpcg, replicate_layer)); auto consumers = get_consumers_of_tensor(mpcg, output_tensor_guid); ASSERT(!consumers.empty()); @@ -64,8 +64,7 @@ static bidict // (discard_copy, machine) pair since replicas are always on different machines bidict result; for (auto const &[consumer_layer, slot_name] : consumers) { - MappedOperatorTaskGroup consumer_mapping = - mpcg.mapped_tasks.at(consumer_layer); + MappedOperatorTaskGroup consumer_mapping = mpcg_get_mapping_for_layer(mpcg, consumer_layer); bidict binding = get_tensor_bindings_for_slot_name(consumer_mapping, slot_name); for (auto const &[p, m] : binding) { @@ -80,14 +79,13 @@ static DynamicNodeInvocation ReplicateAttrs const &attrs, MappedParallelComputationGraph const &mpcg) { auto [input_slot_name, input_tensor_guid] = - get_only(get_incoming_tensors(mpcg.pcg, layer)); - auto incoming = get_incoming_tensors(mpcg.pcg, layer); - ASSERT(!incoming.empty(), - "replicate layer has no incoming tensors — " - "check PCG edge construction in test"); + get_only(mpcg_get_incoming_tensors(mpcg, layer).l_to_r()); + + auto incoming = mpcg_get_incoming_tensors(mpcg, layer); + ASSERT(!incoming.empty(), "Replicate layer has no incoming tensors."); ParallelTensorAttrs input_attrs = - get_parallel_tensor_attrs(mpcg.pcg, input_tensor_guid); + mpcg_get_parallel_tensor_attrs(mpcg, input_tensor_guid); bidict input_mapping = get_input_mapping_for_replicate(mpcg, layer); @@ -101,9 +99,9 @@ static DynamicNodeInvocation }; auto [output_slot_name, output_tensor_guid] = - get_only(get_outgoing_tensors(mpcg.pcg, layer)); + get_only(mpcg_get_outgoing_tensors(mpcg, layer)); ParallelTensorAttrs output_attrs = - get_parallel_tensor_attrs(mpcg.pcg, output_tensor_guid); + mpcg_get_parallel_tensor_attrs(mpcg, output_tensor_guid); DynamicValueAttrs output_value{ /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h new file mode 100644 index 0000000000..b5557e9e49 --- /dev/null +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h @@ -0,0 +1,33 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_KWARG_DATAFLOW_VALUE_USES_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_KWARG_DATAFLOW_VALUE_USES_H + +#include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h" + +namespace FlexFlow { + +template +std::unordered_set> + get_kwarg_dataflow_value_uses( + KwargDataflowGraphView const &g, + KwargDataflowOutput const &v) { + + KwargDataflowEdgeQuery query = + KwargDataflowEdgeQuery{ + /*src_nodes=*/query_set::match_single_value(v.node), + /*src_slots=*/query_set::match_single_value(v.slot_name), + /*dst_nodes=*/query_set::matchall(), + /*dst_slots=*/query_set::matchall(), + }; + + std::unordered_set> edges = + g.query_edges(query); + + return transform( + edges, [&](KwargDataflowEdge const &e) { + return e.dst; + }); +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/many_to_one/many_to_one.h b/lib/utils/include/utils/many_to_one/many_to_one.h index d2f727661c..c73f696172 100644 --- a/lib/utils/include/utils/many_to_one/many_to_one.h +++ b/lib/utils/include/utils/many_to_one/many_to_one.h @@ -19,6 +19,7 @@ #include #include #include +#include "utils/containers/require_same.h" namespace FlexFlow { @@ -106,6 +107,10 @@ struct ManyToOne { return this->m_r_to_l; } + bool empty() const { + return require_same(this->m_l_to_r.empty(), this->m_r_to_l.empty()); + } + private: std::unordered_map m_l_to_r; std::unordered_map> m_r_to_l; diff --git a/lib/utils/include/utils/one_to_many/one_to_many.h b/lib/utils/include/utils/one_to_many/one_to_many.h index 30d84d34c3..7b725fdec1 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many.h +++ b/lib/utils/include/utils/one_to_many/one_to_many.h @@ -23,6 +23,7 @@ #include #include #include +#include "utils/containers/require_same.h" namespace FlexFlow { @@ -114,6 +115,10 @@ struct OneToMany { return this->m_r_to_l; } + bool empty() const { + return require_same(this->m_l_to_r.empty(), this->m_r_to_l.empty()); + } + private: std::unordered_map> m_l_to_r; std::unordered_map m_r_to_l; diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc new file mode 100644 index 0000000000..2e42863e53 --- /dev/null +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc @@ -0,0 +1,14 @@ +#include "utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using SlotName = ordered_value_type<0>; + +template + std::unordered_set> + get_kwarg_dataflow_value_uses( + KwargDataflowGraphView const &, + KwargDataflowOutput const &); + +} // namespace FlexFlow From ac4fffcb307fe1116fde88b3c7aa85599c224a4e Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 15 May 2026 21:39:42 -0700 Subject: [PATCH 06/35] Clean up pass expansion code --- .../op-attrs/pcg_operator_attrs.dtg.toml | 22 +- .../test/src/op-attrs/ops/element_unary.cc | 13 + .../mapped_parallel_computation_graph.h | 7 + .../parallel_tensor_use_t.h | 14 + .../mapped_parallel_computation_graph.cc | 27 +- .../parallel_tensor_use_t.cc | 13 + .../src/realm-execution/pcg_instance.cc | 1 - .../src/realm-execution/test_op_replicate.cc | 587 ++++++------------ .../output_expr_to_result_sub_pcg_mapping.cc | 4 +- .../src/substitutions/pcg_pattern_match.cc | 4 +- .../dynamic_graph/training_operation_attrs.h | 13 + ...mic_open_dataflow_graph_from_mapped_pcg.cc | 220 +++---- .../task-spec/dynamic_graph/pass_expansion.cc | 85 +-- .../dynamic_graph/training_operation_attrs.cc | 21 + .../task-spec/dynamic_graph/pass_expansion.cc | 270 +++++--- .../binary_merge_disjoint_bidicts.h | 37 ++ .../algorithms/merge_disjoint_bidicts.h | 39 +- lib/utils/include/utils/bidict/bidict.h | 8 + .../utils/containers/transform_pairs.h | 46 ++ .../binary_merge_disjoint_bidicts.cc | 12 + .../algorithms/merge_disjoint_bidicts.cc | 10 + .../src/utils/containers/transform_pairs.cc | 17 + ...ts.cc => binary_merge_disjoint_bidicts.cc} | 12 +- 23 files changed, 792 insertions(+), 690 deletions(-) create mode 100644 lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h create mode 100644 lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc create mode 100644 lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h create mode 100644 lib/utils/include/utils/containers/transform_pairs.h create mode 100644 lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc create mode 100644 lib/utils/src/utils/containers/transform_pairs.cc rename lib/utils/test/src/utils/bidict/algorithms/{merge_disjoint_bidicts.cc => binary_merge_disjoint_bidicts.cc} (72%) diff --git a/lib/op-attrs/include/op-attrs/pcg_operator_attrs.dtg.toml b/lib/op-attrs/include/op-attrs/pcg_operator_attrs.dtg.toml index 88a65f75c5..f2dd7c9350 100644 --- a/lib/op-attrs/include/op-attrs/pcg_operator_attrs.dtg.toml +++ b/lib/op-attrs/include/op-attrs/pcg_operator_attrs.dtg.toml @@ -11,13 +11,13 @@ features = [ ] includes = [ - "op-attrs/ops/attention_attrs.dtg.h", - "op-attrs/ops/batch_matmul_attrs.dtg.h", - "op-attrs/ops/batch_norm_attrs.dtg.h", - "op-attrs/ops/broadcast_attrs.dtg.h", - "op-attrs/ops/cast_attrs.dtg.h", - "op-attrs/ops/combine_attrs.dtg.h", - "op-attrs/ops/concat_attrs.dtg.h", + "op-attrs/ops/attention_attrs.dtg.h", + "op-attrs/ops/batch_matmul_attrs.dtg.h", + "op-attrs/ops/batch_norm_attrs.dtg.h", + "op-attrs/ops/broadcast_attrs.dtg.h", + "op-attrs/ops/cast_attrs.dtg.h", + "op-attrs/ops/combine_attrs.dtg.h", + "op-attrs/ops/concat_attrs.dtg.h", "op-attrs/ops/conv_2d_attrs.dtg.h", "op-attrs/ops/dropout_attrs.dtg.h", "op-attrs/ops/element_binary_attrs.dtg.h", @@ -61,7 +61,7 @@ key = "cast" [[values]] type = "::FlexFlow::CombineAttrs" -key = "combine_distributed" +key = "parallel_combine" [[values]] type = "::FlexFlow::ConcatAttrs" @@ -125,15 +125,15 @@ key = "reduce" [[values]] type = "::FlexFlow::ReductionAttrs" -key = "reduce_distributed" +key = "parallel_reduce" [[values]] type = "::FlexFlow::RepartitionAttrs" -key = "partition_distributed" +key = "parallel_partition" [[values]] type = "::FlexFlow::ReplicateAttrs" -key = "replicate_distributed" +key = "parallel_replicate" [[values]] type = "::FlexFlow::ReverseAttrs" diff --git a/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc b/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc index 43b4be06d8..8b2555610e 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc @@ -53,6 +53,19 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } + SUBCASE("discard copy degree > 1") { + positive_int degree = 2_p; + + ParallelTensorShape par_input = make_input( + SumDegree{1_p}, DiscardCopyDegree{degree}, 1_p, 1_p, 1_p); + + tl::expected result = + get_output_shape(attrs, par_input); + tl::expected correct = par_input; + + CHECK(result == correct); + } + SUBCASE("sum degree > 1") { positive_int degree = 2_p; diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h index 984a524c21..6c24d4c1e1 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h @@ -18,6 +18,9 @@ ParallelComputationGraph pcg_from_mpcg(MappedParallelComputationGraph const &); parallel_layer_guid_t mpcg_get_source_layer(MappedParallelComputationGraph const &, parallel_tensor_guid_t const &); +PCGOperatorAttrs mpcg_get_pcg_op_attrs(MappedParallelComputationGraph const &, + parallel_layer_guid_t const &); + ParallelTensorAttrs mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &, parallel_tensor_guid_t const &); @@ -40,6 +43,10 @@ bidict std::unordered_set mpcg_get_edges(MappedParallelComputationGraph const &); +std::unordered_set + mpcg_get_parallel_tensor_uses(MappedParallelComputationGraph const &, + parallel_tensor_guid_t const &); + MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( ParallelComputationGraph const &pcg, std::unordered_map const diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h new file mode 100644 index 0000000000..88f1512149 --- /dev/null +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h @@ -0,0 +1,14 @@ +#ifndef _FLEXFLOW_LIB_PCG_INCLUDE_PCG_PARALLEL_COMPUTATION_GRAPH_PARALLEL_TENSOR_USE_T_H +#define _FLEXFLOW_LIB_PCG_INCLUDE_PCG_PARALLEL_COMPUTATION_GRAPH_PARALLEL_TENSOR_USE_T_H + +#include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" +#include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" + +namespace FlexFlow { + +parallel_layer_guid_t parallel_tensor_use_get_layer(parallel_tensor_use_t const &); +TensorSlotName parallel_tensor_use_get_slot(parallel_tensor_use_t const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 571b89b6dd..3b996ccdab 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -54,22 +54,28 @@ parallel_layer_guid_t mpcg_get_source_layer(MappedParallelComputationGraph const return get_source_layer(pcg_from_mpcg(mpcg), t); } +PCGOperatorAttrs mpcg_get_pcg_op_attrs(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) +{ + return pcg_get_op_attrs(pcg_from_mpcg(mpcg), l); +} + ParallelTensorAttrs mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &mpcg, - parallel_tensor_guid_t const &t) + parallel_tensor_guid_t const &t) { return get_parallel_tensor_attrs(pcg_from_mpcg(mpcg), t); } std::unordered_map mpcg_get_incoming_edges(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) + parallel_layer_guid_t const &l) { return get_incoming_edges(pcg_from_mpcg(mpcg), l); } std::unordered_set mpcg_get_outgoing_edges(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) + parallel_layer_guid_t const &l) { return get_outgoing_edges(pcg_from_mpcg(mpcg), l); } @@ -84,11 +90,24 @@ ManyToOne bidict mpcg_get_outgoing_tensors(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) + parallel_layer_guid_t const &l) { return bidict_from_map(get_outgoing_tensors(pcg_from_mpcg(mpcg), l)); } +std::unordered_set + mpcg_get_edges(MappedParallelComputationGraph const &mpcg) +{ + return get_edges(pcg_from_mpcg(mpcg)); +} + +std::unordered_set + mpcg_get_parallel_tensor_uses(MappedParallelComputationGraph const &mpcg, + parallel_tensor_guid_t const &t) +{ + return pcg_get_parallel_tensor_uses(pcg_from_mpcg(mpcg), t); +} + MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( ParallelComputationGraph const &pcg, std::unordered_map const diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc new file mode 100644 index 0000000000..e93341d312 --- /dev/null +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc @@ -0,0 +1,13 @@ +#include "pcg/parallel_computation_graph/parallel_tensor_use_t.h" + +namespace FlexFlow { + +parallel_layer_guid_t parallel_tensor_use_get_layer(parallel_tensor_use_t const &u) { + return parallel_layer_guid_t{u.raw_dataflow_input.node}; +} + +TensorSlotName parallel_tensor_use_get_slot(parallel_tensor_use_t const &u) { + return u.raw_dataflow_input.slot_name; +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 17a6a383e6..f2edac7f88 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -216,7 +216,6 @@ static Realm::Event spawn_dynamic_node_invocation( precondition); }; - // issue_replicate_bwd lambda auto issue_replicate_bwd = [&]() { DynamicValueAttrs output_grad = get_only( diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index cae5ca1756..2523cae798 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -49,6 +49,190 @@ static bool did_loss_decrease(GenericTensorAccessorR const &first_epoch, compare_tensor_accessors_le(last_epoch, first_epoch, allocator)); } +MappedParallelComputationGraph make_test_mpcg_for_device_type(DeviceType device_type) { + positive_int batch_size = 10_p; + positive_int data_dim = 16_p; + positive_int hidden_dim = 32_p; + positive_int output_dim = 1_p; + + TensorShape output_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + TensorShape label_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + TensorShape input_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + ParallelLayerAddedResult inputs_layer = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input = + require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult inputs_layer_2 = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input_2 = + require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); + + ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ + OperatorType::EW_ADD, + DataType::FLOAT, + false, + false, + }; + + ParallelLayerAddedResult add_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(add_attrs), + { + { + TensorSlotName::LHS_INPUT, + t_input, + }, + { + TensorSlotName::RHS_INPUT, + t_input_2, + }, + }, + /*weights=*/{}); + + parallel_tensor_guid_t t_add_1 = + require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); + + positive_int replicate_degree = 2_p; + ReplicateAttrs repl_attrs = ReplicateAttrs{replicate_degree}; + ParallelLayerAddedResult repl_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(repl_attrs), + { + { + TensorSlotName::INPUT, + t_add_1, + }, + }, + /*weight=*/{}); + + parallel_tensor_guid_t t_repl_1 = + require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult relu_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(make_relu_attrs()), + /*inputs=*/ + { + { + TensorSlotName::INPUT, + t_repl_1, + }, + }, + /*weights=*/{}); + + parallel_tensor_guid_t t_relu_1 = + require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); + + MachineSpaceCoordinate cpu0{0_n, 0_n, device_type}; + MachineSpaceCoordinate cpu1{0_n, 1_n, device_type}; + + ParallelTensorSpaceCoordinate tensor_coord0{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{0_n}}; + ParallelTensorSpaceCoordinate tensor_coord1{ + /*sum_component=*/0_n, + /*discard_copy_component=*/1_n, + /*shard_component=*/FFOrdered{0_n}}; + + MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( + /*pcg=*/pcg, + /*mapped_op_task_groups=*/{ + { + inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + add_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, + }, + { + relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, + }, + }); + + return mpcg; +} + TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("RealmBackend e2e Training Replicate Op (CPU Model Parallelism)") { std::vector fake_args = @@ -61,215 +245,12 @@ TEST_SUITE(FF_TEST_SUITE) { manager.start_controller([](RealmContext &ctx) { Allocator allocator = ctx.get_current_device_allocator(); - positive_int batch_size = 10_p; - positive_int data_dim = 16_p; - positive_int hidden_dim = 32_p; - positive_int output_dim = 1_p; - - // 10,2 - TensorShape output_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - // 10,2 - TensorShape label_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - GenericTensorAccessorW label_tensor = - allocator.allocate_tensor(label_tensor_shape); - - // construct computation graph - ParallelComputationGraph pcg = empty_parallel_computation_graph(); - - // input tensor - // 10, 16 - TensorShape input_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; - - // parallel layer -> input tensor - ParallelLayerAddedResult inputs_layer = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input = - require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> input tensor 2 - ParallelLayerAddedResult inputs_layer_2 = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input_2 = - require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); - - // binary ADD attribute - ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ - OperatorType::EW_ADD, - DataType::FLOAT, - false, - false, - }; - - // parallel layer -> perform add - ParallelLayerAddedResult add_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(add_attrs), - { - { - TensorSlotName::LHS_INPUT, - t_input, - }, - { - TensorSlotName::RHS_INPUT, - t_input_2, - }, - }, - {/* weight */}); - - parallel_tensor_guid_t t_add_1 = - require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> perform replicate - const positive_int replicate_degree = 2_p; - ReplicateAttrs repl_attrs = ReplicateAttrs(replicate_degree); - ParallelLayerAddedResult repl_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(repl_attrs), - { - { - TensorSlotName::INPUT, - t_add_1, - }, - }, - /*weight=*/{}); - // output of replicate layer - parallel_tensor_guid_t t_repl_1 = - require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> perform RelU - ParallelLayerAddedResult relu_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(make_relu_attrs()), - /*inputs=*/ - { - { - TensorSlotName::INPUT, - t_repl_1, - }, - }, - /*weights=*/{}); - // output of relu layer - parallel_tensor_guid_t t_relu_1 = - require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); - - // machine - MachineSpaceCoordinate cpu0{0_n, 0_n, DeviceType::CPU}; - MachineSpaceCoordinate cpu1{0_n, 1_n, DeviceType::CPU}; - - ParallelTensorSpaceCoordinate tensor_coord0{ - /* sum_component */ 0_n, - /* discard_copy_component */ 0_n, - /*shard_component*/ FFOrdered{0_n}}; - ParallelTensorSpaceCoordinate tensor_coord1{ - /* sum_component */ 0_n, - /* discard_copy_component */ 1_n, - /*shard_component*/ FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( - /*pcg=*/pcg, - /*mapped_op_task_groups=*/{ - { - inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - { - { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - }, - }, - }, - { - inputs_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - { - { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - }, - }, - }, - { - add_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { - { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - }, - }, - }, - { - repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { - { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - { - cpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, - }, - }, - }, - }, - { - relu_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { - { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - { - cpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, - }, - }, - }, - }, - }); + MappedParallelComputationGraph mpcg = make_test_mpcg_for_device_type(DeviceType::CPU); - MappedOperatorTaskGroup loss_mapping{ - { - { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}, - }, - }, - }; - // instantiate computation graph - LossAttrs loss_attrs = LossAttrs{ - NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; + std::unordered_map + input_tensors; + OptimizerAttrs optimizer_attrs = OptimizerAttrs{ SGDOptimizerAttrs{ @@ -280,13 +261,11 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_map - input_tensors; - DistributedFfHandle device_handle = create_distributed_ff_handle( ctx, /*workSpaceSize=*/1024 * 1024, /*allowTensorOpMathConversion=*/true); + PCGInstance pcg_instance = create_pcg_instance( /*ctx=*/ctx, /*mpcg=*/mpcg, @@ -324,194 +303,8 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { manager.start_controller([](RealmContext &ctx) { Allocator allocator = ctx.get_current_device_allocator(); - positive_int batch_size = 10_p; - positive_int data_dim = 16_p; - positive_int hidden_dim = 32_p; - positive_int output_dim = 1_p; - - // 10,2 - TensorShape output_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - // 10,2 - TensorShape label_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - GenericTensorAccessorW label_tensor = - allocator.allocate_tensor(label_tensor_shape); - - // construct computation graph - ParallelComputationGraph pcg = empty_parallel_computation_graph(); - - // input tensor - // 10, 16 - TensorShape input_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; - - // parallel layer -> input tensor - ParallelLayerAddedResult inputs_layer = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input = - require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> input tensor 2 - ParallelLayerAddedResult inputs_layer_2 = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input_2 = - require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); - - // binary ADD attribute - ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ - OperatorType::EW_ADD, - DataType::FLOAT, - false, - false, - }; - - // parallel layer -> perform add - ParallelLayerAddedResult add_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(add_attrs), - { - { - TensorSlotName::LHS_INPUT, - t_input, - }, - { - TensorSlotName::RHS_INPUT, - t_input_2, - }, - }, - {/* weight */}); - - parallel_tensor_guid_t t_add_1 = - require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> perform replicate - const positive_int replicate_degree = 2_p; - ReplicateAttrs repl_attrs = ReplicateAttrs(replicate_degree); - ParallelLayerAddedResult repl_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(repl_attrs), - { - { - TensorSlotName::INPUT, - t_add_1, - }, - }, - /*weight=*/{}); - // output of replicate layer - parallel_tensor_guid_t t_repl_1 = - require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); - - // parallel layer -> perform RelU - ParallelLayerAddedResult relu_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(make_relu_attrs()), - /*inputs=*/ - { - { - TensorSlotName::INPUT, - t_repl_1, - }, - }, - /*weights=*/{}); - // output of relu layer - parallel_tensor_guid_t t_relu_1 = - require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); - - // machine - MachineSpaceCoordinate gpu0{0_n, 0_n, DeviceType::GPU}; - MachineSpaceCoordinate gpu1{0_n, 1_n, DeviceType::GPU}; - ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; - ParallelTensorSpaceCoordinate tensor_coord1{0_n, 1_n, FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( - /*pcg=*/pcg, - /*mapped_op_task_groups=*/{ - { - inputs_layer.parallel_layer, - MappedOperatorTaskGroup{{ - { - gpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}, - }, - }}, - }, - { - inputs_layer_2.parallel_layer, - MappedOperatorTaskGroup{{ - { - gpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}, - }}, - }, - }, - { - add_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - }}, - }, - { - repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, - }, - }}, - }, - { - relu_operator_1.parallel_layer, - MappedOperatorTaskGroup{{ - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, - }, - }}, - }, - }); - - MappedOperatorTaskGroup loss_mapping{{ - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}, - }, - }}; + MappedParallelComputationGraph mpcg = make_test_mpcg_for_device_type(DeviceType::GPU); - // instantiate computation graph - LossAttrs loss_attrs = LossAttrs{ - NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; OptimizerAttrs optimizer_attrs = OptimizerAttrs{ SGDOptimizerAttrs{ diff --git a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc index 2ad5b54a17..4374a951f8 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc @@ -2,7 +2,7 @@ #include "substitutions/output_graph/output_graph_expr.h" #include "substitutions/sub_parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_pairs.h" -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include "utils/containers/values.h" #include "utils/containers/zip_values_strict.h" @@ -26,7 +26,7 @@ bidict mapping_for_layer = bidict_from_pairs(values( zip_values_strict(layer_outputs, output_graph_expr_outputs))); - result = merge_disjoint_bidicts(result, mapping_for_layer); + result = binary_merge_disjoint_bidicts(result, mapping_for_layer); } return result; diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index 498fd6c1bf..dbd968d476 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -5,7 +5,7 @@ #include "utils/bidict/algorithms/bidict_from_keys_and_values.h" #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/bidict/algorithms/exhaustive_relational_join.h" -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include "utils/bidict/algorithms/transform_values.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/map_values.h" @@ -34,7 +34,7 @@ bidict exhaustive_relational_join(pattern_node_outputs.reversed(), matched_layer_output_tensors); - result = merge_disjoint_bidicts(result, mapping); + result = binary_merge_disjoint_bidicts(result, mapping); } return result; diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h new file mode 100644 index 0000000000..bb8ca4f840 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h @@ -0,0 +1,13 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_TRAINING_OPERATION_ATTRS_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_TRAINING_OPERATION_ATTRS_H + +#include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" +#include "op-attrs/operator_type.dtg.h" + +namespace FlexFlow { + +bool training_op_attrs_has_op_type(TrainingOperationAttrs const &, OperatorType); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index b23edc0411..664c615a90 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -13,15 +13,21 @@ #include #include #include +#include "pcg/parallel_computation_graph/parallel_tensor_use_t.h" +#include "utils/containers/require_only_key.h" +#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/containers/map_keys_and_values.h" +#include "utils/containers/transform_pairs.h" namespace FlexFlow { + static bidict get_input_mapping_for_replicate( MappedParallelComputationGraph const &mpcg, parallel_layer_guid_t const &replicate_layer) { - // get_incoming_edges returns map - // replicate has exactly one input + ASSERT(mpcg_get_pcg_op_attrs(mpcg, replicate_layer).is_parallel_replicate()); + auto [input_slot_name, input_edge] = get_only(mpcg_get_incoming_edges(mpcg, replicate_layer)); @@ -33,44 +39,32 @@ static bidict /*slot_name=*/producer_slot); } -static std::unordered_map - get_consumers_of_tensor(MappedParallelComputationGraph const &mpcg, - parallel_tensor_guid_t const &tensor) { - parallel_layer_guid_t producer_layer = mpcg_get_source_layer(mpcg, tensor); - - std::unordered_map result; - // get_outgoing_edges returns unordered_set - for (ParallelComputationGraphEdge const &edge : mpcg_get_outgoing_edges(mpcg, producer_layer)) { - if (get_parallel_tensor(edge) == tensor) { - result.insert( - std::pair{get_dst_layer(edge), get_dst_layer_input_slot_name(edge)}); - } - } - return result; -} - static bidict build_replicated_output_mapping( MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &replicate_layer) { + parallel_tensor_guid_t const &output_tensor_guid) { - auto [output_slot_name, output_tensor_guid] = - get_only(mpcg_get_outgoing_tensors(mpcg, replicate_layer)); - - auto consumers = get_consumers_of_tensor(mpcg, output_tensor_guid); + std::unordered_set consumers = mpcg_get_parallel_tensor_uses(mpcg, output_tensor_guid); ASSERT(!consumers.empty()); // union all consumer bindings — each consumer shard maps to a distinct // (discard_copy, machine) pair since replicas are always on different machines - bidict result; - for (auto const &[consumer_layer, slot_name] : consumers) { - MappedOperatorTaskGroup consumer_mapping = mpcg_get_mapping_for_layer(mpcg, consumer_layer); - bidict binding = - get_tensor_bindings_for_slot_name(consumer_mapping, slot_name); - for (auto const &[p, m] : binding) { - result.equate(p, m); - } - } + bidict result = + merge_disjoint_bidicts( + transform(consumers, + [&](parallel_tensor_use_t const &use) + -> bidict + { + parallel_layer_guid_t consumer_layer = parallel_tensor_use_get_layer(use); + TensorSlotName slot_name = parallel_tensor_use_get_slot(use); + + MappedOperatorTaskGroup consumer_mapping = mpcg_get_mapping_for_layer(mpcg, consumer_layer); + bidict binding = + get_tensor_bindings_for_slot_name(consumer_mapping, slot_name); + + return binding; + })); + return result; } @@ -78,14 +72,19 @@ static DynamicNodeInvocation build_replicate_invocation(parallel_layer_guid_t const &layer, ReplicateAttrs const &attrs, MappedParallelComputationGraph const &mpcg) { - auto [input_slot_name, input_tensor_guid] = - get_only(mpcg_get_incoming_tensors(mpcg, layer).l_to_r()); - - auto incoming = mpcg_get_incoming_tensors(mpcg, layer); - ASSERT(!incoming.empty(), "Replicate layer has no incoming tensors."); + ManyToOne incoming = mpcg_get_incoming_tensors(mpcg, layer); + TensorSlotName input_slot_name = TensorSlotName::INPUT; + parallel_tensor_guid_t input_tensor_guid = require_only_key(incoming.l_to_r(), input_slot_name); ParallelTensorAttrs input_attrs = mpcg_get_parallel_tensor_attrs(mpcg, input_tensor_guid); + + bidict outgoing = mpcg_get_outgoing_tensors(mpcg, layer); + TensorSlotName output_slot_name = TensorSlotName::OUTPUT; + parallel_tensor_guid_t output_tensor_guid = require_only_key(outgoing.l_to_r(), output_slot_name); + ParallelTensorAttrs output_attrs = + mpcg_get_parallel_tensor_attrs(mpcg, output_tensor_guid); + bidict input_mapping = get_input_mapping_for_replicate(mpcg, layer); @@ -93,24 +92,20 @@ static DynamicNodeInvocation /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, /*parallel_tensor_shape=*/input_attrs.shape, /*shard_coord=*/std::nullopt, - /*mapping=*/get_input_mapping_for_replicate(mpcg, layer), + /*mapping=*/input_mapping, /*accessor=*/std::nullopt, /*role=*/std::nullopt, }; - auto [output_slot_name, output_tensor_guid] = - get_only(mpcg_get_outgoing_tensors(mpcg, layer)); - ParallelTensorAttrs output_attrs = - mpcg_get_parallel_tensor_attrs(mpcg, output_tensor_guid); - DynamicValueAttrs output_value{ /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, /*parallel_tensor_shape=*/output_attrs.shape, /*shard_coord=*/std::nullopt, - /*mapping=*/build_replicated_output_mapping(mpcg, layer), + /*mapping=*/build_replicated_output_mapping(mpcg, output_tensor_guid), /*accessor=*/std::nullopt, /*role=*/std::nullopt, }; + DynamicNodeAttrs node_attrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, @@ -122,85 +117,92 @@ static DynamicNodeInvocation DynamicNodeInvocation invocation_node{ /*inputs=*/{ - {DynamicTensorSlot{input_slot_name, std::nullopt}, input_value}}, + { + DynamicTensorSlot{input_slot_name, std::nullopt}, + input_value, + }, + }, /*node_attrs=*/node_attrs, - /*outputs=*/ - {{DynamicTensorSlot{output_slot_name, std::nullopt}, output_value}}, + /*outputs=*/{ + { + DynamicTensorSlot{output_slot_name, std::nullopt}, + output_value, + }, + }, }; + return invocation_node; } DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( MappedParallelComputationGraph const &mpcg) { - DynamicOpenDataflowGraph result = make_empty_dynamic_open_dataflow_graph(); ParallelComputationGraph pcg = pcg_from_mpcg(mpcg); - for (auto const &[layer, attrs] : get_parallel_layer_attrs_mapping(pcg)) { - if (attrs.op_attrs.has()) { + auto mk_invocation = [&](parallel_layer_guid_t layer, ParallelLayerAttrs const &attrs) + -> DynamicNodeInvocation + { + if (attrs.op_attrs.is_parallel_replicate()) { // build replicate invocation DynamicNodeInvocation repl_inv = build_replicate_invocation( - layer, attrs.op_attrs.get(), mpcg); - result.invocations.emplace(repl_inv); - continue; - } - - DynamicNodeAttrs result_attrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/mpcg_get_mapping_for_layer(mpcg, layer), - /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs}, - /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, - /*per_device_op_state=*/std::nullopt, + layer, attrs.op_attrs.require_parallel_replicate(), mpcg); + return repl_inv; + } else { + DynamicNodeAttrs result_attrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/mpcg_get_mapping_for_layer(mpcg, layer), + /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs}, + /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, + /*per_device_op_state=*/std::nullopt, + }; + + auto mk_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + }; + }; + + auto mk_value_attrs = [&](parallel_tensor_guid_t const &tensor) -> DynamicValueAttrs + { + ParallelTensorAttrs attrs = + get_parallel_tensor_attrs(pcg, tensor); + + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, + /*parallel_tensor_shape=*/attrs.shape, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + }; + + std::unordered_map result_inputs = + map_keys_and_values(get_incoming_tensors(pcg, layer), + mk_slot, + mk_value_attrs); + + std::unordered_map result_outputs = + map_keys_and_values(get_outgoing_tensors(pcg, layer), + mk_slot, + mk_value_attrs); + + DynamicNodeInvocation invocation = DynamicNodeInvocation{ + /*inputs=*/result_inputs, + /*node_attrs=*/result_attrs, + /*outputs=*/result_outputs, + }; + + return invocation; }; + }; - std::unordered_map result_inputs = - transform(get_incoming_tensors(pcg, layer), - [&](TensorSlotName const &slot_name, - parallel_tensor_guid_t const &tensor) { - ParallelTensorAttrs attrs = - get_parallel_tensor_attrs(pcg, tensor); - return std::pair{ - DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - }, - DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, - /*parallel_tensor_shape=*/attrs.shape, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }, - }; - }); - std::unordered_map result_outputs = - transform(get_outgoing_tensors(pcg, layer), - [&](TensorSlotName const &slot_name, - parallel_tensor_guid_t const &tensor) { - ParallelTensorAttrs attrs = - get_parallel_tensor_attrs(pcg, tensor); - return std::pair{ - DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - }, - DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, - /*parallel_tensor_shape=*/attrs.shape, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }, - }; - }); - - result.invocations.emplace(result_inputs, result_attrs, result_outputs); - } - - return result; + return dynamic_open_dataflow_graph_from_invocation_set( + transform_pairs( + unordered_set_of(get_parallel_layer_attrs_mapping(pcg)), + mk_invocation)); } } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index faa1e186c3..25958b5cb7 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -5,6 +5,7 @@ #include "utils/containers/get_only.h" #include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/transform.h" +#include "task-spec/dynamic_graph/training_operation_attrs.h" namespace FlexFlow { @@ -88,6 +89,8 @@ DynamicNodeInvocation perform_fwd_pass_expansion_for_invocation( DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( DynamicNodeInvocation const &invocation) { + TrainingOperationAttrs op_attrs = assert_unwrap(invocation.node_attrs.op_attrs); + auto to_fwd = [](DynamicTensorSlot const &k, DynamicValueAttrs const &v) { return std::pair{ pass_expand_slot(k, FwbTensorType::FORWARD), @@ -102,56 +105,37 @@ DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( }; }; - return DynamicNodeInvocation{ - /*inputs=*/ - merge_disjoint_maps(std::vector{ - transform(invocation.inputs, to_fwd), - transform(invocation.outputs, to_fwd), - transform(invocation.outputs, to_grad), - }), - /*node_attrs=*/ - pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), - /*outputs=*/ - transform(invocation.inputs, to_grad), - }; -} -static std::unordered_set - perform_pass_expansion_for_replicate( - DynamicNodeInvocation const &invocation) { - - auto const &[input_slot, input] = get_only(invocation.inputs); - auto const &[output_slot, output] = get_only(invocation.outputs); - - // forward: INPUT/FWD → OUTPUT/FWD (copy to replicas) - DynamicNodeInvocation fwd{ - /*inputs=*/{{pass_expand_slot(input_slot, FwbTensorType::FORWARD), - pass_expand_value(input, FwbTensorType::FORWARD)}}, - /*node_attrs=*/ - pass_expand_node(invocation.node_attrs, DynamicTaskType::FWD), - /*outputs=*/ - {{pass_expand_slot(output_slot, FwbTensorType::FORWARD), - pass_expand_value(output, FwbTensorType::FORWARD)}}, - }; - - // backward: OUTPUT/FWD + OUTPUT/GRAD → INPUT/GRAD (reduce gradients) - // The backward node needs the mapping from the output (replicated) - // so it knows which replicas to reduce from - DynamicNodeAttrs bwd_node_attrs = invocation.node_attrs; - bwd_node_attrs.task_type = DynamicTaskType::BWD; - - DynamicNodeInvocation bwd{ - /*inputs=*/{ - {pass_expand_slot(output_slot, FwbTensorType::FORWARD), - pass_expand_value(output, FwbTensorType::FORWARD)}, - {pass_expand_slot(output_slot, FwbTensorType::GRADIENT), - pass_expand_value(output, FwbTensorType::GRADIENT)}, - }, - /*node_attrs=*/bwd_node_attrs, - /*outputs=*/ - {{pass_expand_slot(input_slot, FwbTensorType::GRADIENT), - pass_expand_value(input, FwbTensorType::GRADIENT)}}, + if (training_op_attrs_has_op_type(op_attrs, OperatorType::REPLICATE)) { + auto [input_slot, input] = get_only(invocation.inputs); + auto [output_slot, output] = get_only(invocation.outputs); + + DynamicNodeInvocation bwd{ + /*inputs=*/{ + to_fwd(output_slot, output), + to_grad(output_slot, output), + }, + /*node_attrs=*/ + pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), + /*outputs=*/{ + to_grad(input_slot, input), + }, + }; + + return bwd; + } else { + return DynamicNodeInvocation{ + /*inputs=*/ + merge_disjoint_maps(std::vector{ + transform(invocation.inputs, to_fwd), + transform(invocation.outputs, to_fwd), + transform(invocation.outputs, to_grad), + }), + /*node_attrs=*/ + pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), + /*outputs=*/ + transform(invocation.inputs, to_grad), + }; }; - return {fwd, bwd}; } DynamicOpenDataflowGraph @@ -161,9 +145,6 @@ DynamicOpenDataflowGraph DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [](DynamicNodeInvocation const &invocation) { - if (is_replicate_attrs(invocation.node_attrs)) { - return perform_pass_expansion_for_replicate(invocation); - } if (invocation.inputs.empty()) { return std::unordered_set{ perform_fwd_pass_expansion_for_invocation(invocation), diff --git a/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc new file mode 100644 index 0000000000..d1452242ca --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc @@ -0,0 +1,21 @@ +#include "task-spec/dynamic_graph/training_operation_attrs.h" +#include "op-attrs/pcg_operator_attrs.h" +#include "utils/overload.h" + +namespace FlexFlow { + +bool training_op_attrs_has_op_type(TrainingOperationAttrs const &op_attrs, OperatorType op_type) { + return op_attrs.visit(overload { + [&](PCGOperatorAttrs const &pcg_op_attrs) -> bool { + return pcg_op_attrs_get_op_type(pcg_op_attrs) == op_type; + }, + [](LossAttrs const &) -> bool { + return false; + }, + [](CopyAttrs const &) -> bool { + return false; + }, + }); +} + +} // namespace FlexFlow diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc index fb087f5295..ed22a8cbde 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc @@ -2,6 +2,7 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include +#include "op-attrs/ops/element_unary.h" using namespace ::FlexFlow; @@ -36,6 +37,19 @@ TEST_SUITE(FF_TEST_SUITE) { dynamic_layer_guid_t layer_guid{parallel_layer_guid_t{Node{20}}}; + TrainingOperationAttrs op_attrs = + TrainingOperationAttrs{ + PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/8_p, + /*use_bias=*/true, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, + }, + }; + DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); @@ -46,14 +60,13 @@ TEST_SUITE(FF_TEST_SUITE) { {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, {mk_slot(TensorSlotName::WEIGHT, std::nullopt), v2}, {mk_slot(TensorSlotName::BIAS, std::nullopt), v1}, - {mk_slot(TensorSlotName::SCALE, std::nullopt), v1}, }, /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, - /*op_attrs=*/std::nullopt, + /*op_attrs=*/op_attrs, /*layer_guid=*/layer_guid, /*per_device_op_state=*/std::nullopt, }, @@ -79,14 +92,13 @@ TEST_SUITE(FF_TEST_SUITE) { {mk_slot(TensorSlotName::INPUT, fwd_role), v1_fwd}, {mk_slot(TensorSlotName::WEIGHT, fwd_role), v2_fwd}, {mk_slot(TensorSlotName::BIAS, fwd_role), v1_fwd}, - {mk_slot(TensorSlotName::SCALE, fwd_role), v1_fwd}, }, /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/DynamicTaskType::FWD, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, - /*op_attrs=*/std::nullopt, + /*op_attrs=*/op_attrs, /*layer_guid=*/layer_guid, /*per_device_op_state=*/std::nullopt, }, @@ -130,88 +142,163 @@ TEST_SUITE(FF_TEST_SUITE) { dynamic_layer_guid_t layer_guid{parallel_layer_guid_t{Node{20}}}; - DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { - DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); - DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); - DynamicValueAttrs v3 = mk_value_attrs(2, std::nullopt); - - return DynamicNodeInvocation{ - /*inputs=*/{ - {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, - {mk_slot(TensorSlotName::WEIGHT, std::nullopt), v2}, - {mk_slot(TensorSlotName::BIAS, std::nullopt), v1}, - {mk_slot(TensorSlotName::SCALE, std::nullopt), v1}, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v3}, - }, - }; - }(); - - DynamicNodeInvocation result = - perform_bwd_pass_expansion_for_invocation(invocation); - - DynamicNodeInvocation correct = [&]() -> DynamicNodeInvocation { - DynamicTensorRole fwd_role = DynamicTensorRole{FwbTensorType::FORWARD}; - DynamicTensorRole grad_role = DynamicTensorRole{FwbTensorType::GRADIENT}; - - DynamicValueAttrs v1_fwd = mk_value_attrs(0, fwd_role); - DynamicValueAttrs v2_fwd = mk_value_attrs(1, fwd_role); - DynamicValueAttrs v3_fwd = mk_value_attrs(2, fwd_role); - DynamicValueAttrs v1_grad = mk_value_attrs(0, grad_role); - DynamicValueAttrs v2_grad = mk_value_attrs(1, grad_role); - DynamicValueAttrs v3_grad = mk_value_attrs(2, grad_role); - - return DynamicNodeInvocation{ - /*inputs=*/{ - {mk_slot(TensorSlotName::INPUT, fwd_role), v1_fwd}, - {mk_slot(TensorSlotName::WEIGHT, fwd_role), v2_fwd}, - {mk_slot(TensorSlotName::BIAS, fwd_role), v1_fwd}, - {mk_slot(TensorSlotName::SCALE, fwd_role), v1_fwd}, - {mk_slot(TensorSlotName::OUTPUT, fwd_role), v3_fwd}, - {mk_slot(TensorSlotName::OUTPUT, grad_role), v3_grad}, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*pass_type=*/DynamicTaskType::BWD, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, + DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); + DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); + DynamicValueAttrs v3 = mk_value_attrs(2, std::nullopt); + + DynamicTensorRole fwd_role = DynamicTensorRole{FwbTensorType::FORWARD}; + DynamicTensorRole grad_role = DynamicTensorRole{FwbTensorType::GRADIENT}; + + DynamicValueAttrs v1_fwd = mk_value_attrs(0, fwd_role); + DynamicValueAttrs v2_fwd = mk_value_attrs(1, fwd_role); + DynamicValueAttrs v3_fwd = mk_value_attrs(2, fwd_role); + DynamicValueAttrs v1_grad = mk_value_attrs(0, grad_role); + DynamicValueAttrs v2_grad = mk_value_attrs(1, grad_role); + DynamicValueAttrs v3_grad = mk_value_attrs(2, grad_role); + + SUBCASE("normal operator") { + TrainingOperationAttrs op_attrs = + TrainingOperationAttrs{ + PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/8_p, + /*use_bias=*/true, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }, - /*outputs=*/ - { - {mk_slot(TensorSlotName::INPUT, grad_role), v1_grad}, - {mk_slot(TensorSlotName::WEIGHT, grad_role), v2_grad}, - {mk_slot(TensorSlotName::BIAS, grad_role), v1_grad}, - {mk_slot(TensorSlotName::SCALE, grad_role), v1_grad}, + }; + + DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, + {mk_slot(TensorSlotName::WEIGHT, std::nullopt), v2}, + {mk_slot(TensorSlotName::BIAS, std::nullopt), v1}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v3}, + }, + }; + }(); + + DynamicNodeInvocation result = + perform_bwd_pass_expansion_for_invocation(invocation); + + DynamicNodeInvocation correct = [&]() -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, fwd_role), v1_fwd}, + {mk_slot(TensorSlotName::WEIGHT, fwd_role), v2_fwd}, + {mk_slot(TensorSlotName::BIAS, fwd_role), v1_fwd}, + {mk_slot(TensorSlotName::OUTPUT, fwd_role), v3_fwd}, + {mk_slot(TensorSlotName::OUTPUT, grad_role), v3_grad}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*pass_type=*/DynamicTaskType::BWD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::INPUT, grad_role), v1_grad}, + {mk_slot(TensorSlotName::WEIGHT, grad_role), v2_grad}, + {mk_slot(TensorSlotName::BIAS, grad_role), v1_grad}, + }, + }; + }(); + + ASSERT(result == correct); + } + + SUBCASE("replicate operator optimization") { + TrainingOperationAttrs op_attrs = + TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }, - }; - }(); - - ASSERT(result == correct); + }; + + DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v2}, + }, + }; + }(); + + DynamicNodeInvocation result = + perform_bwd_pass_expansion_for_invocation(invocation); + + DynamicNodeInvocation correct = [&]() -> DynamicNodeInvocation { + DynamicTensorRole fwd_role = DynamicTensorRole{FwbTensorType::FORWARD}; + DynamicTensorRole grad_role = DynamicTensorRole{FwbTensorType::GRADIENT}; + + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::OUTPUT, fwd_role), v2_fwd}, + {mk_slot(TensorSlotName::OUTPUT, grad_role), v2_grad}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*pass_type=*/DynamicTaskType::BWD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::INPUT, grad_role), v1_grad}, + }, + }; + }(); + + ASSERT(result == correct); + } } TEST_CASE("perform_pass_expansion(DynamicOpenDataflowGraph)") { auto mk_node_attrs = [](size_t layer_id, + TrainingOperationAttrs const &op_attrs, std::optional const &pass_type) -> DynamicNodeAttrs { return DynamicNodeAttrs{ /*pass_type=*/pass_type, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, - /*op_attrs=*/std::nullopt, + /*op_attrs=*/op_attrs, /*layer_guid=*/ dynamic_layer_guid_t{parallel_layer_guid_t{Node{layer_id}}}, /*per_device_op_state=*/std::nullopt, @@ -236,9 +323,32 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; + TrainingOperationAttrs input_op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + InputAttrs{ + TensorShape{ + TensorDims{ + FFOrdered{ + 4_p, + 8_p, + }, + }, + DataType::FLOAT, + }, + }, + }, + }; + + TrainingOperationAttrs relu_op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + make_relu_attrs(), + }, + }; + + DynamicOpenDataflowGraph input = [&]() -> DynamicOpenDataflowGraph { - DynamicNodeAttrs n1 = mk_node_attrs(10, std::nullopt); - DynamicNodeAttrs n2 = mk_node_attrs(11, std::nullopt); + DynamicNodeAttrs n1 = mk_node_attrs(10, input_op_attrs, std::nullopt); + DynamicNodeAttrs n2 = mk_node_attrs(11, relu_op_attrs, std::nullopt); DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); @@ -286,10 +396,10 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicOpenDataflowGraph result = perform_pass_expansion(input); DynamicOpenDataflowGraph correct = [&]() -> DynamicOpenDataflowGraph { - DynamicNodeAttrs n1_fwd = mk_node_attrs(10, DynamicTaskType::FWD); - DynamicNodeAttrs n2_fwd = mk_node_attrs(11, DynamicTaskType::FWD); - DynamicNodeAttrs n1_bwd = mk_node_attrs(10, DynamicTaskType::BWD); - DynamicNodeAttrs n2_bwd = mk_node_attrs(11, DynamicTaskType::BWD); + DynamicNodeAttrs n1_fwd = mk_node_attrs(10, input_op_attrs, DynamicTaskType::FWD); + DynamicNodeAttrs n2_fwd = mk_node_attrs(11, relu_op_attrs, DynamicTaskType::FWD); + DynamicNodeAttrs n1_bwd = mk_node_attrs(10, input_op_attrs, DynamicTaskType::BWD); + DynamicNodeAttrs n2_bwd = mk_node_attrs(11, relu_op_attrs, DynamicTaskType::BWD); DynamicValueAttrs v1_activation = mk_value_attrs(0, mk_dynamic_tensor_role_fwd()); diff --git a/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h b/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h new file mode 100644 index 0000000000..5b0bb45910 --- /dev/null +++ b/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h @@ -0,0 +1,37 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BINARY_MERGE_DISJOINT_BIDICTS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BINARY_MERGE_DISJOINT_BIDICTS_H + +#include "utils/bidict/algorithms/left_entries.h" +#include "utils/bidict/algorithms/right_entries.h" +#include "utils/bidict/bidict.h" +#include "utils/containers/are_disjoint.h" +#include "utils/exception.h" + +namespace FlexFlow { + +template +bidict binary_merge_disjoint_bidicts(bidict const &lhs, + bidict const &rhs) { + if (!are_disjoint(left_entries(lhs), left_entries(rhs))) { + throw mk_runtime_error( + fmt::format("Left entries of {} and {} are non-disjoint", lhs, rhs)); + } + if (!are_disjoint(right_entries(lhs), right_entries(rhs))) { + throw mk_runtime_error( + fmt::format("Right entries of {} and {} are non-disjoint", lhs, rhs)); + } + + bidict result; + for (auto const &kv : lhs) { + result.equate_strict(kv.first, kv.second); + } + for (auto const &kv : rhs) { + result.equate_strict(kv.first, kv.second); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h index 97e7334c26..f2104fd113 100644 --- a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h +++ b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h @@ -1,35 +1,22 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H -#include "utils/bidict/algorithms/left_entries.h" -#include "utils/bidict/algorithms/right_entries.h" -#include "utils/bidict/bidict.h" -#include "utils/containers/are_disjoint.h" -#include "utils/exception.h" +#include "utils/containers/foldl.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" namespace FlexFlow { -template -bidict merge_disjoint_bidicts(bidict const &lhs, - bidict const &rhs) { - if (!are_disjoint(left_entries(lhs), left_entries(rhs))) { - throw mk_runtime_error( - fmt::format("Left entries of {} and {} are non-disjoint", lhs, rhs)); - } - if (!are_disjoint(right_entries(lhs), right_entries(rhs))) { - throw mk_runtime_error( - fmt::format("Right entries of {} and {} are non-disjoint", lhs, rhs)); - } - - bidict result; - for (auto const &kv : lhs) { - result.equate(kv.first, kv.second); - } - for (auto const &kv : rhs) { - result.equate(kv.first, kv.second); - } - - return result; +template +bidict merge_disjoint_bidicts(C const &c) { + bidict empty = {}; + return foldl(c, + /*init=*/empty, + [](bidict const &lhs, + bidict const &rhs) { + return binary_merge_disjoint_bidicts(lhs, rhs); + }); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/bidict/bidict.h b/lib/utils/include/utils/bidict/bidict.h index 5dbd1c603d..2d8c5d23a8 100644 --- a/lib/utils/include/utils/bidict/bidict.h +++ b/lib/utils/include/utils/bidict/bidict.h @@ -213,6 +213,14 @@ struct bidict { return this->fwd_map; } + std::unordered_map const &l_to_r() const { + return this->fwd_map; + } + + std::unordered_map const &r_to_l() const { + return this->bwd_map; + } + bidict(std::unordered_map const &fwd_map, std::unordered_map const &bwd_map) : fwd_map(fwd_map), bwd_map(bwd_map) {} diff --git a/lib/utils/include/utils/containers/transform_pairs.h b/lib/utils/include/utils/containers/transform_pairs.h new file mode 100644 index 0000000000..c01b50554f --- /dev/null +++ b/lib/utils/include/utils/containers/transform_pairs.h @@ -0,0 +1,46 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_TRANSFORM_PAIRS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_TRANSFORM_PAIRS_H + +#include "utils/containers/transform.h" + +namespace FlexFlow { + +template > +std::vector transform_pairs(std::vector> const &c, F &&f) { + auto ff = [&](std::pair const &p) -> Out { + return f(p.first, p.second); + }; + + return transform(c, ff); +} + +template > +std::unordered_set transform_pairs(std::unordered_set> const &c, F &&f) { + auto ff = [&](std::pair const &p) -> Out { + return f(p.first, p.second); + }; + + return transform(c, ff); +} + +template > +std::set transform_pairs(std::set> const &c, F &&f) { + auto ff = [&](std::pair const &p) -> Out { + return f(p.first, p.second); + }; + + return transform(c, ff); +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc new file mode 100644 index 0000000000..8650de44f6 --- /dev/null +++ b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc @@ -0,0 +1,12 @@ +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template + bidict binary_merge_disjoint_bidicts(bidict const &, bidict const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc index 754b8d2e90..2c27821d3b 100644 --- a/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc +++ b/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc @@ -1 +1,11 @@ #include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template bidict merge_disjoint_bidicts(std::vector> const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/transform_pairs.cc b/lib/utils/src/utils/containers/transform_pairs.cc new file mode 100644 index 0000000000..241f1ad425 --- /dev/null +++ b/lib/utils/src/utils/containers/transform_pairs.cc @@ -0,0 +1,17 @@ +#include "utils/containers/transform_pairs.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using L = value_type<0>; +using R = value_type<1>; +using Out = value_type<2>; +using F = std::function; + +template + std::vector transform_pairs(std::vector> const &, F &&); + +template + std::unordered_set transform_pairs(std::unordered_set> const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/test/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc b/lib/utils/test/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc similarity index 72% rename from lib/utils/test/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc rename to lib/utils/test/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc index 0a1babd9f9..8a3371b8d8 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc @@ -1,17 +1,17 @@ -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("merge_disjoint_bidicts") { + TEST_CASE("binary_merge_disjoint_bidicts") { SUBCASE("disjoint keys and values") { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{3, "three"}, {4, "four"}}; - bidict result = merge_disjoint_bidicts(bd1, bd2); + bidict result = binary_merge_disjoint_bidicts(bd1, bd2); bidict correct = { {1, "one"}, {2, "two"}, {3, "three"}, {4, "four"}}; @@ -22,21 +22,21 @@ TEST_SUITE(FF_TEST_SUITE) { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{2, "three"}, {3, "four"}}; - CHECK_THROWS(merge_disjoint_bidicts(bd1, bd2)); + CHECK_THROWS(binary_merge_disjoint_bidicts(bd1, bd2)); } SUBCASE("overlapping key, same associated value") { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{2, "two"}, {3, "three"}}; - CHECK_THROWS(merge_disjoint_bidicts(bd1, bd2)); + CHECK_THROWS(binary_merge_disjoint_bidicts(bd1, bd2)); } SUBCASE("overlapping values") { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{3, "two"}, {4, "four"}}; - CHECK_THROWS(merge_disjoint_bidicts(bd1, bd2)); + CHECK_THROWS(binary_merge_disjoint_bidicts(bd1, bd2)); } } } From fdf4fe5e74d4ede4cc21da923ee0aaedf5771351 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 15 May 2026 21:42:09 -0700 Subject: [PATCH 07/35] Remove unnecessary is_replicate_attrs function --- lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc | 5 ----- 1 file changed, 5 deletions(-) diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index 25958b5cb7..f4960fe67a 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -31,11 +31,6 @@ bool graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &g) { g, node_is_pass_expanded, value_is_pass_expanded, slot_is_pass_expanded); } -static bool is_replicate_attrs(DynamicNodeAttrs const &n) { - return n.op_attrs.has_value() && n.op_attrs.value().has() && - n.op_attrs.value().get().has(); -} - DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &s, FwbTensorType tensor_type) { ASSERT(!slot_is_pass_expanded(s)); From dce15e23118a128c65d7c35f7e57da91011329f4 Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 21 May 2026 12:14:33 -0700 Subject: [PATCH 08/35] Format. --- .../test/src/op-attrs/ops/element_unary.cc | 5 +- .../mapped_parallel_computation_graph.h | 20 +-- .../parallel_computation_graph.h | 2 +- .../parallel_tensor_use_t.h | 5 +- .../mapped_parallel_computation_graph.cc | 43 +++--- .../parallel_computation_graph.cc | 9 +- .../parallel_tensor_use_t.cc | 3 +- .../sub_parallel_computation_graph.h | 2 +- .../apply_substitution/apply_substitution.cc | 7 +- .../src/substitutions/pcg_pattern_match.cc | 2 +- .../sub_parallel_computation_graph.cc | 2 +- .../dynamic_graph/training_operation_attrs.h | 5 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 129 +++++++++--------- .../task-spec/dynamic_graph/pass_expansion.cc | 16 ++- .../dynamic_graph/training_operation_attrs.cc | 19 ++- .../task-spec/dynamic_graph/pass_expansion.cc | 95 ++++++------- .../algorithms/merge_disjoint_bidicts.h | 5 +- .../utils/containers/transform_pairs.h | 3 +- .../get_kwarg_dataflow_value_uses.h | 33 ++--- .../include/utils/many_to_one/many_to_one.h | 2 +- .../include/utils/one_to_many/one_to_many.h | 2 +- .../binary_merge_disjoint_bidicts.cc | 4 +- .../src/utils/containers/transform_pairs.cc | 8 +- .../get_kwarg_dataflow_value_uses.cc | 8 +- 24 files changed, 210 insertions(+), 219 deletions(-) diff --git a/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc b/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc index 8b2555610e..09e49a123c 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/element_unary.cc @@ -56,8 +56,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("discard copy degree > 1") { positive_int degree = 2_p; - ParallelTensorShape par_input = make_input( - SumDegree{1_p}, DiscardCopyDegree{degree}, 1_p, 1_p, 1_p); + ParallelTensorShape par_input = + make_input(SumDegree{1_p}, DiscardCopyDegree{degree}, 1_p, 1_p, 1_p); tl::expected result = get_output_shape(attrs, par_input); @@ -74,6 +74,5 @@ TEST_SUITE(FF_TEST_SUITE) { make_input( SumDegree{degree}, DiscardCopyDegree{1_p}, 1_p, 1_p, 1_p))); } - } } diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h index 6c24d4c1e1..a2afdb7914 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h @@ -15,22 +15,24 @@ MappedOperatorTaskGroup ParallelComputationGraph pcg_from_mpcg(MappedParallelComputationGraph const &); -parallel_layer_guid_t mpcg_get_source_layer(MappedParallelComputationGraph const &, - parallel_tensor_guid_t const &); +parallel_layer_guid_t + mpcg_get_source_layer(MappedParallelComputationGraph const &, + parallel_tensor_guid_t const &); PCGOperatorAttrs mpcg_get_pcg_op_attrs(MappedParallelComputationGraph const &, parallel_layer_guid_t const &); -ParallelTensorAttrs mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &, - parallel_tensor_guid_t const &); +ParallelTensorAttrs + mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &, + parallel_tensor_guid_t const &); std::unordered_map - mpcg_get_incoming_edges(MappedParallelComputationGraph const &, - parallel_layer_guid_t const &); + mpcg_get_incoming_edges(MappedParallelComputationGraph const &, + parallel_layer_guid_t const &); std::unordered_set - mpcg_get_outgoing_edges(MappedParallelComputationGraph const &, - parallel_layer_guid_t const &); + mpcg_get_outgoing_edges(MappedParallelComputationGraph const &, + parallel_layer_guid_t const &); ManyToOne mpcg_get_incoming_tensors(MappedParallelComputationGraph const &, @@ -38,7 +40,7 @@ ManyToOne bidict mpcg_get_outgoing_tensors(MappedParallelComputationGraph const &, - parallel_layer_guid_t const &); + parallel_layer_guid_t const &); std::unordered_set mpcg_get_edges(MappedParallelComputationGraph const &); diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h index 1b2d5a0b67..9764e40627 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h @@ -10,8 +10,8 @@ #include "pcg/parallel_computation_graph/parallel_layer_added_result.dtg.h" #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h" -#include #include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" +#include namespace FlexFlow { diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h index 88f1512149..f5e5575632 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_use_t.h @@ -1,12 +1,13 @@ #ifndef _FLEXFLOW_LIB_PCG_INCLUDE_PCG_PARALLEL_COMPUTATION_GRAPH_PARALLEL_TENSOR_USE_T_H #define _FLEXFLOW_LIB_PCG_INCLUDE_PCG_PARALLEL_COMPUTATION_GRAPH_PARALLEL_TENSOR_USE_T_H -#include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" +#include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" namespace FlexFlow { -parallel_layer_guid_t parallel_tensor_use_get_layer(parallel_tensor_use_t const &); +parallel_layer_guid_t + parallel_tensor_use_get_layer(parallel_tensor_use_t const &); TensorSlotName parallel_tensor_use_get_slot(parallel_tensor_use_t const &); } // namespace FlexFlow diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 3b996ccdab..fc1dff504b 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -2,13 +2,13 @@ #include "op-attrs/pcg_operator_attrs.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_attrs.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" +#include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/bidict/algorithms/transform_keys.h" #include "utils/containers/transform.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/materialize_labelled_kwarg_dataflow_graph_view.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/rewrite_labelled_kwarg_dataflow_graph_node_labels.h" -#include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/many_to_one/many_to_one_from_map.h" namespace FlexFlow { @@ -48,63 +48,56 @@ ParallelComputationGraph }; } -parallel_layer_guid_t mpcg_get_source_layer(MappedParallelComputationGraph const &mpcg, - parallel_tensor_guid_t const &t) -{ +parallel_layer_guid_t + mpcg_get_source_layer(MappedParallelComputationGraph const &mpcg, + parallel_tensor_guid_t const &t) { return get_source_layer(pcg_from_mpcg(mpcg), t); } -PCGOperatorAttrs mpcg_get_pcg_op_attrs(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) -{ +PCGOperatorAttrs + mpcg_get_pcg_op_attrs(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) { return pcg_get_op_attrs(pcg_from_mpcg(mpcg), l); } -ParallelTensorAttrs mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &mpcg, - parallel_tensor_guid_t const &t) -{ +ParallelTensorAttrs + mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &mpcg, + parallel_tensor_guid_t const &t) { return get_parallel_tensor_attrs(pcg_from_mpcg(mpcg), t); } std::unordered_map - mpcg_get_incoming_edges(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) -{ + mpcg_get_incoming_edges(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) { return get_incoming_edges(pcg_from_mpcg(mpcg), l); } std::unordered_set - mpcg_get_outgoing_edges(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) -{ + mpcg_get_outgoing_edges(MappedParallelComputationGraph const &mpcg, + parallel_layer_guid_t const &l) { return get_outgoing_edges(pcg_from_mpcg(mpcg), l); } ManyToOne mpcg_get_incoming_tensors(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) -{ + parallel_layer_guid_t const &l) { return many_to_one_from_map(get_incoming_tensors(pcg_from_mpcg(mpcg), l)); } - bidict mpcg_get_outgoing_tensors(MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &l) -{ + parallel_layer_guid_t const &l) { return bidict_from_map(get_outgoing_tensors(pcg_from_mpcg(mpcg), l)); } std::unordered_set - mpcg_get_edges(MappedParallelComputationGraph const &mpcg) -{ + mpcg_get_edges(MappedParallelComputationGraph const &mpcg) { return get_edges(pcg_from_mpcg(mpcg)); } std::unordered_set mpcg_get_parallel_tensor_uses(MappedParallelComputationGraph const &mpcg, - parallel_tensor_guid_t const &t) -{ + parallel_tensor_guid_t const &t) { return pcg_get_parallel_tensor_uses(pcg_from_mpcg(mpcg), t); } diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index 2c5197242d..5098cadafe 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -28,6 +28,7 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.h" +#include "utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_edges_for_node.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h" @@ -36,7 +37,6 @@ #include "utils/graph/node/node.dtg.h" #include "utils/record_formatter.h" #include -#include "utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h" namespace FlexFlow { @@ -209,18 +209,15 @@ std::unordered_map std::unordered_set pcg_get_parallel_tensor_uses(ParallelComputationGraph const &pcg, - parallel_tensor_guid_t const &t) -{ + parallel_tensor_guid_t const &t) { std::unordered_set> raw_uses = - get_kwarg_dataflow_value_uses(pcg.raw_graph, - t.raw_graph_output); + get_kwarg_dataflow_value_uses(pcg.raw_graph, t.raw_graph_output); return transform(raw_uses, [](KwargDataflowInput const &i) { return parallel_tensor_use_t{i}; }); } - std::unordered_set get_initial_layers(ParallelComputationGraph const &pcg) { std::unordered_set raw_sources = get_initial_nodes(pcg.raw_graph); diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc index e93341d312..71a9cadf1c 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_tensor_use_t.cc @@ -2,7 +2,8 @@ namespace FlexFlow { -parallel_layer_guid_t parallel_tensor_use_get_layer(parallel_tensor_use_t const &u) { +parallel_layer_guid_t + parallel_tensor_use_get_layer(parallel_tensor_use_t const &u) { return parallel_layer_guid_t{u.raw_dataflow_input.node}; } diff --git a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h index 26c98e915c..2a3dc8bbb8 100644 --- a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h +++ b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h @@ -49,7 +49,7 @@ std::unordered_set get_subgraph_outgoing_edges( std::unordered_set get_open_parallel_tensor_uses(SubParallelComputationGraph const &, - open_parallel_tensor_guid_t const &); + open_parallel_tensor_guid_t const &); SubParallelComputationGraphData get_sub_pcg_data(SubParallelComputationGraph const &); diff --git a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc index a56555550f..f2686f7cf7 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc @@ -109,9 +109,10 @@ SubParallelComputationGraph apply_substitution_from_output_result( input_parallel_tensor_guid_t output_graph_input = output_expr_to_result_sub_pcg_mapping.input_mapping.at_r( output_expr_input); - std::unordered_set uses = get_open_parallel_tensor_uses( - substitution_output_graph, - open_parallel_tensor_guid_from_input(output_graph_input)); + std::unordered_set uses = + get_open_parallel_tensor_uses( + substitution_output_graph, + open_parallel_tensor_guid_from_input(output_graph_input)); for (parallel_tensor_use_t const &use : uses) { SubParallelComputationGraphEdge new_edge = subpcg_edge_from_tensor_and_use(base_graph_tensor, use); diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index dbd968d476..85a0493e33 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -4,8 +4,8 @@ #include "substitutions/unlabelled/unlabelled_graph_pattern.h" #include "utils/bidict/algorithms/bidict_from_keys_and_values.h" #include "utils/bidict/algorithms/bidict_from_map.h" -#include "utils/bidict/algorithms/exhaustive_relational_join.h" #include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" +#include "utils/bidict/algorithms/exhaustive_relational_join.h" #include "utils/bidict/algorithms/transform_values.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/map_values.h" diff --git a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc index 990975bff9..c0c05ad5b1 100644 --- a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc +++ b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc @@ -132,7 +132,7 @@ std::unordered_set get_subgraph_incoming_edges( std::unordered_set get_open_parallel_tensor_uses(SubParallelComputationGraph const &spcg, - open_parallel_tensor_guid_t const &t) { + open_parallel_tensor_guid_t const &t) { std::unordered_set> raw_uses = get_open_kwarg_dataflow_value_uses(spcg.raw_graph, t.raw_open_dataflow_value); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h index bb8ca4f840..9caea8c341 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h @@ -1,12 +1,13 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_TRAINING_OPERATION_ATTRS_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_TRAINING_OPERATION_ATTRS_H -#include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" #include "op-attrs/operator_type.dtg.h" +#include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" namespace FlexFlow { -bool training_op_attrs_has_op_type(TrainingOperationAttrs const &, OperatorType); +bool training_op_attrs_has_op_type(TrainingOperationAttrs const &, + OperatorType); } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 664c615a90..7a149787b9 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -5,19 +5,19 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" #include "pcg/parallel_computation_graph/parallel_tensor_attrs.dtg.h" +#include "pcg/parallel_computation_graph/parallel_tensor_use_t.h" #include "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" #include "utils/containers/generate_map.h" #include "utils/containers/get_only.h" +#include "utils/containers/map_keys_and_values.h" +#include "utils/containers/require_only_key.h" +#include "utils/containers/transform_pairs.h" #include #include #include -#include "pcg/parallel_computation_graph/parallel_tensor_use_t.h" -#include "utils/containers/require_only_key.h" -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" -#include "utils/containers/map_keys_and_values.h" -#include "utils/containers/transform_pairs.h" namespace FlexFlow { @@ -35,8 +35,8 @@ static bidict TensorSlotName producer_slot = get_src_layer_output_slot_name(input_edge); return get_tensor_bindings_for_slot_name( - /*task_group=*/mpcg_get_mapping_for_layer(mpcg, producer_layer), - /*slot_name=*/producer_slot); + /*task_group=*/mpcg_get_mapping_for_layer(mpcg, producer_layer), + /*slot_name=*/producer_slot); } static bidict @@ -44,26 +44,29 @@ static bidict MappedParallelComputationGraph const &mpcg, parallel_tensor_guid_t const &output_tensor_guid) { - std::unordered_set consumers = mpcg_get_parallel_tensor_uses(mpcg, output_tensor_guid); + std::unordered_set consumers = + mpcg_get_parallel_tensor_uses(mpcg, output_tensor_guid); ASSERT(!consumers.empty()); // union all consumer bindings — each consumer shard maps to a distinct // (discard_copy, machine) pair since replicas are always on different machines bidict result = - merge_disjoint_bidicts( - transform(consumers, - [&](parallel_tensor_use_t const &use) - -> bidict - { - parallel_layer_guid_t consumer_layer = parallel_tensor_use_get_layer(use); - TensorSlotName slot_name = parallel_tensor_use_get_slot(use); - - MappedOperatorTaskGroup consumer_mapping = mpcg_get_mapping_for_layer(mpcg, consumer_layer); - bidict binding = - get_tensor_bindings_for_slot_name(consumer_mapping, slot_name); - - return binding; - })); + merge_disjoint_bidicts(transform( + consumers, + [&](parallel_tensor_use_t const &use) + -> bidict { + parallel_layer_guid_t consumer_layer = + parallel_tensor_use_get_layer(use); + TensorSlotName slot_name = parallel_tensor_use_get_slot(use); + + MappedOperatorTaskGroup consumer_mapping = + mpcg_get_mapping_for_layer(mpcg, consumer_layer); + bidict + binding = get_tensor_bindings_for_slot_name(consumer_mapping, + slot_name); + + return binding; + })); return result; } @@ -73,15 +76,19 @@ static DynamicNodeInvocation ReplicateAttrs const &attrs, MappedParallelComputationGraph const &mpcg) { - ManyToOne incoming = mpcg_get_incoming_tensors(mpcg, layer); + ManyToOne incoming = + mpcg_get_incoming_tensors(mpcg, layer); TensorSlotName input_slot_name = TensorSlotName::INPUT; - parallel_tensor_guid_t input_tensor_guid = require_only_key(incoming.l_to_r(), input_slot_name); + parallel_tensor_guid_t input_tensor_guid = + require_only_key(incoming.l_to_r(), input_slot_name); ParallelTensorAttrs input_attrs = mpcg_get_parallel_tensor_attrs(mpcg, input_tensor_guid); - bidict outgoing = mpcg_get_outgoing_tensors(mpcg, layer); + bidict outgoing = + mpcg_get_outgoing_tensors(mpcg, layer); TensorSlotName output_slot_name = TensorSlotName::OUTPUT; - parallel_tensor_guid_t output_tensor_guid = require_only_key(outgoing.l_to_r(), output_slot_name); + parallel_tensor_guid_t output_tensor_guid = + require_only_key(outgoing.l_to_r(), output_slot_name); ParallelTensorAttrs output_attrs = mpcg_get_parallel_tensor_attrs(mpcg, output_tensor_guid); @@ -117,17 +124,18 @@ static DynamicNodeInvocation DynamicNodeInvocation invocation_node{ /*inputs=*/{ - { - DynamicTensorSlot{input_slot_name, std::nullopt}, - input_value, - }, + { + DynamicTensorSlot{input_slot_name, std::nullopt}, + input_value, + }, }, /*node_attrs=*/node_attrs, - /*outputs=*/{ - { - DynamicTensorSlot{output_slot_name, std::nullopt}, - output_value, - }, + /*outputs=*/ + { + { + DynamicTensorSlot{output_slot_name, std::nullopt}, + output_value, + }, }, }; @@ -139,9 +147,9 @@ DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( ParallelComputationGraph pcg = pcg_from_mpcg(mpcg); - auto mk_invocation = [&](parallel_layer_guid_t layer, ParallelLayerAttrs const &attrs) - -> DynamicNodeInvocation - { + auto mk_invocation = + [&](parallel_layer_guid_t layer, + ParallelLayerAttrs const &attrs) -> DynamicNodeInvocation { if (attrs.op_attrs.is_parallel_replicate()) { // build replicate invocation DynamicNodeInvocation repl_inv = build_replicate_invocation( @@ -159,50 +167,45 @@ DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( auto mk_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, }; }; - auto mk_value_attrs = [&](parallel_tensor_guid_t const &tensor) -> DynamicValueAttrs - { - ParallelTensorAttrs attrs = - get_parallel_tensor_attrs(pcg, tensor); + auto mk_value_attrs = + [&](parallel_tensor_guid_t const &tensor) -> DynamicValueAttrs { + ParallelTensorAttrs attrs = get_parallel_tensor_attrs(pcg, tensor); return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, - /*parallel_tensor_shape=*/attrs.shape, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, + /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, + /*parallel_tensor_shape=*/attrs.shape, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, }; }; std::unordered_map result_inputs = - map_keys_and_values(get_incoming_tensors(pcg, layer), - mk_slot, - mk_value_attrs); + map_keys_and_values( + get_incoming_tensors(pcg, layer), mk_slot, mk_value_attrs); std::unordered_map result_outputs = - map_keys_and_values(get_outgoing_tensors(pcg, layer), - mk_slot, - mk_value_attrs); + map_keys_and_values( + get_outgoing_tensors(pcg, layer), mk_slot, mk_value_attrs); DynamicNodeInvocation invocation = DynamicNodeInvocation{ - /*inputs=*/result_inputs, - /*node_attrs=*/result_attrs, - /*outputs=*/result_outputs, + /*inputs=*/result_inputs, + /*node_attrs=*/result_attrs, + /*outputs=*/result_outputs, }; return invocation; }; }; - return dynamic_open_dataflow_graph_from_invocation_set( - transform_pairs( - unordered_set_of(get_parallel_layer_attrs_mapping(pcg)), - mk_invocation)); + return dynamic_open_dataflow_graph_from_invocation_set(transform_pairs( + unordered_set_of(get_parallel_layer_attrs_mapping(pcg)), mk_invocation)); } } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index f4960fe67a..64fe2df0be 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -1,11 +1,11 @@ #include "task-spec/dynamic_graph/pass_expansion.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include "task-spec/dynamic_graph/training_operation_attrs.h" #include "utils/containers/are_all_same.h" #include "utils/containers/get_only.h" #include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/transform.h" -#include "task-spec/dynamic_graph/training_operation_attrs.h" namespace FlexFlow { @@ -84,7 +84,8 @@ DynamicNodeInvocation perform_fwd_pass_expansion_for_invocation( DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( DynamicNodeInvocation const &invocation) { - TrainingOperationAttrs op_attrs = assert_unwrap(invocation.node_attrs.op_attrs); + TrainingOperationAttrs op_attrs = + assert_unwrap(invocation.node_attrs.op_attrs); auto to_fwd = [](DynamicTensorSlot const &k, DynamicValueAttrs const &v) { return std::pair{ @@ -106,15 +107,16 @@ DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( DynamicNodeInvocation bwd{ /*inputs=*/{ - to_fwd(output_slot, output), - to_grad(output_slot, output), + to_fwd(output_slot, output), + to_grad(output_slot, output), }, /*node_attrs=*/ pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), - /*outputs=*/{ - to_grad(input_slot, input), + /*outputs=*/ + { + to_grad(input_slot, input), }, - }; + }; return bwd; } else { diff --git a/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc index d1452242ca..a9be225ff5 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc @@ -4,17 +4,14 @@ namespace FlexFlow { -bool training_op_attrs_has_op_type(TrainingOperationAttrs const &op_attrs, OperatorType op_type) { - return op_attrs.visit(overload { - [&](PCGOperatorAttrs const &pcg_op_attrs) -> bool { - return pcg_op_attrs_get_op_type(pcg_op_attrs) == op_type; - }, - [](LossAttrs const &) -> bool { - return false; - }, - [](CopyAttrs const &) -> bool { - return false; - }, +bool training_op_attrs_has_op_type(TrainingOperationAttrs const &op_attrs, + OperatorType op_type) { + return op_attrs.visit(overload{ + [&](PCGOperatorAttrs const &pcg_op_attrs) -> bool { + return pcg_op_attrs_get_op_type(pcg_op_attrs) == op_type; + }, + [](LossAttrs const &) -> bool { return false; }, + [](CopyAttrs const &) -> bool { return false; }, }); } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc index ed22a8cbde..bf88d5ec38 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc @@ -1,8 +1,8 @@ #include "task-spec/dynamic_graph/pass_expansion.h" +#include "op-attrs/ops/element_unary.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include -#include "op-attrs/ops/element_unary.h" using namespace ::FlexFlow; @@ -37,18 +37,17 @@ TEST_SUITE(FF_TEST_SUITE) { dynamic_layer_guid_t layer_guid{parallel_layer_guid_t{Node{20}}}; - TrainingOperationAttrs op_attrs = - TrainingOperationAttrs{ + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/8_p, - /*use_bias=*/true, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, - }, + LinearAttrs{ + /*out_channels=*/8_p, + /*use_bias=*/true, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }, - }; + }; DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); @@ -157,18 +156,17 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicValueAttrs v3_grad = mk_value_attrs(2, grad_role); SUBCASE("normal operator") { - TrainingOperationAttrs op_attrs = - TrainingOperationAttrs{ + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/8_p, - /*use_bias=*/true, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, - }, + LinearAttrs{ + /*out_channels=*/8_p, + /*use_bias=*/true, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }, - }; + }; DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { return DynamicNodeInvocation{ @@ -227,14 +225,13 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("replicate operator optimization") { - TrainingOperationAttrs op_attrs = - TrainingOperationAttrs{ + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, - }, + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }, - }; + }; DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { return DynamicNodeInvocation{ @@ -262,7 +259,8 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicNodeInvocation correct = [&]() -> DynamicNodeInvocation { DynamicTensorRole fwd_role = DynamicTensorRole{FwbTensorType::FORWARD}; - DynamicTensorRole grad_role = DynamicTensorRole{FwbTensorType::GRADIENT}; + DynamicTensorRole grad_role = + DynamicTensorRole{FwbTensorType::GRADIENT}; return DynamicNodeInvocation{ /*inputs=*/{ @@ -324,28 +322,27 @@ TEST_SUITE(FF_TEST_SUITE) { }; TrainingOperationAttrs input_op_attrs = TrainingOperationAttrs{ - PCGOperatorAttrs{ - InputAttrs{ - TensorShape{ - TensorDims{ - FFOrdered{ - 4_p, - 8_p, - }, + PCGOperatorAttrs{ + InputAttrs{ + TensorShape{ + TensorDims{ + FFOrdered{ + 4_p, + 8_p, + }, + }, + DataType::FLOAT, + }, }, - DataType::FLOAT, - }, }, - }, }; TrainingOperationAttrs relu_op_attrs = TrainingOperationAttrs{ - PCGOperatorAttrs{ - make_relu_attrs(), - }, + PCGOperatorAttrs{ + make_relu_attrs(), + }, }; - DynamicOpenDataflowGraph input = [&]() -> DynamicOpenDataflowGraph { DynamicNodeAttrs n1 = mk_node_attrs(10, input_op_attrs, std::nullopt); DynamicNodeAttrs n2 = mk_node_attrs(11, relu_op_attrs, std::nullopt); @@ -396,10 +393,14 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicOpenDataflowGraph result = perform_pass_expansion(input); DynamicOpenDataflowGraph correct = [&]() -> DynamicOpenDataflowGraph { - DynamicNodeAttrs n1_fwd = mk_node_attrs(10, input_op_attrs, DynamicTaskType::FWD); - DynamicNodeAttrs n2_fwd = mk_node_attrs(11, relu_op_attrs, DynamicTaskType::FWD); - DynamicNodeAttrs n1_bwd = mk_node_attrs(10, input_op_attrs, DynamicTaskType::BWD); - DynamicNodeAttrs n2_bwd = mk_node_attrs(11, relu_op_attrs, DynamicTaskType::BWD); + DynamicNodeAttrs n1_fwd = + mk_node_attrs(10, input_op_attrs, DynamicTaskType::FWD); + DynamicNodeAttrs n2_fwd = + mk_node_attrs(11, relu_op_attrs, DynamicTaskType::FWD); + DynamicNodeAttrs n1_bwd = + mk_node_attrs(10, input_op_attrs, DynamicTaskType::BWD); + DynamicNodeAttrs n2_bwd = + mk_node_attrs(11, relu_op_attrs, DynamicTaskType::BWD); DynamicValueAttrs v1_activation = mk_value_attrs(0, mk_dynamic_tensor_role_fwd()); diff --git a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h index f2104fd113..0c944bb9bd 100644 --- a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h +++ b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H -#include "utils/containers/foldl.h" #include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" +#include "utils/containers/foldl.h" namespace FlexFlow { @@ -13,8 +13,7 @@ bidict merge_disjoint_bidicts(C const &c) { bidict empty = {}; return foldl(c, /*init=*/empty, - [](bidict const &lhs, - bidict const &rhs) { + [](bidict const &lhs, bidict const &rhs) { return binary_merge_disjoint_bidicts(lhs, rhs); }); } diff --git a/lib/utils/include/utils/containers/transform_pairs.h b/lib/utils/include/utils/containers/transform_pairs.h index c01b50554f..3e421ea445 100644 --- a/lib/utils/include/utils/containers/transform_pairs.h +++ b/lib/utils/include/utils/containers/transform_pairs.h @@ -21,7 +21,8 @@ template > -std::unordered_set transform_pairs(std::unordered_set> const &c, F &&f) { +std::unordered_set + transform_pairs(std::unordered_set> const &c, F &&f) { auto ff = [&](std::pair const &p) -> Out { return f(p.first, p.second); }; diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h index b5557e9e49..52c225d157 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h @@ -7,25 +7,20 @@ namespace FlexFlow { template std::unordered_set> - get_kwarg_dataflow_value_uses( - KwargDataflowGraphView const &g, - KwargDataflowOutput const &v) { - - KwargDataflowEdgeQuery query = - KwargDataflowEdgeQuery{ - /*src_nodes=*/query_set::match_single_value(v.node), - /*src_slots=*/query_set::match_single_value(v.slot_name), - /*dst_nodes=*/query_set::matchall(), - /*dst_slots=*/query_set::matchall(), - }; - - std::unordered_set> edges = - g.query_edges(query); - - return transform( - edges, [&](KwargDataflowEdge const &e) { - return e.dst; - }); + get_kwarg_dataflow_value_uses(KwargDataflowGraphView const &g, + KwargDataflowOutput const &v) { + + KwargDataflowEdgeQuery query = KwargDataflowEdgeQuery{ + /*src_nodes=*/query_set::match_single_value(v.node), + /*src_slots=*/query_set::match_single_value(v.slot_name), + /*dst_nodes=*/query_set::matchall(), + /*dst_slots=*/query_set::matchall(), + }; + + std::unordered_set> edges = g.query_edges(query); + + return transform(edges, + [&](KwargDataflowEdge const &e) { return e.dst; }); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/many_to_one/many_to_one.h b/lib/utils/include/utils/many_to_one/many_to_one.h index c73f696172..2d078eb304 100644 --- a/lib/utils/include/utils/many_to_one/many_to_one.h +++ b/lib/utils/include/utils/many_to_one/many_to_one.h @@ -2,6 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_MANY_TO_ONE_H #include "utils/containers/keys.h" +#include "utils/containers/require_same.h" #include "utils/containers/try_at.h" #include "utils/containers/unordered_set_of.h" #include "utils/containers/values.h" @@ -19,7 +20,6 @@ #include #include #include -#include "utils/containers/require_same.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/one_to_many/one_to_many.h b/lib/utils/include/utils/one_to_many/one_to_many.h index 7b725fdec1..5492ff3f78 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many.h +++ b/lib/utils/include/utils/one_to_many/one_to_many.h @@ -4,6 +4,7 @@ #include "utils/containers/generate_map.h" #include "utils/containers/items.h" #include "utils/containers/keys.h" +#include "utils/containers/require_same.h" #include "utils/containers/transform.h" #include "utils/containers/try_at.h" #include "utils/containers/unordered_set_of.h" @@ -23,7 +24,6 @@ #include #include #include -#include "utils/containers/require_same.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc index 8650de44f6..13a1bcd968 100644 --- a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc +++ b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using K = value_type<0>; using V = value_type<1>; -template - bidict binary_merge_disjoint_bidicts(bidict const &, bidict const &); +template bidict binary_merge_disjoint_bidicts(bidict const &, + bidict const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/transform_pairs.cc b/lib/utils/src/utils/containers/transform_pairs.cc index 241f1ad425..4afda936e4 100644 --- a/lib/utils/src/utils/containers/transform_pairs.cc +++ b/lib/utils/src/utils/containers/transform_pairs.cc @@ -8,10 +8,10 @@ using R = value_type<1>; using Out = value_type<2>; using F = std::function; -template - std::vector transform_pairs(std::vector> const &, F &&); +template std::vector transform_pairs(std::vector> const &, + F &&); -template - std::unordered_set transform_pairs(std::unordered_set> const &, F &&); +template std::unordered_set + transform_pairs(std::unordered_set> const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc index 2e42863e53..b1d2988223 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc @@ -5,10 +5,8 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template - std::unordered_set> - get_kwarg_dataflow_value_uses( - KwargDataflowGraphView const &, - KwargDataflowOutput const &); +template std::unordered_set> + get_kwarg_dataflow_value_uses(KwargDataflowGraphView const &, + KwargDataflowOutput const &); } // namespace FlexFlow From f2b075482f46dff9d5c66fd4a7616bc425f23421 Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 21 May 2026 12:15:38 -0700 Subject: [PATCH 09/35] Format Realm. --- .../src/realm-execution/pcg_instance.cc | 31 ++- .../src/realm-execution/test_op_replicate.cc | 185 +++++++++--------- 2 files changed, 107 insertions(+), 109 deletions(-) diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index f2edac7f88..332669a9dc 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -217,14 +217,11 @@ static Realm::Event spawn_dynamic_node_invocation( }; auto issue_replicate_bwd = [&]() { - - DynamicValueAttrs output_grad = get_only( - values( - filter_keys( - invocation.inputs, - [](DynamicTensorSlot const &s) -> bool { - return s.slot_tensor_role == DynamicTensorRole{FwbTensorType::GRADIENT}; - }))); + DynamicValueAttrs output_grad = get_only(values( + filter_keys(invocation.inputs, [](DynamicTensorSlot const &s) -> bool { + return s.slot_tensor_role == + DynamicTensorRole{FwbTensorType::GRADIENT}; + }))); DynamicValueAttrs input_grad = get_only(values(invocation.outputs)); @@ -246,15 +243,15 @@ static Realm::Event spawn_dynamic_node_invocation( tensor_instance_backing.backing.at(replica_key).first; e = ctx.issue_copy( - /*src_shape=*/assert_unwrap(output_grad.parallel_tensor_shape), - /*src_inst=*/src_inst, - /*dst_shape=*/assert_unwrap(input_grad.parallel_tensor_shape), - /*dst_inst=*/dst_inst, - /*requests=*/Realm::ProfilingRequestSet{}, - /*wait_on=*/e, - /*priority=*/0, - /*redop_id=*/redop_id, - /*exlusive=*/false); + /*src_shape=*/assert_unwrap(output_grad.parallel_tensor_shape), + /*src_inst=*/src_inst, + /*dst_shape=*/assert_unwrap(input_grad.parallel_tensor_shape), + /*dst_inst=*/dst_inst, + /*requests=*/Realm::ProfilingRequestSet{}, + /*wait_on=*/e, + /*priority=*/0, + /*redop_id=*/redop_id, + /*exlusive=*/false); } return e; }; diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index 2523cae798..46d29e2bef 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -13,6 +13,7 @@ #include "op-attrs/tensor_slot_name.dtg.h" #include "pcg/device_type.dtg.h" #include "pcg/machine_space_coordinate.dtg.h" +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" #include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_builder.h" @@ -27,7 +28,6 @@ #include "test/utils/doctest/check_kv.h" #include "utils/containers/require_only_key.h" #include -#include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" namespace test { @@ -49,7 +49,8 @@ static bool did_loss_decrease(GenericTensorAccessorR const &first_epoch, compare_tensor_accessors_le(last_epoch, first_epoch, allocator)); } -MappedParallelComputationGraph make_test_mpcg_for_device_type(DeviceType device_type) { +MappedParallelComputationGraph + make_test_mpcg_for_device_type(DeviceType device_type) { positive_int batch_size = 10_p; positive_int data_dim = 16_p; positive_int hidden_dim = 32_p; @@ -63,8 +64,8 @@ MappedParallelComputationGraph make_test_mpcg_for_device_type(DeviceType device_ ParallelComputationGraph pcg = empty_parallel_computation_graph(); - TensorShape input_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + TensorShape input_tensor_shape = + TensorShape{TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; ParallelLayerAddedResult inputs_layer = pcg_add_input_layer(pcg, input_tensor_shape); @@ -144,91 +145,92 @@ MappedParallelComputationGraph make_test_mpcg_for_device_type(DeviceType device_ /*discard_copy_component=*/1_n, /*shard_component=*/FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( - /*pcg=*/pcg, - /*mapped_op_task_groups=*/{ - { - inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - { + MappedParallelComputationGraph mpcg = + mapped_pcg_from_pcg_and_mapped_op_task_groups( + /*pcg=*/pcg, + /*mapped_op_task_groups=*/{ { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, + inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, }, - }, - }, - }, - { - inputs_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - { { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, + inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, }, - }, - }, - }, - { - add_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, + add_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, }, - }, - }, - }, - { - repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, + repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, }, { - cpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, + relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, }, - }, - }, - }, - { - relu_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { - { - cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - { - cpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, - }, - }, - }, - }, - }); + }); return mpcg; } @@ -245,21 +247,20 @@ TEST_SUITE(FF_TEST_SUITE) { manager.start_controller([](RealmContext &ctx) { Allocator allocator = ctx.get_current_device_allocator(); - MappedParallelComputationGraph mpcg = make_test_mpcg_for_device_type(DeviceType::CPU); - + MappedParallelComputationGraph mpcg = + make_test_mpcg_for_device_type(DeviceType::CPU); std::unordered_map input_tensors; - OptimizerAttrs optimizer_attrs = - OptimizerAttrs{ - SGDOptimizerAttrs{ + OptimizerAttrs optimizer_attrs = OptimizerAttrs{ + SGDOptimizerAttrs{ /*lr=*/0.001, /*momentum=*/0.9, /*nesterov=*/false, /*weight_decay=*/0.001, - }, - }; + }, + }; DistributedFfHandle device_handle = create_distributed_ff_handle( ctx, @@ -303,17 +304,17 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { manager.start_controller([](RealmContext &ctx) { Allocator allocator = ctx.get_current_device_allocator(); - MappedParallelComputationGraph mpcg = make_test_mpcg_for_device_type(DeviceType::GPU); + MappedParallelComputationGraph mpcg = + make_test_mpcg_for_device_type(DeviceType::GPU); - OptimizerAttrs optimizer_attrs = - OptimizerAttrs{ - SGDOptimizerAttrs{ + OptimizerAttrs optimizer_attrs = OptimizerAttrs{ + SGDOptimizerAttrs{ /*lr=*/0.001, /*momentum=*/0.9, /*nesterov=*/false, /*weight_decay=*/0.001, - }, - }; + }, + }; std::unordered_map input_tensors; From 9d03c9766beaa67b2cbc1eea87976d2c44ae6153 Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 21 May 2026 12:16:39 -0700 Subject: [PATCH 10/35] Refactor redop infrastructure and switch to Legion's redops. --- .../redops/realm_redop_registry.h | 16 + .../redops/redop_id_t.dtg.toml | 30 + .../realm-execution/redops/redop_id_t.h | 22 + .../realm-execution/tasks/realm_reduction.h | 154 ---- .../src/realm-execution/realm_manager.cc | 4 +- .../redops/realm_redop_registry.cc | 689 ++++++++++++++++++ .../src/realm-execution/redops/redop_id_t.cc | 32 + .../tasks/realm_task_registry.cc | 10 - 8 files changed, 792 insertions(+), 165 deletions(-) create mode 100644 lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h create mode 100644 lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml create mode 100644 lib/realm-execution/include/realm-execution/redops/redop_id_t.h delete mode 100644 lib/realm-execution/include/realm-execution/tasks/realm_reduction.h create mode 100644 lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc create mode 100644 lib/realm-execution/src/realm-execution/redops/redop_id_t.cc diff --git a/lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h b/lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h new file mode 100644 index 0000000000..a338a38bbf --- /dev/null +++ b/lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h @@ -0,0 +1,16 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H + +#include "realm-execution/realm.h" +#include "realm-execution/redops/redop_id_t.dtg.h" + +namespace FlexFlow { + +/** + * \brief Registers all known reduction operators (redops). + */ +void Realm::Event register_all_redops(Realm::Runtime); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml b/lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml new file mode 100644 index 0000000000..5183ff5e72 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml @@ -0,0 +1,30 @@ +namespace = "FlexFlow" +name = "redop_id_t" +type = "enum" +features = [ + "hash", + "fmt", + "rapidcheck", + "json", +] +docstring = ''' +\brief An enum for identifying reduction operators (redops) for use in the Realm runtime. +''' + +[[values]] +name = "SUM_BOOL_REDOP_ID" + +[[values]] +name = "SUM_INT32_REDOP_ID" + +[[values]] +name = "SUM_INT64_REDOP_ID" + +[[values]] +name = "SUM_HALF_REDOP_ID" + +[[values]] +name = "SUM_FLOAT_REDOP_ID" + +[[values]] +name = "SUM_DOUBLE_REDOP_ID" diff --git a/lib/realm-execution/include/realm-execution/redops/redop_id_t.h b/lib/realm-execution/include/realm-execution/redops/redop_id_t.h new file mode 100644 index 0000000000..b9ef91a05a --- /dev/null +++ b/lib/realm-execution/include/realm-execution/redops/redop_id_t.h @@ -0,0 +1,22 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H + +#include "realm-execution/realm.h" +#include "realm-execution/redops/redop_id_t.dtg.h" + +namespace FlexFlow { + +/** + * \brief Registers all known reduction operators (redops). + */ +Realm::ReductionOpID get_sum_redop_id_for_data_type(DataType); + +/** + * \brief Convert a \ref FlexFlow::redop_id_t into a Realm reduction op ID. + */ +Realm::Processor::ReductionOpID + get_realm_reduction_op_id_for_redop_id(redop_id_t); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h b/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h deleted file mode 100644 index 512e344824..0000000000 --- a/lib/realm-execution/include/realm-execution/tasks/realm_reduction.h +++ /dev/null @@ -1,154 +0,0 @@ -#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDUCTION_H -#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDUCTION_H -#include "op-attrs/datatype.dtg.h" -#include - -namespace FlexFlow { - -/** - * \brief Realm Sum Reduction for Float - * \see https://legion.stanford.edu/tutorial/realm/reductions.html - */ -struct SumReductionFloat { - using LHS = float; - using RHS = float; - - /** \brief Identity element for addition (0.0) */ - static constexpr RHS identity = 0.0f; - - /** - * \brief Apply reduction: lhs += rhs - * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop - * \param lhs Left-hand side accumulator (modified in place) - * \param rhs Value to add - */ - template - static void apply(LHS &lhs, RHS rhs) { - if (EXCLUSIVE) { - lhs += rhs; - } else { - // Atomic float add via CAS loop - union { - float f; - int i; - } old_val, new_val; - do { - old_val.f = lhs; - new_val.f = old_val.f + rhs; - } while ( - !__sync_bool_compare_and_swap((int *)&lhs, old_val.i, new_val.i)); - } - } - - /** - * \brief Fold two RHS values: rhs1 += rhs2 - * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop - * \param rhs1 Accumulator (modified in place) - * \param rhs2 Value to fold in - */ - template - static void fold(RHS &rhs1, RHS rhs2) { - if (EXCLUSIVE) { - rhs1 += rhs2; - } else { - // Atomic float add via CAS loop - union { - float f; - int i; - } old_val, new_val; - do { - old_val.f = rhs1; - new_val.f = old_val.f + rhs2; - } while ( - !__sync_bool_compare_and_swap((int *)&rhs1, old_val.i, new_val.i)); - } - } -}; - -/** - * \brief Realm Sum Reduction for Double - * \see https://legion.stanford.edu/tutorial/realm/reductions.html - */ -struct SumReductionDouble { - using LHS = double; - using RHS = double; - - /** \brief Identity element for addition (0.0) */ - static constexpr RHS identity = 0.0; - - /** - * \brief Apply reduction: lhs += rhs - * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop - * \param lhs Left-hand side accumulator (modified in place) - * \param rhs Value to add - */ - template - static void apply(LHS &lhs, RHS rhs) { - if (EXCLUSIVE) { - lhs += rhs; - } else { - // Atomic double add via CAS loop using long long reinterpretation - union { - double d; - long long i; - } old_val, new_val; - do { - old_val.d = lhs; - new_val.d = old_val.d + rhs; - } while (!__sync_bool_compare_and_swap( - (long long *)&lhs, old_val.i, new_val.i)); - } - } - - /** - * \brief Fold two RHS values: rhs1 += rhs2 - * \tparam EXCLUSIVE If true, direct addition; if false, atomic CAS loop - * \param rhs1 Accumulator (modified in place) - * \param rhs2 Value to fold in - */ - template - static void fold(RHS &rhs1, RHS rhs2) { - if (EXCLUSIVE) { - rhs1 += rhs2; - } else { - // Atomic double add via CAS loop using long long reinterpretation - union { - double d; - long long i; - } old_val, new_val; - do { - old_val.d = rhs1; - new_val.d = old_val.d + rhs2; - } while (!__sync_bool_compare_and_swap( - (long long *)&rhs1, old_val.i, new_val.i)); - } - } -}; - -/** - * \brief Reduction op IDs for sum reductions - * \warning These IDs must not conflict with other registered reduction ops - */ -enum SumReductionOpIDs { - REDOP_SUM_FLOAT = 1, ///< Sum reduction op ID for float - REDOP_SUM_DOUBLE = 2, ///< Sum reduction op ID for double -}; - -/** - * \brief Returns the Realm reduction op ID for a sum reduction over the given datatype - * \param dtype The datatype to look up - * \return The corresponding Realm::ReductionOpID - * \throws PANIC if no sum reduction is registered for the given datatype - */ -inline Realm::ReductionOpID get_sum_reduction_op_id(DataType dtype) { - switch (dtype) { - case DataType::FLOAT: - return REDOP_SUM_FLOAT; - case DataType::DOUBLE: - return REDOP_SUM_DOUBLE; - default: - PANIC("no sum reduction registered for datatype {}", dtype); - } -} -} // namespace FlexFlow -#endif diff --git a/lib/realm-execution/src/realm-execution/realm_manager.cc b/lib/realm-execution/src/realm-execution/realm_manager.cc index e76be7054b..c7136d8a98 100644 --- a/lib/realm-execution/src/realm-execution/realm_manager.cc +++ b/lib/realm-execution/src/realm-execution/realm_manager.cc @@ -1,6 +1,7 @@ #include "realm-execution/realm_manager.h" #include "realm-execution/realm_context.h" #include "realm-execution/tasks/realm_task_registry.h" +#include "realm-execution/redops/realm_redop_registry.h" namespace FlexFlow { @@ -9,8 +10,9 @@ RealmManager::RealmManager(int *argc, char ***argv) bool ok = this->get_runtime().init(argc, argv); ASSERT(ok); - // Register all tasks at initialization time so we don't need to later + // Register all tasks and redops at initialization time so we don't need to later register_all_tasks().wait(); + register_all_redops(this->get_runtime()); } RealmManager::~RealmManager() { diff --git a/lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc b/lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc new file mode 100644 index 0000000000..d10b158463 --- /dev/null +++ b/lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc @@ -0,0 +1,689 @@ +#include "realm-execution/redops/realm_redop_registry.h" + +namespace FlexFlow { + +// Reduction operators and related infrastructure borrowed from Legion. We +// maintain the Legion naming scheme to maximizing compatibility with the +// existing code, despite not otherwise relying or using Legion in any way. +// https://gitlab.com/StanfordLegion/legion/-/blob/5263aeff477fb94239c50d9306d58c4244e9fc38/runtime/legion/api/redop.inl#L31 +#if !defined(__cpp_lib_atomic_ref) || (__cpp_lib_atomic_ref < 201806L) +// We only need this crap if we're using a version of c++ < 20 +// Starting with c++20 we can do all this the right way with atomic_ref +namespace TypePunning { +// The tenth circle of hell is reserved for members of the C++ committee +// that decided to deviate from C's support for type punning unions. +// Add on to it the fact that it took them 9 fucking years to realize +// that they needed std::atomic_ref and it's plain to see they are all +// just a bunch of idiots that should never be allowed near a programming +// language standard ever again. They've clearly never written lock-free +// code in their lives. +template +class Pointer { +public: + Pointer(void *p) : pointer(convert(p)) {} + static inline T *convert(void *p) { + T *ptr = nullptr; + static_assert(sizeof(ptr) == sizeof(p)); + memcpy(&ptr, &p, sizeof(p)); + return ptr; + } + inline operator T *(void) const { + return (T *)pointer; + } + inline T operator*(void) const { + return *pointer; + } + inline T operator[](size_t off) const { + return pointer[off]; + } + +private: + T volatile *const pointer; +}; +template +class AlignedPointer { +public: + AlignedPointer(void *p) : off(align(p)), pointer(convert(p, off)) {} + static inline T *convert(void *p, size_t off) { + uint8_t *p1 = nullptr; + static_assert(sizeof(p1) == sizeof(p)); + memcpy(&p1, &p, sizeof(p)); + p1 = p1 - off; + T *p2 = nullptr; + static_assert(sizeof(p1) == sizeof(p2)); + memcpy(&p2, &p1, sizeof(p1)); + return p2; + } + static inline size_t align(void *p) { + uintptr_t ptr; + static_assert(sizeof(ptr) == sizeof(p)); + memcpy(&ptr, &p, sizeof(ptr)); + return ptr % ALIGNMENT; + } + inline operator T *(void) const { + return (T *)pointer; + } + inline T operator*(void) const { + return *pointer; + } + inline size_t offset(void) const { + return off; + } + +private: + size_t off; + T volatile *const pointer; +}; +template +class Alias { +public: + inline void load(Pointer const &pointer, size_t off = 0) { + T1 value = pointer[off]; + memcpy(buffer, (void *)&value, sizeof(T1)); + } + template + inline void load(AlignedPointer const &pointer) { + T1 value = *pointer; + memcpy(buffer, (void *)&value, sizeof(T1)); + } + inline T1 as_one(void) const { + T1 result; + memcpy((void *)&result, buffer, sizeof(result)); + return result; + } + inline T2 as_two(void) const { + T2 result; + memcpy((void *)&result, buffer, sizeof(result)); + return result; + } + inline Alias &operator=(T2 rhs) { + memcpy(buffer, (void *)&rhs, sizeof(rhs)); + return *this; + } + +private: + // Make this one private so it is can never be called + inline Alias &operator=(T1 rhs) { + memcpy(buffer, (void *)&rhs, sizeof(rhs)); + return *this; + } + static_assert(sizeof(T1) == sizeof(T2)); + uint8_t buffer[sizeof(T1)]; +}; +}; // namespace TypePunning +#endif + +// Define a prefix for annotating functions for CUDA compilation +#if defined(__CUDACC__) || defined(__HIPCC__) +#define __LEGION_CUDA_HD__ __host__ __device__ +#else +#define __LEGION_CUDA_HD__ +#endif + +template <> +class SumReduction { +public: + typedef bool LHS; + typedef bool RHS; + + static constexpr bool identity = false; + static constexpr int REDOP_ID = LEGION_REDOP_OR_BOOL; + + template + __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); + template + __LEGION_CUDA_HD__ static void fold(RHS &rhs1, RHS rhs2); +}; + +template <> +class SumReduction { +public: + typedef int32_t LHS; + typedef int32_t RHS; + + static constexpr int32_t identity = 0; + static constexpr int REDOP_ID = LEGION_REDOP_SUM_INT32; + + template + __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); + template + __LEGION_CUDA_HD__ static void fold(RHS &rhs1, RHS rhs2); +}; + +template <> +class SumReduction { +public: + typedef int64_t LHS; + typedef int64_t RHS; + + static constexpr int64_t identity = 0; + static constexpr int REDOP_ID = LEGION_REDOP_SUM_INT64; + + template + __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); + template + __LEGION_CUDA_HD__ static void fold(RHS &rhs1, RHS rhs2); +}; + +template <> +class SumReduction<__half> { +public: + typedef __half LHS; + typedef __half RHS; + + static inline const __half identity = __half(0, false /*raw*/); + static constexpr int REDOP_ID = LEGION_REDOP_SUM_FLOAT16; + + template + __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); + template + __LEGION_CUDA_HD__ static void fold(RHS &rhs1, RHS rhs2); +}; + +template <> +class SumReduction { +public: + typedef float LHS; + typedef float RHS; + + static constexpr float identity = 0.f; + static constexpr int REDOP_ID = LEGION_REDOP_SUM_FLOAT32; + + template + __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); + template + __LEGION_CUDA_HD__ static void fold(RHS &rhs1, RHS rhs2); +}; + +template <> +class SumReduction { +public: + typedef double LHS; + typedef double RHS; + + static constexpr double identity = 0.0; + static constexpr int REDOP_ID = LEGION_REDOP_SUM_FLOAT64; + + template + __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); + template + __LEGION_CUDA_HD__ static void fold(RHS &rhs1, RHS rhs2); +}; + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { + lhs = lhs || rhs; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + // GPU atomics need 4 byte alignment + const uintptr_t unaligned = reinterpret_cast(&lhs); + unsigned const offset = unaligned % sizeof(unsigned int); + const uintptr_t aligned = unaligned - offset; + unsigned int *ptr = reinterpret_cast(aligned); + unsigned int newval = *ptr, oldval; + do { + RHS previous = __uint2bool(newval, offset); + RHS next = previous || rhs; + oldval = newval; + newval = __bool2uint(newval, next, offset); + newval = atomicCAS(ptr, oldval, newval); + } while (oldval != newval); +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(lhs); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval || rhs; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic logical operations so use compare and swap + TypePunning::Alias oldval, newval; + TypePunning::Pointer pointer((void *)&lhs); + do { + oldval.load(pointer); + newval = oldval.as_two() || rhs; + } while (!__sync_bool_compare_and_swap( + (int8_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { + rhs1 = rhs1 || rhs2; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + // GPU atomics need 4 byte alignment + const uintptr_t unaligned = reinterpret_cast(&rhs1); + unsigned const offset = unaligned % sizeof(unsigned int); + const uintptr_t aligned = unaligned - offset; + unsigned int *ptr = reinterpret_cast(aligned); + unsigned int newval = *ptr, oldval; + do { + RHS previous = __uint2bool(newval, offset); + RHS next = previous || rhs2; + oldval = newval; + newval = __bool2uint(newval, next, offset); + newval = atomicCAS(ptr, oldval, newval); + } while (oldval != newval); +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(rhs1); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval || rhs2; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic logical operations so use compare and swap + TypePunning::Alias oldval, newval; + TypePunning::Pointer pointer((void *)&rhs1); + do { + oldval.load(pointer); + newval = oldval.as_two() || rhs2; + } while (!__sync_bool_compare_and_swap( + (int8_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { + lhs += rhs; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + atomicAdd(&lhs, rhs); +#else + __sync_fetch_and_add(&lhs, rhs); +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { + rhs1 += rhs2; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + atomicAdd(&rhs1, rhs2); +#else + __sync_fetch_and_add(&rhs1, rhs2); +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { + lhs += rhs; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + // Apparently there is no signed 64bit int atomic yet + RHS newval = lhs, oldval; + // Type punning like this is illegal in C++ but the + // CUDA manual has an example just like it so fuck it + unsigned long long int *ptr = (unsigned long long int *)&lhs; + do { + oldval = newval; + newval += rhs; + newval = __ulonglong_as_longlong(atomicCAS( + ptr, __longlong_as_ulonglong(oldval), __longlong_as_ulonglong(newval))); + } while (oldval != newval); +#else + __sync_fetch_and_add(&lhs, rhs); +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { + rhs1 += rhs2; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + // Apparently there is no signed 64bit int atomic yet + RHS newval = rhs1, oldval; + // Type punning like this is illegal in C++ but the + // CUDA manual has an example just like it so fuck it + unsigned long long int *ptr = (unsigned long long int *)&rhs1; + do { + oldval = newval; + newval += rhs2; + newval = __ulonglong_as_longlong(atomicCAS( + ptr, __longlong_as_ulonglong(oldval), __longlong_as_ulonglong(newval))); + } while (oldval != newval); +#else + __sync_fetch_and_add(&rhs1, rhs2); +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction<__half>::apply(LHS &lhs, + RHS rhs) { + lhs = lhs + rhs; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction<__half>::apply(LHS &lhs, + RHS rhs) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) +#if (__CUDA_ARCH__ >= 700) && (__CUDACC_VER_MAJOR__ >= 10) + atomicAdd(&lhs, rhs); +#else + // 16-bit atomics are not supported prior to volta + // 32-bit GPU atomics need 4 byte alignment + const uintptr_t unaligned = reinterpret_cast(&lhs); + unsigned const offset = unaligned % sizeof(unsigned int); + const uintptr_t aligned = unaligned - offset; + unsigned int *ptr = reinterpret_cast(aligned); + RHS newval = lhs, oldval, other; + if (offset == 0) { + other = *((&lhs) + 1); + do { + oldval = newval; + newval = newval + rhs; + unsigned int const result = atomicCAS( + ptr, __hilohalf2uint(other, oldval), __hilohalf2uint(other, newval)); + newval = __uint2lohalf(result); + other = __uint2hihalf(result); + } while (oldval != newval); + } else { + other = *((&lhs) - 1); + do { + oldval = newval; + newval = newval + rhs; + unsigned int const result = atomicCAS( + ptr, __hilohalf2uint(oldval, other), __hilohalf2uint(newval, other)); + other = __uint2lohalf(result); + newval = __uint2hihalf(result); + } while (oldval != newval); + } +#endif +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(lhs); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval + rhs; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic floating point operations so use compare and swap + TypePunning::Alias> oldval, newval; + TypePunning::AlignedPointer pointer((void *)&lhs); + unsigned const offset = pointer.offset() / sizeof(__half); + do { + oldval.load(pointer); + std::array next = oldval.as_two(); + next[offset] = __convert_float_to_halfint( + __convert_halfint_to_float(next[offset]) + float(rhs)); + newval = next; + } while (!__sync_bool_compare_and_swap( + (int32_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction<__half>::fold(RHS &rhs1, + RHS rhs2) { + rhs1 = rhs1 + rhs2; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction<__half>::fold(RHS &rhs1, + RHS rhs2) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) +#if (__CUDA_ARCH__ >= 700) && (__CUDACC_VER_MAJOR__ >= 10) + atomicAdd(&rhs1, rhs2); +#else + // 16-bit atomics are not supported prior to volta + // 32-bit GPU atomics need 4 byte alignment + const uintptr_t unaligned = reinterpret_cast(&rhs1); + unsigned const offset = unaligned % sizeof(unsigned int); + const uintptr_t aligned = unaligned - offset; + unsigned int *ptr = reinterpret_cast(aligned); + RHS newval = rhs1, oldval, other; + if (offset == 0) { + other = *((&rhs1) + 1); + do { + oldval = newval; + newval = newval + rhs2; + unsigned int const result = atomicCAS( + ptr, __hilohalf2uint(other, oldval), __hilohalf2uint(other, newval)); + newval = __uint2lohalf(result); + other = __uint2hihalf(result); + } while (oldval != newval); + } else { + other = *((&rhs1) - 1); + do { + oldval = newval; + newval = newval + rhs2; + unsigned int const result = atomicCAS( + ptr, __hilohalf2uint(oldval, other), __hilohalf2uint(newval, other)); + other = __uint2lohalf(result); + newval = __uint2hihalf(result); + } while (oldval != newval); + } +#endif +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(rhs1); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval + rhs2; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic floating point operations so use compare and swap + TypePunning::Alias> oldval, newval; + TypePunning::AlignedPointer pointer((void *)&rhs1); + unsigned const offset = pointer.offset() / sizeof(__half); + do { + oldval.load(pointer); + std::array next = oldval.as_two(); + next[offset] = __convert_float_to_halfint( + __convert_halfint_to_float(next[offset]) + float(rhs2)); + newval = next; + } while (!__sync_bool_compare_and_swap( + (int32_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { + lhs += rhs; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + atomicAdd(&lhs, rhs); +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(lhs); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval + rhs; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic floating point operations so use compare and swap + TypePunning::Alias oldval, newval; + TypePunning::Pointer pointer((void *)&lhs); + do { + oldval.load(pointer); + newval = oldval.as_two() + rhs; + } while (!__sync_bool_compare_and_swap( + (int32_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { + rhs1 += rhs2; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) + atomicAdd(&rhs1, rhs2); +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(rhs1); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval + rhs2; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic floating point operations so use compare and swap + TypePunning::Alias oldval, newval; + TypePunning::Pointer pointer((void *)&rhs1); + do { + oldval.load(pointer); + newval = oldval.as_two() + rhs2; + } while (!__sync_bool_compare_and_swap( + (int32_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { + lhs += rhs; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, + RHS rhs) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) +#if (__CUDA_ARCH__ >= 600) || defined(__HIP_DEVICE_COMPILE__) + atomicAdd(&lhs, rhs); +#else + RHS newval = lhs, oldval; + // Type punning like this is illegal in C++ but the + // CUDA manual has an example just like it so fuck it + unsigned long long int *ptr = (unsigned long long int *)&lhs; + do { + oldval = newval; + newval += rhs; + newval = __ulonglong_as_double(atomicCAS( + ptr, __double_as_ulonglong(oldval), __double_as_ulonglong(newval))); + } while (oldval != newval); +#endif +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(lhs); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval + rhs; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic floating point operations so use compare and swap + TypePunning::Alias oldval, newval; + TypePunning::Pointer pointer((void *)&lhs); + do { + oldval.load(pointer); + newval = oldval.as_two() + rhs; + } while (!__sync_bool_compare_and_swap( + (int64_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { + rhs1 += rhs2; +} + +template <> +__LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, + RHS rhs2) { +#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) +#if (__CUDA_ARCH__ >= 600) || defined(__HIP_DEVICE_COMPILE__) + atomicAdd(&rhs1, rhs2); +#else + RHS newval = rhs1, oldval; + // Type punning like this is illegal in C++ but the + // CUDA manual has an example just like it so fuck it + unsigned long long int *ptr = (unsigned long long int *)&rhs1; + do { + oldval = newval; + newval += rhs2; + newval = __ulonglong_as_double(atomicCAS( + ptr, __double_as_ulonglong(oldval), __double_as_ulonglong(newval))); + } while (oldval != newval); +#endif +#else +#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) + std::atomic_ref atomic(rhs1); + RHS oldval = atomic.load(); + RHS newval; + do { + newval = oldval + rhs2; + } while (!atomic.compare_exchange_weak(oldval, newval)); +#else + // No atomic floating point operations so use compare and swap + TypePunning::Alias oldval, newval; + TypePunning::Pointer pointer((void *)&rhs1); + do { + oldval.load(pointer); + newval = oldval.as_two() + rhs2; + } while (!__sync_bool_compare_and_swap( + (int64_t *)pointer, oldval.as_one(), newval.as_one())); +#endif +#endif +} + +void Realm::Event register_all_redops(Realm::Runtime rt) { + // Registration is synchronous, so no need to capture events here + rt.register_reduction>( + get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_BOOL_REDOP_ID)); + rt.register_reduction>( + get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_INT32_REDOP_ID)); + rt.register_reduction>( + get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_INT64_REDOP_ID)); + rt.register_reduction>( + get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_HALF_REDOP_ID)); + rt.register_reduction>( + get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_FLOAT_REDOP_ID)); + rt.register_reduction>( + get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_DOUBLE_REDOP_ID)); +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/redops/redop_id_t.cc b/lib/realm-execution/src/realm-execution/redops/redop_id_t.cc new file mode 100644 index 0000000000..702ddd5e97 --- /dev/null +++ b/lib/realm-execution/src/realm-execution/redops/redop_id_t.cc @@ -0,0 +1,32 @@ +#include "realm-execution/redops/redop_id_t.h" + +namespace FlexFlow { + +Realm::ReductionOpID get_sum_redop_id_for_data_type(DataType) { + + switch (dtype) { + case DataType::BOOL: + return redop_id_t::SUM_BOOL_REDOP_ID; + case DataType::INT32: + return redop_id_t::SUM_INT32_REDOP_ID; + case DataType::INT64: + return redop_id_t::SUM_INT64_REDOP_ID; + case DataType::HALF: + return redop_id_t::SUM_HALF_REDOP_ID; + case DataType::FLOAT: + return redop_id_t::SUM_FLOAT_REDOP_ID; + case DataType::DOUBLE: + return redop_id_t::SUM_DOUBLE_REDOP_ID; + default: + PANIC("No known sum reduction for data type {}", dtype); + } +} + +Realm::Processor::ReductionOpID + get_realm_reduction_op_id_for_redop_id(redop_id_t redop_id) { + return static_cast(redop_id); +} + +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc index acafdf59fd..e7a8948f8d 100644 --- a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc +++ b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc @@ -5,7 +5,6 @@ #include "realm-execution/tasks/impl/op_task.h" #include "realm-execution/tasks/impl/per_device_op_state_init_return_task.h" #include "realm-execution/tasks/impl/per_device_op_state_init_task.h" -#include "realm-execution/tasks/realm_reduction.h" #include "realm-execution/tasks/task_id_t.h" #include "utils/exception.h" @@ -31,18 +30,9 @@ Realm::Event register_task(Realm::Processor::Kind target_kind, Realm::ProfilingRequestSet()); } -static void register_reductions() { - // register sum reduction ops - Realm::Runtime rt = Realm::Runtime::get_runtime(); - rt.register_reduction(REDOP_SUM_FLOAT); - rt.register_reduction(REDOP_SUM_DOUBLE); - // register_reduction is synchronous — no event returned -} - Realm::Event register_all_tasks() { std::vector pending_registrations; - register_reductions(); std::vector init_task_ids = { // Init tasks task_id_t::BATCHNORM_INIT_TASK_ID, From 0d589f2056022d1068deb2d6abf745e3dcc048cd Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 21 May 2026 12:39:41 -0700 Subject: [PATCH 11/35] Fix build for reductions. --- .../redops/realm_redop_registry.h | 6 +- .../redops/redop_id_t.dtg.toml | 3 - .../realm-execution/redops/redop_id_t.h | 12 +- .../src/realm-execution/pcg_instance.cc | 7 +- .../src/realm-execution/realm_manager.cc | 2 +- .../redops/realm_redop_registry.cc | 165 +----------------- .../src/realm-execution/redops/redop_id_t.cc | 12 +- 7 files changed, 26 insertions(+), 181 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h b/lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h index a338a38bbf..e7e51326e1 100644 --- a/lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h +++ b/lib/realm-execution/include/realm-execution/redops/realm_redop_registry.h @@ -1,5 +1,5 @@ -#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H -#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_REDOPS_REALM_REDOP_REGISTRY_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_REDOPS_REALM_REDOP_REGISTRY_H #include "realm-execution/realm.h" #include "realm-execution/redops/redop_id_t.dtg.h" @@ -9,7 +9,7 @@ namespace FlexFlow { /** * \brief Registers all known reduction operators (redops). */ -void Realm::Event register_all_redops(Realm::Runtime); +void register_all_redops(Realm::Runtime); } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml b/lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml index 5183ff5e72..44e1f32c59 100644 --- a/lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml +++ b/lib/realm-execution/include/realm-execution/redops/redop_id_t.dtg.toml @@ -20,9 +20,6 @@ name = "SUM_INT32_REDOP_ID" [[values]] name = "SUM_INT64_REDOP_ID" -[[values]] -name = "SUM_HALF_REDOP_ID" - [[values]] name = "SUM_FLOAT_REDOP_ID" diff --git a/lib/realm-execution/include/realm-execution/redops/redop_id_t.h b/lib/realm-execution/include/realm-execution/redops/redop_id_t.h index b9ef91a05a..8565b20b17 100644 --- a/lib/realm-execution/include/realm-execution/redops/redop_id_t.h +++ b/lib/realm-execution/include/realm-execution/redops/redop_id_t.h @@ -1,21 +1,21 @@ -#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H -#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_TASKS_REALM_REDOP_REGISTRY_H +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_REDOPS_REALM_REDOP_ID_T_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_REDOPS_REALM_REDOP_ID_T_H +#include "op-attrs/datatype.dtg.h" #include "realm-execution/realm.h" #include "realm-execution/redops/redop_id_t.dtg.h" namespace FlexFlow { /** - * \brief Registers all known reduction operators (redops). + * \brief Return the sum reduction operator (redop) ID for a given data type. */ -Realm::ReductionOpID get_sum_redop_id_for_data_type(DataType); +redop_id_t get_sum_redop_id_for_data_type(DataType); /** * \brief Convert a \ref FlexFlow::redop_id_t into a Realm reduction op ID. */ -Realm::Processor::ReductionOpID - get_realm_reduction_op_id_for_redop_id(redop_id_t); +Realm::ReductionOpID get_realm_reduction_op_id_for_redop_id(redop_id_t); } // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 332669a9dc..1ac3821142 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -5,8 +5,8 @@ #include "realm-execution/distributed_per_device_op_state_initialization.h" #include "realm-execution/instance_allocation.h" #include "realm-execution/realm_context.h" +#include "realm-execution/redops/redop_id_t.h" #include "realm-execution/tasks/impl/op_task.h" -#include "realm-execution/tasks/realm_reduction.h" #include "realm-execution/tensor_instance_backing.h" #include "task-spec/dynamic_graph/copy_insertion.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" @@ -228,8 +228,9 @@ static Realm::Event spawn_dynamic_node_invocation( Realm::RegionInstance dst_inst = tensor_instance_backing.backing.at(input_grad).first; - Realm::ReductionOpID redop_id = get_sum_reduction_op_id( - assert_unwrap(output_grad.parallel_tensor_shape).data_type); + Realm::ReductionOpID redop_id = + get_realm_reduction_op_id_for_redop_id(get_sum_redop_id_for_data_type( + assert_unwrap(output_grad.parallel_tensor_shape).data_type)); // chain reductions sequentially to avoid write races on dst Realm::Event e = precondition; diff --git a/lib/realm-execution/src/realm-execution/realm_manager.cc b/lib/realm-execution/src/realm-execution/realm_manager.cc index c7136d8a98..5a8f9cbbbb 100644 --- a/lib/realm-execution/src/realm-execution/realm_manager.cc +++ b/lib/realm-execution/src/realm-execution/realm_manager.cc @@ -1,7 +1,7 @@ #include "realm-execution/realm_manager.h" #include "realm-execution/realm_context.h" -#include "realm-execution/tasks/realm_task_registry.h" #include "realm-execution/redops/realm_redop_registry.h" +#include "realm-execution/tasks/realm_task_registry.h" namespace FlexFlow { diff --git a/lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc b/lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc index d10b158463..ab3304836a 100644 --- a/lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc +++ b/lib/realm-execution/src/realm-execution/redops/realm_redop_registry.cc @@ -1,4 +1,5 @@ #include "realm-execution/redops/realm_redop_registry.h" +#include "realm-execution/redops/redop_id_t.h" namespace FlexFlow { @@ -120,6 +121,12 @@ class Alias { #define __LEGION_CUDA_HD__ #endif +template +class SumReduction { + // Empty definition + // Specializations provided for each type +}; + template <> class SumReduction { public: @@ -127,7 +134,6 @@ class SumReduction { typedef bool RHS; static constexpr bool identity = false; - static constexpr int REDOP_ID = LEGION_REDOP_OR_BOOL; template __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); @@ -142,7 +148,6 @@ class SumReduction { typedef int32_t RHS; static constexpr int32_t identity = 0; - static constexpr int REDOP_ID = LEGION_REDOP_SUM_INT32; template __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); @@ -157,22 +162,6 @@ class SumReduction { typedef int64_t RHS; static constexpr int64_t identity = 0; - static constexpr int REDOP_ID = LEGION_REDOP_SUM_INT64; - - template - __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); - template - __LEGION_CUDA_HD__ static void fold(RHS &rhs1, RHS rhs2); -}; - -template <> -class SumReduction<__half> { -public: - typedef __half LHS; - typedef __half RHS; - - static inline const __half identity = __half(0, false /*raw*/); - static constexpr int REDOP_ID = LEGION_REDOP_SUM_FLOAT16; template __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); @@ -187,7 +176,6 @@ class SumReduction { typedef float RHS; static constexpr float identity = 0.f; - static constexpr int REDOP_ID = LEGION_REDOP_SUM_FLOAT32; template __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); @@ -202,7 +190,6 @@ class SumReduction { typedef double RHS; static constexpr double identity = 0.0; - static constexpr int REDOP_ID = LEGION_REDOP_SUM_FLOAT64; template __LEGION_CUDA_HD__ static void apply(LHS &lhs, RHS rhs); @@ -382,140 +369,6 @@ __LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, #endif } -template <> -__LEGION_CUDA_HD__ inline void SumReduction<__half>::apply(LHS &lhs, - RHS rhs) { - lhs = lhs + rhs; -} - -template <> -__LEGION_CUDA_HD__ inline void SumReduction<__half>::apply(LHS &lhs, - RHS rhs) { -#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) -#if (__CUDA_ARCH__ >= 700) && (__CUDACC_VER_MAJOR__ >= 10) - atomicAdd(&lhs, rhs); -#else - // 16-bit atomics are not supported prior to volta - // 32-bit GPU atomics need 4 byte alignment - const uintptr_t unaligned = reinterpret_cast(&lhs); - unsigned const offset = unaligned % sizeof(unsigned int); - const uintptr_t aligned = unaligned - offset; - unsigned int *ptr = reinterpret_cast(aligned); - RHS newval = lhs, oldval, other; - if (offset == 0) { - other = *((&lhs) + 1); - do { - oldval = newval; - newval = newval + rhs; - unsigned int const result = atomicCAS( - ptr, __hilohalf2uint(other, oldval), __hilohalf2uint(other, newval)); - newval = __uint2lohalf(result); - other = __uint2hihalf(result); - } while (oldval != newval); - } else { - other = *((&lhs) - 1); - do { - oldval = newval; - newval = newval + rhs; - unsigned int const result = atomicCAS( - ptr, __hilohalf2uint(oldval, other), __hilohalf2uint(newval, other)); - other = __uint2lohalf(result); - newval = __uint2hihalf(result); - } while (oldval != newval); - } -#endif -#else -#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) - std::atomic_ref atomic(lhs); - RHS oldval = atomic.load(); - RHS newval; - do { - newval = oldval + rhs; - } while (!atomic.compare_exchange_weak(oldval, newval)); -#else - // No atomic floating point operations so use compare and swap - TypePunning::Alias> oldval, newval; - TypePunning::AlignedPointer pointer((void *)&lhs); - unsigned const offset = pointer.offset() / sizeof(__half); - do { - oldval.load(pointer); - std::array next = oldval.as_two(); - next[offset] = __convert_float_to_halfint( - __convert_halfint_to_float(next[offset]) + float(rhs)); - newval = next; - } while (!__sync_bool_compare_and_swap( - (int32_t *)pointer, oldval.as_one(), newval.as_one())); -#endif -#endif -} - -template <> -__LEGION_CUDA_HD__ inline void SumReduction<__half>::fold(RHS &rhs1, - RHS rhs2) { - rhs1 = rhs1 + rhs2; -} - -template <> -__LEGION_CUDA_HD__ inline void SumReduction<__half>::fold(RHS &rhs1, - RHS rhs2) { -#if defined(__CUDA_ARCH__) || defined(__HIP_DEVICE_COMPILE__) -#if (__CUDA_ARCH__ >= 700) && (__CUDACC_VER_MAJOR__ >= 10) - atomicAdd(&rhs1, rhs2); -#else - // 16-bit atomics are not supported prior to volta - // 32-bit GPU atomics need 4 byte alignment - const uintptr_t unaligned = reinterpret_cast(&rhs1); - unsigned const offset = unaligned % sizeof(unsigned int); - const uintptr_t aligned = unaligned - offset; - unsigned int *ptr = reinterpret_cast(aligned); - RHS newval = rhs1, oldval, other; - if (offset == 0) { - other = *((&rhs1) + 1); - do { - oldval = newval; - newval = newval + rhs2; - unsigned int const result = atomicCAS( - ptr, __hilohalf2uint(other, oldval), __hilohalf2uint(other, newval)); - newval = __uint2lohalf(result); - other = __uint2hihalf(result); - } while (oldval != newval); - } else { - other = *((&rhs1) - 1); - do { - oldval = newval; - newval = newval + rhs2; - unsigned int const result = atomicCAS( - ptr, __hilohalf2uint(oldval, other), __hilohalf2uint(newval, other)); - other = __uint2lohalf(result); - newval = __uint2hihalf(result); - } while (oldval != newval); - } -#endif -#else -#if defined(__cpp_lib_atomic_ref) && (__cpp_lib_atomic_ref >= 201806L) - std::atomic_ref atomic(rhs1); - RHS oldval = atomic.load(); - RHS newval; - do { - newval = oldval + rhs2; - } while (!atomic.compare_exchange_weak(oldval, newval)); -#else - // No atomic floating point operations so use compare and swap - TypePunning::Alias> oldval, newval; - TypePunning::AlignedPointer pointer((void *)&rhs1); - unsigned const offset = pointer.offset() / sizeof(__half); - do { - oldval.load(pointer); - std::array next = oldval.as_two(); - next[offset] = __convert_float_to_halfint( - __convert_halfint_to_float(next[offset]) + float(rhs2)); - newval = next; - } while (!__sync_bool_compare_and_swap( - (int32_t *)pointer, oldval.as_one(), newval.as_one())); -#endif -#endif -} - template <> __LEGION_CUDA_HD__ inline void SumReduction::apply(LHS &lhs, RHS rhs) { @@ -670,7 +523,7 @@ __LEGION_CUDA_HD__ inline void SumReduction::fold(RHS &rhs1, #endif } -void Realm::Event register_all_redops(Realm::Runtime rt) { +void register_all_redops(Realm::Runtime rt) { // Registration is synchronous, so no need to capture events here rt.register_reduction>( get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_BOOL_REDOP_ID)); @@ -678,8 +531,6 @@ void Realm::Event register_all_redops(Realm::Runtime rt) { get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_INT32_REDOP_ID)); rt.register_reduction>( get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_INT64_REDOP_ID)); - rt.register_reduction>( - get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_HALF_REDOP_ID)); rt.register_reduction>( get_realm_reduction_op_id_for_redop_id(redop_id_t::SUM_FLOAT_REDOP_ID)); rt.register_reduction>( diff --git a/lib/realm-execution/src/realm-execution/redops/redop_id_t.cc b/lib/realm-execution/src/realm-execution/redops/redop_id_t.cc index 702ddd5e97..f31769419f 100644 --- a/lib/realm-execution/src/realm-execution/redops/redop_id_t.cc +++ b/lib/realm-execution/src/realm-execution/redops/redop_id_t.cc @@ -1,9 +1,9 @@ #include "realm-execution/redops/redop_id_t.h" +#include "utils/exception.h" namespace FlexFlow { -Realm::ReductionOpID get_sum_redop_id_for_data_type(DataType) { - +redop_id_t get_sum_redop_id_for_data_type(DataType dtype) { switch (dtype) { case DataType::BOOL: return redop_id_t::SUM_BOOL_REDOP_ID; @@ -11,8 +11,6 @@ Realm::ReductionOpID get_sum_redop_id_for_data_type(DataType) { return redop_id_t::SUM_INT32_REDOP_ID; case DataType::INT64: return redop_id_t::SUM_INT64_REDOP_ID; - case DataType::HALF: - return redop_id_t::SUM_HALF_REDOP_ID; case DataType::FLOAT: return redop_id_t::SUM_FLOAT_REDOP_ID; case DataType::DOUBLE: @@ -22,11 +20,9 @@ Realm::ReductionOpID get_sum_redop_id_for_data_type(DataType) { } } -Realm::Processor::ReductionOpID +Realm::ReductionOpID get_realm_reduction_op_id_for_redop_id(redop_id_t redop_id) { - return static_cast(redop_id); -} - + return static_cast(redop_id); } } // namespace FlexFlow From 5ef0b070ad9cf21cbae6346e8294074fb81e74ad Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 21 May 2026 14:45:34 -0700 Subject: [PATCH 12/35] Split reduction from copy and put back device op state init code. --- .../include/realm-execution/realm_context.h | 29 ++-- ...uted_per_device_op_state_initialization.cc | 6 +- .../src/realm-execution/pcg_instance.cc | 19 ++- .../src/realm-execution/realm_context.cc | 127 ++++++++++++------ 4 files changed, 112 insertions(+), 69 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/realm_context.h b/lib/realm-execution/include/realm-execution/realm_context.h index eab42d0d79..5b76d52e2c 100644 --- a/lib/realm-execution/include/realm-execution/realm_context.h +++ b/lib/realm-execution/include/realm-execution/realm_context.h @@ -9,6 +9,7 @@ #include "pcg/device_id_t.dtg.h" #include "pcg/machine_space_coordinate.dtg.h" #include "realm-execution/realm.h" +#include "realm-execution/redops/redop_id_t.dtg.h" #include "realm-execution/tasks/task_id_t.dtg.h" #include #include @@ -65,16 +66,24 @@ struct RealmContext { /** \name Data movement and reduction */ ///\{ - Realm::Event - issue_copy(ParallelTensorShape const &src_shape, - Realm::RegionInstance src_inst, - ParallelTensorShape const &dst_shape, - Realm::RegionInstance dst_inst, - Realm::ProfilingRequestSet const &requests, - Realm::Event wait_on = Realm::Event::NO_EVENT, - int priority = 0, - std::optional redop_id = std::nullopt, - bool exclusive = false); + Realm::Event issue_copy(ParallelTensorShape const &src_shape, + Realm::RegionInstance src_inst, + ParallelTensorShape const &dst_shape, + Realm::RegionInstance dst_inst, + Realm::ProfilingRequestSet const &requests, + Realm::Event wait_on = Realm::Event::NO_EVENT, + int priority = 0); + + Realm::Event issue_reduction(ParallelTensorShape const &src_shape, + Realm::RegionInstance src_inst, + ParallelTensorShape const &dst_shape, + Realm::RegionInstance dst_inst, + redop_id_t redop_id, + bool is_fold, + bool exclusive, + Realm::ProfilingRequestSet const &requests, + Realm::Event wait_on = Realm::Event::NO_EVENT, + int priority = 0); ///\} /** \name Instance management */ diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index e7d8647b12..1d517a8fe4 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -31,7 +31,6 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( std::unordered_map *> device_state_map; - std::vector completion_events; for (DynamicNodeInvocation const &invocation : dg.invocations) { Realm::Processor target_proc = ctx.map_device_coord_to_processor( assert_unwrap(invocation.node_attrs.device_coord)); @@ -57,7 +56,6 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( precondition); if (completion_event.has_value()) { - completion_events.push_back(completion_event.value()); device_state_map.insert(std::pair{invocation, device_state_ptr}); } else { // Task doesn't require initialization, clean up and don't store result @@ -65,9 +63,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( } } - // wait for all init tasks — direct write to *result_ptr happens - // before each init task event fires so result is ready after this - Realm::Event::merge_events(completion_events).wait(); + ctx.get_outstanding_events().wait(); auto deref = [](DeviceSpecificPtr *const &p) { return *p; }; std::unordered_map> diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 1ac3821142..aa67110127 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -228,12 +228,11 @@ static Realm::Event spawn_dynamic_node_invocation( Realm::RegionInstance dst_inst = tensor_instance_backing.backing.at(input_grad).first; - Realm::ReductionOpID redop_id = - get_realm_reduction_op_id_for_redop_id(get_sum_redop_id_for_data_type( - assert_unwrap(output_grad.parallel_tensor_shape).data_type)); + redop_id_t redop_id = get_sum_redop_id_for_data_type( + assert_unwrap(output_grad.parallel_tensor_shape).data_type); // chain reductions sequentially to avoid write races on dst - Realm::Event e = precondition; + Realm::Event result = precondition; for (auto const &[p, m] : assert_unwrap(output_grad.mapping)) { DynamicValueAttrs replica_key = output_grad; replica_key.mapping = @@ -243,18 +242,18 @@ static Realm::Event spawn_dynamic_node_invocation( Realm::RegionInstance src_inst = tensor_instance_backing.backing.at(replica_key).first; - e = ctx.issue_copy( + result = ctx.issue_reduction( /*src_shape=*/assert_unwrap(output_grad.parallel_tensor_shape), /*src_inst=*/src_inst, /*dst_shape=*/assert_unwrap(input_grad.parallel_tensor_shape), /*dst_inst=*/dst_inst, - /*requests=*/Realm::ProfilingRequestSet{}, - /*wait_on=*/e, - /*priority=*/0, /*redop_id=*/redop_id, - /*exlusive=*/false); + /*is_fold=*/false, + /*exlusive=*/false, + /*requests=*/Realm::ProfilingRequestSet{}, + /*wait_on=*/result); } - return e; + return result; }; TrainingOperationAttrs op_attrs = diff --git a/lib/realm-execution/src/realm-execution/realm_context.cc b/lib/realm-execution/src/realm-execution/realm_context.cc index a4669bf43e..36dd7c71cc 100644 --- a/lib/realm-execution/src/realm-execution/realm_context.cc +++ b/lib/realm-execution/src/realm-execution/realm_context.cc @@ -7,6 +7,7 @@ #include "pcg/device_id_t.h" #include "pcg/device_type.dtg.h" #include "realm-execution/realm_allocator.h" +#include "realm-execution/redops/redop_id_t.h" #include "realm-execution/tasks/task_id_t.dtg.h" #include "realm-execution/tasks/task_id_t.h" #include "utils/containers/contains_key.h" @@ -154,6 +155,46 @@ static Realm::IndexSpace ispace_from_dims(TensorDims const &dims) { return Realm::IndexSpace{rect}; } +[[nodiscard]] static Realm::Event + issue_copy_for_field(TensorDims const &dims, + Realm::CopySrcDstField const &src_field, + Realm::CopySrcDstField const &dst_field, + Realm::ProfilingRequestSet const &requests, + Realm::Event wait_on, + int priority) { + switch (dims.ff_ordered.num_dims()) { +#if REALM_MAX_DIM >= 1 + case 1: + return ispace_from_dims<1>(dims).copy( + {src_field}, {dst_field}, requests, wait_on, priority); +#endif +#if REALM_MAX_DIM >= 2 + case 2: + return ispace_from_dims<2>(dims).copy( + {src_field}, {dst_field}, requests, wait_on, priority); +#endif +#if REALM_MAX_DIM >= 3 + case 3: + return ispace_from_dims<3>(dims).copy( + {src_field}, {dst_field}, requests, wait_on, priority); +#endif +#if REALM_MAX_DIM >= 4 + case 4: + return ispace_from_dims<4>(dims).copy( + {src_field}, {dst_field}, requests, wait_on, priority); +#endif +#if REALM_MAX_DIM >= 5 + case 5: + return ispace_from_dims<5>(dims).copy( + {src_field}, {dst_field}, requests, wait_on, priority); +#endif + default: + PANIC("TensorShape dims greater than REALM_MAX_DIM: {}", + dims.ff_ordered.num_dims()); + break; + } +} + Realm::Event RealmContext::issue_copy(ParallelTensorShape const &src_shape, Realm::RegionInstance src_inst, @@ -161,9 +202,7 @@ Realm::Event Realm::RegionInstance dst_inst, Realm::ProfilingRequestSet const &requests, Realm::Event wait_on, - int priority, - std::optional redop_id, - bool exclusive) { + int priority) { TensorShape src_piece_shape = get_piece_shape(src_shape); TensorShape dst_piece_shape = get_piece_shape(dst_shape); ASSERT(src_piece_shape == dst_piece_shape); // For now, assume they match @@ -185,48 +224,48 @@ Realm::Event size_of_datatype(src_piece_shape.data_type).int_from_positive_int()), /*subfield_offset=*/0); - // set reduction op on dst field if provided - if (redop_id.has_value()) { - dst_field.set_redop(redop_id.value(), /*is_fold=*/false, exclusive); - } + Realm::Event result = issue_copy_for_field( + src_piece_shape.dims, src_field, dst_field, requests, wait_on, priority); + this->outstanding_events.push_back(result); + return result; +} - Realm::Event result; - switch (src_piece_shape.dims.ff_ordered.num_dims()) { -#if REALM_MAX_DIM >= 1 - case 1: - result = ispace_from_dims<1>(src_piece_shape.dims) - .copy({src_field}, {dst_field}, requests, wait_on, priority); - break; -#endif -#if REALM_MAX_DIM >= 2 - case 2: - result = ispace_from_dims<2>(src_piece_shape.dims) - .copy({src_field}, {dst_field}, requests, wait_on, priority); - break; -#endif -#if REALM_MAX_DIM >= 3 - case 3: - result = ispace_from_dims<3>(src_piece_shape.dims) - .copy({src_field}, {dst_field}, requests, wait_on, priority); - break; -#endif -#if REALM_MAX_DIM >= 4 - case 4: - result = ispace_from_dims<4>(src_piece_shape.dims) - .copy({src_field}, {dst_field}, requests, wait_on, priority); - break; -#endif -#if REALM_MAX_DIM >= 5 - case 5: - result = ispace_from_dims<5>(src_piece_shape.dims) - .copy({src_field}, {dst_field}, requests, wait_on, priority); - break; -#endif - default: - PANIC("TensorShape dims greater than REALM_MAX_DIM: {}", - src_piece_shape.dims.ff_ordered.num_dims()); - break; - } +Realm::Event + RealmContext::issue_reduction(ParallelTensorShape const &src_shape, + Realm::RegionInstance src_inst, + ParallelTensorShape const &dst_shape, + Realm::RegionInstance dst_inst, + redop_id_t redop_id, + bool is_fold, + bool exclusive, + Realm::ProfilingRequestSet const &requests, + Realm::Event wait_on, + int priority) { + TensorShape src_piece_shape = get_piece_shape(src_shape); + TensorShape dst_piece_shape = get_piece_shape(dst_shape); + ASSERT(src_piece_shape == dst_piece_shape); // For now, assume they match + + Realm::CopySrcDstField src_field; + src_field.set_field( + /*inst=*/src_inst, + /*field_id=*/0, + /*size=*/ + static_cast( + size_of_datatype(src_piece_shape.data_type).int_from_positive_int()), + /*subfield_offset=*/0); + Realm::CopySrcDstField dst_field; + dst_field.set_field( + /*inst=*/dst_inst, + /*field_id=*/0, + /*size=*/ + static_cast( + size_of_datatype(src_piece_shape.data_type).int_from_positive_int()), + /*subfield_offset=*/0); + dst_field.set_redop( + get_realm_reduction_op_id_for_redop_id(redop_id), is_fold, exclusive); + + Realm::Event result = issue_copy_for_field( + src_piece_shape.dims, src_field, dst_field, requests, wait_on, priority); this->outstanding_events.push_back(result); return result; } From b73751727c77c92358624fb9bd9997e3f77e73c3 Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 21 May 2026 14:48:55 -0700 Subject: [PATCH 13/35] Replicate is not a task, don't represent it as one. --- .../include/realm-execution/tasks/task_id_t.dtg.toml | 9 --------- .../src/realm-execution/tasks/realm_task_registry.cc | 3 --- 2 files changed, 12 deletions(-) diff --git a/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml index b1e5e07e28..b0bcc23b4d 100644 --- a/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/task_id_t.dtg.toml @@ -327,15 +327,6 @@ name = "COMBINE_FWD_TASK_ID" [[values]] name = "COMBINE_BWD_TASK_ID" -[[values]] -name = "REPLICATE_INIT_TASK_ID" - -[[values]] -name = "REPLICATE_FWD_TASK_ID" - -[[values]] -name = "REPLICATE_BWD_TASK_ID" - [[values]] name = "REDUCTION_INIT_TASK_ID" diff --git a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc index e7a8948f8d..dfdfe72ce0 100644 --- a/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc +++ b/lib/realm-execution/src/realm-execution/tasks/realm_task_registry.cc @@ -49,7 +49,6 @@ Realm::Event register_all_tasks() { task_id_t::REDUCE_INIT_TASK_ID, task_id_t::REDUCTION_INIT_TASK_ID, task_id_t::REPARTITION_INIT_TASK_ID, - task_id_t::REPLICATE_INIT_TASK_ID, task_id_t::SOFTMAX_INIT_TASK_ID, }; @@ -86,7 +85,6 @@ Realm::Event register_all_tasks() { task_id_t::REDUCE_FWD_TASK_ID, task_id_t::REDUCTION_FWD_TASK_ID, task_id_t::REPARTITION_FWD_TASK_ID, - task_id_t::REPLICATE_FWD_TASK_ID, task_id_t::RESHAPE_FWD_TASK_ID, task_id_t::REVERSE_FWD_TASK_ID, task_id_t::SOFTMAX_FWD_TASK_ID, @@ -115,7 +113,6 @@ Realm::Event register_all_tasks() { task_id_t::REDUCE_BWD_TASK_ID, task_id_t::REDUCTION_BWD_TASK_ID, task_id_t::REPARTITION_BWD_TASK_ID, - task_id_t::REPLICATE_BWD_TASK_ID, task_id_t::RESHAPE_BWD_TASK_ID, task_id_t::REVERSE_BWD_TASK_ID, task_id_t::SOFTMAX_BWD_TASK_ID, From c536cc96bfd5b4b88739dcbcfcab9f00a96a7fed Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 21 May 2026 14:49:47 -0700 Subject: [PATCH 14/35] Put back the per device op state return code path. --- .../tasks/impl/per_device_op_state_init_task.cc | 16 +++++----------- 1 file changed, 5 insertions(+), 11 deletions(-) diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc index 0ea51810e4..753fccf74b 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc @@ -66,17 +66,11 @@ void per_device_op_state_init_task_body(void const *args, result_state, ctx.get_current_device_idx())}; DeviceSpecificPtr result_device_specific{ ctx.get_current_device_idx(), result_state_ptr}; - - // replace spawn_per_device_op_state_init_return_task with: - // NOTE: SM/TODO: direct write assumes single-node shared address space - // For multi-node, replace with UserEvent trigger pattern - *task_args.origin_result_ptr = result_device_specific; - - // spawn_per_device_op_state_init_return_task(ctx, - // task_args.origin_proc, - // result_device_specific, - // task_args.origin_result_ptr, - // Realm::Event::NO_EVENT); + spawn_per_device_op_state_init_return_task(ctx, + task_args.origin_proc, + result_device_specific, + task_args.origin_result_ptr, + Realm::Event::NO_EVENT); } std::optional spawn_per_device_op_state_init_task( From 48673bc64d390ea29c412f3d78661e4bd80b7b7e Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Thu, 28 May 2026 03:50:13 -0700 Subject: [PATCH 15/35] Changes for breaking up the replicate PR testing code --- .../pcg_task_graph.dtg.toml | 18 +- .../abstracted_single_tensor_movement.cc | 2 +- ...racted_tensor_set_movement_across_split.cc | 11 +- ...substitution_and_update_machine_mapping.cc | 2 +- .../get_optimal_machine_mapping.cc | 4 +- .../get_tensor_set_movement_across_split.cc | 1 - .../machine_mapping/machine_mapping.cc | 8 +- .../machine_mapping_constraints.cc | 12 +- .../compiler/machine_mapping/machine_view.cc | 4 +- ...get_optimal_machine_mapping_with_memory.cc | 4 +- ...el_layer_guid_oblivious_machine_mapping.cc | 10 +- .../pcg/pcg_binary_sp_decomposition.cc | 2 +- .../task_graph_simulator/pcg_task_graph.cc | 17 +- .../task_graph_simulator/task_simulator.cc | 2 +- .../unity_algorithm/unity_algorithm.cc | 2 +- .../src/local-execution/tensor_allocation.cc | 2 +- .../src/op-attrs/get_incoming_tensor_roles.cc | 2 +- .../op-attrs/parallel_tensor_dim_degrees.cc | 6 +- .../parallel_tensor_space_coordinate.cc | 4 +- lib/op-attrs/src/op-attrs/shape_inference.cc | 2 +- .../mapped_operator_task_group.h | 6 +- .../mapped_parallel_computation_graph.h | 4 + .../mapped_parallel_layer_info.dtg.toml | 31 + ...ed_parallel_layer_invocation_info.dtg.toml | 33 + .../mapped_parallel_layer_invocation_info.h | 17 + .../parallel_computation_graph.h | 8 + .../parallel_layer_info.dtg.toml | 26 + .../parallel_layer_invocation_info.dtg.toml | 33 + .../parallel_tensor_info.dtg.toml | 26 + lib/pcg/src/pcg/computation_graph.cc | 2 +- lib/pcg/src/pcg/computation_graph_builder.cc | 8 +- .../mapped_operator_task_group.cc | 38 +- .../mapped_parallel_computation_graph.cc | 38 + .../mapped_parallel_layer_invocation_info.cc | 22 + .../parallel_computation_graph.cc | 48 +- .../parallel_computation_graph_builder.cc | 8 +- .../mapped_parallel_computation_graph.cc | 64 +- .../parallel_computation_graph_builder.cc | 4 +- .../realm-execution/instance_allocation.cc | 2 +- .../src/realm-execution/pcg_instance.cc | 4 +- .../serializable_tensor_instance_backing.cc | 12 +- .../src/realm-execution/test_op_replicate.cc | 2 + lib/runtime/src/parallel_tensor_uses.cc | 2 +- lib/runtime/src/tensor_uses.cc | 2 +- .../apply_substitution/apply_substitution.cc | 4 +- .../perform_shape_inference.cc | 4 +- .../src/substitutions/pcg_pattern_match.cc | 2 +- .../src/substitutions/substitution_builder.cc | 2 +- .../unlabelled/find_pattern_matches.cc | 7 +- .../unlabelled/pattern_matching.cc | 2 +- .../task-spec/dynamic_graph/copy_insertion.h | 5 + .../dynamic_graph/dynamic_node_invocation.h | 20 + ...mic_node_invocation_sharding_info.dtg.toml | 30 + .../dynamic_open_dataflow_graph.h | 6 + .../dynamic_value_attrs.dtg.toml | 29 +- .../dynamic_graph/dynamic_value_attrs.h | 4 + ...dynamic_value_attrs_sharding_info.dtg.toml | 27 + .../include/task-spec/dynamic_graph/index.dox | 1 + ...amic_open_dataflow_graph_from_mapped_pcg.h | 6 + .../task-spec/dynamic_graph/pass_expansion.h | 5 + .../serializable_dynamic_value_attrs.dtg.toml | 4 +- .../task-spec/dynamic_graph/shard_expansion.h | 35 +- .../dynamic_graph/update_insertion.h | 9 + .../task-spec/dynamic_graph/copy_insertion.cc | 125 ++- .../dynamic_graph/dynamic_node_invocation.cc | 35 + .../dynamic_open_dataflow_graph.cc | 13 + .../dynamic_graph/dynamic_value_attrs.cc | 13 + ...ake_dynamic_open_dataflow_graph_from_cg.cc | 2 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 213 +---- .../task-spec/dynamic_graph/pass_expansion.cc | 18 + .../dynamic_graph/shard_expansion.cc | 347 ++++--- .../task-spec/dynamic_graph/copy_insertion.cc | 873 +++++++++++------- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 396 ++++++++ .../dynamic_graph/shard_expansion.cc | 463 +++++----- .../archetypes/jsonable_ordered_value_type.h | 91 ++ .../{filter_keys.h => bidict_filter_keys.h} | 6 +- ...filter_values.h => bidict_filter_values.h} | 6 +- ...filtrans_keys.h => bidict_filtrans_keys.h} | 6 +- ...rans_values.h => bidict_filtrans_values.h} | 6 +- ...red_set_of.h => bidict_unordered_set_of.h} | 6 +- .../utils/bidict/algorithms/transform_keys.h | 2 +- .../bidict/algorithms/transform_values.h | 2 +- .../unstructured_relation_from_bidict.h | 4 +- lib/utils/include/utils/bidict/bidict.h | 42 +- lib/utils/include/utils/containers/all_of.h | 8 +- .../containers/binary_merge_disjoint_maps.h | 4 +- .../utils/containers/binary_merge_maps_with.h | 10 +- lib/utils/include/utils/containers/filter.h | 7 + .../include/utils/containers/generate_map.h | 6 +- .../utils/containers/generate_unordered_map.h | 28 + .../utils/containers/get_all_assignments.h | 4 +- lib/utils/include/utils/containers/index.dox | 2 +- .../include/utils/containers/is_submapeq_of.h | 4 +- lib/utils/include/utils/containers/items.h | 4 +- lib/utils/include/utils/containers/keys.h | 10 +- .../include/utils/containers/lookup_in_map.h | 16 +- .../utils/containers/map_from_unordered.h | 18 + .../include/utils/containers/map_keys2.h | 2 +- .../utils/containers/map_keys_and_values.h | 3 +- .../include/utils/containers/map_values.h | 13 + .../include/utils/containers/require_all_of.h | 33 + .../utils/containers/require_only_key.h | 9 + .../include/utils/containers/require_same.h | 2 +- .../include/utils/containers/transform.h | 14 +- .../utils/containers/unordered_items.h | 17 + .../include/utils/containers/unordered_keys.h | 30 + .../utils/containers/unordered_map_from_map.h | 18 + .../utils/containers/zip_values_strict.h | 8 +- .../utils/containers/zip_values_strict_with.h | 8 +- lib/utils/include/utils/fmt/map.h | 4 +- .../unordered_set_kwarg_dataflow_graph.h | 4 +- ...ordered_set_labelled_open_dataflow_graph.h | 17 +- ...d_set_labelled_open_kwarg_dataflow_graph.h | 23 +- .../unordered_set_open_kwarg_dataflow_graph.h | 4 +- .../algorithms/get_incoming_slots_for_node.h | 4 +- .../algorithms/get_outgoing_slots_for_node.h | 4 +- .../algorithms/get_graph_data.h | 4 +- .../algorithms/permute_input_ids.h | 4 +- .../algorithms/permute_node_ids.h | 6 +- .../algorithms/rewrite_labels.h | 6 +- ..._labelled_open_kwarg_dataflow_graph_data.h | 6 +- .../labelled_open_kwarg_dataflow_graph_data.h | 4 +- ...lled_open_kwarg_dataflow_graph_input_ids.h | 6 +- ...elled_open_kwarg_dataflow_graph_node_ids.h | 6 +- ...abelled_open_kwarg_dataflow_graph_labels.h | 6 +- .../graph/series_parallel/get_ancestors.h | 1 + .../non_normal_parallel_split.dtg.toml | 9 +- .../non_normal_series_split.dtg.toml | 2 + .../non_normal_sp_decomposition.dtg.toml | 1 + .../series_parallel/parallel_split.dtg.toml | 9 +- .../series_parallel_decomposition.dtg.toml | 1 + .../series_parallel/series_split.dtg.toml | 2 + .../include/utils/json/check_is_jsonable.h | 4 +- .../include/utils/many_to_one/many_to_one.h | 29 +- .../many_to_one_from_unstructured_relation.h | 20 - .../unstructured_relation_from_many_to_one.h | 17 - .../include/utils/nonempty_set/nonempty_set.h | 136 +++ .../include/utils/one_to_many/one_to_many.h | 81 +- ...d_relation.h => one_to_many_filter_keys.h} | 15 +- .../one_to_many/one_to_many_filter_values.h | 22 + .../one_to_many_transform_values.h | 6 +- .../unstructured_relation_from_one_to_many.h | 21 - lib/utils/include/utils/orthotope/dim_coord.h | 8 +- .../include/utils/orthotope/dim_domain.h | 4 +- .../utils/orthotope/dim_projection.dtg.toml | 1 + .../utils/orthotope/down_projection.dtg.toml | 1 + .../include/utils/orthotope/down_projection.h | 2 +- .../utils/orthotope/eq_projection.dtg.toml | 1 + .../utils/orthotope/minimal_dim_domain.h | 8 +- .../utils/orthotope/up_projection.dtg.toml | 1 + .../include/utils/orthotope/up_projection.h | 4 +- .../archetypes/jsonable_ordered_value_type.cc | 7 + .../{filter_keys.cc => bidict_filter_keys.cc} | 4 +- ...lter_values.cc => bidict_filter_values.cc} | 4 +- ...ltrans_keys.cc => bidict_filtrans_keys.cc} | 4 +- ...ns_values.cc => bidict_filtrans_values.cc} | 4 +- .../algorithms/bidict_unordered_set_of.cc | 11 + .../bidict/algorithms/unordered_set_of.cc | 11 - lib/utils/src/utils/cli/cli_parse.cc | 4 +- lib/utils/src/utils/containers/filter.cc | 44 + .../src/utils/containers/generate_map.cc | 12 + .../containers/generate_unordered_map.cc | 13 + lib/utils/src/utils/containers/group_by.cc | 4 +- .../src/utils/containers/is_submapeq_of.cc | 11 + lib/utils/src/utils/containers/items.cc | 13 + lib/utils/src/utils/containers/keys.cc | 12 + .../utils/containers/map_from_unordered.cc | 13 + lib/utils/src/utils/containers/map_values.cc | 8 + .../src/utils/containers/require_all_of.cc | 34 + .../src/utils/containers/require_only_key.cc | 5 + lib/utils/src/utils/containers/transform.cc | 46 + .../src/utils/containers/unordered_items.cc | 15 + .../src/utils/containers/unordered_keys.cc | 13 + .../containers/unordered_map_from_map.cc | 12 + .../algorithms/dataflow_graph_as_dot.cc | 1 - .../digraph/algorithms/get_dominators_map.cc | 4 +- .../algorithms/get_imm_dominators_map.cc | 8 +- .../algorithms/get_imm_post_dominator.cc | 4 +- .../digraph/algorithms/get_incoming_edges.cc | 24 +- .../digraph/algorithms/get_outgoing_edges.cc | 25 +- .../graph/digraph/algorithms/is_acyclic.cc | 4 +- .../graph/instances/adjacency_multidigraph.cc | 18 +- .../algorithms/get_incoming_edges.cc | 7 +- .../get_multidiedge_to_diedge_map.cc | 4 +- .../algorithms/get_outgoing_edges.cc | 7 +- .../algorithms/get_incoming_edges.cc | 4 +- .../balanced_binary_sp_tree_from_nary.cc | 5 +- .../non_normal_sp_decomposition.cc | 10 +- .../normalize_sp_decomposition.cc | 5 +- .../series_parallel_decomposition.cc | 7 +- .../series_parallel_metrics.cc | 2 + .../sp_ization/escribano_algo.cc | 4 +- .../sp_ization/flexible_algo.cc | 6 +- .../sp_ization/naive_stratum_sync.cc | 4 +- .../series_parallel/sp_ization/node_role.cc | 4 +- .../many_to_one/exhaustive_relational_join.cc | 8 +- .../utils/many_to_one/invert_many_to_one.cc | 6 +- .../src/utils/many_to_one/many_to_one.cc | 12 +- .../many_to_one_from_unstructured_relation.cc | 12 - .../unstructured_relation_from_many_to_one.cc | 12 - .../src/utils/nonempty_set/nonempty_set.cc | 23 + .../one_to_many/exhaustive_relational_join.cc | 8 +- .../utils/one_to_many/invert_one_to_many.cc | 6 +- .../src/utils/one_to_many/one_to_many.cc | 20 +- .../one_to_many/one_to_many_filter_keys.cc | 12 + .../one_to_many/one_to_many_filter_values.cc | 13 + .../one_to_many/one_to_many_from_bidict.cc | 6 +- .../one_to_many_from_l_to_r_mapping.cc | 6 +- .../one_to_many_from_unstructured_relation.cc | 12 - .../one_to_many_transform_values.cc | 8 +- .../unstructured_relation_from_one_to_many.cc | 12 - .../src/utils/orthotope/dim_domain_mapping.cc | 14 +- .../src/utils/orthotope/dim_projection.cc | 11 +- .../src/utils/orthotope/down_projection.cc | 12 +- .../src/utils/orthotope/eq_projection.cc | 12 +- .../orthotope/minimal_dim_domain_mapping.cc | 14 +- .../src/utils/orthotope/up_projection.cc | 7 +- .../{filter_keys.cc => bidict_filter_keys.cc} | 7 +- ...lter_values.cc => bidict_filter_values.cc} | 6 +- ...ltrans_keys.cc => bidict_filtrans_keys.cc} | 6 +- ...ns_values.cc => bidict_filtrans_values.cc} | 8 +- ...d_set_of.cc => bidict_unordered_set_of.cc} | 2 +- lib/utils/test/src/utils/bidict/bidict.cc | 2 + .../test/src/utils/containers/enumerate.cc | 4 +- lib/utils/test/src/utils/containers/keys.cc | 8 +- .../algorithms/get_imm_post_dominators_map.cc | 1 - .../sp_ization/work_duplicating_sp_ization.cc | 4 +- .../test/src/utils/many_to_one/many_to_one.cc | 79 ++ .../many_to_one_from_unstructured_relation.cc | 54 -- .../unstructured_relation_from_many_to_one.cc | 25 - .../test/src/utils/one_to_many/one_to_many.cc | 107 ++- .../one_to_many_from_unstructured_relation.cc | 53 -- .../unstructured_relation_from_one_to_many.cc | 25 - 233 files changed, 3602 insertions(+), 1714 deletions(-) create mode 100644 lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_info.dtg.toml create mode 100644 lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.toml create mode 100644 lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h create mode 100644 lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_info.dtg.toml create mode 100644 lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.toml create mode 100644 lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_info.dtg.toml create mode 100644 lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc create mode 100644 lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc create mode 100644 lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h rename lib/utils/include/utils/bidict/algorithms/{filter_keys.h => bidict_filter_keys.h} (53%) rename lib/utils/include/utils/bidict/algorithms/{filter_values.h => bidict_filter_values.h} (53%) rename lib/utils/include/utils/bidict/algorithms/{filtrans_keys.h => bidict_filtrans_keys.h} (64%) rename lib/utils/include/utils/bidict/algorithms/{filtrans_values.h => bidict_filtrans_values.h} (63%) rename lib/utils/include/utils/bidict/algorithms/{unordered_set_of.h => bidict_unordered_set_of.h} (51%) create mode 100644 lib/utils/include/utils/containers/generate_unordered_map.h create mode 100644 lib/utils/include/utils/containers/map_from_unordered.h create mode 100644 lib/utils/include/utils/containers/require_all_of.h create mode 100644 lib/utils/include/utils/containers/unordered_items.h create mode 100644 lib/utils/include/utils/containers/unordered_keys.h create mode 100644 lib/utils/include/utils/containers/unordered_map_from_map.h delete mode 100644 lib/utils/include/utils/many_to_one/many_to_one_from_unstructured_relation.h delete mode 100644 lib/utils/include/utils/many_to_one/unstructured_relation_from_many_to_one.h create mode 100644 lib/utils/include/utils/nonempty_set/nonempty_set.h rename lib/utils/include/utils/one_to_many/{one_to_many_from_unstructured_relation.h => one_to_many_filter_keys.h} (50%) create mode 100644 lib/utils/include/utils/one_to_many/one_to_many_filter_values.h delete mode 100644 lib/utils/include/utils/one_to_many/unstructured_relation_from_one_to_many.h create mode 100644 lib/utils/src/utils/archetypes/jsonable_ordered_value_type.cc rename lib/utils/src/utils/bidict/algorithms/{filter_keys.cc => bidict_filter_keys.cc} (58%) rename lib/utils/src/utils/bidict/algorithms/{filter_values.cc => bidict_filter_values.cc} (57%) rename lib/utils/src/utils/bidict/algorithms/{filtrans_keys.cc => bidict_filtrans_keys.cc} (61%) rename lib/utils/src/utils/bidict/algorithms/{filtrans_values.cc => bidict_filtrans_values.cc} (61%) create mode 100644 lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc delete mode 100644 lib/utils/src/utils/bidict/algorithms/unordered_set_of.cc create mode 100644 lib/utils/src/utils/containers/generate_unordered_map.cc create mode 100644 lib/utils/src/utils/containers/map_from_unordered.cc create mode 100644 lib/utils/src/utils/containers/require_all_of.cc create mode 100644 lib/utils/src/utils/containers/unordered_items.cc create mode 100644 lib/utils/src/utils/containers/unordered_keys.cc create mode 100644 lib/utils/src/utils/containers/unordered_map_from_map.cc delete mode 100644 lib/utils/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc delete mode 100644 lib/utils/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc create mode 100644 lib/utils/src/utils/nonempty_set/nonempty_set.cc create mode 100644 lib/utils/src/utils/one_to_many/one_to_many_filter_keys.cc create mode 100644 lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc delete mode 100644 lib/utils/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc delete mode 100644 lib/utils/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc rename lib/utils/test/src/utils/bidict/algorithms/{filter_keys.cc => bidict_filter_keys.cc} (64%) rename lib/utils/test/src/utils/bidict/algorithms/{filter_values.cc => bidict_filter_values.cc} (61%) rename lib/utils/test/src/utils/bidict/algorithms/{filtrans_keys.cc => bidict_filtrans_keys.cc} (75%) rename lib/utils/test/src/utils/bidict/algorithms/{filtrans_values.cc => bidict_filtrans_values.cc} (69%) rename lib/utils/test/src/utils/bidict/algorithms/{unordered_set_of.cc => bidict_unordered_set_of.cc} (77%) delete mode 100644 lib/utils/test/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc delete mode 100644 lib/utils/test/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc delete mode 100644 lib/utils/test/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc delete mode 100644 lib/utils/test/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc diff --git a/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml index 2c5b5f56fc..31b87feb31 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml @@ -7,19 +7,19 @@ features = [ includes = [ "utils/graph/digraph/digraph_view.h", - "utils/bidict/bidict.h", "compiler/task_graph_simulator/pcg_task.dtg.h", "pcg/device_id_t.dtg.h", + "utils/many_to_one/many_to_one.h", "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", - "", - "" + "", + "" ] src_includes = [ - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h" + "utils/fmt/set.h", + "utils/hash/set.h", + "utils/fmt/map.h", + "utils/hash/map.h" ] [[fields]] @@ -28,8 +28,8 @@ type = "::FlexFlow::DiGraphView" [[fields]] name = "node_to_task" -type = "::FlexFlow::bidict<::FlexFlow::Node, ::FlexFlow::PCGTask>" +type = "::FlexFlow::ManyToOne<::FlexFlow::Node, ::FlexFlow::PCGTask>" [[fields]] name = "node_to_devices" -type = "std::unordered_map<::FlexFlow::Node, std::unordered_set<::FlexFlow::device_id_t>>" +type = "std::map<::FlexFlow::Node, std::set<::FlexFlow::device_id_t>>" diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc index 1aeb83d202..5f9300973f 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc @@ -15,7 +15,7 @@ std::unordered_set abstracted_single_tensor_movement_get_dst_layers( AbstractedSingleTensorMovement const &m) { return transform( - keys(m.edge_to_size), + unordered_keys(m.edge_to_size), [](AbstractedSingleTensorCommunicationEdge const &e) -> BinaryTreePath { return e.dst.operator_tree_path; }); diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc index 151008f65f..6ff261facd 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc @@ -10,10 +10,9 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" -#include "utils/bidict/algorithms/unordered_set_of.h" +#include "utils/bidict/algorithms/bidict_unordered_set_of.h" #include "utils/containers/binary_cartesian_product.h" #include "utils/containers/flatmap.h" -#include "utils/containers/generate_map.h" #include "utils/containers/get_only.h" #include "utils/containers/group_by.h" #include "utils/containers/map_from_pairs.h" @@ -46,7 +45,7 @@ AbstractedSingleTensorMovement get_abstracted_single_tensor_movement_along_edge( std::unordered_map single_comms = map_from_pairs(transform( - unordered_set_of(coord_mapping), + bidict_unordered_set_of(coord_mapping), [&](std::pair const & src_dst) -> std::pair { @@ -101,9 +100,9 @@ AbstractedTensorSetMovement get_abstracted_tensor_set_movement_across_split( }; return AbstractedTensorSetMovement{ - transform(edges_by_tensor.right_groups(), - [&](nonempty_unordered_set const - &edges) { + transform(unordered_set_of(edges_by_tensor.right_groups()), + [&](nonempty_set const &edges) + { return merge_abstracted_single_tensor_movements(transform( unordered_multiset_of(edges.unwrap_as_unordered_set()), to_abstracted_single_tensor_movement)); diff --git a/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc index 7ccab2fac9..4e38750de3 100644 --- a/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc @@ -59,7 +59,7 @@ SearchResult apply_substitution_and_update_machine_mapping( select_random(substituted_machine_views)); } - ASSERT(is_subseteq_of(keys(post_node_data), keys(machine_views))); + ASSERT(is_subseteq_of(unordered_keys(post_node_data), unordered_keys(machine_views))); std::unordered_map post_node_machine_views = diff --git a/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc index 77e50740aa..48f3bd9eed 100644 --- a/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc @@ -23,7 +23,7 @@ #include "utils/containers/contains.h" #include "utils/containers/contains_key.h" #include "utils/containers/flatmap.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/keys.h" #include "utils/containers/set_minus.h" @@ -103,7 +103,7 @@ MachineMappingResult set_minus(boundary_layers, get_constrained_layers(sub_constraints)); std::unordered_map> - allowed = generate_map( + allowed = generate_unordered_map( unconstrained_boundary_layers, [&](BinaryTreePath const &l) -> std::unordered_set { UnmappedRuntimeOnlyOpCostEstimateKey leaf = diff --git a/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc b/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc index 7d1f28337c..f7dbdd1d05 100644 --- a/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc @@ -7,7 +7,6 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" -#include "utils/containers/generate_map.h" #include "utils/containers/keys.h" #include "utils/containers/map_values.h" #include "utils/containers/sum.h" diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc index a2307716ba..861912efef 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc @@ -8,7 +8,7 @@ #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/containers/are_disjoint.h" #include "utils/containers/binary_merge_disjoint_maps.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" namespace FlexFlow { @@ -20,7 +20,7 @@ MappedParallelComputationGraph get_parallel_layers(pcg); std::unordered_set mapped_layers = - keys(mapping.machine_views); + unordered_keys(mapping.machine_views); ASSERT(mapped_layers == pcg_layers); @@ -40,7 +40,7 @@ MappedParallelComputationGraph }; std::unordered_map - mapped_op_task_groups = generate_map(mapped_layers, mapping_for_layer); + mapped_op_task_groups = generate_unordered_map(mapped_layers, mapping_for_layer); return mapped_pcg_from_pcg_and_mapped_op_task_groups(pcg, mapped_op_task_groups); @@ -54,7 +54,7 @@ MachineMapping combine_disjoint_mappings(MachineMapping const &m1, } bool nodes_are_disjoint(MachineMapping const &m1, MachineMapping const &m2) { - return are_disjoint(keys(m1.machine_views), keys(m2.machine_views)); + return are_disjoint(unordered_keys(m1.machine_views), unordered_keys(m2.machine_views)); } std::optional get_machine_mapping_from_machine_mapping_result( diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc index 8278d5511c..fe92c77def 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc @@ -3,7 +3,7 @@ #include "utils/containers/filter_values.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/flatmap.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/keys.h" #include "utils/containers/map_values.h" #include "utils/containers/restrict_keys.h" @@ -14,7 +14,7 @@ namespace FlexFlow { MachineMappingConstraints get_unconstrained_solution_for_layers( std::unordered_set const &layers) { return MachineMappingConstraints{ - generate_map(layers, + generate_unordered_map(layers, [](BinaryTreePath const &) -> std::optional { return std::nullopt; }), @@ -24,7 +24,7 @@ MachineMappingConstraints get_unconstrained_solution_for_layers( std::unordered_set get_unconstrained_layers(MachineMappingConstraints const &constraints) { - return keys(filter_values( + return unordered_keys(filter_values( constraints.machine_views, [](std::optional const &mv) { return !mv.has_value(); })); } @@ -32,14 +32,14 @@ std::unordered_set std::unordered_set get_constrained_layers(MachineMappingConstraints const &constraints) { - return keys(filter_values( + return unordered_keys(filter_values( constraints.machine_views, [](std::optional const &mv) { return mv.has_value(); })); } std::unordered_set get_all_layers(MachineMappingConstraints const &partial_solution) { - return keys(partial_solution.machine_views); + return unordered_keys(partial_solution.machine_views); } std::optional get_machine_view_for_layer( @@ -103,7 +103,7 @@ MachineMappingConstraints with_additional_constraints( std::optional require_only_root(MachineMappingConstraints const &constraints) { - ASSERT(keys(constraints.machine_views) == + ASSERT(unordered_keys(constraints.machine_views) == std::unordered_set{binary_tree_root_path()}, fmt::format("require_only_root expected constraints to have only a " "single key (the root path), but received {}", diff --git a/lib/compiler/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/src/compiler/machine_mapping/machine_view.cc index 090dec5845..7d00707b0d 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_view.cc @@ -221,8 +221,8 @@ static OperatorAtomicTaskShardBinding mappings = get_operator_to_ptensor_mappings(op_attrs, inputs_dim_degrees); std::unordered_map - ptensor_coords = generate_map( - keys(inputs_dim_degrees), + ptensor_coords = generate_unordered_map( + unordered_keys(inputs_dim_degrees), [&](TensorSlotName const &slot_name) -> ParallelTensorSpaceCoordinate { num_ptensor_shard_dims_t num_shard_dims = diff --git a/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc b/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc index a3f2009a60..dad91dd317 100644 --- a/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc +++ b/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc @@ -17,7 +17,7 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "utils/containers/contains.h" #include "utils/containers/flatmap.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/unordered_set_of.h" #include "utils/exception.h" @@ -84,7 +84,7 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( std::unordered_set const &boundary_layers) -> std::unordered_set { std::unordered_map> - allowed = generate_map( + allowed = generate_unordered_map( boundary_layers, [&](BinaryTreePath const &l) -> std::unordered_set { UnmappedRuntimeOnlyOpCostEstimateKey leaf = diff --git a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc index 6e2096afcc..ac39021f6f 100644 --- a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc @@ -45,7 +45,7 @@ std::unordered_map PCGBinarySPDecomposition const &decomposition, ParallelLayerGuidObliviousMachineMapping const &mapping) { std::unordered_set leaf_paths = require_same( - pcg_sp_tree_get_all_leaf_paths(decomposition), keys(mapping.raw_mapping)); + pcg_sp_tree_get_all_leaf_paths(decomposition), unordered_keys(mapping.raw_mapping)); std::unordered_map path_to_op_task_space_map = @@ -54,7 +54,7 @@ std::unordered_map return get_operator_task_space(pcg, l); }); - return generate_map( + return generate_unordered_map( leaf_paths, [&](BinaryTreePath const &p) -> MachineSpaceStencil { return MachineSpaceStencil{ /*operator_task_space=*/path_to_op_task_space_map.at(p), @@ -71,12 +71,12 @@ std::unordered_map> std::unordered_map tree_leaf_map = mm_problem_tree_get_path_to_leaf_map(tree); - std::unordered_set mapping_paths = keys(mapping.raw_mapping); - std::unordered_set tree_paths = keys(tree_leaf_map); + std::unordered_set mapping_paths = unordered_keys(mapping.raw_mapping); + std::unordered_set tree_paths = unordered_keys(tree_leaf_map); ASSERT(is_subseteq_of(mapping_paths, tree_paths)); - return generate_map( + return generate_unordered_map( tree_paths, [&](BinaryTreePath const &p) -> std::optional { if (!contains_key(mapping.raw_mapping, p)) { diff --git a/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc b/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc index cd8e634f2c..4d1c88d9eb 100644 --- a/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc +++ b/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc @@ -152,7 +152,7 @@ SPDecompositionTreeNodeType std::unordered_set pcg_sp_tree_get_all_leaf_paths(PCGBinarySPDecomposition const &tree) { - return keys(pcg_sp_tree_get_path_to_leaf_map(tree)); + return unordered_keys(pcg_sp_tree_get_path_to_leaf_map(tree)); } std::unordered_set diff --git a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc index d4d5a78d6a..b016b106e9 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc @@ -15,6 +15,7 @@ #include "utils/graph/instances/adjacency_digraph.h" #include #include +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -23,21 +24,21 @@ PCGTaskGraph MachineMapping const &machine_mapping, MachineComputeSpecification const &machine_spec) { DiGraph digraph = DiGraph::create(); - bidict node_to_task; + ManyToOne node_to_task; bidict node_to_layer; - std::unordered_map> node_to_devices; + std::map> node_to_devices; for (parallel_layer_guid_t const &layer : get_parallel_layers(pcg)) { MachineView mv = machine_mapping.machine_views.at(layer); RuntimeOnlyOpCostEstimateKey op_key = get_mapped_runtime_only_op_cost_estimate_key_for_layer(pcg, layer, mv); Node node = digraph.add_node(); - node_to_task.equate(node, PCGTask{op_key}); - node_to_layer.equate(node, layer); + node_to_task.insert({node, PCGTask{op_key}}); + node_to_layer.equate_strict(node, layer); node_to_devices[node] = - get_device_ids(get_operator_task_space(pcg, layer), - machine_mapping.machine_views.at(layer), - machine_spec); + set_of(get_device_ids(get_operator_task_space(pcg, layer), + machine_mapping.machine_views.at(layer), + machine_spec)); } for (ParallelComputationGraphEdge const &edge : get_edges(pcg)) { @@ -46,7 +47,7 @@ PCGTaskGraph TensorSetMovement movement = get_tensor_set_movement_from_pcg_edge(edge, pcg, src_mv, dst_mv); Node node = digraph.add_node(); - node_to_task.equate(node, PCGTask{movement}); + node_to_task.insert({node, PCGTask{movement}}); node_to_devices[node] = {}; Node src_node = node_to_layer.at_r(get_src_layer(edge)); Node dst_node = node_to_layer.at_r(get_dst_layer(edge)); diff --git a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc index bc528493a8..3fabfc3966 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc @@ -62,7 +62,7 @@ milliseconds_t task_simulator_estimate_forward_pass_time( std::unordered_set devices_occupied = set_union(transform(in_progress_tasks, get_devices)); - std::unordered_set required_devices = get_devices(task); + std::unordered_set required_devices = unordered_set_of(get_devices(task)); return intersection(devices_occupied, required_devices).empty(); }; diff --git a/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc b/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc index be8c7c4f98..05a049e98f 100644 --- a/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc +++ b/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc @@ -20,7 +20,7 @@ #include "substitutions/sub_parallel_computation_graph.h" #include "substitutions/substitution.h" #include "substitutions/unity_substitution_set.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/deduplicated_priority_queue.h" #include "utils/graph/node/algorithms.h" #include "utils/optional.h" diff --git a/lib/local-execution/src/local-execution/tensor_allocation.cc b/lib/local-execution/src/local-execution/tensor_allocation.cc index bb2a1ba2a4..203345a9af 100644 --- a/lib/local-execution/src/local-execution/tensor_allocation.cc +++ b/lib/local-execution/src/local-execution/tensor_allocation.cc @@ -55,7 +55,7 @@ DynamicOpenDataflowGraph perform_tensor_allocation( Allocator &allocator) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); - for (DynamicValueAttrs const &v : keys(preallocated)) { + for (DynamicValueAttrs const &v : unordered_keys(preallocated)) { ASSERT(v.accessor == std::nullopt); } diff --git a/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc b/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc index eec9ae869c..89936d9b00 100644 --- a/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc +++ b/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc @@ -46,7 +46,7 @@ std::unordered_map }; }, [&](ConcatAttrs const &) { - return generate_map(get_variadic_inputs_slot_name_sequence(), + return generate_unordered_map(get_variadic_inputs_slot_name_sequence(), [](TensorSlotName) -> IncomingTensorRole { return IncomingTensorRole::INPUT; }); diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc index 51d7968033..83a7aded6a 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc @@ -8,7 +8,7 @@ #include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/filtrans.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/map_keys.h" #include "utils/containers/map_values.h" @@ -96,7 +96,7 @@ std::unordered_map }; std::unordered_map shard_dim_degrees = - generate_map(get_idxs(degrees.shard_degrees), [&](ff_dim_t const &dim) { + generate_unordered_map(get_idxs(degrees.shard_degrees), [&](ff_dim_t const &dim) { return degrees.shard_degrees.at(dim); }); @@ -131,7 +131,7 @@ DimDomain ParallelTensorDimDegrees const &dim_degrees) { return DimDomain{ - generate_map(get_parallel_tensor_dim_indices(dim_degrees), + generate_unordered_map(get_parallel_tensor_dim_indices(dim_degrees), [&](parallel_tensor_dim_idx_t idx) { return get_degree_for_parallel_tensor_dim_idx(dim_degrees, idx); diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc index 0c6e157697..79d765b02c 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc @@ -3,7 +3,7 @@ #include "op-attrs/parallel_tensor_dim_idx_t.h" #include "utils/containers/contains_key.h" #include "utils/containers/filtermap_keys.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/unordered_set_of.h" #include "utils/nonnegative_int/num_elements.h" @@ -73,7 +73,7 @@ DimCoord dim_coord_from_parallel_tensor_space_coord( ParallelTensorSpaceCoordinate const &coord) { return DimCoord{ - generate_map(get_dim_idxs_in_ptensor_space_coord(coord), + generate_unordered_map(get_dim_idxs_in_ptensor_space_coord(coord), [&](parallel_tensor_dim_idx_t idx) { return ptensor_coord_component_for_ptensor_dim_idx(coord, idx); diff --git a/lib/op-attrs/src/op-attrs/shape_inference.cc b/lib/op-attrs/src/op-attrs/shape_inference.cc index a3f8066dee..38e116bfd8 100644 --- a/lib/op-attrs/src/op-attrs/shape_inference.cc +++ b/lib/op-attrs/src/op-attrs/shape_inference.cc @@ -52,7 +52,7 @@ static std::vector std::vector expected_slots = slice(slots, 0, v_num_slots.unwrap_nonnegative()); - ASSERT(unordered_set_of(expected_slots) == keys(v)); + ASSERT(unordered_set_of(expected_slots) == unordered_keys(v)); return transform(expected_slots, [&](TensorSlotName const &slot_name) { return v.at(slot_name); diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h index aded1eb657..41aca802e7 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h @@ -7,6 +7,7 @@ #include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.h" #include "utils/bidict/bidict.h" #include +#include "utils/one_to_many/one_to_many.h" namespace FlexFlow { @@ -38,10 +39,13 @@ struct MappedOperatorTaskGroup { friend struct ::std::hash; }; -bidict +OneToMany get_tensor_bindings_for_slot_name(MappedOperatorTaskGroup const &, TensorSlotName const &); +std::set get_slot_names_for_task_group(MappedOperatorTaskGroup const &); + + nlohmann::json mapped_operator_task_group_as_dot_json(MappedOperatorTaskGroup const &); diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h index a2afdb7914..2e789e14c9 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h @@ -3,12 +3,16 @@ #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.h" namespace FlexFlow { std::unordered_set mpcg_get_parallel_layers(MappedParallelComputationGraph const &); +std::set + mpcg_get_invocation_set(MappedParallelComputationGraph const &); + MappedOperatorTaskGroup mpcg_get_mapping_for_layer(MappedParallelComputationGraph const &, parallel_layer_guid_t); diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_info.dtg.toml b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_info.dtg.toml new file mode 100644 index 0000000000..056cbecdde --- /dev/null +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_info.dtg.toml @@ -0,0 +1,31 @@ +namespace = "FlexFlow" +name = "MappedParallelLayerInfo" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", + "pcg/parallel_computation_graph/parallel_layer_attrs.dtg.h", + "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h", +] + +src_includes = [ +] + +[[fields]] +name = "guid" +type = "::FlexFlow::parallel_layer_guid_t" + +[[fields]] +name = "attrs" +type = "::FlexFlow::ParallelLayerAttrs" + +[[fields]] +name = "mapping" +type = "::FlexFlow::MappedOperatorTaskGroup" diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.toml b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.toml new file mode 100644 index 0000000000..b28884d551 --- /dev/null +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.toml @@ -0,0 +1,33 @@ +namespace = "FlexFlow" +name = "MappedParallelLayerInvocationInfo" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_info.dtg.h", + "pcg/parallel_computation_graph/parallel_tensor_info.dtg.h", + "op-attrs/tensor_slot_name.dtg.h", +] + +src_includes = [ + "utils/fmt/map.h", + "utils/hash/map.h", +] + +[[fields]] +name = "incoming" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorInfo>" + +[[fields]] +name = "layer_info" +type = "::FlexFlow::MappedParallelLayerInfo" + +[[fields]] +name = "outgoing" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorInfo>" diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h new file mode 100644 index 0000000000..dcda3a977b --- /dev/null +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h @@ -0,0 +1,17 @@ +#ifndef _FLEXFLOW_LIB_PCG_INCLUDE_PCG_MAPPED_PARALLEL_COMPUTATION_GRAPH_MAPPED_PARALLEL_LAYER_INVOCATION_INFO_H +#define _FLEXFLOW_LIB_PCG_INCLUDE_PCG_MAPPED_PARALLEL_COMPUTATION_GRAPH_MAPPED_PARALLEL_LAYER_INVOCATION_INFO_H + +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.h" +#include "pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.h" +#include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" + +namespace FlexFlow { + +MappedParallelLayerInvocationInfo + mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping( + ParallelLayerInvocationInfo const &, + MappedOperatorTaskGroup const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h index 9764e40627..7c5a825420 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h @@ -12,6 +12,7 @@ #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h" #include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" #include +#include "pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.h" namespace FlexFlow { @@ -38,6 +39,13 @@ ParallelLayerAddedResult OperatorTaskSpace get_operator_task_space(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &layer); +std::set + pcg_get_invocation_info_set(ParallelComputationGraph const &); + +ParallelLayerInvocationInfo + pcg_get_invocation_info_for_layer(ParallelComputationGraph const &, + parallel_layer_guid_t); + std::unordered_set get_pcg_edges_from_layer_to_layer(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &src, diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_info.dtg.toml b/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_info.dtg.toml new file mode 100644 index 0000000000..67107cdbb1 --- /dev/null +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_info.dtg.toml @@ -0,0 +1,26 @@ +namespace = "FlexFlow" +name = "ParallelLayerInfo" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", + "pcg/parallel_computation_graph/parallel_layer_attrs.dtg.h", +] + +src_includes = [ +] + +[[fields]] +name = "guid" +type = "::FlexFlow::parallel_layer_guid_t" + +[[fields]] +name = "attrs" +type = "::FlexFlow::ParallelLayerAttrs" diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.toml b/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.toml new file mode 100644 index 0000000000..6555752bcf --- /dev/null +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.toml @@ -0,0 +1,33 @@ +namespace = "FlexFlow" +name = "ParallelLayerInvocationInfo" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/parallel_computation_graph/parallel_layer_info.dtg.h", + "pcg/parallel_computation_graph/parallel_tensor_info.dtg.h", + "op-attrs/tensor_slot_name.dtg.h", +] + +src_includes = [ + "utils/fmt/map.h", + "utils/hash/map.h", +] + +[[fields]] +name = "incoming" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorInfo>" + +[[fields]] +name = "layer_info" +type = "::FlexFlow::ParallelLayerInfo" + +[[fields]] +name = "outgoing" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorInfo>" diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_info.dtg.toml b/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_info.dtg.toml new file mode 100644 index 0000000000..09f53d6954 --- /dev/null +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_tensor_info.dtg.toml @@ -0,0 +1,26 @@ +namespace = "FlexFlow" +name = "ParallelTensorInfo" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h", + "pcg/parallel_computation_graph/parallel_tensor_attrs.dtg.h", +] + +src_includes = [ +] + +[[fields]] +name = "guid" +type = "::FlexFlow::parallel_tensor_guid_t" + +[[fields]] +name = "attrs" +type = "::FlexFlow::ParallelTensorAttrs" diff --git a/lib/pcg/src/pcg/computation_graph.cc b/lib/pcg/src/pcg/computation_graph.cc index 56bfb98856..35ba0747f0 100644 --- a/lib/pcg/src/pcg/computation_graph.cc +++ b/lib/pcg/src/pcg/computation_graph.cc @@ -196,7 +196,7 @@ static std::unordered_map ASSERT(incoming_tensors.size() == incoming_slot_roles.size()); std::unordered_set slots_with_desired_role = - keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { + unordered_keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { return role == desired_role; })); diff --git a/lib/pcg/src/pcg/computation_graph_builder.cc b/lib/pcg/src/pcg/computation_graph_builder.cc index b687aa11b6..40e72aee9d 100644 --- a/lib/pcg/src/pcg/computation_graph_builder.cc +++ b/lib/pcg/src/pcg/computation_graph_builder.cc @@ -115,10 +115,10 @@ static void check_incoming_tensor_roles( set_union(input_slots, weight_slots)); std::unordered_map current = binary_merge_disjoint_maps( - generate_map( + generate_unordered_map( input_slots, [](TensorSlotName) { return IncomingTensorRole::INPUT; }), - generate_map(weight_slots, [](TensorSlotName) { + generate_unordered_map(weight_slots, [](TensorSlotName) { return IncomingTensorRole::WEIGHT; })); @@ -134,8 +134,8 @@ std::unordered_map &weight_initializers, std::optional> const &outputs) { - ASSERT(are_disjoint(keys(inputs), keys(weight_initializers))); - check_incoming_tensor_roles(layer, keys(inputs), keys(weight_initializers)); + ASSERT(are_disjoint(unordered_keys(inputs), unordered_keys(weight_initializers))); + check_incoming_tensor_roles(layer, unordered_keys(inputs), unordered_keys(weight_initializers)); std::unordered_map input_shapes = map_values( inputs, [&](tensor_guid_t const &t) { return this->get_shape(t); }); diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc index d0fd3300f5..3bb508681d 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc @@ -12,6 +12,14 @@ #include "utils/containers/vector_of.h" #include "utils/hash/tuple.h" #include "utils/nonnegative_int/num_elements.h" +#include "utils/containers/require_all_same1.h" +#include "utils/containers/set_of.h" +#include "utils/containers/keys.h" +#include "utils/containers/contains.h" +#include "utils/bidict/algorithms/right_entries.h" +#include "utils/containers/map_values.h" +#include "utils/containers/unordered_set_of.h" +#include "utils/many_to_one/invert_many_to_one.h" namespace FlexFlow { @@ -23,7 +31,7 @@ MappedOperatorTaskGroup::MappedOperatorTaskGroup( transform(vector_of(shard_bindings.right_values()), [&](OperatorAtomicTaskShardBinding const &s) -> std::unordered_set { - return keys(s.tensor_coords); + return unordered_keys(s.tensor_coords); }); std::unordered_set slot_names = @@ -39,8 +47,6 @@ MappedOperatorTaskGroup::MappedOperatorTaskGroup( return ptensor_space_coord_for_slot_name(signature, slot_name); }); - ASSERT(are_all_distinct(coords_for_key)); - std::vector coord_dims_for_key = transform(coords_for_key, [](ParallelTensorSpaceCoordinate const &c) { return ptensor_coord_num_dims(c); @@ -92,15 +98,27 @@ bidict const & return this->shard_bindings; } -bidict +OneToMany get_tensor_bindings_for_slot_name(MappedOperatorTaskGroup const &task_group, TensorSlotName const &slot_name) { - return transform_values(task_group.get_shard_bindings(), - [&](OperatorAtomicTaskShardBinding const &b) { - return ptensor_space_coord_for_slot_name(b, - slot_name); - }) - .reversed(); + std::set slot_names = get_slot_names_for_task_group(task_group); + ASSERT(contains(slot_names, slot_name)); + + std::unordered_map m = + map_values(task_group.get_shard_bindings().as_unordered_map(), + [&](OperatorAtomicTaskShardBinding const &b) -> ParallelTensorSpaceCoordinate { + return ptensor_space_coord_for_slot_name(b, slot_name); + }); + + return invert_many_to_one(many_to_one_from_unstructured_relation(unordered_set_of(m))); +} + +std::set get_slot_names_for_task_group(MappedOperatorTaskGroup const &g) { + return require_all_same1( + transform(vector_of(right_entries(g.get_shard_bindings())), + [&](OperatorAtomicTaskShardBinding const &shard_bindings) -> std::set { + return keys(shard_bindings.tensor_coords); + })); } nlohmann::json diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index fc1dff504b..1bece17c9a 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -10,6 +10,12 @@ #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/materialize_labelled_kwarg_dataflow_graph_view.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/rewrite_labelled_kwarg_dataflow_graph_node_labels.h" #include "utils/many_to_one/many_to_one_from_map.h" +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h" +#include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" +#include "utils/containers/set_of.h" +#include "utils/containers/set_union.h" +#include "utils/containers/keys.h" +#include "utils/containers/require_all_of.h" namespace FlexFlow { @@ -18,6 +24,22 @@ std::unordered_set return get_parallel_layers(pcg_from_mpcg(mpcg)); } +std::set + mpcg_get_invocation_set(MappedParallelComputationGraph const &mpcg) +{ + auto mk_mapped_invocation = [&](ParallelLayerInvocationInfo const &invocation) + -> MappedParallelLayerInvocationInfo + { + MappedOperatorTaskGroup mapping = mpcg_get_mapping_for_layer(mpcg, invocation.layer_info.guid); + + return mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping(invocation, mapping); + }; + + ParallelComputationGraph pcg = pcg_from_mpcg(mpcg); + + return transform(pcg_get_invocation_info_set(pcg), mk_mapped_invocation); +} + MappedOperatorTaskGroup mpcg_get_mapping_for_layer(MappedParallelComputationGraph const &mpcg, parallel_layer_guid_t l) { @@ -112,6 +134,22 @@ MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( return mapped_op_task_groups.at(l); }; + auto slot_names_for_layer = [&](parallel_layer_guid_t l) -> std::set { + return set_union(keys(get_incoming_tensors(pcg, l)), keys(get_outgoing_tensors(pcg, l))); + }; + + auto slot_names_for_layer_mapping = [&](parallel_layer_guid_t l) -> std::set { + return get_slot_names_for_task_group(mapping_for_layer(l)); + }; + + require_all_of( + get_parallel_layers(pcg), + [&](parallel_layer_guid_t l) -> void { + std::set for_layer = slot_names_for_layer(l); + std::set for_layer_mapping = slot_names_for_layer_mapping(l); + ASSERT(for_layer == for_layer_mapping); + }); + auto mpcg_layer_attrs_from_pcg_layer_attrs = [&](Node const &node, ParallelLayerAttrs const &pcg_layer_attrs) -> MappedParallelLayerAttrs { diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc new file mode 100644 index 0000000000..bda6afb60c --- /dev/null +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc @@ -0,0 +1,22 @@ +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h" + +namespace FlexFlow { + +MappedParallelLayerInvocationInfo + mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping( + ParallelLayerInvocationInfo const &invocation_info, + MappedOperatorTaskGroup const &mapping) +{ + return MappedParallelLayerInvocationInfo{ + /*incoming=*/invocation_info.incoming, + /*layer_info=*/MappedParallelLayerInfo{ + /*guid=*/invocation_info.layer_info.guid, + /*attrs=*/invocation_info.layer_info.attrs, + /*mapping=*/mapping, + }, + /*outgoing=*/invocation_info.outgoing, + }; +} + + +} // namespace FlexFlow diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index 5098cadafe..c4a429a820 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -37,6 +37,7 @@ #include "utils/graph/node/node.dtg.h" #include "utils/record_formatter.h" #include +#include "utils/containers/map_from_unordered.h" namespace FlexFlow { @@ -98,7 +99,7 @@ ParallelLayerAddedResult add_parallel_layer( std::unordered_map output_flags = maybe_output_flags.value_or( - generate_map(keys(output_shapes), + generate_unordered_map(unordered_keys(output_shapes), [](TensorSlotName const &) { return CreateGrad::YES; })); std::unordered_map output_attrs = @@ -164,6 +165,46 @@ OperatorTaskSpace get_operator_task_space(ParallelComputationGraph const &pcg, compgraph_op_attrs_from_pcg_op_attrs(op_attrs).value(), input_degrees); } +std::set + pcg_get_invocation_info_set(ParallelComputationGraph const &pcg) +{ + return transform(set_of(get_parallel_layers(pcg)), + [&](parallel_layer_guid_t l) -> ParallelLayerInvocationInfo { + return pcg_get_invocation_info_for_layer(pcg, l); + }); +} + +ParallelLayerInvocationInfo + pcg_get_invocation_info_for_layer(ParallelComputationGraph const &pcg, + parallel_layer_guid_t l) +{ + ParallelLayerAttrs l_attrs = get_parallel_layer_attrs(pcg, l); + + std::map incoming = + map_from_unordered(get_incoming_tensors(pcg, l)); + + std::map outgoing = + map_from_unordered(get_outgoing_tensors(pcg, l)); + + auto get_parallel_tensor_info = [&](parallel_tensor_guid_t t) -> ParallelTensorInfo { + ParallelTensorAttrs t_attrs = get_parallel_tensor_attrs(pcg, t); + + return ParallelTensorInfo{ + /*guid=*/t, + /*attrs=*/t_attrs, + }; + }; + + return ParallelLayerInvocationInfo{ + /*incoming=*/map_values(incoming, get_parallel_tensor_info), + /*layer_info=*/ParallelLayerInfo{ + /*guid=*/l, + /*attrs=*/l_attrs, + }, + /*outgoing=*/map_values(outgoing, get_parallel_tensor_info), + }; +} + std::unordered_set get_edges(ParallelComputationGraph const &pcg) { return transform(get_all_kwarg_dataflow_edges(pcg.raw_graph), @@ -188,9 +229,10 @@ std::unordered_set get_outgoing_edges(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { std::unordered_set> raw_edges = + unordered_set_of( get_outgoing_kwarg_dataflow_edges_for_node(pcg.raw_graph, l.raw_graph_node) - .right_values(); + .right_values()); return transform(raw_edges, [](KwargDataflowEdge const &e) { return ParallelComputationGraphEdge{e}; }); @@ -308,7 +350,7 @@ static std::unordered_map ASSERT(incoming_tensors.size() == incoming_slot_roles.size()); std::unordered_set slots_with_desired_role = - keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { + unordered_keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { return role == desired_role; })); diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc index 1d6713dcdb..92334cfde9 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc @@ -679,10 +679,10 @@ static void check_incoming_tensor_roles( get_incoming_tensor_roles(layer.op_attrs); std::unordered_map current = binary_merge_disjoint_maps( - generate_map( + generate_unordered_map( input_slots, [](TensorSlotName) { return IncomingTensorRole::INPUT; }), - generate_map(weight_slots, [](TensorSlotName) { + generate_unordered_map(weight_slots, [](TensorSlotName) { return IncomingTensorRole::WEIGHT; })); @@ -698,8 +698,8 @@ std::unordered_map std::unordered_map const &weight_initializers) { - ASSERT(are_disjoint(keys(inputs), keys(weight_initializers))); - check_incoming_tensor_roles(layer, keys(inputs), keys(weight_initializers)); + ASSERT(are_disjoint(unordered_keys(inputs), unordered_keys(weight_initializers))); + check_incoming_tensor_roles(layer, unordered_keys(inputs), unordered_keys(weight_initializers)); std::unordered_map input_shapes = map_values(inputs, [&](parallel_tensor_guid_t const &i) { diff --git a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 7856d89f27..a53c67c336 100644 --- a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -79,18 +79,24 @@ TEST_SUITE(FF_TEST_SUITE) { MappedOperatorTaskGroup partition_mapping = MappedOperatorTaskGroup{ bidict{ - {machine_coord(0_n), - OperatorAtomicTaskShardBinding{ - { - {TensorSlotName::OUTPUT, ptensor_coord(0_n)}, - }, - }}, - {machine_coord(1_n), - OperatorAtomicTaskShardBinding{ - { - {TensorSlotName::OUTPUT, ptensor_coord(1_n)}, - }, - }}, + { + machine_coord(0_n), + OperatorAtomicTaskShardBinding{ + { + {TensorSlotName::INPUT, ptensor_coord(0_n)}, + {TensorSlotName::OUTPUT, ptensor_coord(0_n)}, + }, + }, + }, + { + machine_coord(1_n), + OperatorAtomicTaskShardBinding{ + { + {TensorSlotName::INPUT, ptensor_coord(0_n)}, + {TensorSlotName::OUTPUT, ptensor_coord(1_n)}, + }, + }, + }, }, }; @@ -116,20 +122,26 @@ TEST_SUITE(FF_TEST_SUITE) { MappedOperatorTaskGroup{ bidict{ - {machine_coord(0_n), - OperatorAtomicTaskShardBinding{ - { - {TensorSlotName::LHS_INPUT, ptensor_coord(0_n)}, - {TensorSlotName::RHS_INPUT, ptensor_coord(0_n)}, - }, - }}, - {machine_coord(1_n), - OperatorAtomicTaskShardBinding{ - { - {TensorSlotName::LHS_INPUT, ptensor_coord(1_n)}, - {TensorSlotName::RHS_INPUT, ptensor_coord(1_n)}, - }, - }}, + { + machine_coord(0_n), + OperatorAtomicTaskShardBinding{ + { + {TensorSlotName::LHS_INPUT, ptensor_coord(0_n)}, + {TensorSlotName::RHS_INPUT, ptensor_coord(0_n)}, + {TensorSlotName::OUTPUT, ptensor_coord(0_n)}, + }, + }, + }, + { + machine_coord(1_n), + OperatorAtomicTaskShardBinding{ + { + {TensorSlotName::LHS_INPUT, ptensor_coord(1_n)}, + {TensorSlotName::RHS_INPUT, ptensor_coord(1_n)}, + {TensorSlotName::OUTPUT, ptensor_coord(1_n)}, + }, + }, + }, }, }}, }; diff --git a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc index fd314ebaea..44488e70ea 100644 --- a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc +++ b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc @@ -5,7 +5,7 @@ #include "pcg/parallel_computation_graph/parallel_layer_attrs.h" #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.h" #include "utils/containers/count.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_only.h" #include "utils/containers/items.h" #include "utils/containers/require_only_key.h" @@ -248,7 +248,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*paddingW=*/paddingW); std::unordered_map layers = - generate_map(get_parallel_layers(b.pcg), + generate_unordered_map(get_parallel_layers(b.pcg), [&](parallel_layer_guid_t const &l) { return get_parallel_layer_attrs(b.pcg, l); }); diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index 4ef2919b10..79eef476df 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -42,7 +42,7 @@ TensorInstanceBacking perform_instance_allocation( RealmContext &ctx) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); - for (DynamicValueAttrs const &v : keys(preallocated)) { + for (DynamicValueAttrs const &v : unordered_keys(preallocated)) { ASSERT(v.accessor == std::nullopt); } diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index aa67110127..4b068d70be 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -233,10 +233,10 @@ static Realm::Event spawn_dynamic_node_invocation( // chain reductions sequentially to avoid write races on dst Realm::Event result = precondition; - for (auto const &[p, m] : assert_unwrap(output_grad.mapping)) { + for (auto const &[p, m] : unstructured_relation_from_one_to_many(assert_unwrap(output_grad.mapping))) { DynamicValueAttrs replica_key = output_grad; replica_key.mapping = - bidict{{p, m}}; + OneToMany{{p, {m}}}; replica_key.shard_coord = p; Realm::RegionInstance src_inst = diff --git a/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc b/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc index 79a5176c4f..1d53824c29 100644 --- a/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc +++ b/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc @@ -8,25 +8,29 @@ namespace FlexFlow { SerializableTensorInstanceBacking tensor_instance_backing_to_serializable( TensorInstanceBacking const &backing) { - return SerializableTensorInstanceBacking{/*backing=*/map_keys_and_values( + return SerializableTensorInstanceBacking{ + /*backing=*/map_keys_and_values( backing.backing, dynamic_value_attrs_to_serializable, [](std::pair const &p) { return std::pair{realm_instance_to_serializable(p.first), realm_event_to_serializable(p.second)}; - })}; + }), + }; } TensorInstanceBacking tensor_instance_backing_from_serializable( SerializableTensorInstanceBacking const &backing) { - return TensorInstanceBacking{/*backing=*/map_keys_and_values( + return TensorInstanceBacking{ + /*backing=*/map_keys_and_values( backing.backing, dynamic_value_attrs_from_serializable, [](std::pair const &p) { return std::pair{realm_instance_from_serializable(p.first), realm_event_from_serializable(p.second)}; - })}; + }), + }; } } // namespace FlexFlow diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index 46d29e2bef..6efbb17eb3 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -197,12 +197,14 @@ MappedParallelComputationGraph { cpu0, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, {TensorSlotName::OUTPUT, tensor_coord0}, }}, }, { cpu1, OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, {TensorSlotName::OUTPUT, tensor_coord1}, }}, }, diff --git a/lib/runtime/src/parallel_tensor_uses.cc b/lib/runtime/src/parallel_tensor_uses.cc index 444d3a061a..970a9109dc 100644 --- a/lib/runtime/src/parallel_tensor_uses.cc +++ b/lib/runtime/src/parallel_tensor_uses.cc @@ -31,7 +31,7 @@ Op const *ParallelTensorUses::get_owner(ParallelTensor const &tensor) const { } void ParallelTensorUses::remove(Op const &op) { - for (auto const &k : keys(this->uses)) { + for (auto const &k : unordered_keys(this->uses)) { inplace_filter(this->uses.at(k), [&](ParallelTensorUseDescription const &d) { return d.op->op_guid == op.op_guid; diff --git a/lib/runtime/src/tensor_uses.cc b/lib/runtime/src/tensor_uses.cc index ce4672342d..db2cc942d4 100644 --- a/lib/runtime/src/tensor_uses.cc +++ b/lib/runtime/src/tensor_uses.cc @@ -20,7 +20,7 @@ std::vector TensorUses::at(size_t tensor_guid) const { } void TensorUses::remove(Layer const &layer) { - for (auto const &k : keys(this->uses)) { + for (auto const &k : unordered_keys(this->uses)) { inplace_filter(this->uses.at(k), [&](TensorUseDescription const &d) { return d.layer->layer_guid == layer.layer_guid; }); diff --git a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc index f2686f7cf7..f3ceda7a06 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc @@ -10,7 +10,7 @@ #include "substitutions/sub_parallel_computation_graph_data.h" #include "substitutions/sub_parallel_computation_graph_edge.h" #include "utils/containers/binary_merge_disjoint_maps.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" #include "utils/containers/values.h" @@ -50,7 +50,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( require_sub_parallel_computation_graph_data_is_valid(pre_data); std::unordered_set pre_nodes = - keys(pre_data.node_data); + unordered_keys(pre_data.node_data); std::unordered_set matched_nodes = unordered_set_of(values(match.node_assignment)); std::unordered_set post_nodes_from_original_graph = diff --git a/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc b/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc index 8e1c06b9b5..9ae007ef16 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc @@ -56,12 +56,12 @@ LabelledOpenKwargDataflowGraphView incoming_tensor_roles = get_incoming_tensor_roles(n_attrs.op_attrs); - ASSERT(is_subseteq_of(keys(incoming_shapes), keys(incoming_tensor_roles))); + ASSERT(is_subseteq_of(unordered_keys(incoming_shapes), unordered_keys(incoming_tensor_roles))); auto incoming_shapes_with_role = [&](IncomingTensorRole role) -> std::unordered_map { std::unordered_set slots_with_desired_role = - keys(filter_values(incoming_tensor_roles, + unordered_keys(filter_values(incoming_tensor_roles, [&](IncomingTensorRole r) { return r == role; })); return restrict_keys(incoming_shapes, slots_with_desired_role); diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index 85a0493e33..8a71fe2ad5 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -76,7 +76,7 @@ void assert_pcg_pattern_match_is_valid_for_pattern_and_subpcg( std::unordered_set pattern_inputs = get_inputs(pattern); std::unordered_set match_pattern_inputs = - keys(match.input_assignment); + unordered_keys(match.input_assignment); ASSERT(pattern_inputs == match_pattern_inputs); } diff --git a/lib/substitutions/src/substitutions/substitution_builder.cc b/lib/substitutions/src/substitutions/substitution_builder.cc index f2860326ab..ffda2291b5 100644 --- a/lib/substitutions/src/substitutions/substitution_builder.cc +++ b/lib/substitutions/src/substitutions/substitution_builder.cc @@ -94,7 +94,7 @@ std::unordered_map node_expr, map_values(inputs, raw_open_kwarg_dataflow_value_from_output_graph_expr_value), - generate_map(output_slots, + generate_unordered_map(output_slots, [](TensorSlotName) { return std::monostate{}; })); return map_values( diff --git a/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc b/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc index e9087b5718..a982277d22 100644 --- a/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc +++ b/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc @@ -16,9 +16,6 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.h" #include "utils/many_to_one/invert_many_to_one.h" #include "utils/many_to_one/many_to_one_from_map.h" -#include "utils/many_to_one/many_to_one_from_unstructured_relation.h" -#include "utils/many_to_one/unstructured_relation_from_many_to_one.h" -#include "utils/one_to_many/unstructured_relation_from_one_to_many.h" #include "utils/overload.h" namespace FlexFlow { @@ -46,7 +43,7 @@ static std::optional return OpenKwargDataflowValue{o}; }); - if (keys(pattern_outputs) != keys(graph_outputs)) { + if (unordered_keys(pattern_outputs) != unordered_keys(graph_outputs)) { return std::nullopt; } @@ -64,7 +61,7 @@ static std::optional graph_node_inputs = get_incoming_open_kwarg_dataflow_values_for_node(graph, graph_node); - if (keys(graph_node_inputs) != keys(pattern_node_inputs)) { + if (unordered_keys(graph_node_inputs) != unordered_keys(pattern_node_inputs)) { return std::nullopt; } diff --git a/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc b/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc index 703d651070..25d505b1fa 100644 --- a/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc +++ b/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc @@ -190,7 +190,7 @@ bool unlabelled_pattern_does_match( ASSERT(left_entries(match.node_assignment) == get_pattern_nodes(pattern)); ASSERT( is_subseteq_of(right_entries(match.node_assignment), get_nodes(graph))); - ASSERT(keys(match.input_assignment) == get_pattern_inputs(pattern)); + ASSERT(unordered_keys(match.input_assignment) == get_pattern_inputs(pattern)); ASSERT(is_subseteq_of(matched_by_pattern_inputs, get_all_open_kwarg_dataflow_values(graph))); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h index a1726c2ae1..7a383ee8eb 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h @@ -13,6 +13,11 @@ bool value_is_mapped(DynamicValueAttrs const &); bool no_part_of_graph_is_copy_inserted(DynamicOpenDataflowGraph const &); bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &); +std::unordered_set copies_for_invocation_inputs( + DynamicNodeInvocation const &i, + std::unordered_map const + &unmapped_value_to_mapped_source_value); + std::unordered_set perform_copy_insertion_for_invocation( DynamicNodeInvocation const &i, std::unordered_map const diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h new file mode 100644 index 0000000000..a87f2241c4 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h @@ -0,0 +1,20 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_INVOCATION_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_INVOCATION_H + +#include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" + +namespace FlexFlow { + +bool invocation_fully_satisfies(DynamicNodeInvocation const &, + std::function const &node_condition, + std::function const &value_condition, + std::function const &slot_condition); + +void require_invocation_fully_satisfies(DynamicNodeInvocation const &, + std::function const &require_node_condition, + std::function const &require_value_condition, + std::function const &require_slot_condition); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml new file mode 100644 index 0000000000..a59aba92d7 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml @@ -0,0 +1,30 @@ +namespace = "FlexFlow" +name = "DynamicNodeInvocationShardingInfo" +type = "struct" +#include "task-spec/dynamic_graph/shard_expansion.h" +features = [ + "eq", + "ord", + "hash", + "fmt", + "json", +] + +includes = [ + "pcg/machine_space_coordinate.dtg.h", + "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h", + "task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.h", +] + +src_includes = [ + "utils/hash/map.h", + "utils/fmt/map.h", +] + +[[fields]] +name = "device_coord" +type = "::FlexFlow::MachineSpaceCoordinate" + +[[fields]] +name = "value_sharding" +type = "std::map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::DynamicValueAttrsShardingInfo>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h index 4ca62db5b1..1aba00a675 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h @@ -23,6 +23,12 @@ bool no_part_of_dynamic_graph_satisfies( std::function const &, std::function const &); +void require_full_dynamic_graph_satisfies( + DynamicOpenDataflowGraph const &, + std::function const &, + std::function const &, + std::function const &); + std::unordered_multiset get_dynamic_nodes(DynamicOpenDataflowGraph const &); std::unordered_multiset diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml index 490a51f88d..add72764f1 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml @@ -13,7 +13,7 @@ includes = [ "op-attrs/parallel_tensor_shape.dtg.h", "op-attrs/parallel_tensor_space_coordinate.dtg.h", "pcg/machine_space_coordinate.dtg.h", - "utils/bidict/bidict.h", + "utils/one_to_many/one_to_many.h", "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_role.dtg.h", ] @@ -25,18 +25,38 @@ src_includes = [ [[fields]] name = "tensor_guid" type = "::FlexFlow::dynamic_tensor_guid_t" +docstring = ''' +\brief The \ref tensor_guid_t or \ref parallel_tensor_guid_t of the (usually parallel) tensor this value originates from. Also allows representing tensors for computing the loss that lie outside of the scope of the \ref ComputationGraph or \ref ParallelComputationGraph, e.g., the label tensor. + +For a \ref DynamicOpenDataflowGraph originating from a \ref MapepdParallelComputationGraph, this field is filled in by \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h. +''' [[fields]] name = "parallel_tensor_shape" type = "std::optional<::FlexFlow::ParallelTensorShape>" +docstring = ''' +\brief The \ref ParallelTensorShape of the parallel tensor this value originates from. + +For a \ref DynamicOpenDataflowGraph originating form a \ref MappedParallelComputationGraph, this field is filled in by \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h. +''' [[fields]] name = "shard_coord" type = "std::optional<::FlexFlow::ParallelTensorSpaceCoordinate>" +docstring = ''' +\brief The shard (i.e., \ref ParallelTensorSpaceCoordinate) of the (usually parallel) tensor represented by this value. + +For a \ref DynamicOpenDataflowGraph originating from a \ref MappedParallelComputationGraph, this field is filled in by \ref shard_expansion.h. +''' [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" +type = "std::optional<::FlexFlow::OneToMany<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" +docstring = ''' +\brief The location (i.e., \ref MachineSpaceCoordinate) of each shard (i.e., \ref ParallelTensorSpaceCoordinate) of this (usually parallel) tensor. + +For a \ref DynamicOpenDataflowGraph originating from a \ref MappedParallelComputationGraph, this field is filled in by \ref shard_expansion.h. +''' [[fields]] name = "accessor" @@ -45,3 +65,8 @@ type = "std::optional<::FlexFlow::DynamicTensorAccessor>" [[fields]] name = "role" type = "std::optional<::FlexFlow::DynamicTensorRole>" +docstring = ''' +\brief Identifies the role this tensor plays in training, e.g., a forward tensor, a gradient tensor, an optimizer buffer, a loss tensor, etc. + +For a \ref DynamicOpenDataflowGraph originating from a \ref MappedParallelComputationGraph, this field is filled in by \ref pass_expansion.h. +''' diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h index 9cccc565cc..aa9fbd2874 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h @@ -8,6 +8,10 @@ namespace FlexFlow { DynamicValueAttrs decide_dynamic_value_attrs_role(DynamicValueAttrs const &, DynamicTensorRole); +DynamicValueAttrs decide_dynamic_value_attrs_mapping( + DynamicValueAttrs const &, + OneToMany const &); + } // namespace FlexFlow #endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml new file mode 100644 index 0000000000..5a9234d815 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml @@ -0,0 +1,27 @@ +namespace = "FlexFlow" +name = "DynamicValueAttrsShardingInfo" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "fmt", + "json", +] + +includes = [ + "utils/one_to_many/one_to_many.h", + "op-attrs/parallel_tensor_space_coordinate.dtg.h", + "pcg/machine_space_coordinate.dtg.h", +] + +src_includes = [ +] + +[[fields]] +name = "shard_coord" +type = "::FlexFlow::ParallelTensorSpaceCoordinate" + +[[fields]] +name = "mapping" +type = "::FlexFlow::OneToMany<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/index.dox b/lib/task-spec/include/task-spec/dynamic_graph/index.dox index c48e67f4b3..97b72f8553 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/index.dox +++ b/lib/task-spec/include/task-spec/dynamic_graph/index.dox @@ -7,6 +7,7 @@ namespace FlexFlow { \section task-spec-lowering-passes Lowering Passes +- \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h: Embeds a \ref MappedParallelComputationGraph as a \ref DynamicOpenDataflowGraph. The first of the lowering passes. - \ref pass_expansion.h - \ref shard_expansion.h - \ref update_insertion.h diff --git a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h index 6a269ec3c9..f1d693c975 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h @@ -3,9 +3,15 @@ #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.h" namespace FlexFlow { +DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( + MappedParallelLayerInvocationInfo const &); + +DynamicNodeInvocation build_replicate_invocation(MappedParallelLayerInvocationInfo const &); + DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( MappedParallelComputationGraph const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h b/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h index 6dce8ad514..ad07b2941f 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h @@ -10,8 +10,13 @@ bool node_is_pass_expanded(DynamicNodeAttrs const &); bool value_is_pass_expanded(DynamicValueAttrs const &); bool slot_is_pass_expanded(DynamicTensorSlot const &); +bool node_is_ready_for_pass_expansion(DynamicNodeAttrs const &); +bool value_is_ready_for_pass_expansion(DynamicValueAttrs const &); +bool slot_is_ready_for_pass_expansion(DynamicTensorSlot const &); + bool no_part_of_graph_is_pass_expanded(DynamicOpenDataflowGraph const &); bool graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &); +bool graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &); DynamicNodeInvocation perform_fwd_pass_expansion_for_invocation(DynamicNodeInvocation const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml index 454f1b7e8c..d3cab6ecdb 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml @@ -14,7 +14,7 @@ includes = [ "op-attrs/parallel_tensor_shape.dtg.h", "op-attrs/parallel_tensor_space_coordinate.dtg.h", "pcg/machine_space_coordinate.dtg.h", - "utils/bidict/bidict.h", + "utils/one_to_many/one_to_many.h", "task-spec/dynamic_graph/dynamic_tensor_role.dtg.h", ] @@ -37,7 +37,7 @@ type = "std::optional<::FlexFlow::ParallelTensorSpaceCoordinate>" [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" +type = "std::optional<::FlexFlow::OneToMany<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" [[fields]] name = "role" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h index 4e0db1cd7e..50c713ee3b 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h @@ -4,19 +4,42 @@ #include "task-spec/dynamic_graph/dynamic_node_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.h" namespace FlexFlow { -bool node_is_shard_expanded(DynamicNodeAttrs const &); -bool value_is_shard_expanded(DynamicValueAttrs const &); +[[nodiscard]] bool node_is_shard_expanded(DynamicNodeAttrs const &); +[[nodiscard]] bool value_is_shard_expanded(DynamicValueAttrs const &); +[[nodiscard]] bool invocation_is_fully_shard_expanded(DynamicNodeInvocation const &); -bool no_part_of_graph_is_shard_expanded(DynamicOpenDataflowGraph const &); -bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &); +[[nodiscard]] bool node_is_ready_for_shard_expansion(DynamicNodeAttrs const &); +[[nodiscard]] bool value_is_ready_for_shard_expansion(DynamicValueAttrs const &); +[[nodiscard]] bool invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &); -std::unordered_set +[[nodiscard]] bool no_part_of_graph_is_shard_expanded(DynamicOpenDataflowGraph const &); +[[nodiscard]] bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &); +[[nodiscard]] bool graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &); + +[[nodiscard]] DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( + DynamicNodeAttrs const &, + MachineSpaceCoordinate const &); + +[[nodiscard]] DynamicValueAttrs apply_dynamic_value_attrs_sharding_info( + DynamicValueAttrs const &, + DynamicValueAttrsShardingInfo const &); + +[[nodiscard]] DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( + DynamicNodeInvocation const &, + DynamicNodeInvocationShardingInfo const &); + +[[nodiscard]] std::unordered_set + generate_shard_expansion_for_invocation(DynamicNodeInvocation const &); + +[[nodiscard]] std::unordered_set perform_shard_expansion_for_invocation(DynamicNodeInvocation const &); -DynamicOpenDataflowGraph +[[nodiscard]] DynamicOpenDataflowGraph perform_shard_expansion(DynamicOpenDataflowGraph const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h index 23fb7050a0..9818152b34 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h @@ -6,6 +6,15 @@ namespace FlexFlow { +bool node_has_already_had_update_insertion_performed(DynamicNodeAttrs const &); +bool value_has_already_had_update_insertion_performed(DynamicValueAttrs const &); + +bool node_is_ready_for_update_insertion(DynamicNodeAttrs const &); +bool value_is_ready_for_update_insertion(DynamicValueAttrs const &); + +bool no_part_of_graph_has_had_update_insertion_performed(DynamicOpenDataflowGraph const &); +bool graph_is_ready_for_update_insertion(DynamicOpenDataflowGraph const &); + std::unordered_set perform_update_insertion_for_invocation(DynamicNodeInvocation const &, OptimizerAttrs const &); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index ef41042a51..08ab3b11aa 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -10,7 +10,7 @@ #include "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" #include "utils/bidict/algorithms/bidict_from_pairs.h" -#include "utils/bidict/algorithms/unordered_set_of.h" +#include "utils/bidict/algorithms/bidict_unordered_set_of.h" #include "utils/containers/contains_key.h" #include "utils/containers/flatmap.h" #include "utils/containers/intersection.h" @@ -25,25 +25,13 @@ bool node_is_copy(DynamicNodeAttrs const &n) { return n.op_attrs.has_value() && n.op_attrs.value().is_copy(); } -static bool is_replicate_invocation(DynamicNodeInvocation const &i) { - return i.node_attrs.op_attrs.has_value() && - i.node_attrs.op_attrs.value().has() && - i.node_attrs.op_attrs.value() - .get() - .has(); -} - bool value_is_mapped(DynamicValueAttrs const &n) { return n.mapping.has_value(); } bool no_part_of_graph_is_copy_inserted(DynamicOpenDataflowGraph const &g) { auto slot_is_mapped = [](DynamicTensorSlot const &) -> bool { return false; }; - // check all non-replicate invocations for (DynamicNodeInvocation const &i : g.invocations) { - if (is_replicate_invocation(i)) { - continue; // replicate tensors have mapping set by design - } if (node_is_copy(i.node_attrs)) { return false; } @@ -69,6 +57,26 @@ bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { g, node_is_any, value_is_mapped, slot_is_mapped); } +void require_node_is_ready_for_copy_insertion(DynamicNodeAttrs const &node_attrs) { + ASSERT(node_attrs.mapping.has_value()); +} + +void require_graph_is_ready_for_copy_insertion(DynamicOpenDataflowGraph const &g) { + auto require_slot_is_ready_for_copy_insertion = [](DynamicTensorSlot const &slot) -> void { + return; + }; + + auto require_value_is_ready_for_copy_insertion = [](DynamicValueAttrs const &value_attrs) -> void { + return; + }; + + require_full_dynamic_graph_satisfies( + g, + require_node_is_ready_for_copy_insertion, + require_value_is_ready_for_copy_insertion, + require_slot_is_ready_for_copy_insertion); +} + static DynamicValueAttrs map_dynamic_value_attrs_for_task_group( DynamicTensorSlot const &slot, DynamicValueAttrs const &value, @@ -83,10 +91,10 @@ static std::pair DynamicValueAttrs const &output) { std::unordered_set< std::pair> - input_mapping = unordered_set_of(assert_unwrap(input.mapping)); + input_mapping = unstructured_relation_from_one_to_many(assert_unwrap(input.mapping)); std::unordered_set< std::pair> - output_mapping = unordered_set_of(assert_unwrap(output.mapping)); + output_mapping = unstructured_relation_from_one_to_many(assert_unwrap(output.mapping)); // Exclude the point shared between the input and output mappings, because // those will not result in actual copies once shard expansion is performed @@ -96,25 +104,19 @@ static std::pair DynamicValueAttrs filtered_input = input; filtered_input.mapping = - bidict_from_pairs(set_difference(input_mapping, remove)); + one_to_many_from_unstructured_relation(set_difference(input_mapping, remove)); DynamicValueAttrs filtered_output = output; filtered_output.mapping = - bidict_from_pairs(set_difference(output_mapping, remove)); + one_to_many_from_unstructured_relation(set_difference(output_mapping, remove)); return std::pair{filtered_input, filtered_output}; } -std::unordered_set perform_copy_insertion_for_invocation( - DynamicNodeInvocation const &i, - std::unordered_map const - &unmapped_value_to_mapped_source_value) { - - // replicate nodes have no MappedOperatorTaskGroup — - // pass through unchanged, no copies needed - if (is_replicate_invocation(i)) { - return {i}; - } +std::unordered_set copies_for_invocation_inputs( + DynamicNodeInvocation const &i, + std::unordered_map const &unmapped_value_to_src_mapped_value) +{ MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); auto map_tensor = [&](DynamicTensorSlot const &slot, @@ -124,31 +126,26 @@ std::unordered_set perform_copy_insertion_for_invocation( std::unordered_map mapped_inputs = map_values2(i.inputs, map_tensor); - std::unordered_map mapped_outputs = - map_values2(i.outputs, map_tensor); - std::unordered_set result{DynamicNodeInvocation{ - /*inputs=*/mapped_inputs, - /*node_attrs=*/i.node_attrs, - /*outputs=*/mapped_outputs, - }}; + std::unordered_set result; for (auto const &[slot, input] : i.inputs) { - if (!contains_key(unmapped_value_to_mapped_source_value, input)) { + if (!contains_key(unmapped_value_to_src_mapped_value, input)) { continue; } - DynamicValueAttrs source_value = - unmapped_value_to_mapped_source_value.at(input); - DynamicValueAttrs use_value = mapped_inputs.at(slot); - if (source_value != use_value) { - auto const &[filtered_source, filtered_use] = - filter_mapping_to_avoid_degenerate_copies(source_value, use_value); + DynamicValueAttrs src_mapped_value = unmapped_value_to_src_mapped_value.at(input); + DynamicValueAttrs use_mapped_value = mapped_inputs.at(slot); + + if (src_mapped_value != use_mapped_value) { + auto const &[filtered_source, filtered_use] = filter_mapping_to_avoid_degenerate_copies(src_mapped_value, use_mapped_value); DynamicNodeInvocation copy{ /*inputs=*/{ { - DynamicTensorSlot{TensorSlotName::INPUT, - slot.slot_tensor_role}, + DynamicTensorSlot{ + TensorSlotName::INPUT, + slot.slot_tensor_role, + }, filtered_source, }, }, @@ -179,22 +176,48 @@ std::unordered_set perform_copy_insertion_for_invocation( return result; } +std::unordered_set perform_copy_insertion_for_invocation( + DynamicNodeInvocation const &i, + std::unordered_map const + &unmapped_value_to_mapped_source_value) { + + MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); + + auto map_tensor = [&](DynamicTensorSlot const &slot, + DynamicValueAttrs const &value) { + return map_dynamic_value_attrs_for_task_group(slot, value, mapping); + }; + + DynamicNodeInvocation mapped_i = [&] { + std::unordered_map mapped_inputs = + map_values2(i.inputs, map_tensor); + std::unordered_map mapped_outputs = + map_values2(i.outputs, map_tensor); + + DynamicNodeInvocation r = i; + r.inputs = mapped_inputs; + r.outputs = mapped_outputs; + return r; + }(); + + std::unordered_set result = set_union( + copies_for_invocation_inputs(i, unmapped_value_to_mapped_source_value), + std::unordered_set{ + mapped_i, + }); + + return result; +} + DynamicOpenDataflowGraph perform_copy_insertion(DynamicOpenDataflowGraph const &g) { ASSERT(no_part_of_graph_is_copy_inserted(g)); + require_graph_is_ready_for_copy_insertion(g); std::unordered_map unmapped_value_to_mapped_source_value; for (DynamicNodeInvocation const &i : g.invocations) { - // replicate nodes have no MappedOperatorTaskGroup — - // output mapping already fully set, maps to itself - if (is_replicate_invocation(i)) { - for (auto const &[slot, value] : i.outputs) { - unmapped_value_to_mapped_source_value.insert(std::pair{value, value}); - } - continue; - } for (auto const &[slot, value] : i.outputs) { unmapped_value_to_mapped_source_value.insert( std::pair{value, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc new file mode 100644 index 0000000000..ea50449347 --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc @@ -0,0 +1,35 @@ +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" +#include "utils/containers/values.h" +#include "utils/containers/keys.h" +#include "utils/containers/all_of.h" + +namespace FlexFlow { + +bool invocation_fully_satisfies(DynamicNodeInvocation const &i, + std::function const &node_condition, + std::function const &value_condition, + std::function const &slot_condition) +{ + return node_condition(i.node_attrs) + && all_of(values(i.inputs), value_condition) + && all_of(keys(i.inputs), slot_condition) + && all_of(values(i.outputs), value_condition) + && all_of(keys(i.outputs), slot_condition); +} + +void require_invocation_fully_satisfies(DynamicNodeInvocation const &i, + std::function const &require_node_condition, + std::function const &require_value_condition, + std::function const &require_slot_condition) { + require_node_condition(i.node_attrs); + for (DynamicTensorSlot const &k : keys(i.inputs)) { + require_slot_condition(k); + require_value_condition(i.inputs.at(k)); + } + for (DynamicTensorSlot const &k : keys(i.outputs)) { + require_slot_condition(k); + require_value_condition(i.outputs.at(k)); + } +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index d2a5b653e5..a100c3adfb 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -19,6 +19,7 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/kwarg_dataflow_graph_input.dtg.h" #include "utils/many_to_one/many_to_one.h" +#include "utils/containers/require_all_of.h" namespace FlexFlow { @@ -56,6 +57,18 @@ bool no_part_of_dynamic_graph_satisfies( [&](DynamicTensorSlot const &s) -> bool { return !slot_condition(s); }); } +void require_full_dynamic_graph_satisfies( + DynamicOpenDataflowGraph const &g, + std::function const &node_condition, + std::function const &value_condition, + std::function const &slot_condition) +{ + require_all_of(get_dynamic_nodes(g), node_condition); + require_all_of(get_dynamic_values(g), value_condition); + require_all_of(get_dynamic_tensor_slots(g), slot_condition); +} + + std::unordered_multiset get_dynamic_nodes(DynamicOpenDataflowGraph const &g) { return transform(unordered_multiset_of(g.invocations), diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc index 282279edbe..9a70c5cdd0 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc @@ -13,4 +13,17 @@ DynamicValueAttrs return result; } +DynamicValueAttrs decide_dynamic_value_attrs_mapping( + DynamicValueAttrs const &attrs, + OneToMany const &mapping) +{ + ASSERT(!attrs.mapping.has_value()); + + DynamicValueAttrs result = attrs; + result.mapping = mapping; + + return result; +} + + } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc index 6bfc477e3a..7fe3927fd1 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc @@ -7,7 +7,7 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include #include #include diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 7a149787b9..391ebaff3b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -10,7 +10,6 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "utils/bidict/algorithms/merge_disjoint_bidicts.h" -#include "utils/containers/generate_map.h" #include "utils/containers/get_only.h" #include "utils/containers/map_keys_and_values.h" #include "utils/containers/require_only_key.h" @@ -18,194 +17,64 @@ #include #include #include +#include "utils/containers/unordered_map_from_map.h" +#include "utils/bidict/algorithms/bidict_unordered_set_of.h" namespace FlexFlow { -static bidict - get_input_mapping_for_replicate( - MappedParallelComputationGraph const &mpcg, - parallel_layer_guid_t const &replicate_layer) { - - ASSERT(mpcg_get_pcg_op_attrs(mpcg, replicate_layer).is_parallel_replicate()); - - auto [input_slot_name, input_edge] = - get_only(mpcg_get_incoming_edges(mpcg, replicate_layer)); - - parallel_layer_guid_t producer_layer = get_src_layer(input_edge); - TensorSlotName producer_slot = get_src_layer_output_slot_name(input_edge); - - return get_tensor_bindings_for_slot_name( - /*task_group=*/mpcg_get_mapping_for_layer(mpcg, producer_layer), - /*slot_name=*/producer_slot); -} - -static bidict - build_replicated_output_mapping( - MappedParallelComputationGraph const &mpcg, - parallel_tensor_guid_t const &output_tensor_guid) { - - std::unordered_set consumers = - mpcg_get_parallel_tensor_uses(mpcg, output_tensor_guid); - ASSERT(!consumers.empty()); - - // union all consumer bindings — each consumer shard maps to a distinct - // (discard_copy, machine) pair since replicas are always on different machines - bidict result = - merge_disjoint_bidicts(transform( - consumers, - [&](parallel_tensor_use_t const &use) - -> bidict { - parallel_layer_guid_t consumer_layer = - parallel_tensor_use_get_layer(use); - TensorSlotName slot_name = parallel_tensor_use_get_slot(use); - - MappedOperatorTaskGroup consumer_mapping = - mpcg_get_mapping_for_layer(mpcg, consumer_layer); - bidict - binding = get_tensor_bindings_for_slot_name(consumer_mapping, - slot_name); - - return binding; - })); - - return result; -} - -static DynamicNodeInvocation - build_replicate_invocation(parallel_layer_guid_t const &layer, - ReplicateAttrs const &attrs, - MappedParallelComputationGraph const &mpcg) { - - ManyToOne incoming = - mpcg_get_incoming_tensors(mpcg, layer); - TensorSlotName input_slot_name = TensorSlotName::INPUT; - parallel_tensor_guid_t input_tensor_guid = - require_only_key(incoming.l_to_r(), input_slot_name); - ParallelTensorAttrs input_attrs = - mpcg_get_parallel_tensor_attrs(mpcg, input_tensor_guid); - - bidict outgoing = - mpcg_get_outgoing_tensors(mpcg, layer); - TensorSlotName output_slot_name = TensorSlotName::OUTPUT; - parallel_tensor_guid_t output_tensor_guid = - require_only_key(outgoing.l_to_r(), output_slot_name); - ParallelTensorAttrs output_attrs = - mpcg_get_parallel_tensor_attrs(mpcg, output_tensor_guid); - - bidict input_mapping = - get_input_mapping_for_replicate(mpcg, layer); - - DynamicValueAttrs input_value{ - /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, - /*parallel_tensor_shape=*/input_attrs.shape, - /*shard_coord=*/std::nullopt, - /*mapping=*/input_mapping, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }; - - DynamicValueAttrs output_value{ - /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, - /*parallel_tensor_shape=*/output_attrs.shape, - /*shard_coord=*/std::nullopt, - /*mapping=*/build_replicated_output_mapping(mpcg, output_tensor_guid), - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }; - - DynamicNodeAttrs node_attrs{ +DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( + MappedParallelLayerInvocationInfo const &invocation_info) +{ + DynamicNodeAttrs result_attrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{PCGOperatorAttrs{attrs}}, - /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, + /*mapping=*/invocation_info.layer_info.mapping, + /*op_attrs=*/TrainingOperationAttrs{invocation_info.layer_info.attrs.op_attrs}, + /*pcg_layer_guid=*/dynamic_layer_guid_t{invocation_info.layer_info.guid}, /*per_device_op_state=*/std::nullopt, }; - DynamicNodeInvocation invocation_node{ - /*inputs=*/{ - { - DynamicTensorSlot{input_slot_name, std::nullopt}, - input_value, - }, + auto lift_kv_pair = + [&](TensorSlotName slot_name, + ParallelTensorInfo const &tensor) + -> std::pair + { + return { + DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, }, - /*node_attrs=*/node_attrs, - /*outputs=*/ - { - { - DynamicTensorSlot{output_slot_name, std::nullopt}, - output_value, - }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{tensor.guid}, + /*parallel_tensor_shape=*/tensor.attrs.shape, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, }, + }; }; - return invocation_node; -} - -DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( - MappedParallelComputationGraph const &mpcg) { - - ParallelComputationGraph pcg = pcg_from_mpcg(mpcg); - - auto mk_invocation = - [&](parallel_layer_guid_t layer, - ParallelLayerAttrs const &attrs) -> DynamicNodeInvocation { - if (attrs.op_attrs.is_parallel_replicate()) { - // build replicate invocation - DynamicNodeInvocation repl_inv = build_replicate_invocation( - layer, attrs.op_attrs.require_parallel_replicate(), mpcg); - return repl_inv; - } else { - DynamicNodeAttrs result_attrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/mpcg_get_mapping_for_layer(mpcg, layer), - /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs}, - /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, - /*per_device_op_state=*/std::nullopt, - }; + std::map result_inputs = + transform(invocation_info.incoming, lift_kv_pair); - auto mk_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { - return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - }; - }; + std::map result_outputs = + transform(invocation_info.outgoing, lift_kv_pair); - auto mk_value_attrs = - [&](parallel_tensor_guid_t const &tensor) -> DynamicValueAttrs { - ParallelTensorAttrs attrs = get_parallel_tensor_attrs(pcg, tensor); - - return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, - /*parallel_tensor_shape=*/attrs.shape, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }; - }; - - std::unordered_map result_inputs = - map_keys_and_values( - get_incoming_tensors(pcg, layer), mk_slot, mk_value_attrs); - - std::unordered_map result_outputs = - map_keys_and_values( - get_outgoing_tensors(pcg, layer), mk_slot, mk_value_attrs); + DynamicNodeInvocation invocation = DynamicNodeInvocation{ + /*inputs=*/unordered_map_from_map(result_inputs), + /*node_attrs=*/result_attrs, + /*outputs=*/unordered_map_from_map(result_outputs), + }; - DynamicNodeInvocation invocation = DynamicNodeInvocation{ - /*inputs=*/result_inputs, - /*node_attrs=*/result_attrs, - /*outputs=*/result_outputs, - }; + return invocation; +} - return invocation; - }; - }; +DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( + MappedParallelComputationGraph const &mpcg) { - return dynamic_open_dataflow_graph_from_invocation_set(transform_pairs( - unordered_set_of(get_parallel_layer_attrs_mapping(pcg)), mk_invocation)); + return dynamic_open_dataflow_graph_from_invocation_set( + transform(unordered_set_of(mpcg_get_invocation_set(mpcg)), make_dynamic_node_invocation_from_mapped)); } } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index 64fe2df0be..a348ba77da 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -21,6 +21,18 @@ bool value_is_pass_expanded(DynamicValueAttrs const &v) { return v.role.has_value(); } +bool node_is_ready_for_pass_expansion(DynamicNodeAttrs const &) { + return true; +} + +bool value_is_ready_for_pass_expansion(DynamicValueAttrs const &) { + return true; +} + +bool slot_is_ready_for_pass_expansion(DynamicTensorSlot const &) { + return true; +} + bool no_part_of_graph_is_pass_expanded(DynamicOpenDataflowGraph const &g) { return no_part_of_dynamic_graph_satisfies( g, node_is_pass_expanded, value_is_pass_expanded, slot_is_pass_expanded); @@ -31,6 +43,11 @@ bool graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &g) { g, node_is_pass_expanded, value_is_pass_expanded, slot_is_pass_expanded); } +bool graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &g) { + return full_dynamic_graph_satisfies( + g, node_is_ready_for_pass_expansion, value_is_ready_for_pass_expansion, slot_is_ready_for_pass_expansion); +} + DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &s, FwbTensorType tensor_type) { ASSERT(!slot_is_pass_expanded(s)); @@ -139,6 +156,7 @@ DynamicOpenDataflowGraph perform_pass_expansion(DynamicOpenDataflowGraph const &g) { ASSERT(no_part_of_graph_is_pass_expanded(g)); + ASSERT(graph_is_ready_for_pass_expansion(g)); DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [](DynamicNodeInvocation const &invocation) { diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index d3365ae44c..badd376a8b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -1,12 +1,16 @@ #include "task-spec/dynamic_graph/shard_expansion.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" -#include "utils/bidict/algorithms/filter_keys.h" +#include "utils/bidict/algorithms/bidict_filter_keys.h" #include "utils/containers/get_only.h" #include "utils/containers/map_values2.h" #include "utils/containers/require_same.h" #include "utils/containers/transform.h" #include "utils/optional.h" +#include "utils/containers/binary_merge_disjoint_maps.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" +#include "utils/containers/map_from_unordered.h" +#include "utils/one_to_many/one_to_many_filter_keys.h" namespace FlexFlow { @@ -14,8 +18,74 @@ bool node_is_shard_expanded(DynamicNodeAttrs const &n) { return n.device_coord.has_value(); } +bool node_is_ready_for_shard_expansion(DynamicNodeAttrs const &n) { + if (!n.op_attrs.has_value()) { + return false; + } + + if (n.op_attrs.value().is_pcg_op()) { + if (!n.mapping.has_value()) { + return false; + } + } + + return true; +} + +void require_node_is_ready_for_shard_expansion(DynamicNodeAttrs const &n) { + ASSERT(n.op_attrs.has_value()); + if (n.op_attrs.value().is_pcg_op()) { + ASSERT(n.mapping.has_value()); + } +} + + +bool invocation_is_fully_shard_expanded(DynamicNodeInvocation const &i) { + auto slot_is_shard_expanded = [](DynamicTensorSlot const &) { + return true; + }; + + return invocation_fully_satisfies( + i, + node_is_shard_expanded, + value_is_shard_expanded, + slot_is_shard_expanded); +} + bool value_is_shard_expanded(DynamicValueAttrs const &n) { - return n.shard_coord.has_value(); + return n.shard_coord.has_value() && n.mapping.has_value(); +} + +bool value_is_ready_for_shard_expansion(DynamicValueAttrs const &n) { + return true; +} + +void require_value_is_ready_for_shard_expansion(DynamicValueAttrs const &n) { + return; +} + +bool invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &i) { + auto slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) { + return true; + }; + + return invocation_fully_satisfies( + i, + node_is_ready_for_shard_expansion, + value_is_ready_for_shard_expansion, + slot_is_ready_for_shard_expansion); +} + +void require_invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &i) { + auto require_slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) -> void { + return; + }; + + require_invocation_fully_satisfies( + i, + require_node_is_ready_for_shard_expansion, + require_value_is_ready_for_shard_expansion, + require_slot_is_ready_for_shard_expansion); } bool no_part_of_graph_is_shard_expanded(DynamicOpenDataflowGraph const &g) { @@ -39,20 +109,53 @@ bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { value_is_shard_expanded, slot_is_shard_expanded); } -static bidict + +static OneToMany restrict_tensor_mapping_keys_to_coord( - bidict const + OneToMany const &mapping, ParallelTensorSpaceCoordinate const ¶llel_tensor_coord) { - return filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { + return one_to_many_filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { return p == parallel_tensor_coord; }); } +static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( + DynamicNodeInvocation const &i, + MachineSpaceCoordinate const &machine_coord, + OperatorAtomicTaskShardBinding const &binding) { + + auto shard_expand_value_attrs = + [&](DynamicTensorSlot const &s, DynamicValueAttrs const &v) -> DynamicValueAttrsShardingInfo { + ParallelTensorSpaceCoordinate parallel_tensor_coord = + binding.tensor_coords.at(s.slot_name); + + return DynamicValueAttrsShardingInfo{ + /*shard_coord=*/parallel_tensor_coord, + /*mapping=*/restrict_tensor_mapping_keys_to_coord(v.mapping.value(), parallel_tensor_coord), + }; + }; + + DynamicNodeAttrs expanded_node_attrs = [&]() { + DynamicNodeAttrs result = i.node_attrs; + result.device_coord = machine_coord; + return result; + }(); + + return DynamicNodeInvocationShardingInfo{ + /*device_coord=*/machine_coord, + /*value_sharding=*/map_from_unordered( + map_values2( + binary_merge_disjoint_maps(i.inputs, i.outputs), + shard_expand_value_attrs)), + }; +} + static DynamicNodeInvocation shard_invocation_for_binding( DynamicNodeInvocation const &i, MachineSpaceCoordinate const &machine_coord, OperatorAtomicTaskShardBinding const &binding) { + auto shard_expand_value_attrs = [&](DynamicTensorSlot const &s, DynamicValueAttrs const &v) -> DynamicValueAttrs { @@ -63,8 +166,9 @@ static DynamicNodeInvocation shard_invocation_for_binding( result.shard_coord = parallel_tensor_coord; result.mapping = transform( v.mapping, - [&](bidict const - &mapping) { + [&](OneToMany const &mapping) + -> OneToMany + { return restrict_tensor_mapping_keys_to_coord(mapping, parallel_tensor_coord); }); @@ -84,124 +188,19 @@ static DynamicNodeInvocation shard_invocation_for_binding( }; } -static std::unordered_set - perform_shard_expansion_for_replicate(DynamicNodeInvocation const &i) { - auto const &[input_slot, input] = get_only(i.inputs); - auto const &[output_slot, output] = get_only(i.outputs); - - bidict input_mapping = - assert_unwrap(input.mapping); - bidict output_mapping = - assert_unwrap(output.mapping); - - return transform(output_mapping.left_values(), - [&](ParallelTensorSpaceCoordinate const &p) { - ParallelTensorSpaceCoordinate input_p{ - /*sum_component=*/p.sum_component, - /*discard_copy_component=*/nonnegative_int{0}, - /*shard_components=*/p.shard_components, - }; - return shard_invocation_for_binding( - i, - output_mapping.at_l(p), - OperatorAtomicTaskShardBinding{{ - {input_slot.slot_name, input_p}, - {output_slot.slot_name, p}, - }}); - }); -} - -static std::unordered_set - perform_shard_expansion_for_replicate_bwd(DynamicNodeInvocation const &i) { - - std::optional output_grad_opt; - std::optional output_fwd_opt; - std::optional output_grad_slot_opt; - std::optional output_fwd_slot_opt; - - for (auto const &[slot, value] : i.inputs) { - if (slot.slot_tensor_role == DynamicTensorRole{FwbTensorType::GRADIENT}) { - output_grad_slot_opt = slot; - output_grad_opt = value; - } else { - output_fwd_slot_opt = slot; - output_fwd_opt = value; - } - } - - DynamicValueAttrs output_grad = assert_unwrap(output_grad_opt); - DynamicValueAttrs output_fwd = assert_unwrap(output_fwd_opt); - DynamicTensorSlot output_grad_slot = assert_unwrap(output_grad_slot_opt); - DynamicTensorSlot output_fwd_slot = assert_unwrap(output_fwd_slot_opt); - auto const &[input_grad_slot, input_grad] = get_only(i.outputs); - - bidict - output_grad_mapping = assert_unwrap(output_grad.mapping); - bidict - input_grad_mapping = assert_unwrap(input_grad.mapping); - - std::unordered_map, - std::unordered_set> - by_shard; - for (auto const &p : output_grad_mapping.left_values()) { - by_shard[p.shard_components].insert(p); - } - - std::unordered_set result; - for (auto const &[shard_components, replica_coords] : by_shard) { - ParallelTensorSpaceCoordinate src_p{ - nonnegative_int{0}, nonnegative_int{0}, shard_components}; - MachineSpaceCoordinate src_machine = input_grad_mapping.at_l(src_p); - - bidict - replica_mapping; - for (auto const &p : replica_coords) { - replica_mapping.equate(p, output_grad_mapping.at_l(p)); - } - - DynamicValueAttrs sharded_output_grad = output_grad; - sharded_output_grad.mapping = replica_mapping; - sharded_output_grad.shard_coord = src_p; - - DynamicValueAttrs sharded_output_fwd = output_fwd; - sharded_output_fwd.mapping = replica_mapping; - sharded_output_fwd.shard_coord = src_p; - - DynamicValueAttrs sharded_input_grad = input_grad; - sharded_input_grad.mapping = - bidict{ - {src_p, src_machine}}; - sharded_input_grad.shard_coord = src_p; - - DynamicNodeAttrs sharded_node = i.node_attrs; - sharded_node.device_coord = src_machine; - - result.insert(DynamicNodeInvocation{ - /*inputs=*/{ - {output_fwd_slot, sharded_output_fwd}, - {output_grad_slot, sharded_output_grad}, - }, - /*node_attrs=*/sharded_node, - /*outputs=*/ - { - {input_grad_slot, sharded_input_grad}, - }, - }); - } - return result; -} - -static std::unordered_set - perform_shard_expansion_for_copy(DynamicNodeInvocation const &i) { +static std::set + generate_shard_expansion_for_copy(DynamicNodeInvocation const &i) { auto [input_slot, input] = get_only(i.inputs); auto [output_slot, output] = get_only(i.outputs); - bidict input_mapping = + + OneToMany input_mapping = assert_unwrap(input.mapping); require_same(input_mapping.left_values(), assert_unwrap(output.mapping).left_values()); return transform( - input_mapping.left_values(), [&](ParallelTensorSpaceCoordinate const &p) { + input_mapping.left_values(), + [&](ParallelTensorSpaceCoordinate const &p) -> DynamicNodeInvocationShardingInfo { // The machine coord for a copy is inherently nebulous because it // doesn't strictly run in any single location. Further, Realm has the // flexibility to issue a copy operation from anywhere in the machine, @@ -209,9 +208,9 @@ static std::unordered_set // because we expect this to align with the most efficient way to issue // copies in Realm, although the current Realm backend uses a // centralized controller and thus issues copies all from a single node. - MachineSpaceCoordinate machine_coord = input_mapping.at_l(p); + MachineSpaceCoordinate machine_coord = get_only(input_mapping.at_l(p)); - return shard_invocation_for_binding(i, + return invocation_sharding_info_for_binding(i, machine_coord, OperatorAtomicTaskShardBinding{{ {input_slot.slot_name, p}, @@ -222,28 +221,91 @@ static std::unordered_set std::unordered_set perform_shard_expansion_for_invocation(DynamicNodeInvocation const &i) { - if (i.node_attrs.op_attrs.has_value() && - i.node_attrs.op_attrs.value().is_copy()) { - return perform_shard_expansion_for_copy(i); - } - bool const is_replicate = - i.node_attrs.op_attrs.has_value() && - i.node_attrs.op_attrs.value().has() && - i.node_attrs.op_attrs.value() - .get() - .has(); - - // forward replicate - if (is_replicate && i.node_attrs.task_type.has_value() && - i.node_attrs.task_type.value() == DynamicTaskType::FWD) { - return perform_shard_expansion_for_replicate(i); - } + std::unordered_set + shard_expansion_info = generate_shard_expansion_for_invocation(i); + + return transform( + shard_expansion_info, + [&](DynamicNodeInvocationShardingInfo const &s) + -> DynamicNodeInvocation + { + return apply_dynamic_node_invocation_sharding_info(i, s); + }); +} + +bool graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &g) { + auto slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) -> bool { + return false; + }; + + return full_dynamic_graph_satisfies(g, + node_is_ready_for_shard_expansion, + value_is_ready_for_shard_expansion, + slot_is_ready_for_shard_expansion); +} + + +void require_graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &g) { + auto require_slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) -> void { + return; + }; - // backward replicate - if (is_replicate && i.node_attrs.task_type.has_value() && - i.node_attrs.task_type.value() == DynamicTaskType::BWD) { - return perform_shard_expansion_for_replicate_bwd(i); + return require_full_dynamic_graph_satisfies(g, + require_node_is_ready_for_shard_expansion, + require_value_is_ready_for_shard_expansion, + require_slot_is_ready_for_shard_expansion); +} + +DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( + DynamicNodeAttrs const &node_attrs, + MachineSpaceCoordinate const &device_coord) +{ + DynamicNodeAttrs result = node_attrs; + result.device_coord = device_coord; + + return result; +} + +DynamicValueAttrs apply_dynamic_value_attrs_sharding_info( + DynamicValueAttrs const &value_attrs, + DynamicValueAttrsShardingInfo const &value_sharding_info) +{ + DynamicValueAttrs result = value_attrs; + result.shard_coord = value_sharding_info.shard_coord; + result.mapping = value_sharding_info.mapping; + return result; +} + +DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( + DynamicNodeInvocation const &invocation, + DynamicNodeInvocationShardingInfo const &invocation_sharding_info) +{ + require_invocation_is_ready_for_shard_expansion(invocation); + + auto shard_value = [&](DynamicTensorSlot const &slot, DynamicValueAttrs const &value_attrs) -> DynamicValueAttrs { + DynamicValueAttrsShardingInfo sharding_info = invocation_sharding_info.value_sharding.at(slot); + return apply_dynamic_value_attrs_sharding_info(value_attrs, sharding_info); + }; + + DynamicNodeInvocation result = DynamicNodeInvocation{ + /*inputs=*/map_values2(invocation.inputs, shard_value), + /*node_attrs=*/apply_dynamic_node_attrs_sharding_info(invocation.node_attrs, invocation_sharding_info.device_coord), + /*outputs=*/map_values2(invocation.outputs, shard_value), + }; + + ASSERT(invocation_is_fully_shard_expanded(result)); + return result; +} + +std::unordered_set + generate_shard_expansion_for_invocation(DynamicNodeInvocation const &i) +{ + require_invocation_is_ready_for_shard_expansion(i); + + if (i.node_attrs.op_attrs.has_value() && + i.node_attrs.op_attrs.value().is_copy()) { + return unordered_set_of(generate_shard_expansion_for_copy(i)); } MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); @@ -253,11 +315,11 @@ std::unordered_set return transform( shard_machine_coords, - [&](MachineSpaceCoordinate const &c) -> DynamicNodeInvocation { + [&](MachineSpaceCoordinate const &c) -> DynamicNodeInvocationShardingInfo { OperatorAtomicTaskShardBinding slot_bindings = mapping.get_shard_bindings().at_l(c); - return shard_invocation_for_binding(i, c, slot_bindings); + return invocation_sharding_info_for_binding(i, c, slot_bindings); }); } @@ -265,6 +327,7 @@ DynamicOpenDataflowGraph perform_shard_expansion(DynamicOpenDataflowGraph const &g) { ASSERT(no_part_of_graph_is_shard_expanded(g)); + require_graph_is_ready_for_shard_expansion(g); DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set(g, [&](DynamicNodeInvocation const &i) { diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index 2160f6bf82..fdc705dc54 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -6,11 +6,13 @@ #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" #include "test/utils/doctest/fmt/unordered_set.h" #include +#include "task-spec/dynamic_graph/dynamic_value_attrs.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("perform_copy_insertion_for_invocation") { + TEST_CASE("copies_for_invocation_inputs") { auto mk_machine_coord = [](nonnegative_int node_idx, nonnegative_int device_idx) -> MachineSpaceCoordinate { @@ -21,6 +23,33 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; + auto mk_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + }; + }; + + auto mk_value = [](size_t src_node_id, + TensorSlotName src_slot_name) + -> DynamicValueAttrs { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{src_node_id}, + src_slot_name, + }, + }, + }, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + }; + auto mk_pt_coord = [](nonnegative_int idx1, nonnegative_int idx2, @@ -37,397 +66,567 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - auto mk_input_shard_binding = [&](ParallelTensorSpaceCoordinate const &c) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ + size_t invocation_id = 20; + + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); + MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); + MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); + + SUBCASE("standard operator") { + auto mk_input_shard_binding = [&](ParallelTensorSpaceCoordinate const &c) + -> OperatorAtomicTaskShardBinding { + return OperatorAtomicTaskShardBinding{ + /*tensor_coords=*/{ + { + TensorSlotName::OUTPUT, + c, + }, + }, + }; + }; + + auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, + ParallelTensorSpaceCoordinate const &c2, + ParallelTensorSpaceCoordinate const &c3, + ParallelTensorSpaceCoordinate const &c4) + -> OperatorAtomicTaskShardBinding { + return OperatorAtomicTaskShardBinding{ + /*tensor_coords=*/{ + { + TensorSlotName::INPUT, + c1, + }, + { + TensorSlotName::WEIGHT, + c2, + }, + { + TensorSlotName::OUTPUT_1, + c3, + }, + { + TensorSlotName::OUTPUT_2, + c4, + }, + }, + }; + }; + + ParallelTensorSpaceCoordinate mc1_input_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate mc1_weight_coord = + mk_pt_coord(0_n, 1_n, 2_n, 0_n); + ParallelTensorSpaceCoordinate mc1_output_1_coord = + mk_pt_coord(1_n, 0_n, 0_n, 1_n); + ParallelTensorSpaceCoordinate mc1_output_2_coord = + mk_pt_coord(3_n, 0_n, 0_n, 0_n); + + ParallelTensorSpaceCoordinate mc2_input_coord = + mk_pt_coord(0_n, 1_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate mc2_weight_coord = + mk_pt_coord(0_n, 4_n, 2_n, 0_n); + ParallelTensorSpaceCoordinate mc2_output_1_coord = + mk_pt_coord(1_n, 2_n, 0_n, 1_n); + ParallelTensorSpaceCoordinate mc2_output_2_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + + MappedOperatorTaskGroup input_mapping_same = MappedOperatorTaskGroup{ + bidict{ + { + mc1, + mk_input_shard_binding(mc1_input_coord), + }, { - TensorSlotName::OUTPUT, - c, + mc2, + mk_input_shard_binding(mc2_input_coord), }, }, }; - }; - auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, - ParallelTensorSpaceCoordinate const &c2, - ParallelTensorSpaceCoordinate const &c3, - ParallelTensorSpaceCoordinate const &c4) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ + MappedOperatorTaskGroup weight_mapping_same = MappedOperatorTaskGroup{ + bidict{ { - TensorSlotName::INPUT, - c1, + mc1, + mk_input_shard_binding(mc1_weight_coord), }, { - TensorSlotName::WEIGHT, - c2, + mc2, + mk_input_shard_binding(mc2_weight_coord), }, + }, + }; + + MappedOperatorTaskGroup invocation_mapping = MappedOperatorTaskGroup{ + bidict{ { - TensorSlotName::OUTPUT_1, - c3, + mc1, + mk_shard_binding(mc1_input_coord, + mc1_weight_coord, + mc1_output_1_coord, + mc1_output_2_coord), }, { - TensorSlotName::OUTPUT_2, - c4, + mc2, + mk_shard_binding(mc2_input_coord, + mc2_weight_coord, + mc2_output_1_coord, + mc2_output_2_coord), }, }, }; - }; - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); - MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); - MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); + MappedOperatorTaskGroup invocation_mapping_diff_vs_copy1 = + MappedOperatorTaskGroup{ + bidict{ + { + mc2, + mk_shard_binding(mc2_input_coord, + mc2_weight_coord, + mc2_output_1_coord, + mc2_output_2_coord), + }, + }, + }; - ParallelTensorSpaceCoordinate mc1_input_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc1_weight_coord = - mk_pt_coord(0_n, 1_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc1_output_1_coord = - mk_pt_coord(1_n, 0_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc1_output_2_coord = - mk_pt_coord(3_n, 0_n, 0_n, 0_n); - - ParallelTensorSpaceCoordinate mc2_input_coord = - mk_pt_coord(0_n, 1_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc2_weight_coord = - mk_pt_coord(0_n, 4_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc2_output_1_coord = - mk_pt_coord(1_n, 2_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc2_output_2_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - - MappedOperatorTaskGroup input_mapping_same = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_input_coord), - }, - { - mc2, - mk_input_shard_binding(mc2_input_coord), - }, - }, - }; + DynamicValueAttrs graph_input1 = + mk_value(0, TensorSlotName::OUTPUT); + + DynamicValueAttrs graph_input1_use = + decide_dynamic_value_attrs_mapping( + graph_input1, + get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::INPUT)); + + DynamicValueAttrs graph_input1_use_diff_vs_copy1 = + decide_dynamic_value_attrs_mapping( + graph_input1, + get_tensor_bindings_for_slot_name(invocation_mapping_diff_vs_copy1, TensorSlotName::INPUT)); + + DynamicValueAttrs graph_input2 = + mk_value(1, TensorSlotName::OUTPUT); + + DynamicValueAttrs graph_input2_use = + decide_dynamic_value_attrs_mapping( + graph_input2, + get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::WEIGHT)); + + DynamicValueAttrs invocation_output1 = mk_value(invocation_id, + TensorSlotName::OUTPUT_1); + DynamicValueAttrs invocation_output1_src = + decide_dynamic_value_attrs_mapping( + invocation_output1, + get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT_1)); + + DynamicValueAttrs invocation_output2 = mk_value(invocation_id, + TensorSlotName::OUTPUT_2); + DynamicValueAttrs invocation_output2_src = + decide_dynamic_value_attrs_mapping( + invocation_output2, + get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT_2)); + + DynamicValueAttrs graph_input1_src_same = + decide_dynamic_value_attrs_mapping( + graph_input1, + get_tensor_bindings_for_slot_name(input_mapping_same, TensorSlotName::OUTPUT)); + + DynamicValueAttrs graph_input2_src_same = + decide_dynamic_value_attrs_mapping( + graph_input2, + get_tensor_bindings_for_slot_name(weight_mapping_same, TensorSlotName::OUTPUT)); + + DynamicNodeInvocation input = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + graph_input1, + }, + { + mk_slot(TensorSlotName::WEIGHT), + graph_input2, + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coord=*/std::nullopt, + /*mapping=*/invocation_mapping, + /*op_attrs=*/std::nullopt, + /*layer_guid=*/ + dynamic_layer_guid_t{parallel_layer_guid_t{Node{invocation_id}}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + mk_slot(TensorSlotName::OUTPUT_1), + invocation_output1, + }, + { + mk_slot(TensorSlotName::OUTPUT_2), + invocation_output2, + }, + }, + }; - MappedOperatorTaskGroup weight_mapping_same = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_weight_coord), + auto mk_copy = [&](DynamicValueAttrs const &src, + DynamicValueAttrs const &dst) { + return DynamicNodeInvocation{ + /*inputs=*/{{mk_slot(TensorSlotName::INPUT), src}}, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, }, - { - mc2, - mk_input_shard_binding(mc2_weight_coord), + /*outputs=*/{{mk_slot(TensorSlotName::OUTPUT), dst}}, + }; + }; + + SUBCASE("same mapping, no copies") { + std::unordered_map sources_same{ + {graph_input1, graph_input1_src_same}, + {graph_input2, graph_input2_src_same}, + }; + + std::unordered_set result = + copies_for_invocation_inputs(input, sources_same); + + std::unordered_set correct = {}; + + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } + + SUBCASE("copy one tensor, one point") { + MappedOperatorTaskGroup input_mapping_copy1 = MappedOperatorTaskGroup{ + bidict{ + { + mc1, + mk_input_shard_binding(mc1_input_coord), + }, + { + mc3, + mk_input_shard_binding(mc2_input_coord), + }, }, - }, - }; + }; - MappedOperatorTaskGroup invocation_mapping = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_shard_binding(mc1_input_coord, - mc1_weight_coord, - mc1_output_1_coord, - mc1_output_2_coord), + MappedOperatorTaskGroup input_mapping_copy1_diff_vs_use = + MappedOperatorTaskGroup{ + bidict{ + { + mc3, + mk_input_shard_binding(mc2_input_coord), + }, + }, + }; + + DynamicValueAttrs graph_input1_src_copy1 = + decide_dynamic_value_attrs_mapping( + graph_input1, + get_tensor_bindings_for_slot_name(input_mapping_copy1, TensorSlotName::OUTPUT)); + + DynamicValueAttrs graph_input1_src_copy1_diff_vs_use = + decide_dynamic_value_attrs_mapping( + graph_input1, + get_tensor_bindings_for_slot_name(input_mapping_copy1_diff_vs_use, TensorSlotName::OUTPUT)); + + std::unordered_map sources_copy1{ + {graph_input1, graph_input1_src_copy1}, + {graph_input2, graph_input2_src_same}}; + + std::unordered_set result = + copies_for_invocation_inputs(input, sources_copy1); + + std::unordered_set correct = { + mk_copy(graph_input1_src_copy1_diff_vs_use, graph_input1_use_diff_vs_copy1), + }; + + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } + + SUBCASE("copy two tensors, two points") { + MappedOperatorTaskGroup input_mapping_copy2 = MappedOperatorTaskGroup{ + bidict{ + { + mc3, + mk_input_shard_binding(mc1_input_coord), + }, + { + mc4, + mk_input_shard_binding(mc2_input_coord), + }, }, - { - mc2, - mk_shard_binding(mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), + }; + MappedOperatorTaskGroup weight_mapping_copy2 = MappedOperatorTaskGroup{ + bidict{ + { + mc4, + mk_input_shard_binding(mc1_weight_coord), + }, + { + mc3, + mk_input_shard_binding(mc2_weight_coord), + }, }, - }, - }; + }; - MappedOperatorTaskGroup invocation_mapping_diff_vs_copy1 = - MappedOperatorTaskGroup{ - bidict{ + DynamicValueAttrs graph_input1_src_copy2 = + decide_dynamic_value_attrs_mapping( + graph_input1, + get_tensor_bindings_for_slot_name(input_mapping_copy2, TensorSlotName::OUTPUT)); + + DynamicValueAttrs graph_input2_src_copy2 = + decide_dynamic_value_attrs_mapping( + graph_input2, + get_tensor_bindings_for_slot_name(weight_mapping_copy2, TensorSlotName::OUTPUT)); + + std::unordered_map sources_copy2{ + {graph_input1, graph_input1_src_copy2}, + {graph_input2, graph_input2_src_copy2}}; + + std::unordered_set result = + copies_for_invocation_inputs(input, sources_copy2); + + std::unordered_set correct = { + mk_copy(graph_input1_src_copy2, graph_input1_use), + mk_copy(graph_input2_src_copy2, graph_input2_use), + }; + + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } + } + + SUBCASE("replicate operator") { + + auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, + ParallelTensorSpaceCoordinate const &c2) + -> OperatorAtomicTaskShardBinding { + return OperatorAtomicTaskShardBinding{ + /*tensor_coords=*/{ + { + TensorSlotName::INPUT, + c1, + }, { - mc2, - mk_shard_binding(mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), + TensorSlotName::OUTPUT, + c2, }, }, }; - auto mk_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { - return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), }; - }; - auto mk_value = [&](size_t src_node_id, - TensorSlotName src_slot_name, - MappedOperatorTaskGroup const &mapping, - std::optional const &use_slot_name) - -> DynamicValueAttrs { - return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{src_node_id}, - src_slot_name, + ParallelTensorSpaceCoordinate mc_input_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + + ParallelTensorSpaceCoordinate mc1_output_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate mc2_output_coord = + mk_pt_coord(0_n, 1_n, 0_n, 0_n); + + MappedOperatorTaskGroup invocation_mapping = MappedOperatorTaskGroup{ + bidict{ + { + mc1, + mk_shard_binding(mc_input_coord, + mc1_output_coord), }, - }}, - /*parallel_tensor_shape=*/std::nullopt, - /*shard_coord=*/std::nullopt, - /*mapping=*/ - transform(use_slot_name, - [&](TensorSlotName s) { - return get_tensor_bindings_for_slot_name(mapping, s); - }), - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, + { + mc2, + mk_shard_binding(mc_input_coord, + mc2_output_coord), + }, + }, }; - }; - size_t invocation1_id = 20; - - DynamicValueAttrs graph_input1 = - mk_value(0, TensorSlotName::OUTPUT, invocation_mapping, std::nullopt); - DynamicValueAttrs graph_input1_use = mk_value( - 0, TensorSlotName::OUTPUT, invocation_mapping, TensorSlotName::INPUT); - DynamicValueAttrs graph_input1_use_diff_vs_copy1 = - mk_value(0, - TensorSlotName::OUTPUT, - invocation_mapping_diff_vs_copy1, - TensorSlotName::INPUT); - DynamicValueAttrs graph_input2 = - mk_value(1, TensorSlotName::OUTPUT, invocation_mapping, std::nullopt); - DynamicValueAttrs graph_input2_use = mk_value( - 1, TensorSlotName::OUTPUT, invocation_mapping, TensorSlotName::WEIGHT); - DynamicValueAttrs invocation1_output1 = mk_value(invocation1_id, - TensorSlotName::OUTPUT_1, - invocation_mapping, - std::nullopt); - DynamicValueAttrs invocation1_output1_src = - mk_value(invocation1_id, - TensorSlotName::OUTPUT_1, - invocation_mapping, - TensorSlotName::OUTPUT_1); - DynamicValueAttrs invocation1_output2 = mk_value(invocation1_id, - TensorSlotName::OUTPUT_2, - invocation_mapping, - std::nullopt); - DynamicValueAttrs invocation1_output2_src = - mk_value(invocation1_id, - TensorSlotName::OUTPUT_2, - invocation_mapping, - TensorSlotName::OUTPUT_2); - - DynamicValueAttrs graph_input1_src_same = mk_value( - 0, TensorSlotName::OUTPUT, input_mapping_same, TensorSlotName::OUTPUT); - DynamicValueAttrs graph_input2_src_same = mk_value( - 1, TensorSlotName::OUTPUT, weight_mapping_same, TensorSlotName::OUTPUT); - - DynamicNodeInvocation input = DynamicNodeInvocation{ + DynamicValueAttrs graph_input_unmapped = + mk_value(0, TensorSlotName::OUTPUT); + DynamicValueAttrs graph_input_use_mapped = + decide_dynamic_value_attrs_mapping( + graph_input_unmapped, + get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::INPUT)); + + DynamicValueAttrs invocation_output_unmapped = + mk_value(invocation_id, TensorSlotName::OUTPUT); + DynamicValueAttrs invocation_output_src_mapped = + decide_dynamic_value_attrs_mapping( + invocation_output_unmapped, + get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT)); + + DynamicNodeInvocation input = DynamicNodeInvocation{ /*inputs=*/{ { mk_slot(TensorSlotName::INPUT), - graph_input1, - }, - { - mk_slot(TensorSlotName::WEIGHT), - graph_input2, + graph_input_unmapped, }, }, - /*node_attrs=*/ - DynamicNodeAttrs{ + /*node_attrs=*/DynamicNodeAttrs{ /*task_type=*/DynamicTaskType::FWD, /*device_coord=*/std::nullopt, /*mapping=*/invocation_mapping, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/ - dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT_1), - invocation1_output1, - }, - { - mk_slot(TensorSlotName::OUTPUT_2), - invocation1_output2, - }, - }, - }; - - DynamicNodeInvocation mapped = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - graph_input1_use, + /*op_attrs=*/TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + 2_p, + }, + }, }, - { - mk_slot(TensorSlotName::WEIGHT), - graph_input2_use, + /*layer_guid=*/dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{invocation_id}, + }, }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/invocation_mapping, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/ - dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, /*per_device_op_state=*/std::nullopt, }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT_1), - invocation1_output1_src, - }, + /*outputs=*/{ { - mk_slot(TensorSlotName::OUTPUT_2), - invocation1_output2_src, + mk_slot(TensorSlotName::OUTPUT), + invocation_output_unmapped, }, }, - }; - - auto mk_copy = [&](DynamicValueAttrs const &src, - DynamicValueAttrs const &dst) { - return DynamicNodeInvocation{ - /*inputs=*/{{mk_slot(TensorSlotName::INPUT), src}}, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{{mk_slot(TensorSlotName::OUTPUT), dst}}, }; - }; - SUBCASE("same mapping, no copies") { - std::unordered_map sources_same{ - {graph_input1, graph_input1_src_same}, - {graph_input2, graph_input2_src_same}}; - - std::unordered_set result = - perform_copy_insertion_for_invocation(input, sources_same); - - std::unordered_set correct = {mapped}; - - CHECK(result.size() == correct.size()); - CHECK(result == correct); - } - - SUBCASE("copy one tensor, one point") { - MappedOperatorTaskGroup input_mapping_copy1 = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_input_coord), - }, - { - mc3, - mk_input_shard_binding(mc2_input_coord), - }, + std::unordered_map unmapped_to_mapped_source_value = { + { + graph_input_unmapped, + decide_dynamic_value_attrs_mapping( + graph_input_unmapped, + OneToMany{ + { + mc_input_coord, + {mc3}, + }, + }) }, - }; - MappedOperatorTaskGroup input_mapping_copy1_diff_vs_use = - MappedOperatorTaskGroup{ - bidict{ - { - mc3, - mk_input_shard_binding(mc2_input_coord), - }, - }, - }; + }; - DynamicValueAttrs graph_input1_src_copy1 = - mk_value(0, - TensorSlotName::OUTPUT, - input_mapping_copy1, - TensorSlotName::OUTPUT); - DynamicValueAttrs graph_input1_src_copy1_diff_vs_use = - mk_value(0, - TensorSlotName::OUTPUT, - input_mapping_copy1_diff_vs_use, - TensorSlotName::OUTPUT); - - std::unordered_map sources_copy1{ - {graph_input1, graph_input1_src_copy1}, - {graph_input2, graph_input2_src_same}}; - - std::unordered_set result = - perform_copy_insertion_for_invocation(input, sources_copy1); - - std::unordered_set correct = { - mapped, - mk_copy(graph_input1_src_copy1_diff_vs_use, - graph_input1_use_diff_vs_copy1), - }; + std::unordered_set result = copies_for_invocation_inputs( + input, unmapped_to_mapped_source_value); - CHECK(result.size() == correct.size()); - CHECK(result == correct); - } + std::unordered_set correct = {}; - SUBCASE("copy two tensors, two points") { - MappedOperatorTaskGroup input_mapping_copy2 = MappedOperatorTaskGroup{ - bidict{ - { - mc3, - mk_input_shard_binding(mc1_input_coord), - }, - { - mc4, - mk_input_shard_binding(mc2_input_coord), - }, - }, - }; - MappedOperatorTaskGroup weight_mapping_copy2 = MappedOperatorTaskGroup{ - bidict{ - { - mc4, - mk_input_shard_binding(mc1_weight_coord), - }, - { - mc3, - mk_input_shard_binding(mc2_weight_coord), - }, - }, - }; + nlohmann::json result_j = transform(result, dynamic_node_invocation_to_serializable); + nlohmann::json correct_j = transform(correct, dynamic_node_invocation_to_serializable); - DynamicValueAttrs graph_input1_src_copy2 = - mk_value(0, - TensorSlotName::OUTPUT, - input_mapping_copy2, - TensorSlotName::OUTPUT); - DynamicValueAttrs graph_input2_src_copy2 = - mk_value(1, - TensorSlotName::OUTPUT, - weight_mapping_copy2, - TensorSlotName::OUTPUT); - - std::unordered_map sources_copy2{ - {graph_input1, graph_input1_src_copy2}, - {graph_input2, graph_input2_src_copy2}}; - - std::unordered_set result = - perform_copy_insertion_for_invocation(input, sources_copy2); - - std::unordered_set correct = { - mapped, - mk_copy(graph_input1_src_copy2, graph_input1_use), - mk_copy(graph_input2_src_copy2, graph_input2_use), - }; - - CHECK(result.size() == correct.size()); - CHECK(result == correct); + CHECK(result_j == correct_j); } + + // SUBCASE("reduction operator") { + + // auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, + // ParallelTensorSpaceCoordinate const &c2) + // -> OperatorAtomicTaskShardBinding { + // return OperatorAtomicTaskShardBinding{ + // /*tensor_coords=*/{ + // { + // TensorSlotName::INPUT, + // c1, + // }, + // { + // TensorSlotName::OUTPUT, + // c2, + // }, + // }, + // }; + // }; + + // ParallelTensorSpaceCoordinate mc1_input_coord = + // mk_pt_coord(0_n, 0_n, 0_n, 0_n); + // ParallelTensorSpaceCoordinate mc2_input_coord = + // mk_pt_coord(1_n, 0_n, 0_n, 0_n); + + // ParallelTensorSpaceCoordinate mc_output_coord = + // mk_pt_coord(0_n, 0_n, 0_n, 0_n); + + // MappedOperatorTaskGroup invocation_mapping = MappedOperatorTaskGroup{ + // bidict{ + // { + // mc3, + // mk_shard_binding(mc1_input_coord, + // mc_output_coord), + // }, + // }, + // }; + + // DynamicValueAttrs graph_input_unmapped = + // mk_value(0, TensorSlotName::OUTPUT); + // DynamicValueAttrs graph_input_use_mapped = + // decide_dynamic_value_attrs_mapping( + // graph_input_unmapped, + // get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::INPUT)); + + // DynamicValueAttrs invocation_output_unmapped = + // mk_value(invocation_id, TensorSlotName::OUTPUT); + // DynamicValueAttrs invocation_output_src_mapped = + // decide_dynamic_value_attrs_mapping( + // invocation_output_unmapped, + // get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT)); + + // DynamicNodeInvocation input = DynamicNodeInvocation{ + // /*inputs=*/{ + // { + // mk_slot(TensorSlotName::INPUT), + // graph_input_unmapped, + // }, + // }, + // /*node_attrs=*/DynamicNodeAttrs{ + // /*task_type=*/DynamicTaskType::FWD, + // /*device_coord=*/std::nullopt, + // /*mapping=*/invocation_mapping, + // /*op_attrs=*/TrainingOperationAttrs{ + // PCGOperatorAttrs{ + // ReductionAttrs{ + // /*reduction_degree=*/2_p, + // }, + // }, + // }, + // /*layer_guid=*/dynamic_layer_guid_t{ + // parallel_layer_guid_t{ + // Node{invocation_id}, + // }, + // }, + // /*per_device_op_state=*/std::nullopt, + // }, + // /*outputs=*/{ + // { + // mk_slot(TensorSlotName::OUTPUT), + // invocation_output_unmapped, + // }, + // }, + // }; + + // std::unordered_map unmapped_to_mapped_source_value = { + // { + // graph_input_unmapped, + // decide_dynamic_value_attrs_mapping( + // graph_input_unmapped, + // OneToMany{ + // { + // mc1_input_coord, + // {mc1}, + // }, + // { + // mc2_input_coord, + // {mc2}, + // }, + // }) + // }, + // }; + + // std::unordered_set result = copies_for_invocation_inputs( + // input, unmapped_to_mapped_source_value); + + // std::unordered_set correct = {}; + + // nlohmann::json result_j = transform(result, dynamic_node_invocation_to_serializable); + // nlohmann::json correct_j = transform(correct, dynamic_node_invocation_to_serializable); + + // CHECK(result_j == correct_j); + // } } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc new file mode 100644 index 0000000000..9f8aeee726 --- /dev/null +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -0,0 +1,396 @@ +#include +#include "task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h" +#include "utils/containers/require_only_key.h" +#include "op-attrs/ops/element_unary.h" +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("make_dynamic_node_invocation_from_mapped") { + SUBCASE("Replicate") { + MachineSpaceCoordinate gpu0 = MachineSpaceCoordinate{0_n, 0_n, DeviceType::GPU}; + MachineSpaceCoordinate gpu1 = MachineSpaceCoordinate{0_n, 1_n, DeviceType::GPU}; + + ParallelTensorSpaceCoordinate tensor_coord0 = ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{0_n}, + }; + + ParallelTensorSpaceCoordinate tensor_coord1 = ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/1_n, + /*shard_component=*/FFOrdered{0_n}, + }; + + MappedOperatorTaskGroup mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }; + + ParallelTensorShape input_shape = ParallelTensorShape{ + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{8_p, 2_p}, + ShardParallelDim{5_p, 1_p}, + }, + /*replica_dims=*/ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{1_p}, + }, + }, + /*data_type=*/DataType::FLOAT, + }; + + ParallelTensorShape output_shape = [&] { + ParallelTensorShape shape = input_shape; + shape.dims.replica_dims.discard_copy_degree = DiscardCopyDegree{2_p}; + return shape; + }(); + + PCGOperatorAttrs op_attrs = PCGOperatorAttrs{ + ReplicateAttrs{ + 2_p, + }, + }; + + parallel_layer_guid_t layer_guid = parallel_layer_guid_t{Node{0}}; + parallel_tensor_guid_t input_tensor_guid = parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{5}, + TensorSlotName::OUTPUT, + }, + }; + parallel_tensor_guid_t output_tensor_guid = parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{0}, + TensorSlotName::OUTPUT, + }, + }; + + MappedParallelLayerInvocationInfo input = MappedParallelLayerInvocationInfo{ + /*incoming=*/{ + { + TensorSlotName::INPUT, + ParallelTensorInfo{ + /*guid=*/input_tensor_guid, + /*attrs=*/ParallelTensorAttrs{ + /*shape=*/input_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + }, + /*layer_info=*/MappedParallelLayerInfo{ + /*guid=*/layer_guid, + /*attrs=*/ParallelLayerAttrs{ + /*op_attrs=*/op_attrs, + /*name=*/std::nullopt, + }, + /*mapping=*/mapping, + }, + /*outgoing=*/{ + { + TensorSlotName::OUTPUT, + ParallelTensorInfo{ + /*guid=*/output_tensor_guid, + /*attrs=*/ParallelTensorAttrs{ + /*shape=*/output_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + }, + }; + + DynamicNodeInvocation result = make_dynamic_node_invocation_from_mapped(input); + + DynamicNodeInvocation correct = DynamicNodeInvocation{ + /*inputs=*/{ + { + DynamicTensorSlot{ + TensorSlotName::INPUT, + /*slot_tensor_role=*/std::nullopt, + }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, + /*parallel_tensor_shape=*/input_shape, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }, + }, + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/mapping, + /*op_attrs=*/TrainingOperationAttrs{op_attrs}, + /*layer_guid=*/dynamic_layer_guid_t{layer_guid}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ + { + DynamicTensorSlot{ + TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, + /*parallel_tensor_shape=*/output_shape, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }, + }, + } + }; + + CHECK(result == correct); + } + + // SUBCASE("standard op") { + // + // } + } + + // TEST_CASE("make_dynamic_open_dataflow_graph_from_mapped_pcg") { + // positive_int batch_size = 10_p; + // positive_int data_dim = 16_p; + // positive_int hidden_dim = 32_p; + // positive_int output_dim = 1_p; + + // auto make_layer_attrs = [](auto const &op_attrs) -> ParallelLayerAttrs { + // return ParallelLayerAttrs{ + // /*op_attrs=*/PCGOperatorAttrs{op_attrs}, + // /*name=*/std::nullopt, + // }; + // }; + + + // TensorShape output_tensor_shape = TensorShape{ + // TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + // TensorShape label_tensor_shape = TensorShape{ + // TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + // ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + // TensorShape input_tensor_shape = TensorShape{ + // TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + // ParallelLayerAddedResult inputs_layer = + // pcg_add_input_layer(pcg, input_tensor_shape); + // parallel_tensor_guid_t t_input = + // require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); + + // ParallelLayerAddedResult inputs_layer_2 = + // pcg_add_input_layer(pcg, input_tensor_shape); + // parallel_tensor_guid_t t_input_2 = + // require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); + + // ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ + // OperatorType::EW_ADD, + // DataType::FLOAT, + // false, + // false, + // }; + + // ParallelLayerAddedResult add_operator_1 = + // add_parallel_layer(pcg, + // make_layer_attrs(add_attrs), + // { + // { + // TensorSlotName::LHS_INPUT, + // t_input, + // }, + // { + // TensorSlotName::RHS_INPUT, + // t_input_2, + // }, + // }, + // /*weights=*/{}); + + // parallel_tensor_guid_t t_add_1 = + // require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); + + // positive_int replicate_degree = 2_p; + // ReplicateAttrs repl_attrs = ReplicateAttrs{replicate_degree}; + // ParallelLayerAddedResult repl_operator_1 = + // add_parallel_layer(pcg, + // make_layer_attrs(repl_attrs), + // { + // { + // TensorSlotName::INPUT, + // t_add_1, + // }, + // }, + // /*weight=*/{}); + + // parallel_tensor_guid_t t_repl_1 = + // require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); + + // ParallelLayerAddedResult relu_operator_1 = + // add_parallel_layer(pcg, + // make_layer_attrs(make_relu_attrs()), + // /*inputs=*/ + // { + // { + // TensorSlotName::INPUT, + // t_repl_1, + // }, + // }, + // /*weights=*/{}); + + // parallel_tensor_guid_t t_relu_1 = + // require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); + + // MachineSpaceCoordinate gpu0{0_n, 0_n, DeviceType::GPU}; + // MachineSpaceCoordinate gpu1{0_n, 1_n, DeviceType::GPU}; + + // ParallelTensorSpaceCoordinate tensor_coord0{ + // /*sum_component=*/0_n, + // /*discard_copy_component=*/0_n, + // /*shard_component=*/FFOrdered{0_n}}; + // ParallelTensorSpaceCoordinate tensor_coord1{ + // /*sum_component=*/0_n, + // /*discard_copy_component=*/1_n, + // /*shard_component=*/FFOrdered{0_n}}; + + // MappedOperatorTaskGroup input_1_mapping = MappedOperatorTaskGroup{ + // { + // { + // gpu0, + // OperatorAtomicTaskShardBinding{{ + // {TensorSlotName::OUTPUT, tensor_coord0}, + // }}, + // }, + // }, + // }; + + // MappedOperatorTaskGroup input_2_mapping = MappedOperatorTaskGroup{ + // { + // { + // gpu0, + // OperatorAtomicTaskShardBinding{{ + // {TensorSlotName::OUTPUT, tensor_coord0}, + // }}, + // }, + // }, + // }; + + // MappedOperatorTaskGroup add_operator_1_mapping = MappedOperatorTaskGroup{ + // { + // { + // gpu0, + // OperatorAtomicTaskShardBinding{{ + // {TensorSlotName::LHS_INPUT, tensor_coord0}, + // {TensorSlotName::RHS_INPUT, tensor_coord0}, + // {TensorSlotName::OUTPUT, tensor_coord0}, + // }}, + // }, + // }, + // }; + + // MappedOperatorTaskGroup repl_operator_1_mapping = MappedOperatorTaskGroup{ + // { + // { + // gpu0, + // OperatorAtomicTaskShardBinding{{ + // {TensorSlotName::OUTPUT, tensor_coord0}, + // }}, + // }, + // { + // gpu1, + // OperatorAtomicTaskShardBinding{{ + // {TensorSlotName::OUTPUT, tensor_coord1}, + // }}, + // }, + // }, + // }; + + // MappedOperatorTaskGroup relu_operator_1_mapping = MappedOperatorTaskGroup{ + // { + // { + // gpu0, + // OperatorAtomicTaskShardBinding{{ + // {TensorSlotName::INPUT, tensor_coord0}, + // {TensorSlotName::OUTPUT, tensor_coord0}, + // }}, + // }, + // { + // gpu1, + // OperatorAtomicTaskShardBinding{{ + // {TensorSlotName::INPUT, tensor_coord1}, + // {TensorSlotName::OUTPUT, tensor_coord1}, + // }}, + // }, + // }, + // }; + + // MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( + // /*pcg=*/pcg, + // /*mapped_op_task_groups=*/{ + // { + // inputs_layer.parallel_layer, + // input_1_mapping, + // }, + // { + // inputs_layer_2.parallel_layer, + // input_2_mapping, + // }, + // { + // add_operator_1.parallel_layer, + // add_operator_1_mapping, + // }, + // { + // repl_operator_1.parallel_layer, + // repl_operator_1_mapping, + // }, + // { + // relu_operator_1.parallel_layer, + // relu_operator_1_mapping, + // }, + // }); + + + // DynamicOpenDataflowGraph result = make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg); + + // DynamicNodeInvocation input_1_invocation = DynamicNodeInvocation{ + // DynamicNodeAttrs{ + // /*task_type=*/std::nullopt, + // /*device_coord=*/std::nullopt, + // /*mapping=*/input_1_mapping, + // /*op_attrs=*/TrainingOperationAttrs{ + // /*pcg_layer_guid=*/ + // /*per_device_op_state=*/std::nullopt, + // }, + // }; + + // DynamicNodeInvocation input_2_invocation = + // DynamicNodeInvocation add_operator_1_invocation = + // DynamicNodeInvocation repl_operator_1_invocation = + // DynamicNodeInvocation relu_operator_1_invocation = + + // DynamicOpenDataflowGraph correct = dynamic_open_dataflow_graph_from_invocation_set( + // /*invocations=*/{ + + // }, + // }; + // } +} diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index efe21146db..19c21f5f89 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -4,8 +4,11 @@ #include "task-spec/dynamic_graph/dynamic_copy_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" #include "test/utils/doctest/fmt/unordered_set.h" -#include "utils/bidict/algorithms/filter_keys.h" #include +#include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include "op-attrs/ops/element_unary.h" +#include "utils/one_to_many/one_to_many_filter_keys.h" +#include "utils/one_to_many/one_to_many_filter_values.h" using namespace ::FlexFlow; @@ -43,15 +46,18 @@ DynamicTensorSlot mk_slot(TensorSlotName const &slot_name) { DynamicValueAttrs mk_value(size_t src_node_id, TensorSlotName src_slot_name, - bidict - tensor_binding, - std::optional const &shard_coord) { + OneToMany const &tensor_binding, + std::optional const &shard_coord, + std::optional const &role = std::nullopt) { + + OneToMany mapping = tensor_binding; if (shard_coord.has_value()) { - tensor_binding = filter_keys(tensor_binding, - [&](ParallelTensorSpaceCoordinate const &p) { - return p == shard_coord.value(); - }); + mapping = one_to_many_filter_keys(mapping, + [&](ParallelTensorSpaceCoordinate const &p) { + return p == shard_coord.value(); + }); } + return DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ KwargDataflowOutput{ @@ -61,173 +67,148 @@ DynamicValueAttrs }}, /*parallel_tensor_shape=*/std::nullopt, /*shard_coord=*/shard_coord, - /*mapping=*/ - tensor_binding, + /*mapping=*/mapping, /*accessor=*/std::nullopt, - /*role=*/std::nullopt, + /*role=*/role, }; }; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("perform_shard_expansion_for_invocation") { - auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, - ParallelTensorSpaceCoordinate const &c2, - ParallelTensorSpaceCoordinate const &c3, - ParallelTensorSpaceCoordinate const &c4) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ - { - TensorSlotName::INPUT, - c1, - }, - { - TensorSlotName::WEIGHT, - c2, - }, - { - TensorSlotName::OUTPUT_1, - c3, - }, - { - TensorSlotName::OUTPUT_2, - c4, - }, - }, - }; - }; - - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(2_n, 0_n); - - ParallelTensorSpaceCoordinate mc1_input_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc1_weight_coord = - mk_pt_coord(0_n, 1_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc1_output_1_coord = - mk_pt_coord(1_n, 0_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc1_output_2_coord = - mk_pt_coord(3_n, 0_n, 0_n, 0_n); - - ParallelTensorSpaceCoordinate mc2_input_coord = - mk_pt_coord(0_n, 1_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc2_weight_coord = - mk_pt_coord(0_n, 4_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc2_output_1_coord = - mk_pt_coord(1_n, 2_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc2_output_2_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - - MappedOperatorTaskGroup mapped_task_group = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_shard_binding(mc1_input_coord, - mc1_weight_coord, - mc1_output_1_coord, - mc1_output_2_coord), - }, - { - mc2, - mk_shard_binding(mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), - }, - }, - }; - + TEST_CASE("generate_shard_expansion_for_invocation") { auto mk_op_value = [&](size_t src_node_id, TensorSlotName src_slot_name, TensorSlotName use_slot_name, - std::optional const &shard_coord) + MappedOperatorTaskGroup const &mapped_task_group, + std::optional const &shard_coord, + std::optional const &role = std::nullopt) -> DynamicValueAttrs { - bidict + OneToMany tensor_binding = get_tensor_bindings_for_slot_name(mapped_task_group, use_slot_name); - return mk_value(src_node_id, src_slot_name, tensor_binding, shard_coord); + return mk_value(src_node_id, src_slot_name, tensor_binding, shard_coord, role); }; - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_op_value(0, - TensorSlotName::OUTPUT, - TensorSlotName::INPUT, - std::nullopt), - }, - { - mk_slot(TensorSlotName::WEIGHT), - mk_op_value(1, - TensorSlotName::OUTPUT, - TensorSlotName::WEIGHT, - std::nullopt), - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/mapped_task_group, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/ - dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, - /*per_device_op_state=*/std::nullopt, + auto mk_sharding_info = [&](TensorSlotName slot_name, + ParallelTensorSpaceCoordinate const &shard_coord, + MappedOperatorTaskGroup const &mapped_op_task_group, + MachineSpaceCoordinate const &device_coord) + -> std::pair + { + OneToMany + tensor_binding = get_tensor_bindings_for_slot_name(mapped_op_task_group, + slot_name); + return std::pair{ + mk_slot(slot_name), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/shard_coord, + /*mapping=*/one_to_many_filter_values(tensor_binding, + [&](MachineSpaceCoordinate const &c) -> bool { + return device_coord == c; + }), }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT_1), - mk_op_value(20, - TensorSlotName::OUTPUT_1, - TensorSlotName::OUTPUT_1, - std::nullopt), - }, - { - mk_slot(TensorSlotName::OUTPUT_2), - mk_op_value(20, - TensorSlotName::OUTPUT_2, - TensorSlotName::OUTPUT_2, - std::nullopt), + }; + }; + + SUBCASE("standard operator") { + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(2_n, 0_n); + + auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, + ParallelTensorSpaceCoordinate const &c2, + ParallelTensorSpaceCoordinate const &c3, + ParallelTensorSpaceCoordinate const &c4) + -> OperatorAtomicTaskShardBinding { + return OperatorAtomicTaskShardBinding{ + /*tensor_coords=*/{ + { + TensorSlotName::INPUT, + c1, + }, + { + TensorSlotName::WEIGHT, + c2, + }, + { + TensorSlotName::OUTPUT_1, + c3, + }, + { + TensorSlotName::OUTPUT_2, + c4, + }, }, + }; + }; + + ParallelTensorSpaceCoordinate mc1_input_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate mc1_weight_coord = + mk_pt_coord(0_n, 1_n, 2_n, 0_n); + ParallelTensorSpaceCoordinate mc1_output_1_coord = + mk_pt_coord(1_n, 0_n, 0_n, 1_n); + ParallelTensorSpaceCoordinate mc1_output_2_coord = + mk_pt_coord(3_n, 0_n, 0_n, 0_n); + + ParallelTensorSpaceCoordinate mc2_input_coord = + mk_pt_coord(0_n, 1_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate mc2_weight_coord = + mk_pt_coord(0_n, 4_n, 2_n, 0_n); + ParallelTensorSpaceCoordinate mc2_output_1_coord = + mk_pt_coord(1_n, 2_n, 0_n, 1_n); + ParallelTensorSpaceCoordinate mc2_output_2_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + make_relu_attrs(), }, - }; + }; + + MappedOperatorTaskGroup mapped_task_group = MappedOperatorTaskGroup{ + bidict{ + { + mc1, + mk_shard_binding(mc1_input_coord, + mc1_weight_coord, + mc1_output_1_coord, + mc1_output_2_coord), + }, + { + mc2, + mk_shard_binding(mc2_input_coord, + mc2_weight_coord, + mc2_output_1_coord, + mc2_output_2_coord), + }, + }, + }; - std::unordered_set result = - perform_shard_expansion_for_invocation(input); - - auto mk_invocation_shard = - [&](MachineSpaceCoordinate const &device_coord, - ParallelTensorSpaceCoordinate const &input_shard_coord, - ParallelTensorSpaceCoordinate const &weight_shard_coord, - ParallelTensorSpaceCoordinate const &output_1_shard_coord, - ParallelTensorSpaceCoordinate const &output_2_shard_coord) - -> DynamicNodeInvocation { - return DynamicNodeInvocation{ + DynamicNodeInvocation input = DynamicNodeInvocation{ /*inputs=*/{ { mk_slot(TensorSlotName::INPUT), mk_op_value(0, TensorSlotName::OUTPUT, TensorSlotName::INPUT, - input_shard_coord), + mapped_task_group, + std::nullopt), }, { mk_slot(TensorSlotName::WEIGHT), mk_op_value(1, TensorSlotName::OUTPUT, TensorSlotName::WEIGHT, - weight_shard_coord), + mapped_task_group, + std::nullopt), }, }, /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/std::nullopt, - /*device_coord=*/device_coord, + /*device_coord=*/std::nullopt, /*mapping=*/mapped_task_group, - /*op_attrs=*/std::nullopt, + /*op_attrs=*/op_attrs, /*layer_guid=*/ dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, /*per_device_op_state=*/std::nullopt, @@ -239,112 +220,152 @@ TEST_SUITE(FF_TEST_SUITE) { mk_op_value(20, TensorSlotName::OUTPUT_1, TensorSlotName::OUTPUT_1, - output_1_shard_coord), + mapped_task_group, + std::nullopt), }, { mk_slot(TensorSlotName::OUTPUT_2), mk_op_value(20, TensorSlotName::OUTPUT_2, TensorSlotName::OUTPUT_2, - output_2_shard_coord), + mapped_task_group, + std::nullopt), }, }, }; - }; - std::unordered_set correct = { - mk_invocation_shard(mc1, - mc1_input_coord, - mc1_weight_coord, - mc1_output_1_coord, - mc1_output_2_coord), - mk_invocation_shard(mc2, - mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), - }; + std::unordered_set result = + generate_shard_expansion_for_invocation(input); - CHECK(result.size() == correct.size()); - CHECK(result == correct); - } + auto mk_invocation_shard = + [&](MachineSpaceCoordinate const &device_coord, + ParallelTensorSpaceCoordinate const &input_shard_coord, + ParallelTensorSpaceCoordinate const &weight_shard_coord, + ParallelTensorSpaceCoordinate const &output_1_shard_coord, + ParallelTensorSpaceCoordinate const &output_2_shard_coord) + -> DynamicNodeInvocationShardingInfo { + return DynamicNodeInvocationShardingInfo{ + /*device_coord=*/device_coord, + /*value_sharding=*/{ + mk_sharding_info(TensorSlotName::INPUT, input_shard_coord, mapped_task_group, device_coord), + mk_sharding_info(TensorSlotName::WEIGHT, weight_shard_coord, mapped_task_group, device_coord), + mk_sharding_info(TensorSlotName::OUTPUT_1, output_1_shard_coord, mapped_task_group, device_coord), + mk_sharding_info(TensorSlotName::OUTPUT_2, output_2_shard_coord, mapped_task_group, device_coord), + }, + }; + }; - TEST_CASE("perform_shard_expansion_for_invocation (copy)") { - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); - MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); - MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); + std::unordered_set correct = { + mk_invocation_shard(mc1, + mc1_input_coord, + mc1_weight_coord, + mc1_output_1_coord, + mc1_output_2_coord), + mk_invocation_shard(mc2, + mc2_input_coord, + mc2_weight_coord, + mc2_output_1_coord, + mc2_output_2_coord), + }; - ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); + nlohmann::json result_json = result; + nlohmann::json correct_json = correct; - bidict src_binding{ - {pt1, mc1}, - {pt2, mc2}, - }; - bidict dst_binding{ - {pt1, mc3}, - {pt2, mc4}, - }; + CHECK(result.size() == correct.size()); + CHECK(result_json == correct_json); + CHECK(result == correct); + } - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(0, TensorSlotName::OUTPUT, src_binding, std::nullopt), - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), - }, - }, - }; + SUBCASE("copy operator") { + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); + MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); + MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); + + ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); - std::unordered_set result = - perform_shard_expansion_for_invocation(input); + OneToMany src_binding{ + {pt1, {mc1}}, + {pt2, {mc2}}, + }; + + OneToMany dst_binding{ + {pt1, {mc3}}, + {pt2, {mc4}}, + }; - auto mk_invocation_shard = - [&](MachineSpaceCoordinate const &device_coord, - ParallelTensorSpaceCoordinate const &tensor_shard_coord) - -> DynamicNodeInvocation { - DynamicNodeInvocation result = input; - result.inputs = { + DynamicNodeInvocation input = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(0, TensorSlotName::OUTPUT, src_binding, std::nullopt), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ { - mk_slot(TensorSlotName::INPUT), - mk_value( - 0, TensorSlotName::OUTPUT, src_binding, tensor_shard_coord), + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), + }, }, }; - // See perform_shard_expansion_for_copy in shard_expansion.cc for explanation of the choice of device placement. - result.node_attrs.device_coord = device_coord; - result.outputs = { - { + + std::unordered_set result = + generate_shard_expansion_for_invocation(input); + + auto mk_invocation_shard = + [&](MachineSpaceCoordinate const &device_coord, + ParallelTensorSpaceCoordinate const &tensor_shard_coord) + -> DynamicNodeInvocationShardingInfo { + + return DynamicNodeInvocationShardingInfo{ + /*device_coord=*/device_coord, + /*value_sharding=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + DynamicValueAttrsShardingInfo{ + tensor_shard_coord, + one_to_many_filter_keys(src_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == tensor_shard_coord; + }), + }, + }, + { mk_slot(TensorSlotName::OUTPUT), - mk_value( - 20, TensorSlotName::OUTPUT, dst_binding, tensor_shard_coord), + DynamicValueAttrsShardingInfo{ + tensor_shard_coord, + one_to_many_filter_keys(dst_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == tensor_shard_coord; + }), + }, + }, }, + }; }; - return result; - }; - std::unordered_set correct = { - mk_invocation_shard(mc1, pt1), - mk_invocation_shard(mc2, pt2), - }; + std::unordered_set correct = { + mk_invocation_shard(mc1, pt1), + mk_invocation_shard(mc2, pt2), + }; + + nlohmann::json result_json = result; + nlohmann::json correct_json = correct; - CHECK(result.size() == correct.size()); - CHECK(result == correct); + CHECK(result.size() == correct.size()); + CHECK(result_json == correct_json); + CHECK(result == correct); + } } } diff --git a/lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h b/lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h new file mode 100644 index 0000000000..ad43cd52b6 --- /dev/null +++ b/lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h @@ -0,0 +1,91 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ARCHETYPES_JSONABLE_ORDERED_VALUE_TYPE_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ARCHETYPES_JSONABLE_ORDERED_VALUE_TYPE_H + +#include +#include +#include +#include +#include +#include + +namespace FlexFlow { + +template +struct jsonable_ordered_value_type { + jsonable_ordered_value_type() = delete; + + jsonable_ordered_value_type(jsonable_ordered_value_type const &) { + PANIC(); + } + jsonable_ordered_value_type &operator=(jsonable_ordered_value_type const &) { + PANIC(); + } + + jsonable_ordered_value_type(jsonable_ordered_value_type &&) { + PANIC(); + } + jsonable_ordered_value_type &operator=(jsonable_ordered_value_type &&) { + PANIC(); + } + + bool operator==(jsonable_ordered_value_type const &) const { + PANIC(); + } + + bool operator!=(jsonable_ordered_value_type const &) const { + PANIC(); + } + + bool operator<(jsonable_ordered_value_type const &) const { + PANIC(); + } + bool operator>(jsonable_ordered_value_type const &) const { + PANIC(); + } + bool operator<=(jsonable_ordered_value_type const &) const { + PANIC(); + } + bool operator>=(jsonable_ordered_value_type const &) const { + PANIC(); + } +}; + +template +std::string format_as(jsonable_ordered_value_type const &) { + PANIC(); +} + +template +std::ostream &operator<<(std::ostream &s, jsonable_ordered_value_type const &x) { + PANIC(); +} + +} // namespace FlexFlow + +namespace nlohmann { + +template +struct adl_serializer<::FlexFlow::jsonable_ordered_value_type> { + static ::FlexFlow::jsonable_ordered_value_type from_json(json const &) { + PANIC(); + } + + static void to_json(json &, ::FlexFlow::jsonable_ordered_value_type const &) { + PANIC(); + } +}; + +} // namespace nlohmann + +namespace std { + +template +struct hash<::FlexFlow::jsonable_ordered_value_type> { + size_t operator()(::FlexFlow::jsonable_ordered_value_type const &) const { + PANIC(); + }; +}; + +} // namespace std + +#endif diff --git a/lib/utils/include/utils/bidict/algorithms/filter_keys.h b/lib/utils/include/utils/bidict/algorithms/bidict_filter_keys.h similarity index 53% rename from lib/utils/include/utils/bidict/algorithms/filter_keys.h rename to lib/utils/include/utils/bidict/algorithms/bidict_filter_keys.h index 2734dfaeb5..4c2ceb840b 100644 --- a/lib/utils/include/utils/bidict/algorithms/filter_keys.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_filter_keys.h @@ -1,12 +1,12 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTER_KEYS_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTER_KEYS_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTER_KEYS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTER_KEYS_H #include "utils/bidict/bidict.h" namespace FlexFlow { template -bidict filter_keys(bidict const &m, F &&f) { +bidict bidict_filter_keys(bidict const &m, F &&f) { bidict result; for (auto const &kv : m) { if (f(kv.first)) { diff --git a/lib/utils/include/utils/bidict/algorithms/filter_values.h b/lib/utils/include/utils/bidict/algorithms/bidict_filter_values.h similarity index 53% rename from lib/utils/include/utils/bidict/algorithms/filter_values.h rename to lib/utils/include/utils/bidict/algorithms/bidict_filter_values.h index 5817578e79..cb968f2d02 100644 --- a/lib/utils/include/utils/bidict/algorithms/filter_values.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_filter_values.h @@ -1,12 +1,12 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTER_VALUES_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTER_VALUES_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTER_VALUES_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTER_VALUES_H #include "utils/bidict/bidict.h" namespace FlexFlow { template -bidict filter_values(bidict const &m, F &&f) { +bidict bidict_filter_values(bidict const &m, F &&f) { bidict result; for (auto const &kv : m) { if (f(kv.second)) { diff --git a/lib/utils/include/utils/bidict/algorithms/filtrans_keys.h b/lib/utils/include/utils/bidict/algorithms/bidict_filtrans_keys.h similarity index 64% rename from lib/utils/include/utils/bidict/algorithms/filtrans_keys.h rename to lib/utils/include/utils/bidict/algorithms/bidict_filtrans_keys.h index df6495b400..bd9018bd38 100644 --- a/lib/utils/include/utils/bidict/algorithms/filtrans_keys.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_filtrans_keys.h @@ -1,5 +1,5 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTRANS_KEYS_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTRANS_KEYS_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTRANS_KEYS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTRANS_KEYS_H #include "utils/bidict/bidict.h" @@ -9,7 +9,7 @@ template ::value_type> -bidict filtrans_keys(bidict const &m, F &&f) { +bidict bidict_filtrans_keys(bidict const &m, F &&f) { bidict result; for (auto const &[k, v] : m) { std::optional new_k = f(k); diff --git a/lib/utils/include/utils/bidict/algorithms/filtrans_values.h b/lib/utils/include/utils/bidict/algorithms/bidict_filtrans_values.h similarity index 63% rename from lib/utils/include/utils/bidict/algorithms/filtrans_values.h rename to lib/utils/include/utils/bidict/algorithms/bidict_filtrans_values.h index 11180938b8..592440f7f6 100644 --- a/lib/utils/include/utils/bidict/algorithms/filtrans_values.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_filtrans_values.h @@ -1,5 +1,5 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTRANS_VALUES_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_FILTRANS_VALUES_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTRANS_VALUES_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FILTRANS_VALUES_H #include "utils/bidict/bidict.h" @@ -9,7 +9,7 @@ template ::value_type> -bidict filtrans_values(bidict const &m, F &&f) { +bidict bidict_filtrans_values(bidict const &m, F &&f) { bidict result; for (auto const &[k, v] : m) { std::optional new_v = f(v); diff --git a/lib/utils/include/utils/bidict/algorithms/unordered_set_of.h b/lib/utils/include/utils/bidict/algorithms/bidict_unordered_set_of.h similarity index 51% rename from lib/utils/include/utils/bidict/algorithms/unordered_set_of.h rename to lib/utils/include/utils/bidict/algorithms/bidict_unordered_set_of.h index b3df2514cf..251573d441 100644 --- a/lib/utils/include/utils/bidict/algorithms/unordered_set_of.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_unordered_set_of.h @@ -1,5 +1,5 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_UNORDERED_SET_OF_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_UNORDERED_SET_OF_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_UNORDERED_SET_OF_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_UNORDERED_SET_OF_H #include "utils/bidict/bidict.h" #include "utils/hash/pair.h" @@ -7,7 +7,7 @@ namespace FlexFlow { template -std::unordered_set> unordered_set_of(bidict const &c) { +std::unordered_set> bidict_unordered_set_of(bidict const &c) { std::unordered_set> result; for (auto const &lr : c) { diff --git a/lib/utils/include/utils/bidict/algorithms/transform_keys.h b/lib/utils/include/utils/bidict/algorithms/transform_keys.h index 8ecb10c401..1d82464d17 100644 --- a/lib/utils/include/utils/bidict/algorithms/transform_keys.h +++ b/lib/utils/include/utils/bidict/algorithms/transform_keys.h @@ -12,7 +12,7 @@ template transform_keys(bidict const &m, F &&f) { bidict result; for (auto const &kv : m) { - result.equate(f(kv.first), kv.second); + result.equate_strict(f(kv.first), kv.second); } return result; } diff --git a/lib/utils/include/utils/bidict/algorithms/transform_values.h b/lib/utils/include/utils/bidict/algorithms/transform_values.h index ef5b34ebe9..fc8655594e 100644 --- a/lib/utils/include/utils/bidict/algorithms/transform_values.h +++ b/lib/utils/include/utils/bidict/algorithms/transform_values.h @@ -12,7 +12,7 @@ template transform_values(bidict const &m, F &&f) { bidict result; for (auto const &kv : m) { - result.equate({kv.first, f(kv.second)}); + result.equate_strict({kv.first, f(kv.second)}); } return result; } diff --git a/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h b/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h index 2ceb527b96..63d25332d3 100644 --- a/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h +++ b/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_UNSTRUCTURED_RELATION_FROM_BIDICT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_UNSTRUCTURED_RELATION_FROM_BIDICT_H -#include "utils/bidict/algorithms/unordered_set_of.h" +#include "utils/bidict/algorithms/bidict_unordered_set_of.h" #include "utils/bidict/bidict.h" namespace FlexFlow { @@ -9,7 +9,7 @@ namespace FlexFlow { template std::unordered_set> unstructured_relation_from_bidict(bidict const &b) { - return unordered_set_of(b); + return bidict_unordered_set_of(b); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/bidict/bidict.h b/lib/utils/include/utils/bidict/bidict.h index 2d8c5d23a8..57f8d5e213 100644 --- a/lib/utils/include/utils/bidict/bidict.h +++ b/lib/utils/include/utils/bidict/bidict.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_BIDICT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_BIDICT_H -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/fmt/unordered_map.h" #include "utils/hash/unordered_map.h" @@ -13,6 +13,9 @@ #include #include #include +#include "utils/containers/require_same.h" +#include "utils/containers/values.h" +#include "utils/containers/unordered_set_of.h" namespace FlexFlow { @@ -65,11 +68,15 @@ struct bidict { void equate(L const &l, R const &r) { fwd_map.insert({l, r}); bwd_map.insert({r, l}); + + this->check_invariants(); } void equate(std::pair const &lr) { fwd_map.insert(lr); bwd_map.insert({lr.second, lr.first}); + + this->check_invariants(); } void equate_strict(L const &l, R const &r) { @@ -87,15 +94,17 @@ struct bidict { } bool operator==(bidict const &other) const { - bool result = this->fwd_map == other.fwd_map; - assert(result == (this->bwd_map == other.bwd_map)); - return result; + return require_same( + (this->fwd_map == other.fwd_map), + (this->bwd_map == other.bwd_map) + ); } bool operator!=(bidict const &other) const { - bool result = this->fwd_map != other.fwd_map; - assert(result == (this->bwd_map != other.bwd_map)); - return result; + return require_same( + (this->fwd_map != other.fwd_map), + (this->bwd_map != other.bwd_map) + ); } R const &at_l(L const &l) const { @@ -107,11 +116,11 @@ struct bidict { } std::unordered_set left_values() const { - return keys(this->fwd_map); + return unordered_keys(this->fwd_map); } std::unordered_set right_values() const { - return keys(this->bwd_map); + return unordered_keys(this->bwd_map); } std::size_t size() const { @@ -226,6 +235,21 @@ struct bidict { : fwd_map(fwd_map), bwd_map(bwd_map) {} private: + void check_invariants() const { + std::unordered_set fwd_l_vals = unordered_keys(this->fwd_map); + std::unordered_set bwd_l_vals = unordered_set_of(values(this->bwd_map)); + + std::unordered_set bwd_r_vals = unordered_keys(this->bwd_map); + std::unordered_set fwd_r_vals = unordered_set_of(values(this->fwd_map)); + + ASSERT(fwd_l_vals == bwd_l_vals); + ASSERT(fwd_r_vals == bwd_r_vals); + + for (L const &l : fwd_l_vals) { + ASSERT(bwd_map.at(fwd_map.at(l)) == l); + } + } + friend struct bidict; std::unordered_map fwd_map; diff --git a/lib/utils/include/utils/containers/all_of.h b/lib/utils/include/utils/containers/all_of.h index ef5aac1c41..15ed234511 100644 --- a/lib/utils/include/utils/containers/all_of.h +++ b/lib/utils/include/utils/containers/all_of.h @@ -8,7 +8,7 @@ namespace FlexFlow { template -bool all_of(C const &c, F &&f) { +[[nodiscard]] bool all_of(C const &c, F &&f) { for (auto const &v : c) { if (!f(v)) { return false; @@ -18,7 +18,7 @@ bool all_of(C const &c, F &&f) { } template -bool all_of(std::unordered_map const &m, F &&f) { +[[nodiscard]] bool all_of(std::unordered_map const &m, F &&f) { for (auto const &[k, v] : m) { if (!f(k, v)) { return false; @@ -29,7 +29,7 @@ bool all_of(std::unordered_map const &m, F &&f) { } template -bool all_of(std::map const &m, F &&f) { +[[nodiscard]] bool all_of(std::map const &m, F &&f) { for (auto const &[k, v] : m) { if (!f(k, v)) { return false; @@ -39,7 +39,7 @@ bool all_of(std::map const &m, F &&f) { return true; } -bool all_of(std::vector const &); +[[nodiscard]] bool all_of(std::vector const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h b/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h index 06a42327e1..824fe77b39 100644 --- a/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h +++ b/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h @@ -11,8 +11,8 @@ std::unordered_map binary_merge_disjoint_maps(std::unordered_map const &lhs, std::unordered_map const &rhs) { - std::unordered_set lhs_keys = keys(lhs); - std::unordered_set rhs_keys = keys(rhs); + std::unordered_set lhs_keys = unordered_keys(lhs); + std::unordered_set rhs_keys = unordered_keys(rhs); std::unordered_set shared_keys = intersection(lhs_keys, rhs_keys); ASSERT(shared_keys.empty()); diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with.h b/lib/utils/include/utils/containers/binary_merge_maps_with.h index a7c196d061..2d0b57eb81 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_MAPS_WITH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_MAPS_WITH_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/intersection.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/merge_maps_with_right_dominating.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" @@ -17,8 +17,8 @@ std::unordered_map std::unordered_map const &rhs, F &&f) { - std::unordered_set l_keys = keys(lhs); - std::unordered_set r_keys = keys(rhs); + std::unordered_set l_keys = unordered_keys(lhs); + std::unordered_set r_keys = unordered_keys(rhs); std::unordered_set l_only_keys = set_minus(l_keys, r_keys); std::unordered_set r_only_keys = set_minus(r_keys, l_keys); @@ -27,7 +27,7 @@ std::unordered_map std::unordered_map l_only = restrict_keys(lhs, l_only_keys); std::unordered_map r_only = restrict_keys(rhs, r_only_keys); - std::unordered_map merged = generate_map( + std::unordered_map merged = generate_unordered_map( both_keys, [&](K const &k) { return f(lhs.at(k), rhs.at(k)); }); return merge_maps_with_right_dominating(std::vector{ diff --git a/lib/utils/include/utils/containers/filter.h b/lib/utils/include/utils/containers/filter.h index 07f25dc348..85a413c2c7 100644 --- a/lib/utils/include/utils/containers/filter.h +++ b/lib/utils/include/utils/containers/filter.h @@ -44,6 +44,13 @@ std::map filter(std::map const &m, F const &f) { return result; } +template +std::multiset filter(std::multiset const &m, F const &f) { + std::multiset result; + std::copy_if(m.cbegin(), m.cend(), std::inserter(result, result.begin()), f); + return result; +} + template std::unordered_multiset filter(std::unordered_multiset const &m, F const &f) { diff --git a/lib/utils/include/utils/containers/generate_map.h b/lib/utils/include/utils/containers/generate_map.h index 53b2a590c5..08bfc86350 100644 --- a/lib/utils/include/utils/containers/generate_map.h +++ b/lib/utils/include/utils/containers/generate_map.h @@ -5,7 +5,7 @@ #include "utils/containers/vector_of.h" #include "utils/containers/vector_transform.h" #include "utils/type_traits_core.h" -#include +#include namespace FlexFlow { @@ -13,8 +13,8 @@ template , typename V = std::invoke_result_t> -std::unordered_map generate_map(C const &c, F const &f) { - static_assert(is_hashable_v, "Key type should be hashable (but is not)"); +std::map generate_map(C const &c, F &&f) { + static_assert(is_lt_comparable_v, "Key type should be ordered (but is not)"); auto transformed = vector_transform(vector_of(c), [&](K const &k) -> std::pair { diff --git a/lib/utils/include/utils/containers/generate_unordered_map.h b/lib/utils/include/utils/containers/generate_unordered_map.h new file mode 100644 index 0000000000..d57632b7ae --- /dev/null +++ b/lib/utils/include/utils/containers/generate_unordered_map.h @@ -0,0 +1,28 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GENERATE_UNORDERED_MAP_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GENERATE_UNORDERED_MAP_H + +#include "utils/containers/get_element_type.h" +#include "utils/containers/vector_of.h" +#include "utils/containers/vector_transform.h" +#include "utils/type_traits_core.h" +#include + +namespace FlexFlow { + +template , + typename V = std::invoke_result_t> +std::unordered_map generate_unordered_map(C const &c, F &&f) { + static_assert(is_hashable_v, "Key type should be hashable (but is not)"); + + auto transformed = + vector_transform(vector_of(c), [&](K const &k) -> std::pair { + return {k, f(k)}; + }); + return {transformed.cbegin(), transformed.cend()}; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/get_all_assignments.h b/lib/utils/include/utils/containers/get_all_assignments.h index 9981948f47..8f77ffbc24 100644 --- a/lib/utils/include/utils/containers/get_all_assignments.h +++ b/lib/utils/include/utils/containers/get_all_assignments.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GET_ALL_ASSIGNMENTS_H #include "utils/containers/cartesian_product.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/transform.h" #include "utils/containers/unordered_map_from_pairs.h" #include "utils/containers/unordered_set_of.h" @@ -26,7 +26,7 @@ std::unordered_set> get_all_assignments( return {{}}; } - std::vector ordered_keys = vector_of(keys(options_per_key)); + std::vector ordered_keys = vector_of(unordered_keys(options_per_key)); std::vector> ordered_value_option_sets = transform( ordered_keys, [&](K const &k) { return options_per_key.at(k); }); diff --git a/lib/utils/include/utils/containers/index.dox b/lib/utils/include/utils/containers/index.dox index 9b3865dd78..2ffda1cfdf 100644 --- a/lib/utils/include/utils/containers/index.dox +++ b/lib/utils/include/utils/containers/index.dox @@ -9,7 +9,7 @@ Some of the most commonly-used functions are listed below, but you should ideall - \ref containers/transform.h - \ref containers/filter.h - \ref containers/contains.h -- \ref containers/generate_map.h +- \ref containers/generate_unordered_map.h - \ref containers/get_only.h - \ref containers/slice.h - \ref containers/merge_disjoint_maps.h diff --git a/lib/utils/include/utils/containers/is_submapeq_of.h b/lib/utils/include/utils/containers/is_submapeq_of.h index 03cb5ccd78..e50e85e745 100644 --- a/lib/utils/include/utils/containers/is_submapeq_of.h +++ b/lib/utils/include/utils/containers/is_submapeq_of.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_UTILS_INCLUDE_UTILS_CONTAINERS_IS_SUBMAP_H #define _FLEXFLOW_UTILS_INCLUDE_UTILS_CONTAINERS_IS_SUBMAP_H -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/restrict_keys.h" #include @@ -10,7 +10,7 @@ namespace FlexFlow { template bool is_submapeq_of(std::unordered_map const &sub, std::unordered_map const &m) { - return restrict_keys(m, keys(sub)) == sub; + return restrict_keys(m, unordered_keys(sub)) == sub; } } // namespace FlexFlow diff --git a/lib/utils/include/utils/containers/items.h b/lib/utils/include/utils/containers/items.h index 8e3ba95d6c..13e745b17e 100644 --- a/lib/utils/include/utils/containers/items.h +++ b/lib/utils/include/utils/containers/items.h @@ -1,12 +1,12 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ITEMS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ITEMS_H -#include +#include namespace FlexFlow { template -std::unordered_set> +std::set> items(C const &c) { return {c.begin(), c.end()}; } diff --git a/lib/utils/include/utils/containers/keys.h b/lib/utils/include/utils/containers/keys.h index e14612541e..bd080b7087 100644 --- a/lib/utils/include/utils/containers/keys.h +++ b/lib/utils/include/utils/containers/keys.h @@ -3,13 +3,13 @@ #include #include -#include +#include namespace FlexFlow { template -std::unordered_set keys(std::unordered_map const &c) { - std::unordered_set result; +std::set keys(std::unordered_map const &c) { + std::set result; for (auto const &kv : c) { result.insert(kv.first); } @@ -17,8 +17,8 @@ std::unordered_set keys(std::unordered_map const &c) { } template -std::unordered_set keys(std::map const &c) { - std::unordered_set result; +std::set keys(std::map const &c) { + std::set result; for (auto const &kv : c) { result.insert(kv.first); } diff --git a/lib/utils/include/utils/containers/lookup_in_map.h b/lib/utils/include/utils/containers/lookup_in_map.h index 946fc589db..339b13f042 100644 --- a/lib/utils/include/utils/containers/lookup_in_map.h +++ b/lib/utils/include/utils/containers/lookup_in_map.h @@ -1,24 +1,22 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_LOOKUP_IN_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_LOOKUP_IN_MAP_H -#include "utils/containers/contains.h" -#include "utils/containers/keys.h" -#include "utils/exception.h" #include "utils/fmt/unordered_map.h" +#include "utils/containers/contains_key.h" #include #include #include +#include namespace FlexFlow { template -std::function lookup_in_map(std::unordered_map const &map) { - return [map](K const &key) -> V { - if (!contains(keys(map), key)) { - throw mk_runtime_error(fmt::format( - "Key {} is not present in the underlying map {}", key, map)); +std::function lookup_in_map(std::unordered_map const &m) { + return [m](K const &key) -> V { + if (!contains_key(m, key)) { + PANIC("Key {} is not present in the underlying map {}", key, m); } - return map.at(key); + return m.at(key); }; } diff --git a/lib/utils/include/utils/containers/map_from_unordered.h b/lib/utils/include/utils/containers/map_from_unordered.h new file mode 100644 index 0000000000..1451cd1aa8 --- /dev/null +++ b/lib/utils/include/utils/containers/map_from_unordered.h @@ -0,0 +1,18 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_FROM_UNORDERED_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_FROM_UNORDERED_H + +#include +#include + +namespace FlexFlow { + +template +std::map map_from_unordered(std::unordered_map const &u) { + std::map result{u.cbegin(), u.cend()}; + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/map_keys2.h b/lib/utils/include/utils/containers/map_keys2.h index fd848f18d8..da68fe05b4 100644 --- a/lib/utils/include/utils/containers/map_keys2.h +++ b/lib/utils/include/utils/containers/map_keys2.h @@ -19,7 +19,7 @@ std::unordered_map map_keys2(std::unordered_map const &m, result.insert({f(kv.first, kv.second), kv.second}); } - ASSERT(keys(m).size() == keys(result).size(), + ASSERT(m.size() == result.size(), "keys passed to map_keys must be transformed into distinct keys"); return result; diff --git a/lib/utils/include/utils/containers/map_keys_and_values.h b/lib/utils/include/utils/containers/map_keys_and_values.h index 651ffb2aeb..70b7e17103 100644 --- a/lib/utils/include/utils/containers/map_keys_and_values.h +++ b/lib/utils/include/utils/containers/map_keys_and_values.h @@ -1,7 +1,6 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_KEYS_AND_VALUES_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_KEYS_AND_VALUES_H -#include "utils/containers/keys.h" #include #include @@ -21,7 +20,7 @@ std::unordered_map map_keys_and_values( result.insert({fk(kv.first), fv(kv.second)}); } - ASSERT(keys(m).size() == keys(result).size(), + ASSERT(m.size() == result.size(), "keys passed to map_keys must be transformed into distinct keys"); return result; diff --git a/lib/utils/include/utils/containers/map_values.h b/lib/utils/include/utils/containers/map_values.h index bf377b2c93..575fff977e 100644 --- a/lib/utils/include/utils/containers/map_values.h +++ b/lib/utils/include/utils/containers/map_values.h @@ -3,6 +3,7 @@ #include #include +#include namespace FlexFlow { @@ -18,6 +19,18 @@ std::unordered_map map_values(std::unordered_map const &m, F &&f) { return result; } +template > +std::map map_values(std::map const &m, F &&f) { + std::map result; + for (std::pair const &kv : m) { + result.insert(std::pair{kv.first, f(kv.second)}); + } + return result; +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/require_all_of.h b/lib/utils/include/utils/containers/require_all_of.h new file mode 100644 index 0000000000..c085161659 --- /dev/null +++ b/lib/utils/include/utils/containers/require_all_of.h @@ -0,0 +1,33 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REQUIRE_ALL_OF_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REQUIRE_ALL_OF_H + +#include +#include +#include + +namespace FlexFlow { + +template +void require_all_of(C const &c, F &&f) { + for (auto const &v : c) { + f(v); + } +} + +template +void require_all_of(std::unordered_map const &m, F &&f) { + for (auto const &[k, v] : m) { + f(k, v); + } +} + +template +void require_all_of(std::map const &m, F &&f) { + for (auto const &[k, v] : m) { + f(k, v); + } +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/require_only_key.h b/lib/utils/include/utils/containers/require_only_key.h index c63ff4d440..ef142921ec 100644 --- a/lib/utils/include/utils/containers/require_only_key.h +++ b/lib/utils/include/utils/containers/require_only_key.h @@ -4,6 +4,7 @@ #include "utils/containers/contains_key.h" #include #include +#include namespace FlexFlow { @@ -15,6 +16,14 @@ V require_only_key(std::unordered_map const &m, K const &k) { return m.at(k); } +template +V require_only_key(std::map const &m, K const &k) { + ASSERT(m.size() == 1); + ASSERT(contains_key(m, k)); + + return m.at(k); +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/require_same.h b/lib/utils/include/utils/containers/require_same.h index 2f3439db32..2f6251064c 100644 --- a/lib/utils/include/utils/containers/require_same.h +++ b/lib/utils/include/utils/containers/require_same.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REQUIRE_SAME_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REQUIRE_SAME_H -#include "utils/exception.h" +#include #include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/transform.h b/lib/utils/include/utils/containers/transform.h index 14ef782690..bb34b2b5a5 100644 --- a/lib/utils/include/utils/containers/transform.h +++ b/lib/utils/include/utils/containers/transform.h @@ -2,13 +2,15 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_TRANSFORM_H #include "utils/containers/vector_transform.h" -#include "utils/required_core.h" #include #include #include #include #include #include +#include +#include +#include namespace FlexFlow { @@ -17,12 +19,6 @@ std::vector transform(std::vector const &v, F const &f) { return vector_transform(v, f); } -template -auto transform(req const &c, F const &f) - -> decltype(transform(std::declval(), std::declval())) { - return transform(static_cast(c), f); -} - template > std::unordered_set transform(std::unordered_set const &v, F const &f) { std::unordered_set result; @@ -86,8 +82,8 @@ template ::first_type, typename V2 = typename std::invoke_result_t::second_type> -std::unordered_map transform(std::map const &m, F const &f) { - std::unordered_map result; +std::map transform(std::map const &m, F const &f) { + std::map result; for (auto const &[k, v] : m) { result.insert(f(k, v)); } diff --git a/lib/utils/include/utils/containers/unordered_items.h b/lib/utils/include/utils/containers/unordered_items.h new file mode 100644 index 0000000000..1bd8da1498 --- /dev/null +++ b/lib/utils/include/utils/containers/unordered_items.h @@ -0,0 +1,17 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_ITEMS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_ITEMS_H + +#include +#include "utils/hash/pair.h" + +namespace FlexFlow { + +template +std::unordered_set> + unordered_items(C const &c) { + return {c.begin(), c.end()}; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/unordered_keys.h b/lib/utils/include/utils/containers/unordered_keys.h new file mode 100644 index 0000000000..e4b74a6f5e --- /dev/null +++ b/lib/utils/include/utils/containers/unordered_keys.h @@ -0,0 +1,30 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_KEYS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_KEYS_H + +#include +#include +#include + +namespace FlexFlow { + +template +std::unordered_set unordered_keys(std::unordered_map const &c) { + std::unordered_set result; + for (auto const &kv : c) { + result.insert(kv.first); + } + return result; +} + +template +std::unordered_set unordered_keys(std::map const &c) { + std::unordered_set result; + for (auto const &kv : c) { + result.insert(kv.first); + } + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/unordered_map_from_map.h b/lib/utils/include/utils/containers/unordered_map_from_map.h new file mode 100644 index 0000000000..4410e5a767 --- /dev/null +++ b/lib/utils/include/utils/containers/unordered_map_from_map.h @@ -0,0 +1,18 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_MAP_FROM_MAP_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_MAP_FROM_MAP_H + +#include +#include + +namespace FlexFlow { + +template +std::unordered_map unordered_map_from_map(std::map const &u) { + std::unordered_map result{u.cbegin(), u.cend()}; + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/zip_values_strict.h b/lib/utils/include/utils/containers/zip_values_strict.h index 60a7985bc5..1a3ce95eb1 100644 --- a/lib/utils/include/utils/containers/zip_values_strict.h +++ b/lib/utils/include/utils/containers/zip_values_strict.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_H -#include "utils/containers/generate_map.h" -#include "utils/containers/keys.h" +#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/require_same.h" #include #include @@ -14,9 +14,9 @@ std::unordered_map> zip_values_strict(std::unordered_map const &m1, std::unordered_map const &m2) { - ASSERT(keys(m1) == keys(m2)); + ASSERT(unordered_keys(m1) == unordered_keys(m2)); - return generate_map(require_same(keys(m1), keys(m2)), [&](K const &k) { + return generate_unordered_map(require_same(unordered_keys(m1), unordered_keys(m2)), [&](K const &k) { return std::pair{ m1.at(k), m2.at(k), diff --git a/lib/utils/include/utils/containers/zip_values_strict_with.h b/lib/utils/include/utils/containers/zip_values_strict_with.h index 3b0530db8a..5fc4bb7f5b 100644 --- a/lib/utils/include/utils/containers/zip_values_strict_with.h +++ b/lib/utils/include/utils/containers/zip_values_strict_with.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_WITH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_WITH_H -#include "utils/containers/generate_map.h" -#include "utils/containers/keys.h" +#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/require_same.h" #include #include @@ -19,9 +19,9 @@ std::unordered_map std::unordered_map const &m2, F &&f) { - ASSERT(keys(m1) == keys(m2)); + ASSERT(unordered_keys(m1) == unordered_keys(m2)); - return generate_map(require_same(keys(m1), keys(m2)), + return generate_unordered_map(require_same(unordered_keys(m1), unordered_keys(m2)), [&](K const &k) -> Out { return f(m1.at(k), m2.at(k)); }); } diff --git a/lib/utils/include/utils/fmt/map.h b/lib/utils/include/utils/fmt/map.h index 9225040d4d..5b0d41a9cc 100644 --- a/lib/utils/include/utils/fmt/map.h +++ b/lib/utils/include/utils/fmt/map.h @@ -22,10 +22,8 @@ struct formatter< CHECK_FMTABLE(K); CHECK_FMTABLE(V); - std::vector> items = ::FlexFlow::sorted(m); - std::string result = ::FlexFlow::join_strings( - items.cbegin(), items.cend(), ", ", [](std::pair const &p) { + m.cbegin(), m.cend(), ", ", [](std::pair const &p) { return fmt::to_string(p); }); diff --git a/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h index 418346bb36..4bb8373666 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_KWARG_DATAFLOW_GRAPH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_KWARG_DATAFLOW_GRAPH_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/set_union.h" #include "utils/containers/values.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" @@ -26,7 +26,7 @@ struct UnorderedSetKwargDataflowGraph final Node new_node = this->node_source.new_node(); std::unordered_map> outputs = - generate_map( + generate_unordered_map( output_slots, [&](SlotName const &output_slot) -> KwargDataflowOutput { KwargDataflowOutput output = diff --git a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h index 159778bb6d..f05e6cb58a 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h @@ -4,8 +4,8 @@ #include "utils/containers/count.h" #include "utils/containers/enumerate_vector.h" #include "utils/containers/filter.h" -#include "utils/containers/generate_map.h" -#include "utils/containers/keys.h" +#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/map_keys.h" #include "utils/containers/transform.h" #include "utils/containers/without_nullopts.h" @@ -23,6 +23,7 @@ #include "utils/graph/open_dataflow_graph/dataflow_graph_input_source.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge_query.h" +#include "utils/containers/unordered_keys.h" namespace FlexFlow { @@ -81,7 +82,7 @@ struct UnorderedSetLabelledOpenDataflowGraph final } std::unordered_set query_nodes(NodeQuery const &q) const override { - return filter(keys(this->nodes), + return filter(unordered_keys(this->nodes), [&](Node const &n) { return includes(q.nodes, n); }); } @@ -95,7 +96,7 @@ struct UnorderedSetLabelledOpenDataflowGraph final std::unordered_set query_outputs(DataflowOutputQuery const &q) const override { return without_nullopts(transform( - keys(this->values), + unordered_keys(this->values), [&](OpenDataflowValue const &v) -> std::optional { if (!v.has()) { return std::nullopt; @@ -128,12 +129,12 @@ struct UnorderedSetLabelledOpenDataflowGraph final std::unordered_set outputs = get_all_dataflow_outputs(view); std::unordered_set edges = get_edges(view); std::unordered_map labelled_outputs = - generate_map(outputs, + generate_unordered_map(outputs, [&](DataflowOutput const &o) { return view.at(o); }); this->inputs.clear(); this->nodes = - generate_map(nodes, [&](Node const &n) { return view.at(n); }); + generate_unordered_map(nodes, [&](Node const &n) { return view.at(n); }); this->edges = transform( edges, [](DataflowEdge const &e) { return OpenDataflowEdge{e}; }); this->values = map_keys(labelled_outputs, [](DataflowOutput const &o) { @@ -145,14 +146,14 @@ struct UnorderedSetLabelledOpenDataflowGraph final LabelledOpenDataflowGraphView const &view) override { - std::unordered_map nodes = generate_map( + std::unordered_map nodes = generate_unordered_map( get_nodes(view), [&](Node const &n) { return view.at(n); }); std::unordered_set edges = get_edges(view); std::unordered_set inputs = ::FlexFlow::get_open_dataflow_graph_inputs(view); std::unordered_map values = - generate_map(get_open_dataflow_values(view), + generate_unordered_map(get_open_dataflow_values(view), [&](OpenDataflowValue const &v) { return view.at(v); }); this->inputs = inputs; diff --git a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h index 2b20b94c96..ab60e3d364 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h @@ -4,7 +4,7 @@ #include "utils/containers/contains_key.h" #include "utils/containers/enumerate.h" #include "utils/containers/extend.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/map_values.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" @@ -17,6 +17,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge.h" #include "utils/overload.h" +#include "utils/containers/unordered_keys.h" namespace FlexFlow { @@ -68,8 +69,8 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final } std::unordered_map> outputs = - generate_map( - keys(output_labels), + generate_unordered_map( + unordered_keys(output_labels), [&](SlotName const &output_slot) -> KwargDataflowOutput { ValueLabel value_label = output_labels.at(output_slot); @@ -106,7 +107,7 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final } std::unordered_set query_nodes(NodeQuery const &q) const override { - return filter(keys(this->nodes), + return filter(unordered_keys(this->nodes), [&](Node const &n) { return includes(q.nodes, n); }); } @@ -122,7 +123,7 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final std::unordered_set> query_outputs( KwargDataflowOutputQuery const &q) const override { - return filter(keys(this->outputs), + return filter(unordered_keys(this->outputs), [&](KwargDataflowOutput const &output) { return kwarg_dataflow_output_query_includes(q, output); }); @@ -130,7 +131,7 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final std::unordered_set> get_inputs() const override { - return keys(this->graph_inputs); + return unordered_keys(this->graph_inputs); } NodeLabel at(Node const &n) const override { @@ -159,7 +160,7 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final this->graph_inputs.clear(); this->nodes = - generate_map(view_nodes, [&](Node const &n) { return view.at(n); }); + generate_unordered_map(view_nodes, [&](Node const &n) { return view.at(n); }); this->edges = transform(view_edges, @@ -168,7 +169,7 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final return OpenKwargDataflowEdge{e}; }); this->outputs = - generate_map(view_outputs, [&](KwargDataflowOutput const &o) { + generate_unordered_map(view_outputs, [&](KwargDataflowOutput const &o) { return view.at(o); }); } @@ -186,16 +187,16 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final std::unordered_set> view_outputs = get_all_kwarg_dataflow_outputs(view); - this->graph_inputs = generate_map( + this->graph_inputs = generate_unordered_map( view_inputs, [&](KwargDataflowGraphInput const &i) { return view.at(OpenKwargDataflowValue{i}); }); this->nodes = - generate_map(view_nodes, [&](Node const &n) { return view.at(n); }); + generate_unordered_map(view_nodes, [&](Node const &n) { return view.at(n); }); this->edges = view_edges; this->outputs = - generate_map(view_outputs, [&](KwargDataflowOutput const &o) { + generate_unordered_map(view_outputs, [&](KwargDataflowOutput const &o) { return view.at(OpenKwargDataflowValue{o}); }); } diff --git a/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h index 3c66b2c689..9746533c17 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_OPEN_KWARG_DATAFLOW_GRAPH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_OPEN_KWARG_DATAFLOW_GRAPH_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.h" #include "utils/graph/node/node_source.h" #include "utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h" @@ -36,7 +36,7 @@ struct UnorderedSetOpenKwargDataflowGraph final } std::unordered_map> outputs = - generate_map( + generate_unordered_map( output_slots, [&](SlotName const &output_slot) -> KwargDataflowOutput { KwargDataflowOutput output = diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h index 32848f38a6..87ebab4c2d 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_SLOTS_FOR_NODE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_SLOTS_FOR_NODE_H -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.h" namespace FlexFlow { @@ -10,7 +10,7 @@ template std::unordered_set get_incoming_slots_for_node(KwargDataflowGraphView const &g, Node n) { - return keys(get_incoming_kwarg_dataflow_edges_for_node(g, n)); + return unordered_keys(get_incoming_kwarg_dataflow_edges_for_node(g, n)); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h index 372dfed1e8..6cf2b4b8a4 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_OUTGOING_SLOTS_FOR_NODE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_OUTGOING_SLOTS_FOR_NODE_H -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h" namespace FlexFlow { @@ -10,7 +10,7 @@ template std::unordered_set get_outgoing_slots_for_node(KwargDataflowGraphView const &g, Node n) { - return keys(get_outgoing_kwarg_dataflow_outputs_for_node(g, n)); + return unordered_keys(get_outgoing_kwarg_dataflow_outputs_for_node(g, n)); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h index 2115a03cda..502eeab73b 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h @@ -14,14 +14,14 @@ LabelledOpenDataflowGraphData get_graph_data( LabelledOpenDataflowGraphView const &g) { std::unordered_map node_data = - generate_map(get_nodes(g), [&](Node const &n) { return g.at(n); }); + generate_unordered_map(get_nodes(g), [&](Node const &n) { return g.at(n); }); std::unordered_set edges = get_edges(g); std::unordered_set inputs = g.get_inputs(); std::unordered_map value_data = - generate_map(get_open_dataflow_values(g), + generate_unordered_map(get_open_dataflow_values(g), [&](OpenDataflowValue const &v) { return g.at(v); }); return LabelledOpenDataflowGraphData{ diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h index 88132e0a79..580a35b3f7 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h @@ -30,10 +30,10 @@ LabelledOpenDataflowGraphView permute_input_ids( }; std::unordered_map node_labels = - generate_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); + generate_unordered_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); std::unordered_map value_labels = - generate_map(get_open_dataflow_values(permuted), + generate_unordered_map(get_open_dataflow_values(permuted), [&](OpenDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h index 88950635d2..5119587654 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_NODE_IDS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_NODE_IDS_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h" #include "utils/graph/labelled_open_dataflow_graph/labelled_open_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" @@ -37,12 +37,12 @@ LabelledOpenDataflowGraphView permute_node_ids( }; std::unordered_map node_labels = - generate_map(get_nodes(permuted), [&](Node const &new_node) { + generate_unordered_map(get_nodes(permuted), [&](Node const &new_node) { return g.at(old_node_from_new(new_node)); }); std::unordered_map value_labels = - generate_map(get_open_dataflow_values(permuted), + generate_unordered_map(get_open_dataflow_values(permuted), [&](OpenDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h index 92938d7142..fde90497e7 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELS_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h" #include "utils/graph/labelled_open_dataflow_graph/labelled_open_dataflow_graph_view.h" #include "utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h" @@ -27,9 +27,9 @@ LabelledOpenDataflowGraphView rewrite_labels( }; std::unordered_map node_labels = - generate_map(get_nodes(g), get_new_node_label); + generate_unordered_map(get_nodes(g), get_new_node_label); std::unordered_map value_labels = - generate_map(get_open_dataflow_values(g), get_new_value_label); + generate_unordered_map(get_open_dataflow_values(g), get_new_value_label); return with_labelling(g, node_labels, value_labels); } diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h index d60c396274..e98b858019 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" @@ -28,12 +28,12 @@ LabelledOpenKwargDataflowGraphData{ - /*nodes=*/generate_map( + /*nodes=*/generate_unordered_map( get_nodes(g), [&](Node const &n) -> NodeLabel { return g.at(n); }), /*edges=*/get_all_open_kwarg_dataflow_edges(g), /*inputs=*/get_all_kwarg_dataflow_graph_inputs(g), /*outputs=*/ - generate_map( + generate_unordered_map( get_all_open_kwarg_dataflow_values(g), [&](OpenKwargDataflowValue const &v) -> ValueLabel { return g.at(v); }), diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h index d06b96e37f..b6a4366fc2 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h @@ -21,12 +21,12 @@ OpenKwargDataflowGraphData SlotName> const &labelled_data) { OpenKwargDataflowGraphData result = OpenKwargDataflowGraphData{ - /*nodes=*/keys(labelled_data.node_data), + /*nodes=*/unordered_keys(labelled_data.node_data), /*edges=*/labelled_data.edges, /*inputs=*/labelled_data.inputs, /*outputs=*/ filtrans( - keys(labelled_data.value_data), + unordered_keys(labelled_data.value_data), [](OpenKwargDataflowValue const &v) { return v.try_require_internal(); }), diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_input_ids.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_input_ids.h index 223c2e7673..6e1c8bdc57 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_input_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_input_ids.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_INPUT_IDS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_INPUT_IDS_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" @@ -57,11 +57,11 @@ LabelledOpenKwargDataflowGraphView node_labels = - generate_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); + generate_unordered_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); std::unordered_map, ValueLabel> - value_labels = generate_map( + value_labels = generate_unordered_map( get_all_open_kwarg_dataflow_values(permuted), [&](OpenKwargDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h index 06728949df..d89c935d52 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_NODE_IDS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_NODE_IDS_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms/new_node.dtg.h" @@ -56,13 +56,13 @@ LabelledOpenKwargDataflowGraphView node_labels = - generate_map(get_nodes(permuted), [&](Node const &new_node) { + generate_unordered_map(get_nodes(permuted), [&](Node const &new_node) { return g.at(old_node_from_new(new_node)); }); std::unordered_map, ValueLabel> - value_labels = generate_map( + value_labels = generate_unordered_map( get_all_open_kwarg_dataflow_values(permuted), [&](OpenKwargDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h index a632cd7b64..d5d1435462 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_LABELS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_LABELS_H -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h" #include "utils/graph/node/algorithms.h" @@ -39,10 +39,10 @@ LabelledOpenKwargDataflowGraphView NewValueLabel { return f(v, g.at(v)); }; std::unordered_map node_labels = - generate_map(get_nodes(g), get_new_node_label); + generate_unordered_map(get_nodes(g), get_new_node_label); std::unordered_map, NewValueLabel> - value_labels = generate_map(get_all_open_kwarg_dataflow_values(g), + value_labels = generate_unordered_map(get_all_open_kwarg_dataflow_values(g), get_new_value_label); return open_kwarg_dataflow_graph_view_with_labelling( g, node_labels, value_labels); diff --git a/lib/utils/include/utils/graph/series_parallel/get_ancestors.h b/lib/utils/include/utils/graph/series_parallel/get_ancestors.h index f8e9f52cb5..e658c04060 100644 --- a/lib/utils/include/utils/graph/series_parallel/get_ancestors.h +++ b/lib/utils/include/utils/graph/series_parallel/get_ancestors.h @@ -2,6 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_SERIAL_PARALLLEL_GET_ANCESTORS_H #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/graph/series_parallel/non_normal_parallel_split.dtg.toml b/lib/utils/include/utils/graph/series_parallel/non_normal_parallel_split.dtg.toml index b4b975bb4e..a262de1a6b 100644 --- a/lib/utils/include/utils/graph/series_parallel/non_normal_parallel_split.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/non_normal_parallel_split.dtg.toml @@ -4,6 +4,7 @@ type = "struct" features = [ "eq", "hash", + "ord", "fmt", ] @@ -16,19 +17,19 @@ post_includes = [ ] includes = [ - "", + "", "", "utils/graph/node/node.dtg.h", ] src_includes = [ "utils/fmt/variant.h", - "utils/fmt/unordered_multiset.h", - "utils/hash/unordered_multiset.h", + "utils/fmt/multiset.h", + "utils/hash/multiset.h", ] [[fields]] name = "children" -type = "std::unordered_multiset>" +type = "std::multiset>" indirect = true diff --git a/lib/utils/include/utils/graph/series_parallel/non_normal_series_split.dtg.toml b/lib/utils/include/utils/graph/series_parallel/non_normal_series_split.dtg.toml index 008e58dc3f..8f0de7082d 100644 --- a/lib/utils/include/utils/graph/series_parallel/non_normal_series_split.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/non_normal_series_split.dtg.toml @@ -5,6 +5,7 @@ features = [ "eq", "hash", "fmt", + "ord", ] fwd_decls = [ @@ -23,6 +24,7 @@ includes = [ src_includes = [ "utils/fmt/variant.h", + "utils/ord/vector.h", "utils/fmt/vector.h", "utils/hash/vector.h", ] diff --git a/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.dtg.toml b/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.dtg.toml index c82e771385..d6289f1d75 100644 --- a/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.dtg.toml @@ -3,6 +3,7 @@ name = "NonNormalSPDecomposition" type = "variant" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/series_parallel/parallel_split.dtg.toml b/lib/utils/include/utils/graph/series_parallel/parallel_split.dtg.toml index a3315d506b..eb907d4d43 100644 --- a/lib/utils/include/utils/graph/series_parallel/parallel_split.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/parallel_split.dtg.toml @@ -3,6 +3,7 @@ name = "ParallelSplit" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -16,18 +17,18 @@ post_includes = [ ] includes = [ - "", + "", "", "utils/graph/node/node.dtg.h", ] src_includes = [ "utils/fmt/variant.h", - "utils/fmt/unordered_multiset.h", - "utils/hash/unordered_multiset.h", + "utils/fmt/multiset.h", + "utils/hash/multiset.h", ] [[fields]] name = "children" -type = "std::unordered_multiset>" +type = "std::multiset>" indirect = true diff --git a/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.dtg.toml b/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.dtg.toml index 4635bdd877..b47a8eabaa 100644 --- a/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.dtg.toml @@ -3,6 +3,7 @@ name = "SeriesParallelDecomposition" type = "variant" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/series_parallel/series_split.dtg.toml b/lib/utils/include/utils/graph/series_parallel/series_split.dtg.toml index e37762a059..cb5753627d 100644 --- a/lib/utils/include/utils/graph/series_parallel/series_split.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/series_split.dtg.toml @@ -3,6 +3,7 @@ name = "SeriesSplit" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -23,6 +24,7 @@ includes = [ src_includes = [ "utils/fmt/variant.h", + "utils/ord/vector.h", "utils/fmt/vector.h", "utils/hash/vector.h", ] diff --git a/lib/utils/include/utils/json/check_is_jsonable.h b/lib/utils/include/utils/json/check_is_jsonable.h index 9d4f65f005..8597e11c22 100644 --- a/lib/utils/include/utils/json/check_is_jsonable.h +++ b/lib/utils/include/utils/json/check_is_jsonable.h @@ -7,9 +7,9 @@ namespace FlexFlow { #define CHECK_IS_JSONABLE(...) \ - static_assert(is_json_serializable<__VA_ARGS__>::value, \ + static_assert(::FlexFlow::is_json_serializable<__VA_ARGS__>::value, \ #__VA_ARGS__ " should be json serializeable"); \ - static_assert(is_json_deserializable<__VA_ARGS__>::value, \ + static_assert(::FlexFlow::is_json_deserializable<__VA_ARGS__>::value, \ #__VA_ARGS__ " should be json deserializeable") } // namespace FlexFlow diff --git a/lib/utils/include/utils/many_to_one/many_to_one.h b/lib/utils/include/utils/many_to_one/many_to_one.h index 2d078eb304..a501a0672c 100644 --- a/lib/utils/include/utils/many_to_one/many_to_one.h +++ b/lib/utils/include/utils/many_to_one/many_to_one.h @@ -1,7 +1,6 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_MANY_TO_ONE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_MANY_TO_ONE_H -#include "utils/containers/keys.h" #include "utils/containers/require_same.h" #include "utils/containers/try_at.h" #include "utils/containers/unordered_set_of.h" @@ -20,6 +19,8 @@ #include #include #include +#include "utils/containers/set_of.h" +#include "utils/containers/unordered_keys.h" namespace FlexFlow { @@ -88,7 +89,7 @@ struct ManyToOne { } std::unordered_set left_values() const { - return keys(this->m_l_to_r); + return unordered_keys(this->m_l_to_r); } std::unordered_set> left_groups() const { @@ -96,7 +97,7 @@ struct ManyToOne { } std::unordered_set right_values() const { - return keys(this->m_r_to_l); + return unordered_keys(this->m_r_to_l); } std::unordered_map const &l_to_r() const { @@ -141,6 +142,22 @@ std::ostream &operator<<(std::ostream &s, ManyToOne const &m) { return (s << fmt::to_string(m)); } +template +std::unordered_set> + unstructured_relation_from_many_to_one(ManyToOne const &many_to_one) { + return unordered_set_of(many_to_one.l_to_r()); +} + +template +ManyToOne many_to_one_from_unstructured_relation( + std::unordered_set> const &relation) { + ManyToOne result; + for (auto const &lr : relation) { + result.insert(lr); + } + return result; +} + } // namespace FlexFlow namespace nlohmann { @@ -151,14 +168,16 @@ struct adl_serializer<::FlexFlow::ManyToOne> { CHECK_IS_JSON_DESERIALIZABLE(L); CHECK_IS_JSON_DESERIALIZABLE(R); - NOT_IMPLEMENTED(); + std::unordered_set> s = j; + + return ::FlexFlow::many_to_one_from_unstructured_relation(s); } static void to_json(json &j, ::FlexFlow::ManyToOne const &m) { CHECK_IS_JSON_SERIALIZABLE(L); CHECK_IS_JSON_SERIALIZABLE(R); - NOT_IMPLEMENTED(); + j = ::FlexFlow::set_of(::FlexFlow::unstructured_relation_from_many_to_one(m)); } }; diff --git a/lib/utils/include/utils/many_to_one/many_to_one_from_unstructured_relation.h b/lib/utils/include/utils/many_to_one/many_to_one_from_unstructured_relation.h deleted file mode 100644 index 171c6c15d6..0000000000 --- a/lib/utils/include/utils/many_to_one/many_to_one_from_unstructured_relation.h +++ /dev/null @@ -1,20 +0,0 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_MANY_TO_ONE_FROM_UNSTRUCTURED_RELATION_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_MANY_TO_ONE_FROM_UNSTRUCTURED_RELATION_H - -#include "utils/many_to_one/many_to_one.h" - -namespace FlexFlow { - -template -ManyToOne many_to_one_from_unstructured_relation( - std::unordered_set> const &relation) { - ManyToOne result; - for (auto const &lr : relation) { - result.insert(lr); - } - return result; -} - -} // namespace FlexFlow - -#endif diff --git a/lib/utils/include/utils/many_to_one/unstructured_relation_from_many_to_one.h b/lib/utils/include/utils/many_to_one/unstructured_relation_from_many_to_one.h deleted file mode 100644 index 676c5efa5d..0000000000 --- a/lib/utils/include/utils/many_to_one/unstructured_relation_from_many_to_one.h +++ /dev/null @@ -1,17 +0,0 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_UNSTRUCTURED_RELATION_FROM_MANY_TO_ONE_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_UNSTRUCTURED_RELATION_FROM_MANY_TO_ONE_H - -#include "utils/containers/unordered_set_of.h" -#include "utils/many_to_one/many_to_one.h" - -namespace FlexFlow { - -template -std::unordered_set> - unstructured_relation_from_many_to_one(ManyToOne const &many_to_one) { - return unordered_set_of(many_to_one.l_to_r()); -} - -} // namespace FlexFlow - -#endif diff --git a/lib/utils/include/utils/nonempty_set/nonempty_set.h b/lib/utils/include/utils/nonempty_set/nonempty_set.h new file mode 100644 index 0000000000..93da743592 --- /dev/null +++ b/lib/utils/include/utils/nonempty_set/nonempty_set.h @@ -0,0 +1,136 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_NONEMPTY_SET_NONEMPTY_SET_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_NONEMPTY_SET_NONEMPTY_SET_H + +#include +#include +#include "utils/hash-utils.h" +#include "utils/hash/set.h" +#include "utils/fmt/set.h" +#include "utils/positive_int/positive_int.h" +#include "utils/containers/unordered_set_of.h" + +namespace FlexFlow { + +template +struct nonempty_set { +public: + nonempty_set() = delete; + + nonempty_set(std::initializer_list const &vs) : raw(vs) { + ASSERT(this->raw.size() > 0); + } + + explicit nonempty_set(std::set const &s) : raw(s) { + ASSERT(this->raw.size() > 0); + } + + bool operator==(nonempty_set const &other) const { + return this->unwrap_as_set() == other.unwrap_as_set(); + } + + bool operator!=(nonempty_set const &other) const { + return this->unwrap_as_set() != other.unwrap_as_set(); + } + + bool operator<(nonempty_set const &other) const { + return this->unwrap_as_set() < other.unwrap_as_set(); + } + + bool operator<=(nonempty_set const &other) const { + return this->unwrap_as_set() <= other.unwrap_as_set(); + } + + bool operator>(nonempty_set const &other) const { + return this->unwrap_as_set() > other.unwrap_as_set(); + } + + bool operator>=(nonempty_set const &other) const { + return this->unwrap_as_set() >= other.unwrap_as_set(); + } + + bool operator==(std::set const &other) const { + return this->unwrap_as_set() == other; + } + + bool operator!=(std::set const &other) const { + return this->unwrap_as_set() != other; + } + + void insert(T const &t) { + this->raw.insert(t); + } + + size_t size() const { + return this->raw.size(); + }; + + positive_int num_elements() const { + return positive_int{this->raw.size()}; + }; + + std::set const &unwrap_as_set() const { + return this->raw; + } + + std::unordered_set unwrap_as_unordered_set() const { + return unordered_set_of(this->raw); + } + + using value_type = T; + + typename std::set::const_iterator begin() const { + return this->raw.cbegin(); + } + + typename std::set::const_iterator cbegin() const { + return this->raw.cbegin(); + } + + typename std::set::const_iterator end() const { + return this->raw.cend(); + } + + typename std::set::const_iterator cend() const { + return this->raw.cend(); + } + +private: + std::set raw; +}; + +template +bool operator==(std::set const &lhs, + nonempty_set const &rhs) { + return lhs == rhs.unwrap_as_set(); +} + +template +bool operator!=(std::set const &lhs, + nonempty_set const &rhs) { + return lhs != rhs.unwrap_as_set(); +} + +template +std::set format_as(nonempty_set const &s) { + return s.unwrap_as_set(); +} + +template +std::ostream &operator<<(std::ostream &s, nonempty_set const &m) { + return (s << fmt::to_string(m)); +} + +} // namespace FlexFlow + +namespace std { + +template +struct hash<::FlexFlow::nonempty_set> { + size_t operator()(::FlexFlow::nonempty_set const &x) const { + return ::FlexFlow::get_std_hash(x.unwrap_as_set()); + }; +}; + +} // namespace std + +#endif diff --git a/lib/utils/include/utils/one_to_many/one_to_many.h b/lib/utils/include/utils/one_to_many/one_to_many.h index 5492ff3f78..d57622b950 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many.h +++ b/lib/utils/include/utils/one_to_many/one_to_many.h @@ -7,23 +7,23 @@ #include "utils/containers/require_same.h" #include "utils/containers/transform.h" #include "utils/containers/try_at.h" -#include "utils/containers/unordered_set_of.h" #include "utils/containers/values.h" #include "utils/exception.h" -#include "utils/fmt/unordered_map.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/map.h" +#include "utils/fmt/set.h" #include "utils/hash-utils.h" #include "utils/hash/tuple.h" -#include "utils/hash/unordered_map.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/map.h" +#include "utils/hash/set.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" -#include "utils/nonempty_unordered_set/nonempty_unordered_set.h" +#include "utils/nonempty_set/nonempty_set.h" #include #include #include -#include -#include +#include +#include +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -54,6 +54,22 @@ struct OneToMany { return this->tie() != other.tie(); } + bool operator<(OneToMany const &other) const { + return this->tie() < other.tie(); + } + + bool operator<=(OneToMany const &other) const { + return this->tie() <= other.tie(); + } + + bool operator>(OneToMany const &other) const { + return this->tie() > other.tie(); + } + + bool operator>=(OneToMany const &other) const { + return this->tie() >= other.tie(); + } + void insert(std::pair const &p) { L l = p.first; R r = p.second; @@ -66,7 +82,7 @@ struct OneToMany { if (contains_key(this->m_l_to_r, l)) { this->m_l_to_r.at(l).insert(r); } else { - this->m_l_to_r.insert({l, nonempty_unordered_set{{r}}}); + this->m_l_to_r.insert({l, nonempty_set{{r}}}); } } else if (found_l.value() == l) { return; @@ -80,14 +96,14 @@ struct OneToMany { } } - std::unordered_set> relation() const { + std::set> relation() const { return transform(items(this->m_r_to_l), [](std::pair const &p) -> std::pair { return {p.second, p.first}; }); } - nonempty_unordered_set const &at_l(L const &l) const { + nonempty_set const &at_l(L const &l) const { return this->m_l_to_r.at(l); } @@ -95,23 +111,23 @@ struct OneToMany { return this->m_r_to_l.at(r); } - std::unordered_set left_values() const { + std::set left_values() const { return keys(this->m_l_to_r); } - std::unordered_set right_values() const { + std::set right_values() const { return keys(this->m_r_to_l); } - std::unordered_set> right_groups() const { - return unordered_set_of(values(this->m_l_to_r)); + std::set> right_groups() const { + return set_of(values(this->m_l_to_r)); } - std::unordered_map> const &l_to_r() const { + std::map> const &l_to_r() const { return this->m_l_to_r; } - std::unordered_map const &r_to_l() const { + std::map const &r_to_l() const { return this->m_r_to_l; } @@ -120,8 +136,8 @@ struct OneToMany { } private: - std::unordered_map> m_l_to_r; - std::unordered_map m_r_to_l; + std::map> m_l_to_r; + std::map m_r_to_l; private: std::tuple @@ -133,7 +149,7 @@ struct OneToMany { }; template -std::unordered_map> +std::map> format_as(OneToMany const &m) { return generate_map(m.left_values(), [&](L const &l) { return m.at_l(l); }); } @@ -143,6 +159,25 @@ std::ostream &operator<<(std::ostream &s, OneToMany const &m) { return (s << fmt::to_string(m)); } +template +std::unordered_set> + unstructured_relation_from_one_to_many(OneToMany const &one_to_many) { + return transform(unordered_set_of(one_to_many.r_to_l()), + [](std::pair const &rl) -> std::pair { + return std::pair{rl.second, rl.first}; + }); +} + +template +OneToMany one_to_many_from_unstructured_relation( + std::unordered_set> const &rel) { + OneToMany result; + for (auto const &lr : rel) { + result.insert(lr); + } + return result; +} + } // namespace FlexFlow namespace nlohmann { @@ -153,14 +188,16 @@ struct adl_serializer<::FlexFlow::OneToMany> { CHECK_IS_JSON_DESERIALIZABLE(L); CHECK_IS_JSON_DESERIALIZABLE(R); - NOT_IMPLEMENTED(); + std::unordered_set> s = j; + + return ::FlexFlow::one_to_many_from_unstructured_relation(s); } static void to_json(json &j, ::FlexFlow::OneToMany const &m) { CHECK_IS_JSON_SERIALIZABLE(L); CHECK_IS_JSON_SERIALIZABLE(R); - NOT_IMPLEMENTED(); + j = ::FlexFlow::set_of(::FlexFlow::unstructured_relation_from_one_to_many(m)); } }; diff --git a/lib/utils/include/utils/one_to_many/one_to_many_from_unstructured_relation.h b/lib/utils/include/utils/one_to_many/one_to_many_filter_keys.h similarity index 50% rename from lib/utils/include/utils/one_to_many/one_to_many_from_unstructured_relation.h rename to lib/utils/include/utils/one_to_many/one_to_many_filter_keys.h index 11c0a767d6..5f31fbb584 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many_from_unstructured_relation.h +++ b/lib/utils/include/utils/one_to_many/one_to_many_filter_keys.h @@ -1,16 +1,17 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_ONE_TO_MANY_FROM_UNSTRUCTURED_RELATION_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_ONE_TO_MANY_FROM_UNSTRUCTURED_RELATION_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_ONE_TO_MANY_FILTER_KEYS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_ONE_TO_MANY_FILTER_KEYS_H #include "utils/one_to_many/one_to_many.h" namespace FlexFlow { -template -OneToMany one_to_many_from_unstructured_relation( - std::unordered_set> const &rel) { +template +OneToMany one_to_many_filter_keys(OneToMany const &m, F &&f) { OneToMany result; - for (auto const &lr : rel) { - result.insert(lr); + for (auto const &kv : unstructured_relation_from_one_to_many(m)) { + if (f(kv.first)) { + result.insert(kv); + } } return result; } diff --git a/lib/utils/include/utils/one_to_many/one_to_many_filter_values.h b/lib/utils/include/utils/one_to_many/one_to_many_filter_values.h new file mode 100644 index 0000000000..4694b06d65 --- /dev/null +++ b/lib/utils/include/utils/one_to_many/one_to_many_filter_values.h @@ -0,0 +1,22 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_ONE_TO_MANY_FILTER_VALUES_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_ONE_TO_MANY_FILTER_VALUES_H + +#include "utils/one_to_many/one_to_many.h" + +namespace FlexFlow { + +template +OneToMany one_to_many_filter_values(OneToMany const &m, F &&f) { + OneToMany result; + for (auto const &kv : unstructured_relation_from_one_to_many(m)) { + if (f(kv.second)) { + result.insert(kv); + } + } + return result; +} + +} // namespace FlexFlow + + +#endif diff --git a/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h b/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h index a9afe98988..050f6ad7dd 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h +++ b/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h @@ -2,7 +2,8 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_ONE_TO_MANY_TRANSFORM_VALUES_H #include "utils/containers/transform.h" -#include "utils/one_to_many/one_to_many_from_unstructured_relation.h" +#include "utils/one_to_many/one_to_many.h" + namespace FlexFlow { @@ -13,7 +14,8 @@ template one_to_many_transform_values(OneToMany const &input, F f) { return one_to_many_from_unstructured_relation(transform( - input.relation(), [&](std::pair const &p) -> std::pair { + unordered_set_of(input.relation()), + [&](std::pair const &p) -> std::pair { return {p.first, f(p.second)}; })); } diff --git a/lib/utils/include/utils/one_to_many/unstructured_relation_from_one_to_many.h b/lib/utils/include/utils/one_to_many/unstructured_relation_from_one_to_many.h deleted file mode 100644 index 02ed225610..0000000000 --- a/lib/utils/include/utils/one_to_many/unstructured_relation_from_one_to_many.h +++ /dev/null @@ -1,21 +0,0 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_UNSTRUCTURED_RELATION_FROM_ONE_TO_MANY_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_UNSTRUCTURED_RELATION_FROM_ONE_TO_MANY_H - -#include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" -#include "utils/one_to_many/one_to_many.h" - -namespace FlexFlow { - -template -std::unordered_set> - unstructured_relation_from_one_to_many(OneToMany const &one_to_many) { - return transform(unordered_set_of(one_to_many.r_to_l()), - [](std::pair const &rl) -> std::pair { - return std::pair{rl.second, rl.first}; - }); -} - -} // namespace FlexFlow - -#endif diff --git a/lib/utils/include/utils/orthotope/dim_coord.h b/lib/utils/include/utils/orthotope/dim_coord.h index 87a05a7315..b9a10d6750 100644 --- a/lib/utils/include/utils/orthotope/dim_coord.h +++ b/lib/utils/include/utils/orthotope/dim_coord.h @@ -3,10 +3,10 @@ #include "utils/containers/all_of.h" #include "utils/containers/contains_key.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/is_subseteq_of.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/map_values.h" #include "utils/containers/product.h" @@ -29,7 +29,7 @@ namespace FlexFlow { template std::unordered_set get_coord_dims(DimCoord const &coord) { - return keys(coord.raw); + return unordered_keys(coord.raw); } template @@ -65,7 +65,7 @@ DimCoord lift_dim_coord(DimCoord const &coord, ASSERT(is_subseteq_of(get_coord_dims(coord), lifted_dims)); return DimCoord{ - generate_map(lifted_dims, + generate_unordered_map(lifted_dims, [&](T const &dim) { if (contains_key(coord.raw, dim)) { return coord.raw.at(dim); diff --git a/lib/utils/include/utils/orthotope/dim_domain.h b/lib/utils/include/utils/orthotope/dim_domain.h index c940745a78..6bd63faeae 100644 --- a/lib/utils/include/utils/orthotope/dim_domain.h +++ b/lib/utils/include/utils/orthotope/dim_domain.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ORTHOTOPE_DIM_DOMAIN_H #include "utils/containers/filter.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" @@ -27,7 +27,7 @@ nonnegative_int dim_domain_num_dims(DimDomain const &domain) { template std::unordered_set get_domain_dims(DimDomain const &domain) { - return keys(domain.dims); + return unordered_keys(domain.dims); } template diff --git a/lib/utils/include/utils/orthotope/dim_projection.dtg.toml b/lib/utils/include/utils/orthotope/dim_projection.dtg.toml index a530adac5d..9133c1a03f 100644 --- a/lib/utils/include/utils/orthotope/dim_projection.dtg.toml +++ b/lib/utils/include/utils/orthotope/dim_projection.dtg.toml @@ -3,6 +3,7 @@ name = "DimProjection" type = "variant" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/orthotope/down_projection.dtg.toml b/lib/utils/include/utils/orthotope/down_projection.dtg.toml index 9a642d2b9f..a83fd04e64 100644 --- a/lib/utils/include/utils/orthotope/down_projection.dtg.toml +++ b/lib/utils/include/utils/orthotope/down_projection.dtg.toml @@ -3,6 +3,7 @@ name = "DownProjection" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/orthotope/down_projection.h b/lib/utils/include/utils/orthotope/down_projection.h index f46a0f16c8..8fb1487ad3 100644 --- a/lib/utils/include/utils/orthotope/down_projection.h +++ b/lib/utils/include/utils/orthotope/down_projection.h @@ -48,7 +48,7 @@ DimCoord compute_down_projection(DownProjection const &projection, output_dims_of_down_projection(projection); return DimCoord{ - generate_map( + generate_unordered_map( output_dims, [&](R const &output_dim) { std::unordered_set src_dims = diff --git a/lib/utils/include/utils/orthotope/eq_projection.dtg.toml b/lib/utils/include/utils/orthotope/eq_projection.dtg.toml index 972952f907..456a44b2d3 100644 --- a/lib/utils/include/utils/orthotope/eq_projection.dtg.toml +++ b/lib/utils/include/utils/orthotope/eq_projection.dtg.toml @@ -3,6 +3,7 @@ name = "EqProjection" type = "struct" features = [ "eq", + "ord", "hash", "fmt", "rapidcheck", diff --git a/lib/utils/include/utils/orthotope/minimal_dim_domain.h b/lib/utils/include/utils/orthotope/minimal_dim_domain.h index 3934e2af62..c9d1214278 100644 --- a/lib/utils/include/utils/orthotope/minimal_dim_domain.h +++ b/lib/utils/include/utils/orthotope/minimal_dim_domain.h @@ -4,8 +4,7 @@ #include "utils/containers/are_disjoint.h" #include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/filtermap_values.h" -#include "utils/containers/generate_map.h" -#include "utils/containers/keys.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/map_values.h" #include "utils/containers/restrict_keys.h" @@ -16,6 +15,7 @@ #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/minimal_dim_domain.dtg.h" #include "utils/orthotope/minimal_orthotope.dtg.h" +#include "utils/containers/unordered_keys.h" namespace FlexFlow { @@ -70,14 +70,14 @@ DimDomain dim_domain_from_minimal_dim_domain( map_values( minimal_dim_domain.dims, [](int_ge_two x) { return x.positive_int_from_int_ge_two(); }), - generate_map(trivial_dims, [](T const &) { return 1_p; })), + generate_unordered_map(trivial_dims, [](T const &) { return 1_p; })), }; } template std::unordered_set get_minimal_domain_dims(MinimalDimDomain const &domain) { - return keys(domain.dims); + return unordered_keys(domain.dims); } template diff --git a/lib/utils/include/utils/orthotope/up_projection.dtg.toml b/lib/utils/include/utils/orthotope/up_projection.dtg.toml index c99e6eec93..7f69c9dd0e 100644 --- a/lib/utils/include/utils/orthotope/up_projection.dtg.toml +++ b/lib/utils/include/utils/orthotope/up_projection.dtg.toml @@ -3,6 +3,7 @@ name = "UpProjection" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/orthotope/up_projection.h b/lib/utils/include/utils/orthotope/up_projection.h index e485419fbb..1e241108e2 100644 --- a/lib/utils/include/utils/orthotope/up_projection.h +++ b/lib/utils/include/utils/orthotope/up_projection.h @@ -24,13 +24,13 @@ UpProjection make_empty_up_projection() { template std::unordered_set input_dims_of_up_projection(UpProjection const &projection) { - return projection.dim_mapping.left_values(); + return unordered_set_of(projection.dim_mapping.left_values()); } template std::unordered_set output_dims_of_up_projection(UpProjection const &projection) { - return projection.dim_mapping.right_values(); + return unordered_set_of(projection.dim_mapping.right_values()); } template diff --git a/lib/utils/src/utils/archetypes/jsonable_ordered_value_type.cc b/lib/utils/src/utils/archetypes/jsonable_ordered_value_type.cc new file mode 100644 index 0000000000..b83da321f7 --- /dev/null +++ b/lib/utils/src/utils/archetypes/jsonable_ordered_value_type.cc @@ -0,0 +1,7 @@ +#include "utils/archetypes/jsonable_ordered_value_type.h" + +namespace FlexFlow { + +template struct jsonable_ordered_value_type<0>; + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/filter_keys.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filter_keys.cc similarity index 58% rename from lib/utils/src/utils/bidict/algorithms/filter_keys.cc rename to lib/utils/src/utils/bidict/algorithms/bidict_filter_keys.cc index 57ef4e873d..5c84ec85b1 100644 --- a/lib/utils/src/utils/bidict/algorithms/filter_keys.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filter_keys.cc @@ -1,4 +1,4 @@ -#include "utils/bidict/algorithms/filter_keys.h" +#include "utils/bidict/algorithms/bidict_filter_keys.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -7,6 +7,6 @@ using K = value_type<0>; using V = value_type<1>; using F = std::function; -template bidict filter_keys(bidict const &, F &&); +template bidict bidict_filter_keys(bidict const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/filter_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filter_values.cc similarity index 57% rename from lib/utils/src/utils/bidict/algorithms/filter_values.cc rename to lib/utils/src/utils/bidict/algorithms/bidict_filter_values.cc index 4cf58037ee..7adf808f85 100644 --- a/lib/utils/src/utils/bidict/algorithms/filter_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filter_values.cc @@ -1,4 +1,4 @@ -#include "utils/bidict/algorithms/filter_values.h" +#include "utils/bidict/algorithms/bidict_filter_values.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -7,6 +7,6 @@ using K = value_type<0>; using V = value_type<1>; using F = std::function; -template bidict filter_values(bidict const &, F &&); +template bidict bidict_filter_values(bidict const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/filtrans_keys.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_keys.cc similarity index 61% rename from lib/utils/src/utils/bidict/algorithms/filtrans_keys.cc rename to lib/utils/src/utils/bidict/algorithms/bidict_filtrans_keys.cc index 1e506b8a51..954558a66a 100644 --- a/lib/utils/src/utils/bidict/algorithms/filtrans_keys.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_keys.cc @@ -1,4 +1,4 @@ -#include "utils/bidict/algorithms/filtrans_keys.h" +#include "utils/bidict/algorithms/bidict_filtrans_keys.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,6 +8,6 @@ using V = value_type<1>; using K2 = value_type<2>; using F = std::function(K)>; -template bidict filtrans_keys(bidict const &, F &&); +template bidict bidict_filtrans_keys(bidict const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/filtrans_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_values.cc similarity index 61% rename from lib/utils/src/utils/bidict/algorithms/filtrans_values.cc rename to lib/utils/src/utils/bidict/algorithms/bidict_filtrans_values.cc index 8d1352196c..303d40575a 100644 --- a/lib/utils/src/utils/bidict/algorithms/filtrans_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_values.cc @@ -1,4 +1,4 @@ -#include "utils/bidict/algorithms/filtrans_values.h" +#include "utils/bidict/algorithms/bidict_filtrans_values.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,6 +8,6 @@ using V = value_type<1>; using V2 = value_type<2>; using F = std::function(V)>; -template bidict filtrans_values(bidict const &, F &&); +template bidict bidict_filtrans_values(bidict const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc b/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc new file mode 100644 index 0000000000..425c1d099c --- /dev/null +++ b/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc @@ -0,0 +1,11 @@ +#include "utils/bidict/algorithms/bidict_unordered_set_of.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +std::unordered_set> bidict_unordered_set_of(bidict const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/unordered_set_of.cc b/lib/utils/src/utils/bidict/algorithms/unordered_set_of.cc deleted file mode 100644 index a0bfa1525e..0000000000 --- a/lib/utils/src/utils/bidict/algorithms/unordered_set_of.cc +++ /dev/null @@ -1,11 +0,0 @@ -#include "utils/bidict/algorithms/unordered_set_of.h" -#include "utils/archetypes/value_type.h" - -namespace FlexFlow { - -using K = value_type<0>; -using V = value_type<1>; - -std::unordered_set> unordered_set_of(bidict const &); - -} // namespace FlexFlow diff --git a/lib/utils/src/utils/cli/cli_parse.cc b/lib/utils/src/utils/cli/cli_parse.cc index 36d5837f9c..8f5f81324c 100644 --- a/lib/utils/src/utils/cli/cli_parse.cc +++ b/lib/utils/src/utils/cli/cli_parse.cc @@ -2,7 +2,7 @@ #include "utils/cli/cli_spec.h" #include "utils/containers/contains.h" #include "utils/containers/enumerate.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" namespace FlexFlow { @@ -27,7 +27,7 @@ tl::expected cli_parse_flag(CLISpec const &cli, tl::expected cli_parse(CLISpec const &cli, std::vector const &args) { CLIParseResult result = CLIParseResult{ - generate_map(cli_get_flag_keys(cli), + generate_unordered_map(cli_get_flag_keys(cli), [](CLIFlagKey const &) { return false; }), {}, }; diff --git a/lib/utils/src/utils/containers/filter.cc b/lib/utils/src/utils/containers/filter.cc index dc11d0dffa..4931d97704 100644 --- a/lib/utils/src/utils/containers/filter.cc +++ b/lib/utils/src/utils/containers/filter.cc @@ -1 +1,45 @@ #include "utils/containers/filter.h" +#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using VT0 = value_type<0>; +using VT1 = value_type<1>; +using OVT0 = ordered_value_type<0>; + +template + std::vector filter(std::vector const &, + std::function const &); + +template + std::unordered_set filter(std::unordered_set const &, + std::function const &); + +template + std::unordered_map + filter(std::unordered_map const &, + std::function const &)> const &); + +template + std::set filter( + std::set const &, + std::function const &); + +template + std::map + filter(std::map const &, + std::function const &)> const &); + +template + std::multiset + filter(std::multiset const &, + std::function const &); + +template + std::unordered_multiset + filter(std::unordered_multiset const &, + std::function const &); + + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/generate_map.cc b/lib/utils/src/utils/containers/generate_map.cc index 54bbe13dc9..57caa3d062 100644 --- a/lib/utils/src/utils/containers/generate_map.cc +++ b/lib/utils/src/utils/containers/generate_map.cc @@ -1 +1,13 @@ #include "utils/containers/generate_map.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = value_type<1>; +using F = std::function; + +template std::map generate_map(std::vector const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/generate_unordered_map.cc b/lib/utils/src/utils/containers/generate_unordered_map.cc new file mode 100644 index 0000000000..2287e45632 --- /dev/null +++ b/lib/utils/src/utils/containers/generate_unordered_map.cc @@ -0,0 +1,13 @@ +#include "utils/containers/generate_unordered_map.h" +#include "utils/archetypes/value_type.h" +#include + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; +using F = std::function; + +template std::unordered_map generate_unordered_map(std::unordered_set const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/group_by.cc b/lib/utils/src/utils/containers/group_by.cc index efd8c2032b..a41ab4dc62 100644 --- a/lib/utils/src/utils/containers/group_by.cc +++ b/lib/utils/src/utils/containers/group_by.cc @@ -4,8 +4,8 @@ namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; using F = std::function; template OneToMany group_by(std::unordered_set const &, F &&); diff --git a/lib/utils/src/utils/containers/is_submapeq_of.cc b/lib/utils/src/utils/containers/is_submapeq_of.cc index 567d94fac5..f8fd627b3d 100644 --- a/lib/utils/src/utils/containers/is_submapeq_of.cc +++ b/lib/utils/src/utils/containers/is_submapeq_of.cc @@ -1 +1,12 @@ #include "utils/containers/is_submapeq_of.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +bool is_submapeq_of(std::unordered_map const &, std::unordered_map const &); + + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/items.cc b/lib/utils/src/utils/containers/items.cc index 3b1e80452a..193cec9ca8 100644 --- a/lib/utils/src/utils/containers/items.cc +++ b/lib/utils/src/utils/containers/items.cc @@ -1 +1,14 @@ #include "utils/containers/items.h" +#include "utils/archetypes/ordered_value_type.h" +#include +#include + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; + +template std::set> items(std::unordered_map const &); +template std::set> items(std::map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/keys.cc b/lib/utils/src/utils/containers/keys.cc index 6c6abadd56..96db33f4c7 100644 --- a/lib/utils/src/utils/containers/keys.cc +++ b/lib/utils/src/utils/containers/keys.cc @@ -1 +1,13 @@ #include "utils/containers/keys.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = value_type<1>; + +template std::set keys(std::unordered_map const &); +template std::set keys(std::map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_from_unordered.cc b/lib/utils/src/utils/containers/map_from_unordered.cc new file mode 100644 index 0000000000..11558af765 --- /dev/null +++ b/lib/utils/src/utils/containers/map_from_unordered.cc @@ -0,0 +1,13 @@ +#include "utils/containers/map_from_unordered.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = value_type<1>; + +template + std::map map_from_unordered(std::unordered_map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_values.cc b/lib/utils/src/utils/containers/map_values.cc index e26035e8b1..e850ecf31e 100644 --- a/lib/utils/src/utils/containers/map_values.cc +++ b/lib/utils/src/utils/containers/map_values.cc @@ -1,5 +1,6 @@ #include "utils/containers/map_values.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -11,4 +12,11 @@ using F = std::function; template std::unordered_map map_values(std::unordered_map const &, F &&); +using KO = ordered_value_type<0>; +using VO = value_type<1>; +using VO2 = value_type<2>; +using FO = std::function; + +template std::map map_values(std::map const &, FO &&); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/require_all_of.cc b/lib/utils/src/utils/containers/require_all_of.cc new file mode 100644 index 0000000000..7fbf48c54f --- /dev/null +++ b/lib/utils/src/utils/containers/require_all_of.cc @@ -0,0 +1,34 @@ +#include "utils/containers/require_all_of.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" +#include +#include + +namespace FlexFlow { + +using T1 = value_type<0>; +using F1 = std::function; + +template void require_all_of(std::vector const &, F1 &&); +template void require_all_of(std::unordered_set const &, F1 &&); +template void require_all_of(std::unordered_multiset const &, F1 &&); + +using T2 = ordered_value_type<0>; +using F2 = std::function; + +template void require_all_of(std::set const &, F2 &&); +template void require_all_of(std::multiset const &, F2 &&); + +using K3 = value_type<0>; +using V3 = value_type<1>; +using F3 = std::function; + +template void require_all_of(std::unordered_map const &, F3 &&); + +using K4 = ordered_value_type<0>; +using V4 = ordered_value_type<1>; +using F4 = std::function; + +template void require_all_of(std::map const &, F4 &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/require_only_key.cc b/lib/utils/src/utils/containers/require_only_key.cc index ac8c201303..26ec81528a 100644 --- a/lib/utils/src/utils/containers/require_only_key.cc +++ b/lib/utils/src/utils/containers/require_only_key.cc @@ -1,5 +1,6 @@ #include "utils/containers/require_only_key.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -8,4 +9,8 @@ using V = value_type<1>; template V require_only_key(std::unordered_map const &, K const &); +using K2 = ordered_value_type<0>; + +template V require_only_key(std::map const &, K2 const &); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/transform.cc b/lib/utils/src/utils/containers/transform.cc index 7cd5a56ed4..55255cbc1e 100644 --- a/lib/utils/src/utils/containers/transform.cc +++ b/lib/utils/src/utils/containers/transform.cc @@ -1 +1,47 @@ #include "utils/containers/transform.h" +#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using In = value_type<0>; +using Out = value_type<1>; +using F = std::function; + +template std::vector transform(std::vector const &, F const &); + +template std::unordered_set transform(std::unordered_set const &, F const &); + +template std::unordered_multiset transform(std::unordered_multiset const &, F const &); + +using In2 = ordered_value_type<0>; +using Out2 = ordered_value_type<1>; +using F2 = std::function; + +template std::set transform(std::set const &, F2 const &); + +template std::multiset transform(std::multiset const &v, F2 const &f); + +using F3 = std::function; + +template std::string transform(std::string const &, F3 const &); + +using K = value_type<3>; +using V = value_type<4>; +using K2 = value_type<5>; +using V2 = value_type<6>; + +template std::unordered_map transform(std::unordered_map const &, + std::function(K const &, V const &)> const &); + +using K3 = ordered_value_type<3>; +using V3 = value_type<4>; +using K4 = ordered_value_type<5>; +using V4 = value_type<6>; + +template std::map transform(std::map const &, + std::function(K3 const &, V3 const &)> const &); + +template std::optional transform(std::optional const &o, + std::function const &); +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_items.cc b/lib/utils/src/utils/containers/unordered_items.cc new file mode 100644 index 0000000000..9b58cfd18e --- /dev/null +++ b/lib/utils/src/utils/containers/unordered_items.cc @@ -0,0 +1,15 @@ +#include "utils/containers/unordered_items.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" +#include +#include + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = value_type<1>; + +template std::unordered_set> unordered_items(std::unordered_map const &); +template std::unordered_set> unordered_items(std::map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_keys.cc b/lib/utils/src/utils/containers/unordered_keys.cc new file mode 100644 index 0000000000..e850b5f460 --- /dev/null +++ b/lib/utils/src/utils/containers/unordered_keys.cc @@ -0,0 +1,13 @@ +#include "utils/containers/unordered_keys.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = value_type<1>; + +template std::unordered_set unordered_keys(std::unordered_map const &); +std::unordered_set unordered_keys(std::map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_map_from_map.cc b/lib/utils/src/utils/containers/unordered_map_from_map.cc new file mode 100644 index 0000000000..a0ffa034b5 --- /dev/null +++ b/lib/utils/src/utils/containers/unordered_map_from_map.cc @@ -0,0 +1,12 @@ +#include "utils/containers/unordered_map_from_map.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = value_type<0>; + +template std::unordered_map unordered_map_from_map(std::map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc index f617c52593..c9889a49e7 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc @@ -1,5 +1,4 @@ #include "utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.h" -#include "utils/containers/generate_map.h" #include "utils/containers/map_keys.h" #include "utils/dot/dot_file.h" #include "utils/dot/dot_html_from_json.h" diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc index 2da1b208f4..5a9ca06c5e 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/get_dominators_map.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/transform.h" #include "utils/containers/values.h" @@ -25,7 +25,7 @@ std::unordered_map> } std::unordered_map> result = - generate_map(get_nodes(g), [&](Node const &) { return get_nodes(g); }); + generate_unordered_map(get_nodes(g), [&](Node const &) { return get_nodes(g); }); while (!queue.empty()) { Node n = queue.front(); queue.pop(); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc index 34cc7fcc6f..42a9830837 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc @@ -1,7 +1,7 @@ #include "utils/graph/digraph/algorithms/get_imm_dominators_map.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/filter_values.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_element_counts.h" #include "utils/containers/get_only.h" #include "utils/containers/keys.h" @@ -27,14 +27,14 @@ std::unordered_map> })); std::unordered_map dominator_counts = get_element_counts(recursive_dominator_list); - std::unordered_set imm_dominators = keys( + std::unordered_set imm_dominators = unordered_keys( filter_values(dominator_counts, [](int count) { return count <= 1; })); - assert(imm_dominators.size() <= 1); + ASSERT(imm_dominators.size() <= 1); return maybe_get_only(imm_dominators); }; - return generate_map(get_nodes(g), get_imm_dominator); + return generate_unordered_map(get_nodes(g), get_imm_dominator); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc index 39523f2ec1..899ed385c7 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/get_imm_post_dominator.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_one_of.h" #include "utils/containers/get_only.h" #include "utils/containers/intersection.h" @@ -33,7 +33,7 @@ std::optional Node contracted_node = get_one_of(nodes); std::unordered_map contraction = - generate_map(nodes, [&](Node const &) { return contracted_node; }); + generate_unordered_map(nodes, [&](Node const &) { return contracted_node; }); return get_imm_post_dominator(apply_contraction(g, contraction), contracted_node); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc index db09dd07d6..e6c9d5e557 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc @@ -2,6 +2,8 @@ #include "utils/containers/group_by.h" #include "utils/containers/map_values.h" #include "utils/containers/set_of.h" +#include "utils/nonempty_unordered_set/nonempty_unordered_set.h" +#include "utils/containers/unordered_map_from_map.h" namespace FlexFlow { @@ -16,14 +18,18 @@ std::unordered_set get_incoming_edges(DiGraphView const &g, std::unordered_map> get_incoming_edges(DiGraphView const &g, std::unordered_set const &ns) { - std::unordered_map> result = - map_values(group_by(g.query_edges(DirectedEdgeQuery{ - query_set::matchall(), - query_set::match_values_in(set_of(ns)), - }), - [](DirectedEdge const &e) { return e.dst; }) - .l_to_r(), - [](nonempty_unordered_set const &s) + + std::map> by_dst = + group_by(g.query_edges(DirectedEdgeQuery{ + query_set::matchall(), + query_set::match_values_in(set_of(ns)), + }), + [](DirectedEdge const &e) { return e.dst; }) + .l_to_r(); + + std::map> result = + map_values(by_dst, + [](nonempty_set const &s) -> std::unordered_set { return s.unwrap_as_unordered_set(); }); @@ -32,7 +38,7 @@ std::unordered_map> result[n]; } - return result; + return unordered_map_from_map(result); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc index c2057472cf..883d8e3725 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc @@ -2,21 +2,26 @@ #include "utils/containers/group_by.h" #include "utils/containers/map_values.h" #include "utils/containers/set_of.h" +#include "utils/nonempty_unordered_set/nonempty_unordered_set.h" +#include "utils/containers/unordered_map_from_map.h" namespace FlexFlow { std::unordered_map> get_outgoing_edges(DiGraphView const &g, std::unordered_set const &ns) { - std::unordered_map> result = - map_values(group_by(g.query_edges(DirectedEdgeQuery{ - query_set::match_values_in(set_of(ns)), - query_set::matchall(), - }), - [](DirectedEdge const &e) { return e.src; }) - .l_to_r(), - [](nonempty_unordered_set const &s) - -> std::unordered_set { + + std::map> by_src = + group_by(g.query_edges(DirectedEdgeQuery{ + query_set::match_values_in(set_of(ns)), + query_set::matchall(), + }), + [](DirectedEdge const &e) { return e.src; }) + .l_to_r(); + + std::map> result = + map_values(by_src, + [](nonempty_set const &s) -> std::unordered_set { return s.unwrap_as_unordered_set(); }); @@ -24,7 +29,7 @@ std::unordered_map> result[n]; } - return result; + return unordered_map_from_map(result); } std::unordered_set get_outgoing_edges(DiGraphView const &g, diff --git a/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc b/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc index 096efd49e9..66c04ec59c 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/is_acyclic.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/digraph/algorithms/get_successors.h" #include "utils/graph/node/algorithms.h" #include @@ -11,7 +11,7 @@ enum class ExplorationStatus { NOT_EXPLORED, BEING_EXPLORED, FULLY_EXPLORED }; bool is_acyclic(DiGraphView const &g) { std::unordered_map status = - generate_map(get_nodes(g), [](Node const &n) { + generate_unordered_map(get_nodes(g), [](Node const &n) { return ExplorationStatus::NOT_EXPLORED; }); diff --git a/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc b/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc index 941c8e8e3e..903ba0c589 100644 --- a/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc +++ b/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc @@ -1,12 +1,12 @@ #include "utils/graph/instances/adjacency_multidigraph.h" #include "utils/containers/contains_key.h" #include "utils/containers/extend.h" -#include "utils/containers/generate_map.h" -#include "utils/containers/keys.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/values.h" #include "utils/graph/multidigraph/algorithms/get_edges.h" #include "utils/graph/node/algorithms.h" #include "utils/hash/unordered_set.h" +#include "utils/containers/unordered_keys.h" namespace FlexFlow { @@ -26,8 +26,8 @@ AdjacencyMultiDiGraph::AdjacencyMultiDiGraph( Node AdjacencyMultiDiGraph::add_node() { Node new_node = this->node_source.new_node(); std::unordered_set all_nodes = - set_union(keys(this->adjacency), {new_node}); - this->adjacency[new_node] = generate_map(all_nodes, [](Node const &) { + set_union(unordered_keys(this->adjacency), {new_node}); + this->adjacency[new_node] = generate_unordered_map(all_nodes, [](Node const &) { return std::unordered_set{}; }); @@ -77,15 +77,15 @@ void AdjacencyMultiDiGraph::remove_edge(MultiDiEdge const &e) { std::unordered_set AdjacencyMultiDiGraph::query_nodes(NodeQuery const &q) const { - return apply_query(q.nodes, keys(this->adjacency)); + return apply_query(q.nodes, unordered_keys(this->adjacency)); } std::unordered_set AdjacencyMultiDiGraph::query_edges(MultiDiEdgeQuery const &q) const { std::unordered_set result; - std::unordered_set srcs = apply_query(q.srcs, keys(this->adjacency)); - std::unordered_set dsts = apply_query(q.dsts, keys(this->adjacency)); + std::unordered_set srcs = apply_query(q.srcs, unordered_keys(this->adjacency)); + std::unordered_set dsts = apply_query(q.dsts, unordered_keys(this->adjacency)); for (Node const &src : srcs) { for (Node const &dst : dsts) { extend(result, this->adjacency.at(src).at(dst)); @@ -108,8 +108,8 @@ void AdjacencyMultiDiGraph::inplace_materialize_from( std::unordered_set nodes = get_nodes(g); std::unordered_set edges = get_edges(g); - this->adjacency = generate_map(nodes, [&](Node const &) { - return generate_map( + this->adjacency = generate_unordered_map(nodes, [&](Node const &) { + return generate_unordered_map( nodes, [&](Node const &) { return std::unordered_set{}; }); }); this->edge_nodes.clear(); diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc index db181fbe73..ddc3f4f7ae 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc @@ -7,6 +7,7 @@ #include "utils/graph/multidigraph/multidiedge_query.dtg.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/query_set.h" +#include "utils/containers/unordered_map_from_map.h" namespace FlexFlow { @@ -28,11 +29,11 @@ std::unordered_map> query_set::match_values_in(set_of(ns)), }; - std::unordered_map> result = map_values( + std::map> result = map_values( group_by(g.query_edges(query), [&](MultiDiEdge const &e) { return g.get_multidiedge_dst(e); }) .l_to_r(), - [](nonempty_unordered_set const &s) + [](nonempty_set const &s) -> std::unordered_set { return s.unwrap_as_unordered_set(); }); @@ -41,7 +42,7 @@ std::unordered_map> result[n]; } - return result; + return unordered_map_from_map(result); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc index 826c03f476..466bb903d9 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc @@ -1,5 +1,5 @@ #include "utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/multidigraph/algorithms/get_directed_edge.h" #include "utils/graph/multidigraph/algorithms/get_edges.h" @@ -7,7 +7,7 @@ namespace FlexFlow { std::unordered_map get_multidiedge_to_diedge_map(MultiDiGraphView const &g) { - return generate_map(get_edges(g), [&](MultiDiEdge const &e) { + return generate_unordered_map(get_edges(g), [&](MultiDiEdge const &e) { return get_directed_edge(g, e); }); } diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc index 28e181ebb9..143e59b3db 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc @@ -5,6 +5,7 @@ #include "utils/graph/multidigraph/algorithms/get_edges.h" #include "utils/graph/node/algorithms.h" #include +#include "utils/containers/unordered_map_from_map.h" namespace FlexFlow { @@ -26,11 +27,11 @@ std::unordered_map> query_set::matchall(), }; - std::unordered_map> result = map_values( + std::map> result = map_values( group_by(g.query_edges(query), [&](MultiDiEdge const &e) { return g.get_multidiedge_src(e); }) .l_to_r(), - [](nonempty_unordered_set const &s) + [](nonempty_set const &s) -> std::unordered_set { return s.unwrap_as_unordered_set(); }); @@ -39,7 +40,7 @@ std::unordered_map> result[n]; } - return result; + return unordered_map_from_map(result); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc index 0228fdd8e9..b68aa4b1d8 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc @@ -1,5 +1,5 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/sorted_by.h" #include "utils/containers/transform.h" #include "utils/graph/dataflow_graph/dataflow_edge_query.h" @@ -45,7 +45,7 @@ std::vector get_incoming_edges(OpenDataflowGraphView const &g, std::unordered_map> get_incoming_edges(OpenDataflowGraphView const &g, std::unordered_set const &ns) { - return generate_map(ns, + return generate_unordered_map(ns, [&](Node const &n) { return get_incoming_edges(g, n); }); } diff --git a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc index 6f1c2ace68..398abb3faf 100644 --- a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc +++ b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc @@ -13,6 +13,7 @@ #include "utils/overload.h" #include #include +#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -43,8 +44,8 @@ BinarySPDecompositionTree from_parallel_child(children[0]), from_parallel_child(children[1])}}; } - auto s1 = unordered_multiset_of(slice(children, 0, children.size() / 2)); - auto s2 = unordered_multiset_of( + auto s1 = multiset_of(slice(children, 0, children.size() / 2)); + auto s2 = multiset_of( slice(children, children.size() / 2, std::nullopt)); return BinarySPDecompositionTree{BinaryParallelSplit{ diff --git a/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc index 6105dda704..b5a07a5a4d 100644 --- a/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc @@ -11,6 +11,8 @@ #include "utils/graph/series_parallel/series_split.dtg.h" #include "utils/overload.h" #include "utils/variant.h" +#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -43,7 +45,7 @@ NonNormalSPDecomposition non_normal_parallel_composition( for (NonNormalSPDecomposition const &sp_comp : sp_compositions) { if (sp_comp.has()) { composition = multiset_union( - composition, sp_comp.get().get_children()); + composition, unordered_multiset_of(sp_comp.get().get_children())); } else if (sp_comp.has()) { composition.insert(sp_comp.get()); } else { @@ -51,7 +53,7 @@ NonNormalSPDecomposition non_normal_parallel_composition( composition.insert(sp_comp.get()); } } - return NonNormalSPDecomposition(NonNormalParallelSplit{composition}); + return NonNormalSPDecomposition(NonNormalParallelSplit{multiset_of(composition)}); } static Node as_non_normal(Node const &n) { @@ -70,11 +72,11 @@ static NonNormalSeriesSplit as_non_normal(SeriesSplit const &s) { static NonNormalParallelSplit as_non_normal(ParallelSplit const &p) { return non_normal_parallel_composition( - transform(p.get_children(), + unordered_multiset_of(transform(p.get_children(), [](std::variant const &child) { return as_non_normal( widen(child)); - })) + }))) .get(); } diff --git a/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc index 3851ca38d9..5eda579f81 100644 --- a/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc @@ -6,6 +6,7 @@ #include "utils/graph/series_parallel/non_normal_sp_decomposition.h" #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/variant.h" +#include "utils/containers/unordered_multiset_of.h" namespace FlexFlow { @@ -41,7 +42,7 @@ static SeriesParallelDecomposition static SeriesParallelDecomposition normalize_sp_decomposition(NonNormalParallelSplit const ¶llel) { - std::unordered_multiset normalized_children = + std::multiset normalized_children = transform(filter_empty(parallel.get_children()), [](std::variant const &child) { return normalize_sp_decomposition( @@ -54,7 +55,7 @@ static SeriesParallelDecomposition if (normalized_children.size() == 1) { return get_only(normalized_children); } - return parallel_composition(normalized_children); + return parallel_composition(unordered_multiset_of(normalized_children)); } SeriesParallelDecomposition diff --git a/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc index e0075c2584..8c9f655fe5 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc @@ -16,6 +16,7 @@ #include "utils/nonnegative_int/nonnegative_int.h" #include "utils/variant.h" #include +#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -31,7 +32,7 @@ struct ToFinalAST { .value(); })}; } else { - return ParallelSplit{unordered_multiset_of(transform( + return ParallelSplit{multiset_of(transform( node.children, [](std::variant const &s) { return narrow>( @@ -137,7 +138,7 @@ SeriesParallelDecomposition parallel_composition( for (SeriesParallelDecomposition const &sp_comp : sp_compositions) { if (sp_comp.has()) { composition = multiset_union(composition, - sp_comp.get().get_children()); + unordered_multiset_of(sp_comp.get().get_children())); } else if (sp_comp.has()) { composition.insert(sp_comp.get()); } else { @@ -145,7 +146,7 @@ SeriesParallelDecomposition parallel_composition( composition.insert(sp_comp.get()); } } - return SeriesParallelDecomposition(ParallelSplit{composition}); + return SeriesParallelDecomposition(ParallelSplit{multiset_of(composition)}); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc b/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc index fc7cad225a..590912e93f 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc @@ -5,6 +5,7 @@ #include "utils/containers/values.h" #include "utils/containers/vector_of.h" #include "utils/fmt/unordered_multiset.h" +#include "utils/fmt/multiset.h" #include "utils/graph/digraph/algorithms/get_edges.h" #include "utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h" #include "utils/graph/digraph/digraph_view.h" @@ -15,6 +16,7 @@ #include "utils/nonnegative_int/nonnegative_int.h" #include "utils/variant.h" #include + namespace FlexFlow { static std::unordered_map diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc index 36b8e7294b..e3e24294b9 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc @@ -147,7 +147,7 @@ static std::unordered_set return filter_out_sync_nodes(forest, node_roles); } -static std::pair, nonempty_unordered_set> +static std::pair, nonempty_set> get_up_and_down_sets( DiGraph const &g, std::unordered_set const &forest, @@ -229,7 +229,7 @@ SeriesParallelDecomposition escribano_sp_ization(DiGraph g) { std::unordered_set forest = get_forest_escribano(sp, handle, component, node_roles); - std::pair, nonempty_unordered_set> + std::pair, nonempty_set> up_down_sets = get_up_and_down_sets(sp, forest, depth_map); std::unordered_set up = up_down_sets.first.unwrap_as_unordered_set(); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc index 7206ec5cda..b22d7d3811 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc @@ -3,7 +3,7 @@ #include "utils/containers/argmin.h" #include "utils/containers/contains.h" #include "utils/containers/filter.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_only.h" #include "utils/containers/intersection.h" #include "utils/containers/is_subseteq_of.h" @@ -233,7 +233,7 @@ static std::unordered_set ASSERT(!candidate_nodes.empty()); std::unordered_map critical_path_costs = - generate_map(candidate_nodes, [&](Node const &node) { + generate_unordered_map(candidate_nodes, [&](Node const &node) { std::unordered_set preds = get_predecessors(g, node); float max_parent_cost = maximum(transform(preds, [&](Node const &pred) { return sp_longest_paths.at(pred); @@ -253,7 +253,7 @@ static std::unordered_set static bool cost_map_is_valid(DiGraphView const &g, std::unordered_map const &cost_map) { - bool has_correct_nodes = get_nodes(g) == keys(cost_map); + bool has_correct_nodes = (get_nodes(g) == unordered_keys(cost_map)); bool has_nonnegative_costs = all_of(values(cost_map), [&](float const &cost) { return cost >= 0.0f; }); return has_correct_nodes && has_nonnegative_costs; diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc index 3c38b23f2b..4ebaf89756 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc @@ -1,6 +1,5 @@ #include "utils/graph/series_parallel/sp_ization/naive_stratum_sync.h" #include "utils/containers/group_by.h" -#include "utils/containers/keys.h" #include "utils/containers/maximum.h" #include "utils/containers/range.h" #include "utils/containers/transform.h" @@ -13,6 +12,7 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h" #include +#include "utils/containers/unordered_keys.h" namespace FlexFlow { @@ -21,7 +21,7 @@ std::vector> std::unordered_map node_to_stratum = get_longest_path_lengths_from_root(g); - std::unordered_set nodes = keys(node_to_stratum); + std::unordered_set nodes = unordered_keys(node_to_stratum); OneToMany strata_to_nodes = group_by(nodes, [&](Node const &n) { return node_to_stratum.at(n); }); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc index 97b8f11ec3..a6d4183a23 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/sp_ization/node_role.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_predecessors.h" #include "utils/graph/digraph/algorithms/get_successors.h" @@ -10,7 +10,7 @@ namespace FlexFlow { std::unordered_map get_initial_node_role_map(DiGraphView const &g) { - return generate_map(get_nodes(g), + return generate_unordered_map(get_nodes(g), [](Node const &) { return NodeRole::PURE; }); } diff --git a/lib/utils/src/utils/many_to_one/exhaustive_relational_join.cc b/lib/utils/src/utils/many_to_one/exhaustive_relational_join.cc index a12de02656..b987153285 100644 --- a/lib/utils/src/utils/many_to_one/exhaustive_relational_join.cc +++ b/lib/utils/src/utils/many_to_one/exhaustive_relational_join.cc @@ -1,11 +1,11 @@ #include "utils/many_to_one/exhaustive_relational_join.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T1 = value_type<0>; -using T2 = value_type<1>; -using T3 = value_type<2>; +using T1 = ordered_value_type<0>; +using T2 = ordered_value_type<1>; +using T3 = ordered_value_type<2>; template ManyToOne exhaustive_relational_join(ManyToOne const &, diff --git a/lib/utils/src/utils/many_to_one/invert_many_to_one.cc b/lib/utils/src/utils/many_to_one/invert_many_to_one.cc index 92570a1c7f..adb63f5fd8 100644 --- a/lib/utils/src/utils/many_to_one/invert_many_to_one.cc +++ b/lib/utils/src/utils/many_to_one/invert_many_to_one.cc @@ -1,10 +1,10 @@ #include "utils/many_to_one/invert_many_to_one.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template OneToMany invert_many_to_one(ManyToOne const &); diff --git a/lib/utils/src/utils/many_to_one/many_to_one.cc b/lib/utils/src/utils/many_to_one/many_to_one.cc index bbb3bfcc14..52ae52153b 100644 --- a/lib/utils/src/utils/many_to_one/many_to_one.cc +++ b/lib/utils/src/utils/many_to_one/many_to_one.cc @@ -1,5 +1,5 @@ #include "utils/many_to_one/many_to_one.h" -#include "utils/archetypes/jsonable_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/rapidcheckable_value_type.h" #include "utils/archetypes/value_type.h" @@ -17,12 +17,18 @@ template std::unordered_map, R> template std::ostream &operator<<(std::ostream &, ManyToOne const &); +template std::unordered_set> + unstructured_relation_from_many_to_one(ManyToOne const &); + +template ManyToOne many_to_one_from_unstructured_relation( + std::unordered_set> const &); + } // namespace FlexFlow namespace nlohmann { -using L = ::FlexFlow::jsonable_value_type<0>; -using R = ::FlexFlow::jsonable_value_type<1>; +using L = ::FlexFlow::jsonable_ordered_value_type<0>; +using R = ::FlexFlow::jsonable_ordered_value_type<1>; template struct adl_serializer<::FlexFlow::ManyToOne>; diff --git a/lib/utils/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc b/lib/utils/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc deleted file mode 100644 index dc03030f20..0000000000 --- a/lib/utils/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc +++ /dev/null @@ -1,12 +0,0 @@ -#include "utils/many_to_one/many_to_one_from_unstructured_relation.h" -#include "utils/archetypes/value_type.h" - -namespace FlexFlow { - -using L = value_type<0>; -using R = value_type<1>; - -template ManyToOne many_to_one_from_unstructured_relation( - std::unordered_set> const &); - -} // namespace FlexFlow diff --git a/lib/utils/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc b/lib/utils/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc deleted file mode 100644 index d89df51b79..0000000000 --- a/lib/utils/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc +++ /dev/null @@ -1,12 +0,0 @@ -#include "utils/many_to_one/unstructured_relation_from_many_to_one.h" -#include "utils/archetypes/value_type.h" - -namespace FlexFlow { - -using L = value_type<0>; -using R = value_type<1>; - -template std::unordered_set> - unstructured_relation_from_many_to_one(ManyToOne const &); - -} // namespace FlexFlow diff --git a/lib/utils/src/utils/nonempty_set/nonempty_set.cc b/lib/utils/src/utils/nonempty_set/nonempty_set.cc new file mode 100644 index 0000000000..1af2951f10 --- /dev/null +++ b/lib/utils/src/utils/nonempty_set/nonempty_set.cc @@ -0,0 +1,23 @@ +#include "utils/nonempty_set/nonempty_set.h" +#include "utils/archetypes/ordered_value_type.h" + +using T = ::FlexFlow::ordered_value_type<0>; + +namespace FlexFlow { + +template struct nonempty_set; + +template bool operator==(std::set const &, nonempty_set const &); + +template bool operator!=(std::set const &, nonempty_set const &); + +template std::set format_as(nonempty_set const &); +template std::ostream &operator<<(std::ostream &, nonempty_set const &); + +} // namespace FlexFlow + +namespace std { + +template struct hash<::FlexFlow::nonempty_set>; + +} // namespace std diff --git a/lib/utils/src/utils/one_to_many/exhaustive_relational_join.cc b/lib/utils/src/utils/one_to_many/exhaustive_relational_join.cc index 6d237732e9..ae8fdc0d99 100644 --- a/lib/utils/src/utils/one_to_many/exhaustive_relational_join.cc +++ b/lib/utils/src/utils/one_to_many/exhaustive_relational_join.cc @@ -1,11 +1,11 @@ #include "utils/one_to_many/exhaustive_relational_join.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T1 = value_type<0>; -using T2 = value_type<1>; -using T3 = value_type<2>; +using T1 = ordered_value_type<0>; +using T2 = ordered_value_type<1>; +using T3 = ordered_value_type<2>; template OneToMany exhaustive_relational_join(OneToMany const &, diff --git a/lib/utils/src/utils/one_to_many/invert_one_to_many.cc b/lib/utils/src/utils/one_to_many/invert_one_to_many.cc index cb911ff60a..45edf29b3c 100644 --- a/lib/utils/src/utils/one_to_many/invert_one_to_many.cc +++ b/lib/utils/src/utils/one_to_many/invert_one_to_many.cc @@ -1,10 +1,10 @@ #include "utils/one_to_many/invert_one_to_many.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template ManyToOne invert_one_to_many(OneToMany const &); diff --git a/lib/utils/src/utils/one_to_many/one_to_many.cc b/lib/utils/src/utils/one_to_many/one_to_many.cc index 158d2e10c9..ce6220c509 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many.cc @@ -1,28 +1,32 @@ #include "utils/one_to_many/one_to_many.h" #include "utils/archetypes/jsonable_value_type.h" #include "utils/archetypes/rapidcheckable_value_type.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" using namespace ::FlexFlow; namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template struct OneToMany; -template std::unordered_map> +template std::map> format_as(OneToMany const &); template std::ostream &operator<<(std::ostream &, OneToMany const &); +template std::unordered_set> + unstructured_relation_from_one_to_many(OneToMany const &); + } // namespace FlexFlow namespace nlohmann { -using L = ::FlexFlow::jsonable_value_type<0>; -using R = ::FlexFlow::jsonable_value_type<1>; +using L = ::FlexFlow::jsonable_ordered_value_type<0>; +using R = ::FlexFlow::jsonable_ordered_value_type<1>; template struct adl_serializer<::FlexFlow::OneToMany>; @@ -40,8 +44,8 @@ template struct Arbitrary<::FlexFlow::OneToMany>; namespace std { -using L = ::FlexFlow::value_type<0>; -using R = ::FlexFlow::value_type<1>; +using L = ::FlexFlow::ordered_value_type<0>; +using R = ::FlexFlow::ordered_value_type<1>; template struct hash>; diff --git a/lib/utils/src/utils/one_to_many/one_to_many_filter_keys.cc b/lib/utils/src/utils/one_to_many/one_to_many_filter_keys.cc new file mode 100644 index 0000000000..a8855e1857 --- /dev/null +++ b/lib/utils/src/utils/one_to_many/one_to_many_filter_keys.cc @@ -0,0 +1,12 @@ +#include "utils/one_to_many/one_to_many_filter_keys.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; +using F = std::function; + +template OneToMany one_to_many_filter_keys(OneToMany const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc b/lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc new file mode 100644 index 0000000000..fa11ccf56f --- /dev/null +++ b/lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc @@ -0,0 +1,13 @@ +#include "utils/one_to_many/one_to_many_filter_values.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +using F = std::function; + +template OneToMany one_to_many_filter_values(OneToMany const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/one_to_many/one_to_many_from_bidict.cc b/lib/utils/src/utils/one_to_many/one_to_many_from_bidict.cc index bd6c976488..3dad3d6a1d 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many_from_bidict.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many_from_bidict.cc @@ -1,10 +1,10 @@ #include "utils/one_to_many/one_to_many_from_bidict.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template OneToMany one_to_many_from_bidict(bidict const &); diff --git a/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc b/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc index 76f0a221c5..124adb20c3 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc @@ -1,10 +1,10 @@ #include "utils/one_to_many/one_to_many_from_l_to_r_mapping.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template OneToMany one_to_many_from_l_to_r_mapping( std::unordered_map> const &); diff --git a/lib/utils/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc b/lib/utils/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc deleted file mode 100644 index 4fdd52ef2e..0000000000 --- a/lib/utils/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc +++ /dev/null @@ -1,12 +0,0 @@ -#include "utils/one_to_many/one_to_many_from_unstructured_relation.h" -#include "utils/archetypes/value_type.h" - -namespace FlexFlow { - -using L = value_type<0>; -using R = value_type<1>; - -template OneToMany one_to_many_from_unstructured_relation( - std::unordered_set> const &); - -} // namespace FlexFlow diff --git a/lib/utils/src/utils/one_to_many/one_to_many_transform_values.cc b/lib/utils/src/utils/one_to_many/one_to_many_transform_values.cc index 141db7f1da..47a4652ce2 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many_transform_values.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many_transform_values.cc @@ -1,11 +1,11 @@ #include "utils/one_to_many/one_to_many_transform_values.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R1 = value_type<1>; -using R2 = value_type<2>; +using L = ordered_value_type<0>; +using R1 = ordered_value_type<1>; +using R2 = ordered_value_type<2>; using F = std::function; template OneToMany one_to_many_transform_values(OneToMany const &, diff --git a/lib/utils/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc b/lib/utils/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc deleted file mode 100644 index 9a48510a3a..0000000000 --- a/lib/utils/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc +++ /dev/null @@ -1,12 +0,0 @@ -#include "utils/one_to_many/unstructured_relation_from_one_to_many.h" -#include "utils/archetypes/value_type.h" - -namespace FlexFlow { - -using L = value_type<0>; -using R = value_type<1>; - -template std::unordered_set> - unstructured_relation_from_one_to_many(OneToMany const &); - -} // namespace FlexFlow diff --git a/lib/utils/src/utils/orthotope/dim_domain_mapping.cc b/lib/utils/src/utils/orthotope/dim_domain_mapping.cc index 03762dadca..bd0f46e3dc 100644 --- a/lib/utils/src/utils/orthotope/dim_domain_mapping.cc +++ b/lib/utils/src/utils/orthotope/dim_domain_mapping.cc @@ -1,9 +1,9 @@ #include "utils/orthotope/dim_domain_mapping.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" -using ::FlexFlow::value_type; -using L = value_type<0>; -using R = value_type<1>; +using ::FlexFlow::ordered_value_type; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; namespace FlexFlow { @@ -32,9 +32,9 @@ template DimDomainMapping DimOrdering const &, DimOrdering const &); -using T1 = value_type<2>; -using T2 = value_type<3>; -using T3 = value_type<4>; +using T1 = ordered_value_type<2>; +using T2 = ordered_value_type<3>; +using T3 = ordered_value_type<4>; template DimDomainMapping compose_dim_domain_mappings(DimDomainMapping const &, diff --git a/lib/utils/src/utils/orthotope/dim_projection.cc b/lib/utils/src/utils/orthotope/dim_projection.cc index fdf0472a36..9fa250fa47 100644 --- a/lib/utils/src/utils/orthotope/dim_projection.cc +++ b/lib/utils/src/utils/orthotope/dim_projection.cc @@ -1,10 +1,11 @@ #include "utils/orthotope/dim_projection.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template DimProjection dim_projection_identity_map(DimDomain const &, @@ -27,9 +28,9 @@ template DimCoord compute_dim_projection(DimProjection const &, DimOrdering const &, DimOrdering const &); -using T1 = value_type<2>; -using T2 = value_type<3>; -using T3 = value_type<4>; +using T1 = ordered_value_type<2>; +using T2 = ordered_value_type<3>; +using T3 = ordered_value_type<4>; template DimProjection right_compose_eq_projection(DimProjection const &, diff --git a/lib/utils/src/utils/orthotope/down_projection.cc b/lib/utils/src/utils/orthotope/down_projection.cc index 73842ecc11..684521a5de 100644 --- a/lib/utils/src/utils/orthotope/down_projection.cc +++ b/lib/utils/src/utils/orthotope/down_projection.cc @@ -1,10 +1,10 @@ #include "utils/orthotope/down_projection.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template DownProjection make_empty_down_projection(); @@ -26,9 +26,9 @@ template void project_dims(DownProjection &, template UpProjection invert_down_projection(DownProjection const &); -using T1 = value_type<2>; -using T2 = value_type<3>; -using T3 = value_type<4>; +using T1 = ordered_value_type<2>; +using T2 = ordered_value_type<3>; +using T3 = ordered_value_type<4>; template DownProjection compose_down_projections(DownProjection const &, diff --git a/lib/utils/src/utils/orthotope/eq_projection.cc b/lib/utils/src/utils/orthotope/eq_projection.cc index 877b3f93ae..a6965dc36e 100644 --- a/lib/utils/src/utils/orthotope/eq_projection.cc +++ b/lib/utils/src/utils/orthotope/eq_projection.cc @@ -1,10 +1,10 @@ #include "utils/orthotope/eq_projection.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template EqProjection make_empty_eq_projection(); @@ -18,9 +18,9 @@ template void project_dims(EqProjection &, L const &, R const &); template EqProjection invert_eq_projection(EqProjection const &); -using T1 = value_type<0>; -using T2 = value_type<1>; -using T3 = value_type<2>; +using T1 = ordered_value_type<0>; +using T2 = ordered_value_type<1>; +using T3 = ordered_value_type<2>; template EqProjection compose_eq_projections(EqProjection const &, diff --git a/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc b/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc index a867abfd0a..5d4c47b491 100644 --- a/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc +++ b/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc @@ -1,9 +1,9 @@ #include "utils/orthotope/minimal_dim_domain_mapping.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" -using ::FlexFlow::value_type; -using L = value_type<0>; -using R = value_type<1>; +using ::FlexFlow::ordered_value_type; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; namespace FlexFlow { @@ -40,9 +40,9 @@ template MinimalDimDomainMapping DimOrdering const &, DimOrdering const &); -using T1 = value_type<2>; -using T2 = value_type<3>; -using T3 = value_type<4>; +using T1 = ordered_value_type<2>; +using T2 = ordered_value_type<3>; +using T3 = ordered_value_type<4>; template MinimalDimDomainMapping compose_minimal_dim_domain_mappings( MinimalDimDomainMapping const &, diff --git a/lib/utils/src/utils/orthotope/up_projection.cc b/lib/utils/src/utils/orthotope/up_projection.cc index 604cec08ed..0c8909dffd 100644 --- a/lib/utils/src/utils/orthotope/up_projection.cc +++ b/lib/utils/src/utils/orthotope/up_projection.cc @@ -1,12 +1,11 @@ #include "utils/orthotope/up_projection.h" #include "utils/archetypes/ordered_value_type.h" -#include "utils/archetypes/value_type.h" namespace FlexFlow { -using T1 = value_type<0>; -using T2 = value_type<1>; -using T3 = value_type<2>; +using T1 = ordered_value_type<0>; +using T2 = ordered_value_type<1>; +using T3 = ordered_value_type<2>; template UpProjection compose_up_projections(UpProjection const &, diff --git a/lib/utils/test/src/utils/bidict/algorithms/filter_keys.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_keys.cc similarity index 64% rename from lib/utils/test/src/utils/bidict/algorithms/filter_keys.cc rename to lib/utils/test/src/utils/bidict/algorithms/bidict_filter_keys.cc index 3c3097fa9b..2a0806324b 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/filter_keys.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_keys.cc @@ -1,20 +1,21 @@ -#include "utils/bidict/algorithms/filter_keys.h" +#include "utils/bidict/algorithms/bidict_filter_keys.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("filter_keys(bidict, F)") { + TEST_CASE("bidict_filter_keys(bidict, F)") { bidict dict = { {1, "one"}, {2, "two"}, }; bidict result = - filter_keys(dict, [](int k) { return k == 1; }); + bidict_filter_keys(dict, [](int k) { return k == 1; }); bidict correct = { {1, "one"}, }; + CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/bidict/algorithms/filter_values.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_values.cc similarity index 61% rename from lib/utils/test/src/utils/bidict/algorithms/filter_values.cc rename to lib/utils/test/src/utils/bidict/algorithms/bidict_filter_values.cc index 54d0bad199..5c67e11c4f 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/filter_values.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_values.cc @@ -1,17 +1,17 @@ -#include "utils/bidict/algorithms/filter_values.h" +#include "utils/bidict/algorithms/bidict_filter_values.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("filter_values(bidict, F") { + TEST_CASE("bidict_filter_values(bidict, F") { bidict dict = { {1, "one"}, {2, "two"}, }; bidict result = - filter_values(dict, [](std::string const &v) { return v == "two"; }); + bidict_filter_values(dict, [](std::string const &v) { return v == "two"; }); bidict correct = { {2, "two"}, }; diff --git a/lib/utils/test/src/utils/bidict/algorithms/filtrans_keys.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_keys.cc similarity index 75% rename from lib/utils/test/src/utils/bidict/algorithms/filtrans_keys.cc rename to lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_keys.cc index 300918f978..a2774440e5 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/filtrans_keys.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_keys.cc @@ -1,17 +1,17 @@ -#include "utils/bidict/algorithms/filtrans_keys.h" +#include "utils/bidict/algorithms/bidict_filtrans_keys.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("filtrans_keys(bidict, F)") { + TEST_CASE("bidict_filtrans_keys") { bidict dict = { {1, "one"}, {2, "two"}, }; bidict result = - filtrans_keys(dict, [](int k) -> std::optional { + bidict_filtrans_keys(dict, [](int k) -> std::optional { if (k == 1) { return std::nullopt; } else { diff --git a/lib/utils/test/src/utils/bidict/algorithms/filtrans_values.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_values.cc similarity index 69% rename from lib/utils/test/src/utils/bidict/algorithms/filtrans_values.cc rename to lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_values.cc index 99aaef114d..8687d539bc 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/filtrans_values.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_values.cc @@ -1,26 +1,28 @@ -#include "utils/bidict/algorithms/filtrans_values.h" +#include "utils/bidict/algorithms/bidict_filtrans_values.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("filtrans_values(bidict, F)") { + TEST_CASE("bidict_filtrans_values") { bidict dict = { {1, "one"}, {2, "two"}, }; bidict result = - filtrans_values(dict, [](std::string const &v) -> std::optional { + bidict_filtrans_values(dict, [](std::string const &v) -> std::optional { if (v == "two") { return std::nullopt; } else { return v.size() + 1; } }); + bidict correct = { {1, 4}, }; + CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/bidict/algorithms/unordered_set_of.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_unordered_set_of.cc similarity index 77% rename from lib/utils/test/src/utils/bidict/algorithms/unordered_set_of.cc rename to lib/utils/test/src/utils/bidict/algorithms/bidict_unordered_set_of.cc index b88b6df0ca..d44a2fe62b 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/unordered_set_of.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_unordered_set_of.cc @@ -1,4 +1,4 @@ -#include "utils/bidict/algorithms/unordered_set_of.h" +#include "utils/bidict/algorithms/bidict_unordered_set_of.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/bidict/bidict.cc b/lib/utils/test/src/utils/bidict/bidict.cc index 1365d04027..f15f15b0fe 100644 --- a/lib/utils/test/src/utils/bidict/bidict.cc +++ b/lib/utils/test/src/utils/bidict/bidict.cc @@ -113,11 +113,13 @@ TEST_SUITE(FF_TEST_SUITE) { bidict deserialized = bidict{ {2, "hello"}, {3, "goodbye"}, + {4, "yes"}, }; nlohmann::json serialized = std::vector>{ {2, "hello"}, {3, "goodbye"}, + {4, "yes"}, }; SUBCASE("to_json") { diff --git a/lib/utils/test/src/utils/containers/enumerate.cc b/lib/utils/test/src/utils/containers/enumerate.cc index 2fdb2e481e..22bae1f613 100644 --- a/lib/utils/test/src/utils/containers/enumerate.cc +++ b/lib/utils/test/src/utils/containers/enumerate.cc @@ -4,7 +4,7 @@ #include "test/utils/doctest/fmt/unordered_multiset.h" #include "test/utils/doctest/fmt/unordered_set.h" #include "test/utils/doctest/fmt/vector.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/unordered_multiset_of.h" #include "utils/containers/values.h" #include "utils/containers/vector_of.h" @@ -50,7 +50,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::unordered_multiset correct_values = {"A", "B", "C", "D"}; std::map result = enumerate(input); - CHECK(keys(result) == correct_keys); + CHECK(unordered_keys(result) == correct_keys); CHECK(unordered_multiset_of(values(result)) == correct_values); } } diff --git a/lib/utils/test/src/utils/containers/keys.cc b/lib/utils/test/src/utils/containers/keys.cc index 5bdaef6d08..d2ac3dbfba 100644 --- a/lib/utils/test/src/utils/containers/keys.cc +++ b/lib/utils/test/src/utils/containers/keys.cc @@ -1,5 +1,5 @@ #include "utils/containers/keys.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include #include #include @@ -9,10 +9,10 @@ using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("keys") { - std::unordered_map m = { + std::map m = { {1, "one"}, {2, "two"}, {3, "three"}}; - std::unordered_set result = keys(m); - std::unordered_set expected = {1, 2, 3}; + std::set result = keys(m); + std::set expected = {1, 2, 3}; CHECK(result == expected); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc index 4435ccc26c..e92a37169b 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc @@ -1,5 +1,4 @@ #include "utils/graph/digraph/algorithms/get_imm_post_dominators_map.h" -#include "utils/containers/generate_map.h" #include "utils/graph/algorithms.h" #include "utils/graph/instances/adjacency_digraph.h" #include diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc index 8e276daca7..85c4d66f6d 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc @@ -1,6 +1,6 @@ #include "utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.h" #include "test/utils/rapidcheck.h" -#include "utils/containers/generate_map.h" +#include "utils/containers/generate_unordered_map.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_initial_nodes.h" #include "utils/graph/digraph/algorithms/get_terminal_nodes.h" @@ -46,7 +46,7 @@ static std::pair> } std::unordered_map cost_map = - generate_map(get_nodes(g), [](Node const &) { + generate_unordered_map(get_nodes(g), [](Node const &) { return static_cast(*rc::gen::inRange(1, 101)); }); diff --git a/lib/utils/test/src/utils/many_to_one/many_to_one.cc b/lib/utils/test/src/utils/many_to_one/many_to_one.cc index 13f88fab7c..ce219676e1 100644 --- a/lib/utils/test/src/utils/many_to_one/many_to_one.cc +++ b/lib/utils/test/src/utils/many_to_one/many_to_one.cc @@ -96,6 +96,37 @@ TEST_SUITE(FF_TEST_SUITE) { } } + TEST_CASE("adl_serializer>") { + ManyToOne deserialized = ManyToOne{ + {{2, 20}, {"two"}}, + {{3}, "three"}, + {{4, 40, 400}, "four"}, + }; + + nlohmann::json serialized = std::set>{ + {2, "two"}, + {3, "three"}, + {4, "four"}, + {20, "two"}, + {40, "four"}, + {400, "four"}, + }; + + SUBCASE("to_json") { + nlohmann::json result = deserialized; + nlohmann::json correct = serialized; + + CHECK(result == correct); + } + + SUBCASE("from_json") { + ManyToOne result = serialized; + ManyToOne correct = deserialized; + + CHECK(result == correct); + } + } + TEST_CASE("fmt::to_string(ManyToOne)") { ManyToOne input = ManyToOne{ {{1, 10, 100}, "one"}, @@ -107,4 +138,52 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(multiset_of(result) == multiset_of(correct)); } + + TEST_CASE("many_to_one_from_unstructured_relation") { + SUBCASE("relation is many-to-one") { + std::unordered_set> input = { + {1, "odd"}, + {2, "even"}, + {3, "odd"}, + }; + + ManyToOne result = + many_to_one_from_unstructured_relation(input); + ManyToOne correct = { + {{1, 3}, "odd"}, + {{2}, "even"}, + }; + + CHECK(result == correct); + } + + SUBCASE("relation is one-to-one") { + std::unordered_set> input = { + {1, "one"}, + {2, "two"}, + {3, "three"}, + }; + + ManyToOne result = + many_to_one_from_unstructured_relation(input); + ManyToOne correct = { + {{1}, "one"}, + {{2}, "two"}, + {{3}, "three"}, + }; + + CHECK(result == correct); + } + + SUBCASE("relation is not many-to-one") { + std::unordered_set> input = { + {1, "one"}, + {1, "ODD"}, + {2, "two"}, + {3, "ODD"}, + }; + + CHECK_THROWS(many_to_one_from_unstructured_relation(input)); + } + } } diff --git a/lib/utils/test/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc b/lib/utils/test/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc deleted file mode 100644 index c9d0e866d4..0000000000 --- a/lib/utils/test/src/utils/many_to_one/many_to_one_from_unstructured_relation.cc +++ /dev/null @@ -1,54 +0,0 @@ -#include "utils/many_to_one/many_to_one_from_unstructured_relation.h" -#include - -using namespace ::FlexFlow; - -TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("many_to_one_from_unstructured_relation") { - SUBCASE("relation is many-to-one") { - std::unordered_set> input = { - {1, "odd"}, - {2, "even"}, - {3, "odd"}, - }; - - ManyToOne result = - many_to_one_from_unstructured_relation(input); - ManyToOne correct = { - {{1, 3}, "odd"}, - {{2}, "even"}, - }; - - CHECK(result == correct); - } - - SUBCASE("relation is one-to-one") { - std::unordered_set> input = { - {1, "one"}, - {2, "two"}, - {3, "three"}, - }; - - ManyToOne result = - many_to_one_from_unstructured_relation(input); - ManyToOne correct = { - {{1}, "one"}, - {{2}, "two"}, - {{3}, "three"}, - }; - - CHECK(result == correct); - } - - SUBCASE("relation is not many-to-one") { - std::unordered_set> input = { - {1, "one"}, - {1, "ODD"}, - {2, "two"}, - {3, "ODD"}, - }; - - CHECK_THROWS(many_to_one_from_unstructured_relation(input)); - } - } -} diff --git a/lib/utils/test/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc b/lib/utils/test/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc deleted file mode 100644 index 15a8f5b390..0000000000 --- a/lib/utils/test/src/utils/many_to_one/unstructured_relation_from_many_to_one.cc +++ /dev/null @@ -1,25 +0,0 @@ -#include "utils/many_to_one/unstructured_relation_from_many_to_one.h" -#include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_set.h" -#include - -using namespace ::FlexFlow; - -TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("unstructured_relation_from_many_to_one") { - ManyToOne input = { - {{1, 3}, "odd"}, - {{2}, "even"}, - }; - - std::unordered_set> result = - unstructured_relation_from_many_to_one(input); - std::unordered_set> correct = { - {1, "odd"}, - {2, "even"}, - {3, "odd"}, - }; - - CHECK(result == correct); - } -} diff --git a/lib/utils/test/src/utils/one_to_many/one_to_many.cc b/lib/utils/test/src/utils/one_to_many/one_to_many.cc index d2ea7d6a0b..de149ce609 100644 --- a/lib/utils/test/src/utils/one_to_many/one_to_many.cc +++ b/lib/utils/test/src/utils/one_to_many/one_to_many.cc @@ -1,8 +1,11 @@ #include "utils/one_to_many/one_to_many.h" #include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/containers/multiset_of.h" #include "utils/one_to_many/one_to_many_from_l_to_r_mapping.h" +#include "test/utils/doctest/fmt/pair.h" +#include "test/utils/doctest/fmt/unordered_set.h" #include using namespace ::FlexFlow; @@ -31,9 +34,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("at_l") { - nonempty_unordered_set result = m.at_l(1); + nonempty_set result = m.at_l(1); - nonempty_unordered_set correct = {"one", "One", "ONE"}; + nonempty_set correct = {"one", "One", "ONE"}; CHECK(result == correct); } @@ -47,17 +50,17 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("left_values") { - std::unordered_set result = m.left_values(); + std::set result = m.left_values(); - std::unordered_set correct = {1, 2}; + std::set correct = {1, 2}; CHECK(result == correct); } SUBCASE("right_values") { - std::unordered_set result = m.right_values(); + std::set result = m.right_values(); - std::unordered_set correct = {"one", "One", "ONE", "two"}; + std::set correct = {"one", "One", "ONE", "two"}; CHECK(result == correct); } @@ -90,6 +93,35 @@ TEST_SUITE(FF_TEST_SUITE) { } } + TEST_CASE("adl_serializer>") { + OneToMany deserialized = OneToMany{ + {2, {"two", "TWO"}}, + {3, {"three"}}, + {4, {"four"}}, + }; + + nlohmann::json serialized = std::set>{ + {2, "two"}, + {2, "TWO"}, + {3, "three"}, + {4, "four"}, + }; + + SUBCASE("to_json") { + nlohmann::json result = deserialized; + nlohmann::json correct = serialized; + + CHECK(result == correct); + } + + SUBCASE("from_json") { + OneToMany result = serialized; + OneToMany correct = deserialized; + + CHECK(result == correct); + } + } + TEST_CASE("fmt::to_string(OneToMany)") { OneToMany input = one_to_many_from_l_to_r_mapping( @@ -100,4 +132,67 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(multiset_of(result) == multiset_of(correct)); } + + TEST_CASE("unstructured_relation_from_one_to_many") { + OneToMany input = { + {1, {"one", "ONE"}}, + {2, {"two"}}, + }; + + std::unordered_set> result = + unstructured_relation_from_one_to_many(input); + std::unordered_set> correct = { + {1, "one"}, + {1, "ONE"}, + {2, "two"}, + }; + + CHECK(result == correct); + } + + TEST_CASE("one_to_many_from_unstructured_relation") { + SUBCASE("relation is one-to-many") { + std::unordered_set> input = { + {1, "one"}, + {1, "ONE"}, + {2, "two"}, + }; + + OneToMany result = + one_to_many_from_unstructured_relation(input); + OneToMany correct = { + {1, {"one", "ONE"}}, + {2, {"two"}}, + }; + + CHECK(result == correct); + } + + SUBCASE("relation is one-to-one") { + std::unordered_set> input = { + {1, "one"}, + {2, "two"}, + }; + + OneToMany result = + one_to_many_from_unstructured_relation(input); + OneToMany correct = { + {1, {"one"}}, + {2, {"two"}}, + }; + + CHECK(result == correct); + } + + SUBCASE("relation is not one-to-many") { + std::unordered_set> input = { + {1, "one"}, + {1, "ONE"}, + {2, "two"}, + {3, "ONE"}, + }; + + CHECK_THROWS(one_to_many_from_unstructured_relation(input)); + } + } } diff --git a/lib/utils/test/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc b/lib/utils/test/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc deleted file mode 100644 index 023c556c46..0000000000 --- a/lib/utils/test/src/utils/one_to_many/one_to_many_from_unstructured_relation.cc +++ /dev/null @@ -1,53 +0,0 @@ -#include "utils/one_to_many/one_to_many_from_unstructured_relation.h" -#include -#include - -using namespace ::FlexFlow; - -TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("one_to_many_from_unstructured_relation") { - SUBCASE("relation is one-to-many") { - std::unordered_set> input = { - {1, "one"}, - {1, "ONE"}, - {2, "two"}, - }; - - OneToMany result = - one_to_many_from_unstructured_relation(input); - OneToMany correct = { - {1, {"one", "ONE"}}, - {2, {"two"}}, - }; - - CHECK(result == correct); - } - - SUBCASE("relation is one-to-one") { - std::unordered_set> input = { - {1, "one"}, - {2, "two"}, - }; - - OneToMany result = - one_to_many_from_unstructured_relation(input); - OneToMany correct = { - {1, {"one"}}, - {2, {"two"}}, - }; - - CHECK(result == correct); - } - - SUBCASE("relation is not one-to-many") { - std::unordered_set> input = { - {1, "one"}, - {1, "ONE"}, - {2, "two"}, - {3, "ONE"}, - }; - - CHECK_THROWS(one_to_many_from_unstructured_relation(input)); - } - } -} diff --git a/lib/utils/test/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc b/lib/utils/test/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc deleted file mode 100644 index 06a5d5ff2e..0000000000 --- a/lib/utils/test/src/utils/one_to_many/unstructured_relation_from_one_to_many.cc +++ /dev/null @@ -1,25 +0,0 @@ -#include "utils/one_to_many/unstructured_relation_from_one_to_many.h" -#include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_set.h" -#include - -using namespace ::FlexFlow; - -TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("unstructured_relation_from_one_to_many") { - OneToMany input = { - {1, {"one", "ONE"}}, - {2, {"two"}}, - }; - - std::unordered_set> result = - unstructured_relation_from_one_to_many(input); - std::unordered_set> correct = { - {1, "one"}, - {1, "ONE"}, - {2, "two"}, - }; - - CHECK(result == correct); - } -} From 164004d4da20c4c262bd44f9baa471d63cf8d8e0 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Thu, 28 May 2026 04:00:09 -0700 Subject: [PATCH 16/35] Pass replicate copy insertion test case --- .../src/task-spec/dynamic_graph/copy_insertion.cc | 7 +++++++ .../test/src/task-spec/dynamic_graph/copy_insertion.cc | 7 ++++++- 2 files changed, 13 insertions(+), 1 deletion(-) diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 08ab3b11aa..f24dd27da2 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -18,6 +18,7 @@ #include "utils/containers/set_difference.h" #include "utils/containers/transform.h" #include "utils/optional.h" +#include "task-spec/dynamic_graph/training_operation_attrs.h" namespace FlexFlow { @@ -117,6 +118,12 @@ std::unordered_set copies_for_invocation_inputs( DynamicNodeInvocation const &i, std::unordered_map const &unmapped_value_to_src_mapped_value) { + if (training_op_attrs_has_op_type(assert_unwrap(i.node_attrs.op_attrs), OperatorType::REPLICATE)) { + // copies should not be inserted before a replicate, as the replicate + // implicitly includes the copy operations + return {}; + } + MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); auto map_tensor = [&](DynamicTensorSlot const &slot, diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index fdc705dc54..31de844555 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -8,6 +8,7 @@ #include #include "task-spec/dynamic_graph/dynamic_value_attrs.h" #include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" +#include "op-attrs/ops/element_unary.h" using namespace ::FlexFlow; @@ -250,7 +251,11 @@ TEST_SUITE(FF_TEST_SUITE) { /*task_type=*/DynamicTaskType::FWD, /*device_coord=*/std::nullopt, /*mapping=*/invocation_mapping, - /*op_attrs=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{ + PCGOperatorAttrs{ + make_relu_attrs(), + }, + }, /*layer_guid=*/ dynamic_layer_guid_t{parallel_layer_guid_t{Node{invocation_id}}}, /*per_device_op_state=*/std::nullopt, From 032f6ba992d6e3f157e3115e241762a6ad17b3df Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 29 May 2026 02:21:15 -0700 Subject: [PATCH 17/35] Pass shard expansion test for fwd replicate --- .../abstracted_single_tensor_movement.cc | 8 +- .../abstracted_tensor_set_movement.cc | 4 +- ...racted_tensor_set_movement_across_split.cc | 4 +- .../machine_mapping/machine_mapping.cc | 4 +- .../machine_mapping_constraints.cc | 2 +- ...el_layer_guid_oblivious_machine_mapping.cc | 4 +- lib/kernels/include/kernels/accessor.h | 10 + lib/kernels/src/kernels/accessor.cc | 32 ++++ ...space_to_parallel_tensor_space_mappings.cc | 4 +- .../op-attrs/parallel_tensor_dim_degrees.cc | 4 +- lib/pcg/src/pcg/computation_graph.cc | 4 +- lib/pcg/src/pcg/computation_graph_builder.cc | 3 +- .../parallel_computation_graph.cc | 3 +- .../parallel_computation_graph_builder.cc | 3 +- .../apply_substitution/apply_substitution.cc | 8 +- .../perform_shape_inference.cc | 4 +- .../output_operator_attrs_assignment.cc | 4 +- .../include/task-spec/device_specific.h | 16 ++ ...vice_specific_per_device_op_state.dtg.toml | 1 + .../dynamic_layer_guid_t.dtg.toml | 1 + .../dynamic_graph/dynamic_node_attrs.dtg.toml | 8 +- .../dynamic_node_invocation.dtg.toml | 9 +- ...mic_node_invocation_sharding_info.dtg.toml | 6 +- .../dynamic_tensor_accessor.dtg.toml | 1 + .../dynamic_tensor_guid_t.dtg.toml | 1 + .../dynamic_tensor_slot.dtg.toml | 11 ++ .../dynamic_value_attrs.dtg.toml | 1 + .../serializable_dynamic_node_attrs.dtg.toml | 6 +- ...ializable_dynamic_node_invocation.dtg.toml | 9 +- .../serializable_dynamic_value_attrs.dtg.toml | 1 + .../training_operation_attrs.dtg.toml | 1 + .../task-spec/dynamic_graph/copy_insertion.cc | 14 +- .../dynamic_open_dataflow_graph.cc | 8 +- .../task-spec/dynamic_graph/loss_insertion.cc | 37 +++- .../dynamic_graph/machine_slicing.cc | 4 +- ...ake_dynamic_open_dataflow_graph_from_cg.cc | 6 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 5 +- .../serializable_dynamic_node_attrs.cc | 4 +- .../dynamic_graph/shard_expansion.cc | 154 ++++++++++++++-- .../dynamic_graph/update_insertion.cc | 2 +- .../task-spec/dynamic_graph/copy_insertion.cc | 1 + .../dynamic_open_dataflow_graph.cc | 42 +++-- .../dynamic_graph/machine_slicing.cc | 7 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 8 +- .../task-spec/dynamic_graph/pass_expansion.cc | 55 +++--- .../dynamic_graph/shard_expansion.cc | 173 +++++++++++++++++- .../containers/binary_merge_disjoint_maps.h | 14 +- .../binary_merge_disjoint_unordered_maps.h | 28 +++ .../utils/containers/binary_merge_maps_with.h | 26 +-- .../binary_merge_maps_with_left_dominating.h | 6 +- .../binary_merge_maps_with_right_dominating.h | 6 +- .../binary_merge_unordered_maps_with.h | 42 +++++ ...erge_unordered_maps_with_left_dominating.h | 19 ++ ...rge_unordered_maps_with_right_dominating.h | 19 ++ lib/utils/include/utils/containers/flatmap.h | 4 +- lib/utils/include/utils/containers/get_only.h | 8 +- .../include/utils/containers/map_from_pairs.h | 15 +- lib/utils/include/utils/containers/map_keys.h | 27 ++- .../include/utils/containers/map_values2.h | 15 ++ .../utils/containers/merge_disjoint_maps.h | 8 +- .../merge_disjoint_unordered_maps.h | 24 +++ .../include/utils/containers/merge_in_map.h | 6 +- .../utils/containers/merge_in_unordered_map.h | 23 +++ .../utils/containers/merge_maps_with.h | 10 +- .../merge_maps_with_right_dominating.h | 6 +- .../containers/merge_unordered_maps_with.h | 25 +++ ...rge_unordered_maps_with_right_dominating.h | 23 +++ .../include/utils/containers/restrict_keys.h | 13 ++ .../utils/containers/zip_values_strict.h | 18 ++ .../full_binary_tree/get_path_to_leaf_map.h | 4 +- .../include/utils/nonempty_set/nonempty_set.h | 34 +++- .../require_one_to_many_is_bijection.h | 23 +++ .../utils/orthotope/minimal_dim_domain.h | 4 +- .../containers/binary_merge_disjoint_maps.cc | 9 +- .../binary_merge_disjoint_unordered_maps.cc | 13 ++ .../containers/binary_merge_maps_with.cc | 7 +- .../binary_merge_maps_with_left_dominating.cc | 9 +- ...binary_merge_maps_with_right_dominating.cc | 9 +- .../binary_merge_unordered_maps_with.cc | 13 ++ ...rge_unordered_maps_with_left_dominating.cc | 13 ++ ...ge_unordered_maps_with_right_dominating.cc | 13 ++ .../src/utils/containers/map_from_pairs.cc | 17 +- lib/utils/src/utils/containers/map_keys.cc | 21 +++ lib/utils/src/utils/containers/map_values2.cc | 19 +- .../utils/containers/merge_disjoint_maps.cc | 7 +- .../merge_disjoint_unordered_maps.cc | 12 ++ .../src/utils/containers/merge_in_map.cc | 10 +- .../containers/merge_in_unordered_map.cc | 12 ++ .../src/utils/containers/merge_maps_with.cc | 7 +- .../merge_maps_with_right_dominating.cc | 7 +- .../containers/merge_unordered_maps_with.cc | 14 ++ ...ge_unordered_maps_with_right_dominating.cc | 12 ++ .../src/utils/containers/restrict_keys.cc | 20 ++ .../src/utils/containers/zip_values_strict.cc | 21 ++- .../src/utils/nonempty_set/nonempty_set.cc | 8 + .../require_one_to_many_is_bijection.cc | 12 ++ 96 files changed, 1191 insertions(+), 241 deletions(-) create mode 100644 lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h create mode 100644 lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h create mode 100644 lib/utils/include/utils/containers/binary_merge_unordered_maps_with_left_dominating.h create mode 100644 lib/utils/include/utils/containers/binary_merge_unordered_maps_with_right_dominating.h create mode 100644 lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h create mode 100644 lib/utils/include/utils/containers/merge_in_unordered_map.h create mode 100644 lib/utils/include/utils/containers/merge_unordered_maps_with.h create mode 100644 lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h create mode 100644 lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h create mode 100644 lib/utils/src/utils/containers/binary_merge_disjoint_unordered_maps.cc create mode 100644 lib/utils/src/utils/containers/binary_merge_unordered_maps_with.cc create mode 100644 lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc create mode 100644 lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc create mode 100644 lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc create mode 100644 lib/utils/src/utils/containers/merge_in_unordered_map.cc create mode 100644 lib/utils/src/utils/containers/merge_unordered_maps_with.cc create mode 100644 lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc create mode 100644 lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc index 5f9300973f..e8bd602289 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc @@ -1,13 +1,13 @@ #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h" #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.h" #include "utils/containers/filtermap_keys.h" -#include "utils/containers/map_from_pairs.h" #include "utils/containers/map_keys_with_value_merging.h" -#include "utils/containers/merge_maps_with.h" #include "utils/containers/require_all_same1.h" #include "utils/containers/require_same.h" #include "utils/containers/transform.h" #include "utils/containers/values.h" +#include "utils/containers/merge_unordered_maps_with.h" +#include "utils/containers/unordered_map_from_pairs.h" namespace FlexFlow { @@ -34,7 +34,7 @@ AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( return AbstractedSingleTensorMovement{ /*src_op_tree_path=*/require_all_same1(src_paths), /*edge_to_size=*/ - merge_maps_with(transform(vector_of(movements), + merge_unordered_maps_with(transform(vector_of(movements), [](AbstractedSingleTensorMovement const &m) { return m.edge_to_size; }), @@ -51,7 +51,7 @@ AbstractedSingleTensorMovement return AbstractedSingleTensorMovement{ /*src_op_tree_path=*/src_op_tree_path, /*edge_to_size=*/ - map_from_pairs( + unordered_map_from_pairs( transform(communications, [](AbstractedSingleTensorCommunication const &c) { return std::pair{c.edge, c.size}; diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc index 98a7d9b0b2..37bf62029f 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc @@ -3,13 +3,13 @@ #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.dtg.h" #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h" #include "compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.h" -#include "utils/containers/binary_merge_maps_with.h" #include "utils/containers/flatmap.h" #include "utils/containers/map_keys_with_value_merging.h" #include "utils/containers/merge_maps_with.h" #include "utils/containers/transform.h" #include "utils/containers/unordered_set_of.h" #include "utils/hash/unordered_map.h" +#include "utils/containers/binary_merge_unordered_maps_with.h" namespace FlexFlow { @@ -63,7 +63,7 @@ TensorSetMovement concretize_abstracted_tensor_set_movement( [](TensorSetMovement const &lhs, TensorSetMovement const &rhs) -> TensorSetMovement { return TensorSetMovement{ - binary_merge_maps_with( + binary_merge_unordered_maps_with( lhs.edge_to_size, rhs.edge_to_size, [](num_bytes_t l, num_bytes_t r) { return l + r; }), diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc index 6ff261facd..192ada3fb6 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc @@ -44,7 +44,7 @@ AbstractedSingleTensorMovement get_abstracted_single_tensor_movement_along_edge( op_to_op_get_coord_mapping(mapping); std::unordered_map - single_comms = map_from_pairs(transform( + single_comms = unordered_map_from_pairs(transform( bidict_unordered_set_of(coord_mapping), [&](std::pair const & src_dst) -> std::pair const &edges) + [&](nonempty_set const &edges) { return merge_abstracted_single_tensor_movements(transform( unordered_multiset_of(edges.unwrap_as_unordered_set()), diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc index 861912efef..c7b068d121 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc @@ -7,8 +7,8 @@ #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/containers/are_disjoint.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/unordered_keys.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -49,7 +49,7 @@ MappedParallelComputationGraph MachineMapping combine_disjoint_mappings(MachineMapping const &m1, MachineMapping const &m2) { return MachineMapping{ - binary_merge_disjoint_maps(m1.machine_views, m2.machine_views), + binary_merge_disjoint_unordered_maps(m1.machine_views, m2.machine_views), }; } diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc index fe92c77def..f77d424795 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc @@ -4,7 +4,7 @@ #include "utils/containers/filtermap_keys.h" #include "utils/containers/flatmap.h" #include "utils/containers/generate_unordered_map.h" -#include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/map_values.h" #include "utils/containers/restrict_keys.h" #include "utils/full_binary_tree/binary_tree_path.h" diff --git a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc index ac39021f6f..97814288ab 100644 --- a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc @@ -5,11 +5,11 @@ #include "op-attrs/get_operator_task_space.h" #include "op-attrs/parallel_tensor_shape.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/map_keys.h" #include "utils/containers/require_same.h" #include "utils/containers/try_at.h" #include "utils/full_binary_tree/binary_tree_path.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -17,7 +17,7 @@ ParallelLayerGuidObliviousMachineMapping binary_combine_mappings( ParallelLayerGuidObliviousMachineMapping const &lhs, ParallelLayerGuidObliviousMachineMapping const &rhs) { return ParallelLayerGuidObliviousMachineMapping{ - binary_merge_disjoint_maps( + binary_merge_disjoint_unordered_maps( map_keys(lhs.raw_mapping, nest_inside_left_child), map_keys(rhs.raw_mapping, nest_inside_right_child)), }; diff --git a/lib/kernels/include/kernels/accessor.h b/lib/kernels/include/kernels/accessor.h index 27d3693e7e..cacc56046f 100644 --- a/lib/kernels/include/kernels/accessor.h +++ b/lib/kernels/include/kernels/accessor.h @@ -43,6 +43,11 @@ class GenericTensorAccessorR { bool operator==(GenericTensorAccessorR const &) const; bool operator!=(GenericTensorAccessorR const &) const; + bool operator<(GenericTensorAccessorR const &) const; + bool operator<=(GenericTensorAccessorR const &) const; + bool operator>(GenericTensorAccessorR const &) const; + bool operator>=(GenericTensorAccessorR const &) const; + template real_type_t
const &at(TensorDimsCoord const &indices) const { ASSERT(this->device_type == DeviceType::CPU, @@ -97,6 +102,11 @@ class GenericTensorAccessorW { bool operator==(GenericTensorAccessorW const &) const; bool operator!=(GenericTensorAccessorW const &) const; + bool operator<(GenericTensorAccessorW const &) const; + bool operator<=(GenericTensorAccessorW const &) const; + bool operator>(GenericTensorAccessorW const &) const; + bool operator>=(GenericTensorAccessorW const &) const; + operator GenericTensorAccessorR() const; template diff --git a/lib/kernels/src/kernels/accessor.cc b/lib/kernels/src/kernels/accessor.cc index a3f8ead17f..75f144a57a 100644 --- a/lib/kernels/src/kernels/accessor.cc +++ b/lib/kernels/src/kernels/accessor.cc @@ -98,6 +98,22 @@ bool GenericTensorAccessorW::operator!=( return this->tie() != other.tie(); } +bool GenericTensorAccessorW::operator<(GenericTensorAccessorW const &other) const { + return this->tie() < other.tie(); +} + +bool GenericTensorAccessorW::operator<=(GenericTensorAccessorW const &other) const { + return this->tie() <= other.tie(); +} + +bool GenericTensorAccessorW::operator>(GenericTensorAccessorW const &other) const { + return this->tie() > other.tie(); +} + +bool GenericTensorAccessorW::operator>=(GenericTensorAccessorW const &other) const { + return this->tie() >= other.tie(); +} + int32_t *GenericTensorAccessorW::get_int32_ptr() const { return this->get(); } @@ -150,6 +166,22 @@ bool GenericTensorAccessorR::operator!=( return this->tie() != other.tie(); } +bool GenericTensorAccessorR::operator<(GenericTensorAccessorR const &other) const { + return this->tie() < other.tie(); +} + +bool GenericTensorAccessorR::operator<=(GenericTensorAccessorR const &other) const { + return this->tie() <= other.tie(); +} + +bool GenericTensorAccessorR::operator>(GenericTensorAccessorR const &other) const { + return this->tie() > other.tie(); +} + +bool GenericTensorAccessorR::operator>=(GenericTensorAccessorR const &other) const { + return this->tie() >= other.tie(); +} + int32_t const *GenericTensorAccessorR::get_int32_ptr() const { return this->get(); } diff --git a/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc b/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc index 618eb533ff..1a97f8b38b 100644 --- a/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc +++ b/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc @@ -8,11 +8,11 @@ #include "op-attrs/ops/weight.h" #include "utils/containers/filtrans.h" #include "utils/containers/get_only.h" -#include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/require_only_key.h" #include "utils/containers/require_two_keys.h" #include "utils/containers/zip_values_strict.h" #include "utils/overload.h" +#include "utils/containers/merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -281,7 +281,7 @@ std::unordered_map ComputationGraphOpAttrs const &attrs, std::unordered_map const &inputs_degrees) { - return merge_disjoint_maps(std::vector{ + return merge_disjoint_unordered_maps(std::vector{ get_operator_to_input_mappings(attrs, inputs_degrees), get_operator_to_weight_mappings(attrs, inputs_degrees), get_operator_to_output_mappings(attrs, inputs_degrees), diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc index 83a7aded6a..5b5d0b514f 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc @@ -5,7 +5,6 @@ #include "op-attrs/parallel_tensor_dim_idx_t.dtg.h" #include "op-attrs/parallel_tensor_dim_idx_t.h" #include "op-attrs/parallel_tensor_space_coordinate.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/filtrans.h" #include "utils/containers/generate_unordered_map.h" @@ -19,6 +18,7 @@ #include "utils/nonnegative_int/nonnegative_range.h" #include "utils/nonnegative_int/num_elements.h" #include "utils/orthotope/minimal_dim_domain.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -100,7 +100,7 @@ std::unordered_map return degrees.shard_degrees.at(dim); }); - return binary_merge_disjoint_maps( + return binary_merge_disjoint_unordered_maps( /*lhs=*/replica_dim_degrees, /*rhs=*/map_keys(shard_dim_degrees, [](ff_dim_t const &dim) { return parallel_tensor_dim_idx_t{dim}; diff --git a/lib/pcg/src/pcg/computation_graph.cc b/lib/pcg/src/pcg/computation_graph.cc index 35ba0747f0..1b0b3b3204 100644 --- a/lib/pcg/src/pcg/computation_graph.cc +++ b/lib/pcg/src/pcg/computation_graph.cc @@ -2,7 +2,6 @@ #include "op-attrs/computation_graph_op_attrs.h" #include "op-attrs/get_incoming_tensor_roles.h" #include "op-attrs/shape_inference.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/filter_values.h" #include "utils/containers/filtrans.h" @@ -35,6 +34,7 @@ #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/node/algorithms.h" #include "utils/record_formatter.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -101,7 +101,7 @@ LayerAddedResult add_layer( KwargNodeAddedResult added = computation_graph.raw_graph.add_node( layer_attrs, - binary_merge_disjoint_maps(raw_inputs, raw_weights), + binary_merge_disjoint_unordered_maps(raw_inputs, raw_weights), output_attrs); return LayerAddedResult{ diff --git a/lib/pcg/src/pcg/computation_graph_builder.cc b/lib/pcg/src/pcg/computation_graph_builder.cc index 40e72aee9d..2eb140fa58 100644 --- a/lib/pcg/src/pcg/computation_graph_builder.cc +++ b/lib/pcg/src/pcg/computation_graph_builder.cc @@ -48,6 +48,7 @@ #include "utils/fmt/set.h" #include "utils/stack_vector/stack_vector_of.h" #include +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -114,7 +115,7 @@ static void check_incoming_tensor_roles( restrict_keys(get_incoming_tensor_roles(layer.op_attrs), set_union(input_slots, weight_slots)); std::unordered_map current = - binary_merge_disjoint_maps( + binary_merge_disjoint_unordered_maps( generate_unordered_map( input_slots, [](TensorSlotName) { return IncomingTensorRole::INPUT; }), diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index c4a429a820..40e8cb9e5e 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -38,6 +38,7 @@ #include "utils/record_formatter.h" #include #include "utils/containers/map_from_unordered.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -112,7 +113,7 @@ ParallelLayerAddedResult add_parallel_layer( KwargNodeAddedResult op_added = pcg.raw_graph.add_node( layer_attrs, - binary_merge_disjoint_maps(unwrapped_inputs, unwrapped_weights), + binary_merge_disjoint_unordered_maps(unwrapped_inputs, unwrapped_weights), output_attrs); return ParallelLayerAddedResult{ diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc index 92334cfde9..c01b17c285 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc @@ -34,6 +34,7 @@ #include "utils/containers/transform.h" #include "utils/containers/zip_values_strict_with.h" #include "utils/containers/zip_with.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -678,7 +679,7 @@ static void check_incoming_tensor_roles( std::unordered_map correct = get_incoming_tensor_roles(layer.op_attrs); std::unordered_map current = - binary_merge_disjoint_maps( + binary_merge_disjoint_unordered_maps( generate_unordered_map( input_slots, [](TensorSlotName) { return IncomingTensorRole::INPUT; }), diff --git a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc index f3ceda7a06..b8140440b7 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc @@ -9,11 +9,11 @@ #include "substitutions/sub_parallel_computation_graph_data.dtg.h" #include "substitutions/sub_parallel_computation_graph_data.h" #include "substitutions/sub_parallel_computation_graph_edge.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/unordered_keys.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" #include "utils/containers/values.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -64,7 +64,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( std::unordered_map post_node_data_from_sub = output_graph_data.node_data; - return binary_merge_disjoint_maps(post_node_data_from_orig, + return binary_merge_disjoint_unordered_maps(post_node_data_from_orig, post_node_data_from_sub); }(); @@ -168,8 +168,8 @@ SubParallelComputationGraph apply_substitution_from_output_result( std::unordered_map post_value_data_from_sub = output_graph_data.value_data; - return binary_merge_disjoint_maps(post_value_data_from_orig, - post_value_data_from_sub); + return binary_merge_disjoint_unordered_maps(post_value_data_from_orig, + post_value_data_from_sub); }(); SubParallelComputationGraphData post_data = SubParallelComputationGraphData{ diff --git a/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc b/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc index 9ae007ef16..d3ad4ca246 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc @@ -1,7 +1,6 @@ #include "substitutions/apply_substitution/perform_shape_inference.h" #include "op-attrs/get_incoming_tensor_roles.h" #include "op-attrs/shape_inference.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/filter_values.h" #include "utils/containers/filtrans.h" #include "utils/containers/is_subseteq_of.h" @@ -20,6 +19,7 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.h" #include "utils/nonnegative_int/num_elements.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -72,7 +72,7 @@ LabelledOpenKwargDataflowGraphView weight_shapes = incoming_shapes_with_role(IncomingTensorRole::WEIGHT); - ASSERT(binary_merge_disjoint_maps(input_shapes, weight_shapes) == + ASSERT(binary_merge_disjoint_unordered_maps(input_shapes, weight_shapes) == incoming_shapes); std::unordered_map diff --git a/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc b/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc index 647362ee4d..2755716e44 100644 --- a/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc +++ b/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc @@ -2,9 +2,9 @@ #include "substitutions/operator_pattern/get_attribute_map.h" #include "substitutions/output_graph/materialize_operator_from_attrs_map.h" #include "substitutions/output_graph/output_operator_attribute_expr.h" -#include "utils/containers/binary_merge_maps_with_right_dominating.h" #include "utils/containers/map_values.h" #include "utils/exception.h" +#include "utils/containers/binary_merge_unordered_maps_with_right_dominating.h" namespace FlexFlow { @@ -36,7 +36,7 @@ PCGOperatorAttrs materialize_output_operator_from_attrs_assignment( }); std::unordered_map - joined_attrs_map = binary_merge_maps_with_right_dominating( + joined_attrs_map = binary_merge_unordered_maps_with_right_dominating( template_attrs_map, assignments_attrs_map); return materialize_operator_from_attrs_map(joined_attrs_map); diff --git a/lib/task-spec/include/task-spec/device_specific.h b/lib/task-spec/include/task-spec/device_specific.h index 2055888b1b..834378f415 100644 --- a/lib/task-spec/include/task-spec/device_specific.h +++ b/lib/task-spec/include/task-spec/device_specific.h @@ -25,6 +25,22 @@ struct DeviceSpecific { return this->tie() != other.tie(); } + bool operator<(DeviceSpecific const &other) const { + return this->tie() < other.tie(); + } + + bool operator<=(DeviceSpecific const &other) const { + return this->tie() <= other.tie(); + } + + bool operator>(DeviceSpecific const &other) const { + return this->tie() > other.tie(); + } + + bool operator>=(DeviceSpecific const &other) const { + return this->tie() >= other.tie(); + } + T const *get(device_id_t curr_device_idx) const { ASSERT(curr_device_idx == this->device_idx); return (T const *)this->ptr.get(); diff --git a/lib/task-spec/include/task-spec/device_specific_per_device_op_state.dtg.toml b/lib/task-spec/include/task-spec/device_specific_per_device_op_state.dtg.toml index 4435a472ce..0a32b29b25 100644 --- a/lib/task-spec/include/task-spec/device_specific_per_device_op_state.dtg.toml +++ b/lib/task-spec/include/task-spec/device_specific_per_device_op_state.dtg.toml @@ -4,6 +4,7 @@ type = "variant" features = [ "eq", "hash", + "ord", "fmt", ] diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.toml index 8def0ec5fb..5200bfc6a6 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.toml @@ -3,6 +3,7 @@ name = "dynamic_layer_guid_t" type = "variant" features = [ "eq", + "ord", "hash", "fmt", "json", diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml index 73c023fd40..110c963383 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml @@ -4,6 +4,7 @@ type = "struct" features = [ "eq", "hash", + "ord", "fmt", ] @@ -11,6 +12,7 @@ includes = [ "", "task-spec/dynamic_graph/dynamic_task_type.dtg.h", "pcg/machine_space_coordinate.dtg.h", + "utils/nonempty_set/nonempty_set.h", "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h", "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h", "task-spec/dynamic_graph/training_operation_attrs.dtg.h", @@ -26,10 +28,10 @@ name = "task_type" type = "std::optional<::FlexFlow::DynamicTaskType>" [[fields]] -name = "device_coord" -type = "std::optional<::FlexFlow::MachineSpaceCoordinate>" +name = "device_coords" +type = "std::optional<::FlexFlow::nonempty_set<::FlexFlow::MachineSpaceCoordinate>>" docstring = ''' -\note Right now the \c device_coord for a copy node is sort of meaningless +\note Right now the \c device_coords for a copy node is sort of meaningless because we have one controller issuing all copies for the entire graph, no matter where they are. However the intention is this to be the "owner" or "issuer" of the copy, which matters a lot more down the road once we write the diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml index 07060106c0..4d6d27444a 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml @@ -3,6 +3,7 @@ name = "DynamicNodeInvocation" type = "struct" features = [ "eq", + "ord", "fmt", "hash", ] @@ -15,13 +16,13 @@ includes = [ ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] name = "inputs" -type = "std::unordered_map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::DynamicValueAttrs>" +type = "std::map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::DynamicValueAttrs>" [[fields]] name = "node_attrs" @@ -29,4 +30,4 @@ type = "::FlexFlow::DynamicNodeAttrs" [[fields]] name = "outputs" -type = "std::unordered_map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::DynamicValueAttrs>" +type = "std::map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::DynamicValueAttrs>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml index a59aba92d7..00d98a2e6c 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml @@ -1,7 +1,6 @@ namespace = "FlexFlow" name = "DynamicNodeInvocationShardingInfo" type = "struct" -#include "task-spec/dynamic_graph/shard_expansion.h" features = [ "eq", "ord", @@ -14,6 +13,7 @@ includes = [ "pcg/machine_space_coordinate.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h", "task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.h", + "utils/nonempty_set/nonempty_set.h", ] src_includes = [ @@ -22,8 +22,8 @@ src_includes = [ ] [[fields]] -name = "device_coord" -type = "::FlexFlow::MachineSpaceCoordinate" +name = "device_coords" +type = "::FlexFlow::nonempty_set<::FlexFlow::MachineSpaceCoordinate>" [[fields]] name = "value_sharding" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.toml index 85f8f299a4..bfe30a7a5e 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.toml @@ -3,6 +3,7 @@ name = "DynamicTensorAccessor" type = "variant" features = [ "eq", + "ord", "fmt", "hash", ] diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_guid_t.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_guid_t.dtg.toml index c9171b928b..56f8ae4359 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_guid_t.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_guid_t.dtg.toml @@ -3,6 +3,7 @@ name = "dynamic_tensor_guid_t" type = "variant" features = [ "eq", + "ord", "hash", "fmt", "json", diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.dtg.toml index 378582f428..851b8dc83d 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.dtg.toml @@ -12,6 +12,7 @@ features = [ includes = [ "op-attrs/tensor_slot_name.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_role.dtg.h", + "pcg/machine_space_coordinate.dtg.h", "", ] @@ -27,3 +28,13 @@ type = "::FlexFlow::TensorSlotName" [[fields]] name = "slot_tensor_role" type = "std::optional<::FlexFlow::DynamicTensorRole>" + +[[fields]] +name = "task_shard" +type = "std::optional<::FlexFlow::MachineSpaceCoordinate>" +docstring = ''' +\brief For representing parallel operators such as \ref ReplicateAttrs as a single operator with multiple outputs instead of fully shard expanding. + +This is done for the convenience of the runtime, as the \ref ReplicateAttrs is ultimately executed as a single NCCL task/call, so shard-expanding this operator is ultimately counter-productive. +Since the output values, however, do need to be shard-expanded for the rest of the system to work, we make Replicate a single operator with multiple outputs, each with the same \ref TensorSlotName and \ref DynamicTensorRole, but with different \ref MachineSpaceCoordinate ""s to identify which node is ultimately producing that value. +''' diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml index add72764f1..1e13736fc7 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml @@ -3,6 +3,7 @@ name = "DynamicValueAttrs" type = "struct" features = [ "eq", + "ord", "fmt", "hash", ] diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml index 3c43e1d637..c8ea415ee2 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml @@ -3,6 +3,7 @@ name = "SerializableDynamicNodeAttrs" type = "struct" features = [ "eq", + "ord", "hash", "fmt", "json", @@ -15,6 +16,7 @@ includes = [ "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h", "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h", "task-spec/dynamic_graph/training_operation_attrs.dtg.h", + "utils/nonempty_set/nonempty_set.h", ] src_includes = [ @@ -27,8 +29,8 @@ name = "task_type" type = "std::optional<::FlexFlow::DynamicTaskType>" [[fields]] -name = "device_coord" -type = "std::optional<::FlexFlow::MachineSpaceCoordinate>" +name = "device_coords" +type = "std::optional<::FlexFlow::nonempty_set<::FlexFlow::MachineSpaceCoordinate>>" [[fields]] name = "mapping" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml index 01f4cc8876..6051a49876 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml @@ -3,6 +3,7 @@ name = "SerializableDynamicNodeInvocation" type = "struct" features = [ "eq", + "ord", "fmt", "hash", "json", @@ -16,13 +17,13 @@ includes = [ ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] name = "inputs" -type = "std::unordered_map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::SerializableDynamicValueAttrs>" +type = "std::map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::SerializableDynamicValueAttrs>" [[fields]] name = "node_attrs" @@ -30,4 +31,4 @@ type = "::FlexFlow::SerializableDynamicNodeAttrs" [[fields]] name = "outputs" -type = "std::unordered_map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::SerializableDynamicValueAttrs>" +type = "std::map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::SerializableDynamicValueAttrs>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml index d3cab6ecdb..d05d8e011e 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml @@ -3,6 +3,7 @@ name = "SerializableDynamicValueAttrs" type = "struct" features = [ "eq", + "ord", "hash", "fmt", "json", diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml index 8f8f6467c8..2c4e739571 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.dtg.toml @@ -3,6 +3,7 @@ name = "TrainingOperationAttrs" type = "variant" features = [ "eq", + "ord", "hash", "fmt", "json", diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index f24dd27da2..ebf84f2d81 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -131,7 +131,7 @@ std::unordered_set copies_for_invocation_inputs( return map_dynamic_value_attrs_for_task_group(slot, value, mapping); }; - std::unordered_map mapped_inputs = + std::map mapped_inputs = map_values2(i.inputs, map_tensor); std::unordered_set result; @@ -152,6 +152,7 @@ std::unordered_set copies_for_invocation_inputs( DynamicTensorSlot{ TensorSlotName::INPUT, slot.slot_tensor_role, + /*task_shard=*/std::nullopt, }, filtered_source, }, @@ -170,8 +171,11 @@ std::unordered_set copies_for_invocation_inputs( /*outputs=*/ { { - DynamicTensorSlot{TensorSlotName::OUTPUT, - slot.slot_tensor_role}, + DynamicTensorSlot{ + TensorSlotName::OUTPUT, + slot.slot_tensor_role, + /*task_shard=*/std::nullopt, + }, filtered_use, }, }, @@ -196,9 +200,9 @@ std::unordered_set perform_copy_insertion_for_invocation( }; DynamicNodeInvocation mapped_i = [&] { - std::unordered_map mapped_inputs = + std::map mapped_inputs = map_values2(i.inputs, map_tensor); - std::unordered_map mapped_outputs = + std::map mapped_outputs = map_values2(i.outputs, map_tensor); DynamicNodeInvocation r = i; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index a100c3adfb..38d49ef183 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -20,6 +20,8 @@ #include "utils/graph/open_kwarg_dataflow_graph/kwarg_dataflow_graph_input.dtg.h" #include "utils/many_to_one/many_to_one.h" #include "utils/containers/require_all_of.h" +#include "utils/containers/unordered_map_from_map.h" +#include "utils/containers/map_from_unordered.h" namespace FlexFlow { @@ -216,16 +218,16 @@ std::pair void { KwargNodeAddedResult added = result.add_node( invocation.node_attrs, - map_values(invocation.inputs, + map_values(unordered_map_from_map(invocation.inputs), [&](DynamicValueAttrs const &input) -> OpenKwargDataflowValue { return value_map.at_r(input); }), - invocation.outputs); + unordered_map_from_map(invocation.outputs)); node_map.equate(added.node, invocation); for (auto const &[k, v] : - zip_values_strict(invocation.outputs, added.outputs)) { + zip_values_strict(invocation.outputs, map_from_unordered(added.outputs))) { DynamicValueAttrs invocation_output = v.first; KwargDataflowOutput graph_output = v.second; value_map.equate( diff --git a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc index 8066926262..8ff24b51ad 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc @@ -18,6 +18,7 @@ LossInsertionResult perform_loss_insertion( LossAttrs const &loss_attrs, dynamic_tensor_guid_t logit_tensor, std::optional const &loss_mapping) { + DynamicValueAttrs logit_value = assert_unwrap( find_output_value_attrs(dg, logit_tensor, mk_dynamic_tensor_role_fwd())); @@ -29,6 +30,7 @@ LossInsertionResult perform_loss_insertion( /*accessor=*/std::nullopt, /*role=*/mk_dynamic_tensor_role_loss(), }; + DynamicValueAttrs logit_grad_value{ /*tensor_guid=*/logit_value.tensor_guid, /*parallel_tensor_shape=*/logit_value.parallel_tensor_shape, @@ -37,14 +39,25 @@ LossInsertionResult perform_loss_insertion( /*accessor=*/std::nullopt, /*role=*/mk_dynamic_tensor_role_bwd(), }; + DynamicNodeInvocation loss_invocation{ /*inputs=*/{ - {DynamicTensorSlot{/*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/label_value.role}, - label_value}, - {DynamicTensorSlot{/*slot_name=*/TensorSlotName::LOGIT, - /*slot_tensor_role=*/logit_value.role}, - logit_value}, + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/label_value.role, + /*task_shard=*/std::nullopt, + }, + label_value, + }, + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::LOGIT, + /*slot_tensor_role=*/logit_value.role, + /*task_shard=*/std::nullopt, + }, + logit_value, + }, }, /*node_attrs=*/ DynamicNodeAttrs{ @@ -57,11 +70,17 @@ LossInsertionResult perform_loss_insertion( }, /*outputs=*/ { - {DynamicTensorSlot{/*slot_name=*/TensorSlotName::LOGIT, - /*slot_tensor_role=*/logit_grad_value.role}, - logit_grad_value}, + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::LOGIT, + /*slot_tensor_role=*/logit_grad_value.role, + /*task_shard=*/std::nullopt, + }, + logit_grad_value, + }, }, }; + DynamicOpenDataflowGraph result = dg; result.invocations.insert(loss_invocation); return LossInsertionResult{result, label_value, logit_grad_value}; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc index 0a22015ddf..6dd73bed7d 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc @@ -8,9 +8,9 @@ std::unordered_set DynamicNodeInvocation const &invocation, MachineSpaceCoordinate const &device_coord) { - ASSERT(invocation.node_attrs.device_coord.has_value()); + ASSERT(invocation.node_attrs.device_coords.has_value()); - if (invocation.node_attrs.device_coord.value() == device_coord) { + if (contains(invocation.node_attrs.device_coords.value(), device_coord)) { return {invocation}; } else { return {}; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc index 7fe3927fd1..740a93415e 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc @@ -11,6 +11,7 @@ #include #include #include +#include "utils/containers/map_from_unordered.h" namespace FlexFlow { @@ -39,6 +40,7 @@ DynamicOpenDataflowGraph DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, @@ -50,6 +52,7 @@ DynamicOpenDataflowGraph }, }; }); + std::unordered_map result_outputs = transform( get_outgoing_tensors(cg, layer), @@ -59,6 +62,7 @@ DynamicOpenDataflowGraph DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, @@ -71,7 +75,7 @@ DynamicOpenDataflowGraph }; }); - result.invocations.emplace(result_inputs, result_attrs, result_outputs); + result.invocations.emplace(map_from_unordered(result_inputs), result_attrs, map_from_unordered(result_outputs)); } return result; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 391ebaff3b..9c7638440f 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -43,6 +43,7 @@ DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{tensor.guid}, @@ -62,9 +63,9 @@ DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( transform(invocation_info.outgoing, lift_kv_pair); DynamicNodeInvocation invocation = DynamicNodeInvocation{ - /*inputs=*/unordered_map_from_map(result_inputs), + /*inputs=*/result_inputs, /*node_attrs=*/result_attrs, - /*outputs=*/unordered_map_from_map(result_outputs), + /*outputs=*/result_outputs, }; return invocation; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc index d613194d14..67ae7b58f3 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc @@ -7,7 +7,7 @@ SerializableDynamicNodeAttrs dynamic_node_attrs_to_serializable(DynamicNodeAttrs const &attrs) { return SerializableDynamicNodeAttrs{ /*task_type=*/attrs.task_type, - /*device_coord=*/attrs.device_coord, + /*device_coords=*/attrs.device_coords, /*mapping=*/attrs.mapping, /*op_attrs=*/attrs.op_attrs, /*layer_guid=*/attrs.layer_guid, @@ -18,7 +18,7 @@ DynamicNodeAttrs dynamic_node_attrs_from_serializable( SerializableDynamicNodeAttrs const &attrs) { return DynamicNodeAttrs{ /*task_type=*/attrs.task_type, - /*device_coord=*/attrs.device_coord, + /*device_coords=*/attrs.device_coords, /*mapping=*/attrs.mapping, /*op_attrs=*/attrs.op_attrs, /*layer_guid=*/attrs.layer_guid, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index badd376a8b..db9484e361 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -11,11 +11,19 @@ #include "task-spec/dynamic_graph/dynamic_node_invocation.h" #include "utils/containers/map_from_unordered.h" #include "utils/one_to_many/one_to_many_filter_keys.h" +#include "task-spec/dynamic_graph/training_operation_attrs.h" +#include "utils/one_to_many/require_one_to_many_is_bijection.h" +#include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.h" +#include "utils/bidict/algorithms/bidict_filter_values.h" +#include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include "utils/containers/merge_disjoint_maps.h" +#include "utils/containers/map_keys.h" +#include "utils/containers/require_only_key.h" namespace FlexFlow { bool node_is_shard_expanded(DynamicNodeAttrs const &n) { - return n.device_coord.has_value(); + return n.device_coords.has_value(); } bool node_is_ready_for_shard_expansion(DynamicNodeAttrs const &n) { @@ -138,16 +146,15 @@ static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( DynamicNodeAttrs expanded_node_attrs = [&]() { DynamicNodeAttrs result = i.node_attrs; - result.device_coord = machine_coord; + result.device_coords = nonempty_set{machine_coord}; return result; }(); return DynamicNodeInvocationShardingInfo{ - /*device_coord=*/machine_coord, - /*value_sharding=*/map_from_unordered( - map_values2( + /*device_coord=*/nonempty_set{machine_coord}, + /*value_sharding=*/map_values2( binary_merge_disjoint_maps(i.inputs, i.outputs), - shard_expand_value_attrs)), + shard_expand_value_attrs), }; } @@ -177,7 +184,7 @@ static DynamicNodeInvocation shard_invocation_for_binding( DynamicNodeAttrs expanded_node_attrs = [&]() { DynamicNodeAttrs result = i.node_attrs; - result.device_coord = machine_coord; + result.device_coords = nonempty_set{machine_coord}; return result; }(); @@ -219,6 +226,121 @@ static std::set }); } +static std::set + generate_shard_expansion_for_fwd_replicate(DynamicNodeInvocation const &i) { + ASSERT(i.node_attrs.task_type == DynamicTaskType::FWD); + + MappedOperatorTaskGroup node_mapping = assert_unwrap(i.node_attrs.mapping); + + DynamicTensorSlot expected_input_slot = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }; + + DynamicValueAttrs input = require_only_key(i.inputs, expected_input_slot); + + DynamicTensorSlot expected_output_slot = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }; + + DynamicValueAttrs output = require_only_key(i.outputs, expected_output_slot); + + bidict + input_value_mapping = require_one_to_many_is_bijection( + assert_unwrap(input.mapping)); + + std::set input_tensor_shards = set_of(input_value_mapping.left_values()); + + bidict + output_value_mapping = require_one_to_many_is_bijection( + assert_unwrap(output.mapping)); + + auto get_task_shard_machine_coords_for_input_tensor_shard + = [&](ParallelTensorSpaceCoordinate const &input_tensor_shard) + -> nonempty_set + { + bidict dependent_on_input_tensor_shard + = bidict_filter_values( + node_mapping.get_shard_bindings(), + [&](OperatorAtomicTaskShardBinding const &b) -> bool { + return ptensor_space_coord_for_slot_name(b, TensorSlotName::INPUT) == input_tensor_shard; + }); + + return nonempty_set(set_of(dependent_on_input_tensor_shard.left_values())); + }; + + auto invocation_sharding_info_for_input_tensor_shard = [&](ParallelTensorSpaceCoordinate const &c) + -> DynamicNodeInvocationShardingInfo + { + nonempty_set task_shard_machine_coords = + get_task_shard_machine_coords_for_input_tensor_shard(c); + + std::map output_sharding_infos = + generate_map(task_shard_machine_coords.unwrap_as_set(), + [&](MachineSpaceCoordinate const &mc) + -> DynamicValueAttrsShardingInfo + { + ParallelTensorSpaceCoordinate pc = output_value_mapping.at_r(mc); + + return DynamicValueAttrsShardingInfo{ + /*shard_coord=*/pc, + /*mapping=*/OneToMany{ + { + pc, + {mc}, + }, + }, + }; + }); + + std::map keyed_output_sharding_infos = + map_keys(output_sharding_infos, + [&](MachineSpaceCoordinate const &mc) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/mc, + }; + }); + + DynamicTensorSlot input_slot = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }; + + DynamicValueAttrsShardingInfo input_sharding_info = DynamicValueAttrsShardingInfo{ + /*shard_coord=*/c, + /*mapping=*/OneToMany{ + { + c, + {input_value_mapping.at_l(c)}, + }, + }, + }; + + std::map sharding_infos = + binary_merge_disjoint_maps( + keyed_output_sharding_infos, + std::map{ + { + input_slot, + input_sharding_info, + }, + }); + + return DynamicNodeInvocationShardingInfo{ + /*device_coords=*/task_shard_machine_coords, + /*value_sharding=*/sharding_infos, + }; + }; + + return transform(input_tensor_shards, invocation_sharding_info_for_input_tensor_shard); +} + std::unordered_set perform_shard_expansion_for_invocation(DynamicNodeInvocation const &i) { @@ -259,10 +381,10 @@ void require_graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const & DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( DynamicNodeAttrs const &node_attrs, - MachineSpaceCoordinate const &device_coord) + nonempty_set const &device_coords) { DynamicNodeAttrs result = node_attrs; - result.device_coord = device_coord; + result.device_coords = device_coords; return result; } @@ -283,14 +405,17 @@ DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( { require_invocation_is_ready_for_shard_expansion(invocation); - auto shard_value = [&](DynamicTensorSlot const &slot, DynamicValueAttrs const &value_attrs) -> DynamicValueAttrs { + auto shard_value = [&](DynamicTensorSlot const &slot, DynamicValueAttrs const &value_attrs) + -> DynamicValueAttrs + { DynamicValueAttrsShardingInfo sharding_info = invocation_sharding_info.value_sharding.at(slot); return apply_dynamic_value_attrs_sharding_info(value_attrs, sharding_info); }; DynamicNodeInvocation result = DynamicNodeInvocation{ /*inputs=*/map_values2(invocation.inputs, shard_value), - /*node_attrs=*/apply_dynamic_node_attrs_sharding_info(invocation.node_attrs, invocation_sharding_info.device_coord), + /*node_attrs=*/apply_dynamic_node_attrs_sharding_info( + invocation.node_attrs, invocation_sharding_info.device_coords), /*outputs=*/map_values2(invocation.outputs, shard_value), }; @@ -303,11 +428,14 @@ std::unordered_set { require_invocation_is_ready_for_shard_expansion(i); - if (i.node_attrs.op_attrs.has_value() && - i.node_attrs.op_attrs.value().is_copy()) { + if (i.node_attrs.op_attrs.value().is_copy()) { return unordered_set_of(generate_shard_expansion_for_copy(i)); } + if (training_op_attrs_has_op_type(i.node_attrs.op_attrs.value(), OperatorType::REPLICATE)) { + return unordered_set_of(generate_shard_expansion_for_fwd_replicate(i)); + } + MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); std::unordered_set shard_machine_coords = diff --git a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc index 58a32db6c1..51a79cff59 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc @@ -79,7 +79,7 @@ static DynamicNodeInvocation get_update_invocation_for_invocation( /*inputs=*/map_from_pairs( transform(tensor_roles, create_binding_for_role)), /*node_attrs=*/update_node_attrs, - /*outputs=*/std::unordered_map{}, + /*outputs=*/std::map{}, }; } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index 31de844555..05324c2195 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -28,6 +28,7 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }; }; diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index 49b8d4a77a..96523b6c31 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -58,33 +58,40 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ - {DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/std::nullopt, - }, - value_1}, + /*inputs=*/std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + value_1, + }, }, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ - {DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - }, - value_2}, + std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + value_2, + }, }, }; DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{}, + /*inputs=*/std::map{}, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, value_3, }, @@ -92,11 +99,12 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicNodeInvocation invocation_3 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::INPUT, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, value_1, }, @@ -104,6 +112,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::WEIGHT, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, value_2, }, @@ -111,12 +120,13 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::BIAS, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, value_1, }, }, /*node_attrs=*/node_attrs, - /*outputs=*/std::unordered_map{}, + /*outputs=*/std::map{}, }; std::unordered_set invocation_set = { diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc index 40b3460ee5..789d89a676 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc @@ -58,6 +58,7 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }; }; @@ -112,7 +113,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/std::nullopt, - /*device_coord=*/mc2, + /*device_coords=*/nonempty_set{mc2}, /*mapping=*/std::nullopt, /*op_attrs=*/std::nullopt, /*layer_guid=*/ @@ -139,7 +140,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/std::nullopt, - /*device_coord=*/mc1, + /*device_coord=*/nonempty_set{mc1}, /*mapping=*/std::nullopt, /*op_attrs=*/std::nullopt, /*layer_guid=*/ @@ -173,7 +174,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/std::nullopt, - /*device_coord=*/mc2, + /*device_coord=*/nonempty_set{mc2}, /*mapping=*/std::nullopt, /*op_attrs=*/std::nullopt, /*layer_guid=*/ diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 9f8aeee726..2e14c88654 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -11,7 +11,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Replicate") { MachineSpaceCoordinate gpu0 = MachineSpaceCoordinate{0_n, 0_n, DeviceType::GPU}; MachineSpaceCoordinate gpu1 = MachineSpaceCoordinate{0_n, 1_n, DeviceType::GPU}; - + ParallelTensorSpaceCoordinate tensor_coord0 = ParallelTensorSpaceCoordinate{ /*sum_component=*/0_n, /*discard_copy_component=*/0_n, @@ -78,7 +78,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargDataflowOutput{ Node{0}, TensorSlotName::OUTPUT, - }, + }, }; MappedParallelLayerInvocationInfo input = MappedParallelLayerInvocationInfo{ @@ -124,6 +124,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot{ TensorSlotName::INPUT, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, @@ -148,6 +149,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot{ TensorSlotName::OUTPUT, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, @@ -165,7 +167,7 @@ TEST_SUITE(FF_TEST_SUITE) { } // SUBCASE("standard op") { - // + // // } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc index bf88d5ec38..90fbdec5f7 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc @@ -32,6 +32,7 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/role, + /*task_shard=*/std::nullopt, }; }; @@ -136,6 +137,7 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/role, + /*task_shard=*/std::nullopt, }; }; @@ -352,37 +354,42 @@ TEST_SUITE(FF_TEST_SUITE) { std::unordered_set invocation_set = { DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{}, + /*inputs=*/std::map{}, /*node_attrs=*/n1, /*outputs=*/ - std::unordered_map{ + std::map{ { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, v1, }, }, }, DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ - {DynamicTensorSlot{ + /*inputs=*/std::map{ + { + DynamicTensorSlot{ /*slot_name=*/TensorSlotName::INPUT, /*slot_tensor_role=*/std::nullopt, - }, - v1}, + /*task_shard=*/std::nullopt, + }, + v1, + }, }, /*node_attrs=*/n2, /*outputs=*/ - std::unordered_map{ - {DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - }, - v2}, + std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + v2, + }, }, }, }; @@ -413,48 +420,51 @@ TEST_SUITE(FF_TEST_SUITE) { std::unordered_set invocation_set = { DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{}, + /*inputs=*/std::map{}, /*node_attrs=*/n1_fwd, /*outputs=*/ - std::unordered_map{ + std::map{ std::pair{ DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, v1_activation, }, }, }, DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ std::pair{ DynamicTensorSlot{ TensorSlotName::INPUT, mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, v1_activation, }, }, /*node_attrs=*/n2_fwd, /*outputs=*/ - std::unordered_map{ + std::map{ std::pair{ DynamicTensorSlot{ TensorSlotName::OUTPUT, mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, v2_activation, }, }, }, DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ std::pair{ DynamicTensorSlot{ TensorSlotName::INPUT, mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, v1_activation, }, @@ -462,6 +472,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot{ TensorSlotName::OUTPUT, mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, v2_activation, }, @@ -469,17 +480,19 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot{ TensorSlotName::OUTPUT, mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }, v2_gradient, }, }, /*node_attrs=*/n2_bwd, /*outputs=*/ - std::unordered_map{ + std::map{ std::pair{ DynamicTensorSlot{ TensorSlotName::INPUT, mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }, v1_gradient, }, diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index 19c21f5f89..9204e60c58 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -9,6 +9,8 @@ #include "op-attrs/ops/element_unary.h" #include "utils/one_to_many/one_to_many_filter_keys.h" #include "utils/one_to_many/one_to_many_filter_values.h" +#include "utils/containers/map_from_pairs.h" +#include "utils/containers/binary_merge_disjoint_maps.h" using namespace ::FlexFlow; @@ -36,10 +38,12 @@ static ParallelTensorSpaceCoordinate mk_pt_coord(nonnegative_int idx1, }; }; -DynamicTensorSlot mk_slot(TensorSlotName const &slot_name) { +DynamicTensorSlot mk_slot(TensorSlotName const &slot_name, + std::optional const &task_shard = std::nullopt) { return DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/task_shard, }; }; @@ -245,7 +249,7 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelTensorSpaceCoordinate const &output_2_shard_coord) -> DynamicNodeInvocationShardingInfo { return DynamicNodeInvocationShardingInfo{ - /*device_coord=*/device_coord, + /*device_coord=*/nonempty_set{device_coord}, /*value_sharding=*/{ mk_sharding_info(TensorSlotName::INPUT, input_shard_coord, mapped_task_group, device_coord), mk_sharding_info(TensorSlotName::WEIGHT, weight_shard_coord, mapped_task_group, device_coord), @@ -329,7 +333,7 @@ TEST_SUITE(FF_TEST_SUITE) { -> DynamicNodeInvocationShardingInfo { return DynamicNodeInvocationShardingInfo{ - /*device_coord=*/device_coord, + /*device_coord=*/nonempty_set{device_coord}, /*value_sharding=*/std::map{ { mk_slot(TensorSlotName::INPUT), @@ -360,6 +364,169 @@ TEST_SUITE(FF_TEST_SUITE) { mk_invocation_shard(mc2, pt2), }; + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } + + SUBCASE("replicate operator") { + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); + MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); + MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); + + ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 0_n, 0_n, 1_n); + ParallelTensorSpaceCoordinate pt3 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate pt4 = mk_pt_coord(0_n, 1_n, 0_n, 1_n); + + OneToMany src_binding{ + {pt1, {mc1}}, + {pt2, {mc2}}, + }; + + OneToMany dst_binding{ + {pt1, {mc1}}, + {pt2, {mc2}}, + {pt3, {mc3}}, + {pt4, {mc4}}, + }; + + auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, + ParallelTensorSpaceCoordinate const &c2) + -> OperatorAtomicTaskShardBinding { + return OperatorAtomicTaskShardBinding{ + /*tensor_coords=*/{ + { + TensorSlotName::INPUT, + c1, + }, + { + TensorSlotName::OUTPUT, + c2, + }, + }, + }; + }; + + MappedOperatorTaskGroup mapped_task_group = MappedOperatorTaskGroup{ + bidict{ + { + mc1, + mk_shard_binding(pt1, pt1), + }, + { + mc2, + mk_shard_binding(pt1, pt2), + }, + { + mc3, + mk_shard_binding(pt2, pt3), + }, + { + mc4, + mk_shard_binding(pt2, pt4), + }, + }, + }; + + DynamicNodeInvocation input = DynamicNodeInvocation{ + /*inputs=*/{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + mk_value(0, TensorSlotName::OUTPUT, src_binding, std::nullopt), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coords=*/std::nullopt, + /*mapping=*/mapped_task_group, + /*op_attrs=*/TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, + }, + }, + /*layer_guid=*/dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), + }, + }, + }; + + std::unordered_set result = + generate_shard_expansion_for_invocation(input); + + + auto mk_output_binding = [&](MachineSpaceCoordinate const &mc) + -> std::pair + { + return { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/mc, + }, + DynamicValueAttrsShardingInfo{ + dst_binding.at_r(mc), + one_to_many_filter_keys(dst_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == dst_binding.at_r(mc); + }), + }, + }; + }; + + auto mk_invocation_shard = + [&](nonempty_set const &device_coords, + ParallelTensorSpaceCoordinate const &input_shard_coord, + std::unordered_set const &output_task_shards) + -> DynamicNodeInvocationShardingInfo { + + return DynamicNodeInvocationShardingInfo{ + /*device_coords=*/device_coords, + /*value_sharding=*/ + binary_merge_disjoint_maps( + std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrsShardingInfo{ + input_shard_coord, + one_to_many_filter_keys( + src_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == input_shard_coord; + }), + }, + }, + }, + map_from_pairs(transform(output_task_shards, mk_output_binding))), + }; + }; + + std::unordered_set correct = { + mk_invocation_shard(nonempty_set{mc1, mc2}, pt1, {mc1, mc2}), + mk_invocation_shard(nonempty_set{mc3, mc4}, pt2, {mc3, mc4}), + }; + nlohmann::json result_json = result; nlohmann::json correct_json = correct; diff --git a/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h b/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h index 824fe77b39..85ff3d75fa 100644 --- a/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h +++ b/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h @@ -3,18 +3,20 @@ #include "utils/containers/binary_merge_maps_with.h" #include +#include "utils/containers/keys.h" +#include "utils/containers/intersection.h" namespace FlexFlow { template -std::unordered_map - binary_merge_disjoint_maps(std::unordered_map const &lhs, - std::unordered_map const &rhs) { +std::map + binary_merge_disjoint_maps(std::map const &lhs, + std::map const &rhs) { - std::unordered_set lhs_keys = unordered_keys(lhs); - std::unordered_set rhs_keys = unordered_keys(rhs); + std::set lhs_keys = keys(lhs); + std::set rhs_keys = keys(rhs); - std::unordered_set shared_keys = intersection(lhs_keys, rhs_keys); + std::set shared_keys = intersection(lhs_keys, rhs_keys); ASSERT(shared_keys.empty()); return binary_merge_maps_with( diff --git a/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h b/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h new file mode 100644 index 0000000000..536d402be0 --- /dev/null +++ b/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h @@ -0,0 +1,28 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_DISJOINT_UNORDERED_MAPS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_DISJOINT_UNORDERED_MAPS_H + +#include +#include "utils/containers/binary_merge_unordered_maps_with.h" +#include "utils/containers/unordered_keys.h" +#include "utils/containers/intersection.h" + +namespace FlexFlow { + +template +std::unordered_map + binary_merge_disjoint_unordered_maps(std::unordered_map const &lhs, + std::unordered_map const &rhs) { + + std::unordered_set lhs_keys = unordered_keys(lhs); + std::unordered_set rhs_keys = unordered_keys(rhs); + + std::unordered_set shared_keys = intersection(lhs_keys, rhs_keys); + ASSERT(shared_keys.empty()); + + return binary_merge_unordered_maps_with( + lhs, rhs, [](V const &, V const &) -> V { PANIC(); }); +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with.h b/lib/utils/include/utils/containers/binary_merge_maps_with.h index 2d0b57eb81..3c3b556830 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_MAPS_WITH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_MAPS_WITH_H -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/intersection.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/containers/merge_maps_with_right_dominating.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" @@ -12,22 +12,22 @@ namespace FlexFlow { template -std::unordered_map - binary_merge_maps_with(std::unordered_map const &lhs, - std::unordered_map const &rhs, +std::map + binary_merge_maps_with(std::map const &lhs, + std::map const &rhs, F &&f) { - std::unordered_set l_keys = unordered_keys(lhs); - std::unordered_set r_keys = unordered_keys(rhs); + std::set l_keys = keys(lhs); + std::set r_keys = keys(rhs); - std::unordered_set l_only_keys = set_minus(l_keys, r_keys); - std::unordered_set r_only_keys = set_minus(r_keys, l_keys); - std::unordered_set both_keys = intersection(r_keys, l_keys); + std::set l_only_keys = set_minus(l_keys, r_keys); + std::set r_only_keys = set_minus(r_keys, l_keys); + std::set both_keys = intersection(r_keys, l_keys); - std::unordered_map l_only = restrict_keys(lhs, l_only_keys); - std::unordered_map r_only = restrict_keys(rhs, r_only_keys); + std::map l_only = restrict_keys(lhs, l_only_keys); + std::map r_only = restrict_keys(rhs, r_only_keys); - std::unordered_map merged = generate_unordered_map( + std::map merged = generate_map( both_keys, [&](K const &k) { return f(lhs.at(k), rhs.at(k)); }); return merge_maps_with_right_dominating(std::vector{ diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h b/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h index f6e23af11c..25be62d9c3 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h @@ -6,9 +6,9 @@ namespace FlexFlow { template -std::unordered_map binary_merge_maps_with_left_dominating( - std::unordered_map const &lhs, std::unordered_map const &rhs) { - std::unordered_map result; +std::map binary_merge_maps_with_left_dominating( + std::map const &lhs, std::map const &rhs) { + std::map result; merge_in_map(rhs, result); merge_in_map(lhs, result); return result; diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h b/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h index e5e29dfcb9..e4bfdd6d29 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h @@ -6,9 +6,9 @@ namespace FlexFlow { template -std::unordered_map binary_merge_maps_with_right_dominating( - std::unordered_map const &lhs, std::unordered_map const &rhs) { - std::unordered_map result; +std::map binary_merge_maps_with_right_dominating( + std::map const &lhs, std::map const &rhs) { + std::map result; merge_in_map(lhs, result); merge_in_map(rhs, result); return result; diff --git a/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h new file mode 100644 index 0000000000..2f8be45802 --- /dev/null +++ b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h @@ -0,0 +1,42 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_H + +#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/intersection.h" +#include "utils/containers/unordered_keys.h" +#include "utils/containers/merge_unordered_maps_with_right_dominating.h" +#include "utils/containers/restrict_keys.h" +#include "utils/containers/set_minus.h" +#include + +namespace FlexFlow { + +template +std::unordered_map + binary_merge_unordered_maps_with(std::unordered_map const &lhs, + std::unordered_map const &rhs, + F &&f) { + + std::unordered_set l_keys = unordered_keys(lhs); + std::unordered_set r_keys = unordered_keys(rhs); + + std::unordered_set l_only_keys = set_minus(l_keys, r_keys); + std::unordered_set r_only_keys = set_minus(r_keys, l_keys); + std::unordered_set both_keys = intersection(r_keys, l_keys); + + std::unordered_map l_only = restrict_keys(lhs, l_only_keys); + std::unordered_map r_only = restrict_keys(rhs, r_only_keys); + + std::unordered_map merged = generate_unordered_map( + both_keys, [&](K const &k) { return f(lhs.at(k), rhs.at(k)); }); + + return merge_unordered_maps_with_right_dominating(std::vector{ + l_only, + r_only, + merged, + }); +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/binary_merge_unordered_maps_with_left_dominating.h b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with_left_dominating.h new file mode 100644 index 0000000000..0d71a1b7cc --- /dev/null +++ b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with_left_dominating.h @@ -0,0 +1,19 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_LEFT_DOMINATING_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_LEFT_DOMINATING_H + +#include "utils/containers/merge_in_unordered_map.h" + +namespace FlexFlow { + +template +std::unordered_map binary_merge_unordered_maps_with_left_dominating( + std::unordered_map const &lhs, std::unordered_map const &rhs) { + std::unordered_map result; + merge_in_unordered_map(rhs, result); + merge_in_unordered_map(lhs, result); + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/binary_merge_unordered_maps_with_right_dominating.h b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with_right_dominating.h new file mode 100644 index 0000000000..6a3581a635 --- /dev/null +++ b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with_right_dominating.h @@ -0,0 +1,19 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_RIGHT_DOMINATING_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_RIGHT_DOMINATING_H + +#include "utils/containers/merge_in_unordered_map.h" + +namespace FlexFlow { + +template +std::unordered_map binary_merge_unordered_maps_with_right_dominating( + std::unordered_map const &lhs, std::unordered_map const &rhs) { + std::unordered_map result; + merge_in_unordered_map(lhs, result); + merge_in_unordered_map(rhs, result); + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/flatmap.h b/lib/utils/include/utils/containers/flatmap.h index 70de6b5020..e0440cc791 100644 --- a/lib/utils/include/utils/containers/flatmap.h +++ b/lib/utils/include/utils/containers/flatmap.h @@ -1,12 +1,12 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_FLATMAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_FLATMAP_H -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/extend.h" #include "utils/containers/get_element_type.h" #include #include #include +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -77,7 +77,7 @@ std::unordered_map flatmap(std::unordered_map const &m, std::unordered_map result; for (auto const &[k, v] : m) { - result = binary_merge_disjoint_maps(result, f(k, v)); + result = binary_merge_disjoint_unordered_maps(result, f(k, v)); } return result; diff --git a/lib/utils/include/utils/containers/get_only.h b/lib/utils/include/utils/containers/get_only.h index ed44b26c36..d5e2faf67a 100644 --- a/lib/utils/include/utils/containers/get_only.h +++ b/lib/utils/include/utils/containers/get_only.h @@ -2,17 +2,15 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GET_ONLY_H #include "utils/containers/maybe_get_only.h" -#include "utils/exception.h" +#include #include "utils/optional.h" namespace FlexFlow { template typename C::value_type get_only(C const &c) { - return unwrap(maybe_get_only(c), [&] { - throw mk_runtime_error(fmt::format( - "Encountered container with size {} in get_only", c.size())); - }); + ASSERT(c.size() == 1); + return maybe_get_only(c).value(); } template diff --git a/lib/utils/include/utils/containers/map_from_pairs.h b/lib/utils/include/utils/containers/map_from_pairs.h index 7c470d4d3e..f5f2dad415 100644 --- a/lib/utils/include/utils/containers/map_from_pairs.h +++ b/lib/utils/include/utils/containers/map_from_pairs.h @@ -1,18 +1,15 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_FROM_PAIRS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_FROM_PAIRS_H -#include -#include +#include namespace FlexFlow { -template -std::unordered_map - map_from_pairs(std::unordered_set> const &pairs) { - - std::unordered_map result(pairs.cbegin(), pairs.cend()); - - return result; +template +std::map map_from_pairs(C const &c) { + return std::map(c.cbegin(), c.cend()); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/containers/map_keys.h b/lib/utils/include/utils/containers/map_keys.h index 5cd44d8a5d..ff41248a30 100644 --- a/lib/utils/include/utils/containers/map_keys.h +++ b/lib/utils/include/utils/containers/map_keys.h @@ -2,11 +2,13 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_KEYS_H #include "utils/containers/keys.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/transform.h" #include "utils/containers/unordered_multiset_of.h" #include "utils/exception.h" #include #include +#include namespace FlexFlow { @@ -19,14 +21,35 @@ template > std::unordered_map map_keys(std::unordered_map const &m, - F const &f) { + F &&f) { std::unordered_map result; for (auto const &kv : m) { result.insert({f(kv.first), kv.second}); } - ASSERT(keys(m).size() == keys(result).size(), + ASSERT(m.size() == result.size(), + "keys passed to map_keys must be transformed into distinct keys"); + + return result; +} + +/** + * @brief Applies the given function to all the keys within the given map and + * returns the updated map. + */ +template > +std::map map_keys(std::map const &m, F &&f) { + + std::map result; + for (auto const &kv : m) { + result.insert({f(kv.first), kv.second}); + } + + ASSERT(m.size() == result.size(), "keys passed to map_keys must be transformed into distinct keys"); return result; diff --git a/lib/utils/include/utils/containers/map_values2.h b/lib/utils/include/utils/containers/map_values2.h index 752a8babd3..dd943b02bb 100644 --- a/lib/utils/include/utils/containers/map_values2.h +++ b/lib/utils/include/utils/containers/map_values2.h @@ -3,6 +3,7 @@ #include #include +#include namespace FlexFlow { @@ -19,6 +20,20 @@ std::unordered_map map_values2(std::unordered_map const &m, return result; } +template > +std::map map_values2(std::map const &m, + F &&f) { + std::map result; + for (std::pair const &kv : m) { + result.insert(std::pair{kv.first, f(kv.first, kv.second)}); + } + return result; +} + + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/merge_disjoint_maps.h b/lib/utils/include/utils/containers/merge_disjoint_maps.h index eccb06180a..b541fdbd53 100644 --- a/lib/utils/include/utils/containers/merge_disjoint_maps.h +++ b/lib/utils/include/utils/containers/merge_disjoint_maps.h @@ -9,12 +9,12 @@ namespace FlexFlow { template -std::unordered_map merge_disjoint_maps(C const &c) { - std::unordered_map empty = {}; +std::map merge_disjoint_maps(C const &c) { + std::map empty = {}; return foldl(c, /*init=*/empty, - [](std::unordered_map const &lhs, - std::unordered_map const &rhs) { + [](std::map const &lhs, + std::map const &rhs) { return binary_merge_disjoint_maps(lhs, rhs); }); } diff --git a/lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h b/lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h new file mode 100644 index 0000000000..1bd7fb019b --- /dev/null +++ b/lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h @@ -0,0 +1,24 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_DISJOINT_UNORDERED_MAPS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_DISJOINT_UNORDERED_MAPS_H + +#include "utils/containers/foldl.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" + +namespace FlexFlow { + +template +std::unordered_map merge_disjoint_unordered_maps(C const &c) { + std::unordered_map empty = {}; + return foldl(c, + /*init=*/empty, + [](std::unordered_map const &lhs, + std::unordered_map const &rhs) { + return binary_merge_disjoint_unordered_maps(lhs, rhs); + }); +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/merge_in_map.h b/lib/utils/include/utils/containers/merge_in_map.h index edae4b8a6a..e41c1a6826 100644 --- a/lib/utils/include/utils/containers/merge_in_map.h +++ b/lib/utils/include/utils/containers/merge_in_map.h @@ -1,13 +1,13 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_IN_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_IN_MAP_H -#include +#include namespace FlexFlow { template -void merge_in_map(std::unordered_map const &m, - std::unordered_map &result) { +void merge_in_map(std::map const &m, + std::map &result) { for (auto const &[k, v] : m) { auto it = result.find(k); if (it != result.end()) { diff --git a/lib/utils/include/utils/containers/merge_in_unordered_map.h b/lib/utils/include/utils/containers/merge_in_unordered_map.h new file mode 100644 index 0000000000..7c2b31b8fc --- /dev/null +++ b/lib/utils/include/utils/containers/merge_in_unordered_map.h @@ -0,0 +1,23 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_IN_UNORDERED_MAPS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_IN_UNORDERED_MAPS_H + +#include + +namespace FlexFlow { + +template +void merge_in_unordered_map(std::unordered_map const &m, + std::unordered_map &result) { + for (auto const &[k, v] : m) { + auto it = result.find(k); + if (it != result.end()) { + it->second = v; + } else { + result.insert({k, v}); + } + } +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/merge_maps_with.h b/lib/utils/include/utils/containers/merge_maps_with.h index 2f5a09e26e..eef6c9af67 100644 --- a/lib/utils/include/utils/containers/merge_maps_with.h +++ b/lib/utils/include/utils/containers/merge_maps_with.h @@ -9,13 +9,13 @@ namespace FlexFlow { template -std::unordered_map - merge_maps_with(std::vector> const &to_merge, +std::map + merge_maps_with(std::vector> const &to_merge, F &&f) { return foldl(to_merge, - std::unordered_map{}, - [&](std::unordered_map const &accum, - std::unordered_map const &m) { + std::map{}, + [&](std::map const &accum, + std::map const &m) { return binary_merge_maps_with(accum, m, f); }); } diff --git a/lib/utils/include/utils/containers/merge_maps_with_right_dominating.h b/lib/utils/include/utils/containers/merge_maps_with_right_dominating.h index 1d4f8536d8..6271cef2c0 100644 --- a/lib/utils/include/utils/containers/merge_maps_with_right_dominating.h +++ b/lib/utils/include/utils/containers/merge_maps_with_right_dominating.h @@ -8,10 +8,10 @@ namespace FlexFlow { template -std::unordered_map merge_maps_with_right_dominating(C const &c) { - std::unordered_map result; +std::map merge_maps_with_right_dominating(C const &c) { + std::map result; - for (std::unordered_map const &m : c) { + for (std::map const &m : c) { merge_in_map(m, result); } diff --git a/lib/utils/include/utils/containers/merge_unordered_maps_with.h b/lib/utils/include/utils/containers/merge_unordered_maps_with.h new file mode 100644 index 0000000000..fee7fa2fa4 --- /dev/null +++ b/lib/utils/include/utils/containers/merge_unordered_maps_with.h @@ -0,0 +1,25 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_UNORDERED_MAPS_WITH_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_UNORDERED_MAPS_WITH_H + +#include "utils/containers/binary_merge_unordered_maps_with.h" +#include "utils/containers/foldl.h" +#include +#include + +namespace FlexFlow { + +template +std::unordered_map + merge_unordered_maps_with(std::vector> const &to_merge, + F &&f) { + return foldl(to_merge, + std::unordered_map{}, + [&](std::unordered_map const &accum, + std::unordered_map const &m) { + return binary_merge_unordered_maps_with(accum, m, f); + }); +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h b/lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h new file mode 100644 index 0000000000..1323378019 --- /dev/null +++ b/lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h @@ -0,0 +1,23 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_UNORDERED_MAPS_WITH_RIGHT_DOMINATING_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_UNORDERED_MAPS_WITH_RIGHT_DOMINATING_H + +#include "utils/containers/merge_in_unordered_map.h" + +namespace FlexFlow { + +template +std::unordered_map merge_unordered_maps_with_right_dominating(C const &c) { + std::unordered_map result; + + for (std::unordered_map const &m : c) { + merge_in_unordered_map(m, result); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/restrict_keys.h b/lib/utils/include/utils/containers/restrict_keys.h index bedcc4ed8e..353b2ec237 100644 --- a/lib/utils/include/utils/containers/restrict_keys.h +++ b/lib/utils/include/utils/containers/restrict_keys.h @@ -4,6 +4,8 @@ #include "utils/containers/contains.h" #include #include +#include +#include namespace FlexFlow { @@ -19,6 +21,17 @@ std::unordered_map restrict_keys(std::unordered_map const &m, return result; } +template +std::map restrict_keys(std::map const &m, + std::set const &mask) { + std::map result; + for (auto const &kv : m) { + if (contains(mask, kv.first)) { + result.insert(kv); + } + } + return result; +} } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/zip_values_strict.h b/lib/utils/include/utils/containers/zip_values_strict.h index 1a3ce95eb1..b891490e39 100644 --- a/lib/utils/include/utils/containers/zip_values_strict.h +++ b/lib/utils/include/utils/containers/zip_values_strict.h @@ -6,6 +6,9 @@ #include "utils/containers/require_same.h" #include #include +#include +#include "utils/containers/keys.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { @@ -24,6 +27,21 @@ std::unordered_map> }); } +template +std::map> + zip_values_strict(std::map const &m1, + std::map const &m2) { + + ASSERT(keys(m1) == keys(m2)); + + return generate_map(require_same(keys(m1), keys(m2)), [&](K const &k) { + return std::pair{ + m1.at(k), + m2.at(k), + }; + }); +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h b/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h index fd77509e4d..e3947605c0 100644 --- a/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h +++ b/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_FULL_BINARY_TREE_GET_PATH_TO_LEAF_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_FULL_BINARY_TREE_GET_PATH_TO_LEAF_MAP_H -#include "utils/containers/binary_merge_disjoint_maps.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" #include "utils/containers/map_keys.h" #include "utils/containers/multiset_union.h" #include "utils/full_binary_tree/binary_tree_path.dtg.h" @@ -30,7 +30,7 @@ std::unordered_map get_path_to_leaf_map( get_path_to_leaf_map(impl.get_right_child(parent), impl), [](BinaryTreePath const &p) { return nest_inside_right_child(p); }); - return binary_merge_disjoint_maps(left_map, right_map); + return binary_merge_disjoint_unordered_maps(left_map, right_map); }, [](Leaf const &leaf) -> std::unordered_map { return std::unordered_map{ diff --git a/lib/utils/include/utils/nonempty_set/nonempty_set.h b/lib/utils/include/utils/nonempty_set/nonempty_set.h index 93da743592..fe4b152bd5 100644 --- a/lib/utils/include/utils/nonempty_set/nonempty_set.h +++ b/lib/utils/include/utils/nonempty_set/nonempty_set.h @@ -8,6 +8,8 @@ #include "utils/fmt/set.h" #include "utils/positive_int/positive_int.h" #include "utils/containers/unordered_set_of.h" +#include "utils/json/check_is_json_deserializable.h" +#include "utils/json/check_is_json_serializable.h" namespace FlexFlow { @@ -76,21 +78,24 @@ struct nonempty_set { return unordered_set_of(this->raw); } + using const_iterator = typename std::set::const_iterator; using value_type = T; + using reference = value_type &; + using const_reference = value_type const &; - typename std::set::const_iterator begin() const { + const_iterator begin() const { return this->raw.cbegin(); } - typename std::set::const_iterator cbegin() const { + const_iterator cbegin() const { return this->raw.cbegin(); } - typename std::set::const_iterator end() const { + const_iterator end() const { return this->raw.cend(); } - typename std::set::const_iterator cend() const { + const_iterator cend() const { return this->raw.cend(); } @@ -122,6 +127,27 @@ std::ostream &operator<<(std::ostream &s, nonempty_set const &m) { } // namespace FlexFlow +namespace nlohmann { + +template +struct adl_serializer<::FlexFlow::nonempty_set> { + static ::FlexFlow::nonempty_set from_json(json const &j) { + CHECK_IS_JSON_DESERIALIZABLE(T); + + std::set s = j; + + return ::FlexFlow::nonempty_set{s}; + } + + static void to_json(json &j, ::FlexFlow::nonempty_set const &s) { + CHECK_IS_JSON_SERIALIZABLE(T); + + j = s.unwrap_as_set(); + } +}; + +} // namespace nlohmann + namespace std { template diff --git a/lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h b/lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h new file mode 100644 index 0000000000..9d31dc968c --- /dev/null +++ b/lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h @@ -0,0 +1,23 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_REQUIRE_ONE_TO_MANY_IS_BIJECTION_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_REQUIRE_ONE_TO_MANY_IS_BIJECTION_H + +#include "utils/bidict/algorithms/bidict_from_map.h" +#include "utils/containers/map_values.h" +#include "utils/containers/get_only.h" +#include "utils/one_to_many/one_to_many.h" +#include "utils/nonempty_set/nonempty_set.h" + +namespace FlexFlow { + +template +bidict require_one_to_many_is_bijection(OneToMany const &otm) { + return bidict_from_map( + map_values(otm.l_to_r(), + [](nonempty_set const &s) -> R { + return get_only(s.unwrap_as_set()); + })); +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/orthotope/minimal_dim_domain.h b/lib/utils/include/utils/orthotope/minimal_dim_domain.h index c9d1214278..f9bf5fa979 100644 --- a/lib/utils/include/utils/orthotope/minimal_dim_domain.h +++ b/lib/utils/include/utils/orthotope/minimal_dim_domain.h @@ -2,7 +2,6 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ORTHOTOPE_MINIMAL_DIM_DOMAIN_H #include "utils/containers/are_disjoint.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/filtermap_values.h" #include "utils/containers/generate_unordered_map.h" #include "utils/containers/map_from_keys_and_values.h" @@ -16,6 +15,7 @@ #include "utils/orthotope/minimal_dim_domain.dtg.h" #include "utils/orthotope/minimal_orthotope.dtg.h" #include "utils/containers/unordered_keys.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -66,7 +66,7 @@ DimDomain dim_domain_from_minimal_dim_domain( ASSERT(are_disjoint(nontrivial_dims, trivial_dims)); return DimDomain{ - /*dims=*/binary_merge_disjoint_maps( + /*dims=*/binary_merge_disjoint_unordered_maps( map_values( minimal_dim_domain.dims, [](int_ge_two x) { return x.positive_int_from_int_ge_two(); }), diff --git a/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc b/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc index 0569b3ed0b..95d19083ac 100644 --- a/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc +++ b/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc @@ -1,13 +1,14 @@ #include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; -template std::unordered_map - binary_merge_disjoint_maps(std::unordered_map const &, - std::unordered_map const &); +template std::map + binary_merge_disjoint_maps(std::map const &, + std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_disjoint_unordered_maps.cc b/lib/utils/src/utils/containers/binary_merge_disjoint_unordered_maps.cc new file mode 100644 index 0000000000..60ccf3a7e5 --- /dev/null +++ b/lib/utils/src/utils/containers/binary_merge_disjoint_unordered_maps.cc @@ -0,0 +1,13 @@ +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template std::unordered_map + binary_merge_disjoint_unordered_maps(std::unordered_map const &, + std::unordered_map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_maps_with.cc b/lib/utils/src/utils/containers/binary_merge_maps_with.cc index 35f771f60c..4679d21227 100644 --- a/lib/utils/src/utils/containers/binary_merge_maps_with.cc +++ b/lib/utils/src/utils/containers/binary_merge_maps_with.cc @@ -1,13 +1,14 @@ #include "utils/containers/binary_merge_maps_with.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; using F = std::function; -template std::unordered_map binary_merge_maps_with( - std::unordered_map const &, std::unordered_map const &, F &&); +template std::map binary_merge_maps_with( + std::map const &, std::map const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc b/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc index c459e82061..d5b4f6cebe 100644 --- a/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc @@ -1,13 +1,14 @@ #include "utils/containers/binary_merge_maps_with_left_dominating.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; -template std::unordered_map - binary_merge_maps_with_left_dominating(std::unordered_map const &, - std::unordered_map const &); +template std::map + binary_merge_maps_with_left_dominating(std::map const &, + std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc index df934387d2..bbc799150e 100644 --- a/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc @@ -1,13 +1,14 @@ #include "utils/containers/binary_merge_maps_with_right_dominating.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; -template std::unordered_map - binary_merge_maps_with_right_dominating(std::unordered_map const &, - std::unordered_map const &); +template std::map + binary_merge_maps_with_right_dominating(std::map const &, + std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with.cc b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with.cc new file mode 100644 index 0000000000..de8c484d0c --- /dev/null +++ b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with.cc @@ -0,0 +1,13 @@ +#include "utils/containers/binary_merge_unordered_maps_with.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; +using F = std::function; + +template std::unordered_map binary_merge_unordered_maps_with( + std::unordered_map const &, std::unordered_map const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc new file mode 100644 index 0000000000..d777eb0e29 --- /dev/null +++ b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc @@ -0,0 +1,13 @@ +#include "utils/containers/binary_merge_unordered_maps_with_left_dominating.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template std::unordered_map + binary_merge_unordered_maps_with_left_dominating(std::unordered_map const &, + std::unordered_map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc new file mode 100644 index 0000000000..f5586cec6b --- /dev/null +++ b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc @@ -0,0 +1,13 @@ +#include "utils/containers/binary_merge_unordered_maps_with_right_dominating.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template + std::unordered_map binary_merge_unordered_maps_with_right_dominating( + std::unordered_map const &, std::unordered_map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_from_pairs.cc b/lib/utils/src/utils/containers/map_from_pairs.cc index ba0eed8c15..8dc8ffa29c 100644 --- a/lib/utils/src/utils/containers/map_from_pairs.cc +++ b/lib/utils/src/utils/containers/map_from_pairs.cc @@ -1,12 +1,21 @@ #include "utils/containers/map_from_pairs.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" +#include +#include +#include namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; -template std::unordered_map +template std::map + map_from_pairs(std::set> const &); + +template std::map map_from_pairs(std::unordered_set> const &); +template std::map + map_from_pairs(std::vector> const &); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_keys.cc b/lib/utils/src/utils/containers/map_keys.cc index 7473c7e16d..daf9ed25d4 100644 --- a/lib/utils/src/utils/containers/map_keys.cc +++ b/lib/utils/src/utils/containers/map_keys.cc @@ -1 +1,22 @@ #include "utils/containers/map_keys.h" +#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using VT0 = value_type<0>; +using VT1 = value_type<1>; +using VT2 = value_type<2>; + +template + std::unordered_map map_keys(std::unordered_map const &, + std::function &&); + +using OV0 = ordered_value_type<0>; +using OV1 = ordered_value_type<1>; + +template + std::map map_keys(std::map const &m, + std::function &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_values2.cc b/lib/utils/src/utils/containers/map_values2.cc index 6aba8f4db0..8840f15aee 100644 --- a/lib/utils/src/utils/containers/map_values2.cc +++ b/lib/utils/src/utils/containers/map_values2.cc @@ -1,14 +1,21 @@ #include "utils/containers/map_values2.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; -using V2 = value_type<2>; -using F = std::function; +using VT0 = value_type<0>; +using VT1 = value_type<1>; +using VT2 = value_type<2>; -template std::unordered_map map_values2(std::unordered_map const &, - F &&); +template std::unordered_map map_values2( + std::unordered_map const &, + std::function &&); + +using OT0 = ordered_value_type<0>; + +template std::map map_values2( + std::map const &, + std::function &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_disjoint_maps.cc b/lib/utils/src/utils/containers/merge_disjoint_maps.cc index dec8ee0618..e810b6b4f0 100644 --- a/lib/utils/src/utils/containers/merge_disjoint_maps.cc +++ b/lib/utils/src/utils/containers/merge_disjoint_maps.cc @@ -1,12 +1,13 @@ #include "utils/containers/merge_disjoint_maps.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; -using C = std::vector>; +using C = std::vector>; -template std::unordered_map merge_disjoint_maps(C const &); +template std::map merge_disjoint_maps(C const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc b/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc new file mode 100644 index 0000000000..1ef6af0877 --- /dev/null +++ b/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc @@ -0,0 +1,12 @@ +#include "utils/containers/merge_disjoint_unordered_maps.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; +using C = std::vector>; + +template std::unordered_map merge_disjoint_unordered_maps(C const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_in_map.cc b/lib/utils/src/utils/containers/merge_in_map.cc index ada1a803ad..618128ff3b 100644 --- a/lib/utils/src/utils/containers/merge_in_map.cc +++ b/lib/utils/src/utils/containers/merge_in_map.cc @@ -1,12 +1,12 @@ #include "utils/containers/merge_in_map.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; -template void merge_in_map(std::unordered_map const &, - std::unordered_map &); +template void merge_in_map(std::map const &, + std::map &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_in_unordered_map.cc b/lib/utils/src/utils/containers/merge_in_unordered_map.cc new file mode 100644 index 0000000000..95228dc5f2 --- /dev/null +++ b/lib/utils/src/utils/containers/merge_in_unordered_map.cc @@ -0,0 +1,12 @@ +#include "utils/containers/merge_in_unordered_map.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template void merge_in_unordered_map(std::unordered_map const &, + std::unordered_map &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_maps_with.cc b/lib/utils/src/utils/containers/merge_maps_with.cc index 0375b16bc4..b5b471a1f8 100644 --- a/lib/utils/src/utils/containers/merge_maps_with.cc +++ b/lib/utils/src/utils/containers/merge_maps_with.cc @@ -1,13 +1,14 @@ #include "utils/containers/merge_maps_with.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; using F = std::function; -template std::unordered_map - merge_maps_with(std::vector> const &, F &&); +template std::map + merge_maps_with(std::vector> const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc index d8c269d6e9..f33b46780c 100644 --- a/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc +++ b/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc @@ -1,12 +1,13 @@ #include "utils/containers/merge_maps_with_right_dominating.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; -using C = std::vector>; +using C = std::vector>; -template std::unordered_map merge_maps_with_right_dominating(C const &); +template std::map merge_maps_with_right_dominating(C const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_unordered_maps_with.cc b/lib/utils/src/utils/containers/merge_unordered_maps_with.cc new file mode 100644 index 0000000000..60218312f3 --- /dev/null +++ b/lib/utils/src/utils/containers/merge_unordered_maps_with.cc @@ -0,0 +1,14 @@ +#include "utils/containers/merge_unordered_maps_with.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; +using F = std::function; + +std::unordered_map + merge_unordered_maps_with(std::vector> const &, + F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc new file mode 100644 index 0000000000..1dd7da70a3 --- /dev/null +++ b/lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc @@ -0,0 +1,12 @@ +#include "utils/containers/merge_unordered_maps_with_right_dominating.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; +using C = std::vector>; + +template std::unordered_map merge_unordered_maps_with_right_dominating(C const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/restrict_keys.cc b/lib/utils/src/utils/containers/restrict_keys.cc index d2749b7ea2..13584abec1 100644 --- a/lib/utils/src/utils/containers/restrict_keys.cc +++ b/lib/utils/src/utils/containers/restrict_keys.cc @@ -1 +1,21 @@ #include "utils/containers/restrict_keys.h" +#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using VT0 = value_type<0>; +using VT1 = value_type<1>; + +template + std::unordered_map restrict_keys(std::unordered_map const &, + std::unordered_set const &); + +using OV0 = ordered_value_type<0>; + +template + std::map restrict_keys(std::map const &, + std::set const &); + + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/zip_values_strict.cc b/lib/utils/src/utils/containers/zip_values_strict.cc index b9bed29a1b..c1710129a5 100644 --- a/lib/utils/src/utils/containers/zip_values_strict.cc +++ b/lib/utils/src/utils/containers/zip_values_strict.cc @@ -1,14 +1,23 @@ #include "utils/containers/zip_values_strict.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V1 = value_type<1>; -using V2 = value_type<2>; +using VT0 = value_type<0>; +using VT1 = value_type<1>; +using VT2 = value_type<2>; + +template std::unordered_map> + zip_values_strict(std::unordered_map const &, + std::unordered_map const &); + +using OV0 = ordered_value_type<0>; + +template std::map> + zip_values_strict(std::map const &, + std::map const &); + -template std::unordered_map> - zip_values_strict(std::unordered_map const &, - std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/nonempty_set/nonempty_set.cc b/lib/utils/src/utils/nonempty_set/nonempty_set.cc index 1af2951f10..a7bd7f7c5a 100644 --- a/lib/utils/src/utils/nonempty_set/nonempty_set.cc +++ b/lib/utils/src/utils/nonempty_set/nonempty_set.cc @@ -1,7 +1,9 @@ #include "utils/nonempty_set/nonempty_set.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" using T = ::FlexFlow::ordered_value_type<0>; +using J = ::FlexFlow::jsonable_ordered_value_type<0>; namespace FlexFlow { @@ -16,6 +18,12 @@ template std::ostream &operator<<(std::ostream &, nonempty_set const &); } // namespace FlexFlow +namespace nlohmann { + +template struct adl_serializer<::FlexFlow::nonempty_set>; + +} // namespace nlohmann + namespace std { template struct hash<::FlexFlow::nonempty_set>; diff --git a/lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc b/lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc new file mode 100644 index 0000000000..60a78bb123 --- /dev/null +++ b/lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc @@ -0,0 +1,12 @@ +#include "utils/one_to_many/require_one_to_many_is_bijection.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template + bidict require_one_to_many_is_bijection(OneToMany const &); + +} // namespace FlexFlow From 9bea615de80e2ce0d1fb7f5b9c397b274ad36c5a Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 29 May 2026 03:19:00 -0700 Subject: [PATCH 18/35] Pass shard expansion test for bwd replicate --- .../dynamic_graph/shard_expansion.cc | 129 +++++++- .../dynamic_graph/shard_expansion.cc | 312 ++++++++++++------ lib/utils/include/utils/bidict/bidict.h | 3 + 3 files changed, 343 insertions(+), 101 deletions(-) diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index db9484e361..f910c6c6c2 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -226,6 +226,10 @@ static std::set }); } +// TODO(@lockshaw): There is a lot of code duplication between +// generate_shard_expansion_for_fwd_replicate and +// generate_shard_expansion_for_bwd_replicate that should eventually be +// factored out. static std::set generate_shard_expansion_for_fwd_replicate(DynamicNodeInvocation const &i) { ASSERT(i.node_attrs.task_type == DynamicTaskType::FWD); @@ -341,6 +345,121 @@ static std::set return transform(input_tensor_shards, invocation_sharding_info_for_input_tensor_shard); } +static std::set + generate_shard_expansion_for_bwd_replicate(DynamicNodeInvocation const &i) { + ASSERT(i.node_attrs.task_type == DynamicTaskType::BWD); + + MappedOperatorTaskGroup node_mapping = assert_unwrap(i.node_attrs.mapping); + + DynamicTensorSlot expected_output_grad_slot = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }; + + DynamicValueAttrs output_grad = require_only_key(i.inputs, expected_output_grad_slot); + + DynamicTensorSlot expected_input_grad_slot = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }; + + DynamicValueAttrs input_grad = require_only_key(i.outputs, expected_input_grad_slot); + + bidict + output_grad_value_mapping = require_one_to_many_is_bijection( + assert_unwrap(output_grad.mapping)); + + bidict + input_grad_value_mapping = require_one_to_many_is_bijection( + assert_unwrap(input_grad.mapping)); + + std::set input_grad_tensor_shards = set_of(input_grad_value_mapping.left_values()); + + auto get_task_shard_machine_coords_for_input_grad_tensor_shard + = [&](ParallelTensorSpaceCoordinate const &input_grad_tensor_shard) + -> nonempty_set + { + bidict produce_input_grad_tensor_shard + = bidict_filter_values( + node_mapping.get_shard_bindings(), + [&](OperatorAtomicTaskShardBinding const &b) -> bool { + return ptensor_space_coord_for_slot_name(b, TensorSlotName::INPUT) == input_grad_tensor_shard; + }); + + return nonempty_set(set_of(produce_input_grad_tensor_shard.left_values())); + }; + + auto invocation_sharding_info_for_input_grad_tensor_shard = [&](ParallelTensorSpaceCoordinate const &c) + -> DynamicNodeInvocationShardingInfo + { + nonempty_set task_shard_machine_coords = + get_task_shard_machine_coords_for_input_grad_tensor_shard(c); + + std::map output_grad_sharding_infos = + generate_map(task_shard_machine_coords.unwrap_as_set(), + [&](MachineSpaceCoordinate const &mc) + -> DynamicValueAttrsShardingInfo + { + ParallelTensorSpaceCoordinate pc = output_grad_value_mapping.at_r(mc); + + return DynamicValueAttrsShardingInfo{ + /*shard_coord=*/pc, + /*mapping=*/OneToMany{ + { + pc, + {mc}, + }, + }, + }; + }); + + std::map keyed_output_grad_sharding_infos = + map_keys(output_grad_sharding_infos, + [&](MachineSpaceCoordinate const &mc) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/mc, + }; + }); + + DynamicTensorSlot input_grad_slot = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }; + + DynamicValueAttrsShardingInfo input_grad_sharding_info = DynamicValueAttrsShardingInfo{ + /*shard_coord=*/c, + /*mapping=*/OneToMany{ + { + c, + {input_grad_value_mapping.at_l(c)}, + }, + }, + }; + + std::map sharding_infos = + binary_merge_disjoint_maps( + keyed_output_grad_sharding_infos, + std::map{ + { + input_grad_slot, + input_grad_sharding_info, + }, + }); + + return DynamicNodeInvocationShardingInfo{ + /*device_coords=*/task_shard_machine_coords, + /*value_sharding=*/sharding_infos, + }; + }; + + return transform(input_grad_tensor_shards, invocation_sharding_info_for_input_grad_tensor_shard); +} + std::unordered_set perform_shard_expansion_for_invocation(DynamicNodeInvocation const &i) { @@ -433,7 +552,15 @@ std::unordered_set } if (training_op_attrs_has_op_type(i.node_attrs.op_attrs.value(), OperatorType::REPLICATE)) { - return unordered_set_of(generate_shard_expansion_for_fwd_replicate(i)); + DynamicTaskType task_type = assert_unwrap(i.node_attrs.task_type); + switch (task_type) { + case DynamicTaskType::FWD: + return unordered_set_of(generate_shard_expansion_for_fwd_replicate(i)); + case DynamicTaskType::BWD: + return unordered_set_of(generate_shard_expansion_for_bwd_replicate(i)); + default: + PANIC("Unexpected task type for Replicate: {}", task_type); + } } MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index 9204e60c58..7f23053943 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -379,18 +379,6 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelTensorSpaceCoordinate pt3 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); ParallelTensorSpaceCoordinate pt4 = mk_pt_coord(0_n, 1_n, 0_n, 1_n); - OneToMany src_binding{ - {pt1, {mc1}}, - {pt2, {mc2}}, - }; - - OneToMany dst_binding{ - {pt1, {mc1}}, - {pt2, {mc2}}, - {pt3, {mc3}}, - {pt4, {mc4}}, - }; - auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, ParallelTensorSpaceCoordinate const &c2) -> OperatorAtomicTaskShardBinding { @@ -429,110 +417,234 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ - { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - mk_value(0, TensorSlotName::OUTPUT, src_binding, std::nullopt), - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coords=*/std::nullopt, - /*mapping=*/mapped_task_group, - /*op_attrs=*/TrainingOperationAttrs{ - PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, + SUBCASE("fwd") { + OneToMany src_binding{ + {pt1, {mc1}}, + {pt2, {mc2}}, + }; + + OneToMany dst_binding{ + {pt1, {mc1}}, + {pt2, {mc2}}, + {pt3, {mc3}}, + {pt4, {mc4}}, + }; + + DynamicNodeInvocation input = DynamicNodeInvocation{ + /*inputs=*/{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + mk_value(0, TensorSlotName::OUTPUT, src_binding, std::nullopt), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coords=*/std::nullopt, + /*mapping=*/mapped_task_group, + /*op_attrs=*/TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }, }, - }, - /*layer_guid=*/dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, + /*layer_guid=*/dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), + }, + }, + }; + + std::unordered_set result = + generate_shard_expansion_for_invocation(input); + + + auto mk_output_binding = [&](MachineSpaceCoordinate const &mc) + -> std::pair + { + return { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/mc, + }, + DynamicValueAttrsShardingInfo{ + dst_binding.at_r(mc), + one_to_many_filter_keys(dst_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == dst_binding.at_r(mc); + }), + }, + }; + }; + + auto mk_invocation_shard = + [&](nonempty_set const &device_coords, + ParallelTensorSpaceCoordinate const &input_shard_coord, + std::unordered_set const &output_task_shards) + -> DynamicNodeInvocationShardingInfo { + + return DynamicNodeInvocationShardingInfo{ + /*device_coords=*/device_coords, + /*value_sharding=*/ + binary_merge_disjoint_maps( + std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrsShardingInfo{ + input_shard_coord, + one_to_many_filter_keys( + src_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == input_shard_coord; + }), + }, }, - mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), - }, - }, - }; + }, + map_from_pairs(transform(output_task_shards, mk_output_binding))), + }; + }; - std::unordered_set result = - generate_shard_expansion_for_invocation(input); + std::unordered_set correct = { + mk_invocation_shard(nonempty_set{mc1, mc2}, pt1, {mc1, mc2}), + mk_invocation_shard(nonempty_set{mc3, mc4}, pt2, {mc3, mc4}), + }; + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } - auto mk_output_binding = [&](MachineSpaceCoordinate const &mc) - -> std::pair - { - return { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/mc, - }, - DynamicValueAttrsShardingInfo{ - dst_binding.at_r(mc), - one_to_many_filter_keys(dst_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == dst_binding.at_r(mc); - }), - }, + SUBCASE("bwd") { + OneToMany output_grad_binding{ + {pt1, {mc1}}, + {pt2, {mc2}}, + {pt3, {mc3}}, + {pt4, {mc4}}, }; - }; - auto mk_invocation_shard = - [&](nonempty_set const &device_coords, - ParallelTensorSpaceCoordinate const &input_shard_coord, - std::unordered_set const &output_task_shards) - -> DynamicNodeInvocationShardingInfo { + OneToMany input_grad_binding{ + {pt1, {mc1}}, + {pt2, {mc2}}, + }; - return DynamicNodeInvocationShardingInfo{ - /*device_coords=*/device_coords, - /*value_sharding=*/ - binary_merge_disjoint_maps( - std::map{ + DynamicNodeInvocation input = DynamicNodeInvocation{ + /*inputs=*/{ { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrsShardingInfo{ - input_shard_coord, - one_to_many_filter_keys( - src_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == input_shard_coord; - }), + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }, + mk_value(0, TensorSlotName::OUTPUT, output_grad_binding, std::nullopt), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::BWD, + /*device_coords=*/std::nullopt, + /*mapping=*/mapped_task_group, + /*op_attrs=*/TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }, }, - }, - map_from_pairs(transform(output_task_shards, mk_output_binding))), + /*layer_guid=*/dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }, + mk_value(20, TensorSlotName::INPUT, input_grad_binding, std::nullopt), + }, + }, }; - }; - std::unordered_set correct = { - mk_invocation_shard(nonempty_set{mc1, mc2}, pt1, {mc1, mc2}), - mk_invocation_shard(nonempty_set{mc3, mc4}, pt2, {mc3, mc4}), - }; + std::unordered_set result = + generate_shard_expansion_for_invocation(input); + + auto mk_output_grad_binding = [&](MachineSpaceCoordinate const &mc) + -> std::pair + { + return { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/mc, + }, + DynamicValueAttrsShardingInfo{ + output_grad_binding.at_r(mc), + one_to_many_filter_keys(output_grad_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == output_grad_binding.at_r(mc); + }), + }, + }; + }; - nlohmann::json result_json = result; - nlohmann::json correct_json = correct; + auto mk_invocation_shard = + [&](nonempty_set const &device_coords, + std::unordered_set const &output_grad_task_shards, + ParallelTensorSpaceCoordinate const &input_grad_shard_coord) + -> DynamicNodeInvocationShardingInfo { + + return DynamicNodeInvocationShardingInfo{ + /*device_coords=*/device_coords, + /*value_sharding=*/ + binary_merge_disjoint_maps( + std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrsShardingInfo{ + input_grad_shard_coord, + one_to_many_filter_keys( + input_grad_binding, + [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { + return pt_coord == input_grad_shard_coord; + }), + }, + }, + }, + map_from_pairs(transform(output_grad_task_shards, mk_output_grad_binding))), + }; + }; - CHECK(result.size() == correct.size()); - CHECK(result_json == correct_json); - CHECK(result == correct); + std::unordered_set correct = { + mk_invocation_shard(nonempty_set{mc1, mc2}, {mc1, mc2}, pt1), + mk_invocation_shard(nonempty_set{mc3, mc4}, {mc3, mc4}, pt2), + }; + + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } } } } diff --git a/lib/utils/include/utils/bidict/bidict.h b/lib/utils/include/utils/bidict/bidict.h index 57f8d5e213..7fcc59f116 100644 --- a/lib/utils/include/utils/bidict/bidict.h +++ b/lib/utils/include/utils/bidict/bidict.h @@ -16,6 +16,7 @@ #include "utils/containers/require_same.h" #include "utils/containers/values.h" #include "utils/containers/unordered_set_of.h" +#include "utils/containers/contains_key.h" namespace FlexFlow { @@ -108,10 +109,12 @@ struct bidict { } R const &at_l(L const &l) const { + ASSERT(contains_key(this->fwd_map, l)); return fwd_map.at(l); } L const &at_r(R const &r) const { + ASSERT(contains_key(this->bwd_map, r)); return bwd_map.at(r); } From d7cce600fc2c5c7a4e4615f442e56c4de4737ee5 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 29 May 2026 03:41:35 -0700 Subject: [PATCH 19/35] Fix other test suites --- .../src/local-execution/task_execution.cc | 5 +- .../utils/containers/map_keys_and_values.h | 21 ++++ .../utils/containers/map_keys_and_values.cc | 8 ++ lib/utils/test/src/utils/bidict/bidict.cc | 4 +- .../containers/binary_merge_disjoint_maps.cc | 10 +- .../binary_merge_disjoint_unordered_maps.cc | 34 ++++++ .../containers/binary_merge_maps_with.cc | 42 +++---- .../binary_merge_maps_with_left_dominating.cc | 10 +- ...binary_merge_maps_with_right_dominating.cc | 10 +- .../binary_merge_unordered_maps_with.cc | 110 ++++++++++++++++++ ...rge_unordered_maps_with_left_dominating.cc | 31 +++++ ...ge_unordered_maps_with_right_dominating.cc | 31 +++++ .../src/utils/containers/map_from_pairs.cc | 15 ++- .../utils/containers/merge_disjoint_maps.cc | 24 ++-- .../merge_disjoint_unordered_maps.cc | 78 +++++++++++++ .../src/utils/containers/merge_maps_with.cc | 36 +++--- .../containers/merge_unordered_maps_with.cc | 100 ++++++++++++++++ 17 files changed, 491 insertions(+), 78 deletions(-) create mode 100644 lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc create mode 100644 lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc create mode 100644 lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc create mode 100644 lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc create mode 100644 lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc create mode 100644 lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc diff --git a/lib/local-execution/src/local-execution/task_execution.cc b/lib/local-execution/src/local-execution/task_execution.cc index c96c834d4a..68a4c579f8 100644 --- a/lib/local-execution/src/local-execution/task_execution.cc +++ b/lib/local-execution/src/local-execution/task_execution.cc @@ -14,6 +14,7 @@ #include "utils/optional.h" #include "utils/overload.h" #include +#include "utils/containers/unordered_map_from_map.h" namespace FlexFlow { @@ -58,9 +59,9 @@ TaskArgumentAccessor make_task_argument_accessor_for_invocation( return assert_unwrap(value.accessor); }; std::unordered_map - tensor_slots_backing = binary_merge_disjoint_maps( + tensor_slots_backing = unordered_map_from_map(binary_merge_disjoint_maps( map_keys_and_values(invocation.inputs, make_param, get_accessor), - map_keys_and_values(invocation.outputs, make_param, get_accessor)); + map_keys_and_values(invocation.outputs, make_param, get_accessor))); return TaskArgumentAccessor::create( /*allocator=*/allocator, diff --git a/lib/utils/include/utils/containers/map_keys_and_values.h b/lib/utils/include/utils/containers/map_keys_and_values.h index 70b7e17103..1873421e17 100644 --- a/lib/utils/include/utils/containers/map_keys_and_values.h +++ b/lib/utils/include/utils/containers/map_keys_and_values.h @@ -3,6 +3,7 @@ #include #include +#include namespace FlexFlow { @@ -26,6 +27,26 @@ std::unordered_map map_keys_and_values( return result; } +template , + typename V2 = std::invoke_result_t> +std::map map_keys_and_values( + std::map const &m, FK const &fk, FV const &fv) { + + std::map result; + for (auto const &kv : m) { + result.insert({fk(kv.first), fv(kv.second)}); + } + + ASSERT(m.size() == result.size(), + "keys passed to map_keys must be transformed into distinct keys"); + + return result; +} + } // namespace FlexFlow #endif diff --git a/lib/utils/src/utils/containers/map_keys_and_values.cc b/lib/utils/src/utils/containers/map_keys_and_values.cc index b3b306988e..95608a09bd 100644 --- a/lib/utils/src/utils/containers/map_keys_and_values.cc +++ b/lib/utils/src/utils/containers/map_keys_and_values.cc @@ -1,5 +1,6 @@ #include "utils/containers/map_keys_and_values.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -13,4 +14,11 @@ using FV = std::function; template std::unordered_map map_keys_and_values( std::unordered_map const &, FK const &, FV const &); +using OK = ordered_value_type<0>; +using OK2 = ordered_value_type<1>; +using OFK = std::function; + +template std::map map_keys_and_values( + std::map const &, OFK const &, FV const &); + } // namespace FlexFlow diff --git a/lib/utils/test/src/utils/bidict/bidict.cc b/lib/utils/test/src/utils/bidict/bidict.cc index f15f15b0fe..ead45fe86a 100644 --- a/lib/utils/test/src/utils/bidict/bidict.cc +++ b/lib/utils/test/src/utils/bidict/bidict.cc @@ -63,14 +63,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("bidict::erase_l") { dict.erase_l(1); CHECK(dict.size() == 1); - CHECK_THROWS_AS(dict.at_l(1), std::out_of_range); + CHECK_THROWS(dict.at_l(1)); CHECK(dict.at_r("two") == 2); } SUBCASE("bidict::erase_r") { dict.erase_r("one"); CHECK(dict.size() == 1); - CHECK_THROWS_AS(dict.at_r("one"), std::out_of_range); + CHECK_THROWS(dict.at_r("one")); CHECK(dict.at_l(2) == "two"); } diff --git a/lib/utils/test/src/utils/containers/binary_merge_disjoint_maps.cc b/lib/utils/test/src/utils/containers/binary_merge_disjoint_maps.cc index bcc7b4149f..d4487343f2 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_disjoint_maps.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_disjoint_maps.cc @@ -1,27 +1,27 @@ #include "utils/containers/binary_merge_disjoint_maps.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("binary_merge_disjoint_maps") { - std::unordered_map l_map = { + std::map l_map = { {1, "one"}, {2, "two"}, }; - std::unordered_map r_map = { + std::map r_map = { {3, "three"}, }; - std::unordered_map correct = { + std::map correct = { {1, "one"}, {2, "two"}, {3, "three"}, }; SUBCASE("maps are disjoint") { - std::unordered_map result = + std::map result = binary_merge_disjoint_maps(l_map, r_map); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc new file mode 100644 index 0000000000..250d1c7f69 --- /dev/null +++ b/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc @@ -0,0 +1,34 @@ +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "test/utils/doctest/fmt/unordered_map.h" +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("binary_merge_disjoint_unordered_maps") { + std::unordered_map l_map = { + {1, "one"}, + {2, "two"}, + }; + + std::unordered_map r_map = { + {3, "three"}, + }; + + std::unordered_map correct = { + {1, "one"}, + {2, "two"}, + {3, "three"}, + }; + SUBCASE("maps are disjoint") { + std::unordered_map result = + binary_merge_disjoint_unordered_maps(l_map, r_map); + + CHECK(result == correct); + } + + SUBCASE("maps are not disjoint") { + CHECK_THROWS(binary_merge_disjoint_unordered_maps(l_map, l_map)); + } + } +} diff --git a/lib/utils/test/src/utils/containers/binary_merge_maps_with.cc b/lib/utils/test/src/utils/containers/binary_merge_maps_with.cc index 55b9c428bf..6a848565e1 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_maps_with.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_maps_with.cc @@ -1,5 +1,5 @@ #include "utils/containers/binary_merge_maps_with.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include @@ -11,47 +11,47 @@ TEST_SUITE(FF_TEST_SUITE) { std::string const &) -> std::string { PANIC(); }; SUBCASE("lhs and rhs do not overlap") { - std::unordered_map lhs = { + std::map lhs = { {1, "lhs_one."}, {4, "lhs_four."}, }; - std::unordered_map rhs = { + std::map rhs = { {2, "rhs_two."}, {5, "rhs_five."}, }; - std::unordered_map correct = { + std::map correct = { {1, "lhs_one."}, {2, "rhs_two."}, {4, "lhs_four."}, {5, "rhs_five."}, }; - std::unordered_map result = + std::map result = binary_merge_maps_with(lhs, rhs, fail_if_called); CHECK(result == correct); } SUBCASE("lhs and rhs overlap") { - std::unordered_map lhs = { + std::map lhs = { {1, "lhs_one."}, {4, "lhs_four."}, }; - std::unordered_map rhs = { + std::map rhs = { {2, "rhs_two."}, {4, "rhs_four."}, {5, "rhs_five."}, }; - std::unordered_map result = binary_merge_maps_with( + std::map result = binary_merge_maps_with( lhs, rhs, [](std::string const &l, std::string const &r) { return l + r; }); - std::unordered_map correct = { + std::map correct = { {1, "lhs_one."}, {2, "rhs_two."}, {4, "lhs_four.rhs_four."}, @@ -62,47 +62,47 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("lhs is empty") { - std::unordered_map lhs = {}; + std::map lhs = {}; - std::unordered_map rhs = { + std::map rhs = { {2, "rhs_two."}, {4, "rhs_four."}, {5, "rhs_five."}, }; - std::unordered_map result = + std::map result = binary_merge_maps_with(lhs, rhs, fail_if_called); - std::unordered_map correct = rhs; + std::map correct = rhs; CHECK(result == correct); } SUBCASE("rhs is empty") { - std::unordered_map lhs = { + std::map lhs = { {1, "lhs_one."}, {4, "lhs_four."}, }; - std::unordered_map rhs = {}; + std::map rhs = {}; - std::unordered_map result = + std::map result = binary_merge_maps_with(lhs, rhs, fail_if_called); - std::unordered_map correct = lhs; + std::map correct = lhs; CHECK(result == correct); } SUBCASE("both lhs and rhs are empty") { - std::unordered_map lhs = {}; + std::map lhs = {}; - std::unordered_map rhs = {}; + std::map rhs = {}; - std::unordered_map result = + std::map result = binary_merge_maps_with(lhs, rhs, fail_if_called); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/binary_merge_maps_with_left_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_maps_with_left_dominating.cc index 27a389d400..fe152dd832 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_maps_with_left_dominating.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_maps_with_left_dominating.cc @@ -1,5 +1,5 @@ #include "utils/containers/binary_merge_maps_with_left_dominating.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include @@ -7,23 +7,23 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("binary_merge_maps_with_left_dominating") { - std::unordered_map l_map = { + std::map l_map = { {1, "one"}, {2, "left_two"}, }; - std::unordered_map r_map = { + std::map r_map = { {2, "right_two"}, {3, "three"}, }; - std::unordered_map correct = { + std::map correct = { {1, "one"}, {2, "left_two"}, {3, "three"}, }; - std::unordered_map result = + std::map result = binary_merge_maps_with_left_dominating(l_map, r_map); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/containers/binary_merge_maps_with_right_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_maps_with_right_dominating.cc index 153266989e..c107f2b7ff 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_maps_with_right_dominating.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_maps_with_right_dominating.cc @@ -1,5 +1,5 @@ #include "utils/containers/binary_merge_maps_with_right_dominating.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include @@ -7,23 +7,23 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("binary_merge_maps_with_right_dominating") { - std::unordered_map l_map = { + std::map l_map = { {1, "one"}, {2, "left_two"}, }; - std::unordered_map r_map = { + std::map r_map = { {2, "right_two"}, {3, "three"}, }; - std::unordered_map correct = { + std::map correct = { {1, "one"}, {2, "right_two"}, {3, "three"}, }; - std::unordered_map result = + std::map result = binary_merge_maps_with_right_dominating(l_map, r_map); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc new file mode 100644 index 0000000000..4e825b99e2 --- /dev/null +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc @@ -0,0 +1,110 @@ +#include "utils/containers/binary_merge_unordered_maps_with.h" +#include "test/utils/doctest/fmt/unordered_map.h" +#include +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("binary_merge_unordered_maps_with") { + auto fail_if_called = [](std::string const &, + std::string const &) -> std::string { PANIC(); }; + + SUBCASE("lhs and rhs do not overlap") { + std::unordered_map lhs = { + {1, "lhs_one."}, + {4, "lhs_four."}, + }; + + std::unordered_map rhs = { + {2, "rhs_two."}, + {5, "rhs_five."}, + }; + + std::unordered_map correct = { + {1, "lhs_one."}, + {2, "rhs_two."}, + {4, "lhs_four."}, + {5, "rhs_five."}, + }; + + std::unordered_map result = + binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + + CHECK(result == correct); + } + + SUBCASE("lhs and rhs overlap") { + std::unordered_map lhs = { + {1, "lhs_one."}, + {4, "lhs_four."}, + }; + + std::unordered_map rhs = { + {2, "rhs_two."}, + {4, "rhs_four."}, + {5, "rhs_five."}, + }; + + std::unordered_map result = binary_merge_unordered_maps_with( + lhs, rhs, [](std::string const &l, std::string const &r) { + return l + r; + }); + + std::unordered_map correct = { + {1, "lhs_one."}, + {2, "rhs_two."}, + {4, "lhs_four.rhs_four."}, + {5, "rhs_five."}, + }; + + CHECK(result == correct); + } + + SUBCASE("lhs is empty") { + std::unordered_map lhs = {}; + + std::unordered_map rhs = { + {2, "rhs_two."}, + {4, "rhs_four."}, + {5, "rhs_five."}, + }; + + std::unordered_map result = + binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + + std::unordered_map correct = rhs; + + CHECK(result == correct); + } + + SUBCASE("rhs is empty") { + std::unordered_map lhs = { + {1, "lhs_one."}, + {4, "lhs_four."}, + }; + + std::unordered_map rhs = {}; + + std::unordered_map result = + binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + + std::unordered_map correct = lhs; + + CHECK(result == correct); + } + + SUBCASE("both lhs and rhs are empty") { + std::unordered_map lhs = {}; + + std::unordered_map rhs = {}; + + std::unordered_map result = + binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + + std::unordered_map correct = {}; + + CHECK(result == correct); + } + } +} diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc new file mode 100644 index 0000000000..d857cf2d91 --- /dev/null +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc @@ -0,0 +1,31 @@ +#include "utils/containers/binary_merge_unordered_maps_with_left_dominating.h" +#include "test/utils/doctest/fmt/unordered_map.h" +#include +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("binary_merge_unordered_maps_with_left_dominating") { + std::unordered_map l_map = { + {1, "one"}, + {2, "left_two"}, + }; + + std::unordered_map r_map = { + {2, "right_two"}, + {3, "three"}, + }; + + std::unordered_map correct = { + {1, "one"}, + {2, "left_two"}, + {3, "three"}, + }; + + std::unordered_map result = + binary_merge_unordered_maps_with_left_dominating(l_map, r_map); + + CHECK(result == correct); + } +} diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc new file mode 100644 index 0000000000..71f50c4dac --- /dev/null +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc @@ -0,0 +1,31 @@ +#include "utils/containers/binary_merge_unordered_maps_with_right_dominating.h" +#include "test/utils/doctest/fmt/unordered_map.h" +#include +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("binary_merge_unordered_maps_with_right_dominating") { + std::unordered_map l_map = { + {1, "one"}, + {2, "left_two"}, + }; + + std::unordered_map r_map = { + {2, "right_two"}, + {3, "three"}, + }; + + std::unordered_map correct = { + {1, "one"}, + {2, "right_two"}, + {3, "three"}, + }; + + std::unordered_map result = + binary_merge_unordered_maps_with_right_dominating(l_map, r_map); + + CHECK(result == correct); + } +} diff --git a/lib/utils/test/src/utils/containers/map_from_pairs.cc b/lib/utils/test/src/utils/containers/map_from_pairs.cc index fc387f1d1a..48e8b9fe05 100644 --- a/lib/utils/test/src/utils/containers/map_from_pairs.cc +++ b/lib/utils/test/src/utils/containers/map_from_pairs.cc @@ -1,24 +1,23 @@ #include "utils/containers/map_from_pairs.h" -#include "test/utils/doctest/fmt/unordered_map.h" -#include "utils/hash/pair.h" +#include "test/utils/doctest/fmt/map.h" #include #include +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("map_from_pairs") { - std::unordered_set> input = - std::unordered_set>{ + std::set> input = + std::set>{ {1, "one"}, {2, "two"}, }; - std::unordered_map result = map_from_pairs(input); - - std::unordered_map correct = - std::unordered_map{ + std::map result = map_from_pairs(input); + std::map correct = + std::map{ {1, "one"}, {2, "two"}, }; diff --git a/lib/utils/test/src/utils/containers/merge_disjoint_maps.cc b/lib/utils/test/src/utils/containers/merge_disjoint_maps.cc index 24e8d548ae..bf4b2202d7 100644 --- a/lib/utils/test/src/utils/containers/merge_disjoint_maps.cc +++ b/lib/utils/test/src/utils/containers/merge_disjoint_maps.cc @@ -1,37 +1,37 @@ #include "utils/containers/merge_disjoint_maps.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("merge_disjoint_maps") { - std::unordered_map m1 = { + std::map m1 = { {4, "four"}, {2, "two"}, }; - std::unordered_map m2 = { + std::map m2 = { {3, "four"}, }; - std::unordered_map m3 = { + std::map m3 = { {1, "one"}, }; - std::unordered_map m4 = {}; + std::map m4 = {}; SUBCASE("maps are disjoint") { - std::vector> input = { + std::vector> input = { m1, m2, m3, m4, }; - std::unordered_map result = merge_disjoint_maps(input); + std::map result = merge_disjoint_maps(input); - std::unordered_map correct = { + std::map correct = { {4, "four"}, {2, "two"}, {3, "four"}, @@ -42,12 +42,12 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("maps are not disjoint") { - std::unordered_map m5 = { + std::map m5 = { {4, "five"}, {6, "six"}, }; - std::vector> input = { + std::vector> input = { m1, m2, m3, @@ -59,12 +59,12 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("maps are not disjoint but have identical values") { - std::unordered_map m5 = { + std::map m5 = { {4, "four"}, {6, "six"}, }; - std::vector> input = { + std::vector> input = { m1, m2, m3, diff --git a/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc new file mode 100644 index 0000000000..ceecea96c1 --- /dev/null +++ b/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc @@ -0,0 +1,78 @@ +#include "utils/containers/merge_disjoint_unordered_maps.h" +#include "test/utils/doctest/fmt/unordered_map.h" +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("merge_disjoint_unordered_maps") { + std::unordered_map m1 = { + {4, "four"}, + {2, "two"}, + }; + + std::unordered_map m2 = { + {3, "four"}, + }; + + std::unordered_map m3 = { + {1, "one"}, + }; + + std::unordered_map m4 = {}; + + SUBCASE("maps are disjoint") { + std::vector> input = { + m1, + m2, + m3, + m4, + }; + + std::unordered_map result = merge_disjoint_unordered_maps(input); + + std::unordered_map correct = { + {4, "four"}, + {2, "two"}, + {3, "four"}, + {1, "one"}, + }; + + CHECK(result == correct); + } + + SUBCASE("maps are not disjoint") { + std::unordered_map m5 = { + {4, "five"}, + {6, "six"}, + }; + + std::vector> input = { + m1, + m2, + m3, + m4, + m5, + }; + + CHECK_THROWS(merge_disjoint_unordered_maps(input)); + } + + SUBCASE("maps are not disjoint but have identical values") { + std::unordered_map m5 = { + {4, "four"}, + {6, "six"}, + }; + + std::vector> input = { + m1, + m2, + m3, + m4, + m5, + }; + + CHECK_THROWS(merge_disjoint_unordered_maps(input)); + } + } +} diff --git a/lib/utils/test/src/utils/containers/merge_maps_with.cc b/lib/utils/test/src/utils/containers/merge_maps_with.cc index ec5b31abf3..fd73a9345e 100644 --- a/lib/utils/test/src/utils/containers/merge_maps_with.cc +++ b/lib/utils/test/src/utils/containers/merge_maps_with.cc @@ -1,5 +1,5 @@ #include "utils/containers/merge_maps_with.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include "test/utils/rapidcheck.h" #include "utils/containers/binary_merge_maps_with.h" #include @@ -14,37 +14,37 @@ TEST_SUITE(FF_TEST_SUITE) { RC_SUBCASE( "with two inputs, matches binary_merge_maps_with", - [&](std::unordered_map const &lhs, - std::unordered_map const &rhs) { - std::unordered_map from_merge_maps_with = + [&](std::map const &lhs, + std::map const &rhs) { + std::map from_merge_maps_with = merge_maps_with(std::vector{lhs, rhs}, string_concat); - std::unordered_map from_binary_merge_maps_with = + std::map from_binary_merge_maps_with = binary_merge_maps_with(lhs, rhs, string_concat); CHECK(from_merge_maps_with == from_binary_merge_maps_with); }); SUBCASE("maps overlap") { - std::unordered_map map1 = { + std::map map1 = { {1, "map1_one."}, {4, "map1_four."}, }; - std::unordered_map map2 = { + std::map map2 = { {2, "map2_two."}, {4, "map2_four."}, {5, "map2_five."}, }; - std::unordered_map map3 = { + std::map map3 = { {1, "map3_one."}, }; - std::unordered_map result = + std::map result = merge_maps_with(std::vector{map1, map2, map3}, string_concat); - std::unordered_map correct = { + std::map correct = { {1, "map1_one.map3_one."}, {2, "map2_two."}, {4, "map1_four.map2_four."}, @@ -58,24 +58,24 @@ TEST_SUITE(FF_TEST_SUITE) { std::string const &) -> std::string { PANIC(); }; SUBCASE("maps do not overlap") { - std::unordered_map map1 = { + std::map map1 = { {8, "map1_eight."}, {4, "map1_four."}, }; - std::unordered_map map2 = { + std::map map2 = { {2, "map2_two."}, {5, "map2_five."}, }; - std::unordered_map map3 = { + std::map map3 = { {1, "map3_one."}, }; - std::unordered_map result = + std::map result = merge_maps_with(std::vector{map1, map2, map3}, fail_if_called); - std::unordered_map correct = { + std::map correct = { {1, "map3_one."}, {2, "map2_two."}, {4, "map1_four."}, @@ -87,12 +87,12 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("no maps are provided") { - std::vector> maps = {}; + std::vector> maps = {}; - std::unordered_map result = + std::map result = merge_maps_with(maps, fail_if_called); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc b/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc new file mode 100644 index 0000000000..66827ca453 --- /dev/null +++ b/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc @@ -0,0 +1,100 @@ +#include "utils/containers/merge_unordered_maps_with.h" +#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/rapidcheck.h" +#include "utils/containers/binary_merge_unordered_maps_with.h" +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("merge_unordered_maps_with") { + auto string_concat = [](std::string const &l, std::string const &r) { + return l + r; + }; + + RC_SUBCASE( + "with two inputs, matches binary_merge_unordered_maps_with", + [&](std::unordered_map const &lhs, + std::unordered_map const &rhs) { + std::unordered_map from_merge_unordered_maps_with = + merge_unordered_maps_with(std::vector{lhs, rhs}, string_concat); + + std::unordered_map from_binary_merge_unordered_maps_with = + binary_merge_unordered_maps_with(lhs, rhs, string_concat); + + CHECK(from_merge_unordered_maps_with == from_binary_merge_unordered_maps_with); + }); + + SUBCASE("maps overlap") { + std::unordered_map map1 = { + {1, "map1_one."}, + {4, "map1_four."}, + }; + + std::unordered_map map2 = { + {2, "map2_two."}, + {4, "map2_four."}, + {5, "map2_five."}, + }; + + std::unordered_map map3 = { + {1, "map3_one."}, + }; + + std::unordered_map result = + merge_unordered_maps_with(std::vector{map1, map2, map3}, string_concat); + + std::unordered_map correct = { + {1, "map1_one.map3_one."}, + {2, "map2_two."}, + {4, "map1_four.map2_four."}, + {5, "map2_five."}, + }; + + CHECK(result == correct); + } + + auto fail_if_called = [](std::string const &, + std::string const &) -> std::string { PANIC(); }; + + SUBCASE("maps do not overlap") { + std::unordered_map map1 = { + {8, "map1_eight."}, + {4, "map1_four."}, + }; + + std::unordered_map map2 = { + {2, "map2_two."}, + {5, "map2_five."}, + }; + + std::unordered_map map3 = { + {1, "map3_one."}, + }; + + std::unordered_map result = + merge_unordered_maps_with(std::vector{map1, map2, map3}, fail_if_called); + + std::unordered_map correct = { + {1, "map3_one."}, + {2, "map2_two."}, + {4, "map1_four."}, + {5, "map2_five."}, + {8, "map1_eight."}, + }; + + CHECK(result == correct); + } + + SUBCASE("no maps are provided") { + std::vector> maps = {}; + + std::unordered_map result = + merge_unordered_maps_with(maps, fail_if_called); + + std::unordered_map correct = {}; + + CHECK(result == correct); + } + } +} From 30d0320055ff29689dd2ba4af75a380b5c777210 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Thu, 4 Jun 2026 15:50:24 -0700 Subject: [PATCH 20/35] Change OneToMany DynamicValueAttrs mapping back to bidict --- .../mapped_operator_task_group.h | 3 +- .../mapped_operator_task_group.cc | 12 +- .../dynamic_value_attrs.dtg.toml | 4 +- .../dynamic_graph/dynamic_value_attrs.h | 2 +- ...dynamic_value_attrs_sharding_info.dtg.toml | 6 +- .../serializable_dynamic_value_attrs.dtg.toml | 4 +- .../task-spec/dynamic_graph/copy_insertion.cc | 10 +- .../dynamic_graph/dynamic_value_attrs.cc | 2 +- .../dynamic_graph/shard_expansion.cc | 87 ++++--------- .../task-spec/dynamic_graph/copy_insertion.cc | 121 +----------------- .../dynamic_graph/shard_expansion.cc | 112 +++++++--------- 11 files changed, 94 insertions(+), 269 deletions(-) diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h index 41aca802e7..89d09b9132 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h @@ -7,7 +7,6 @@ #include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.h" #include "utils/bidict/bidict.h" #include -#include "utils/one_to_many/one_to_many.h" namespace FlexFlow { @@ -39,7 +38,7 @@ struct MappedOperatorTaskGroup { friend struct ::std::hash; }; -OneToMany +bidict get_tensor_bindings_for_slot_name(MappedOperatorTaskGroup const &, TensorSlotName const &); diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc index 3bb508681d..72b3aac5dd 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc @@ -19,7 +19,7 @@ #include "utils/bidict/algorithms/right_entries.h" #include "utils/containers/map_values.h" #include "utils/containers/unordered_set_of.h" -#include "utils/many_to_one/invert_many_to_one.h" +#include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" namespace FlexFlow { @@ -98,7 +98,7 @@ bidict const & return this->shard_bindings; } -OneToMany +bidict get_tensor_bindings_for_slot_name(MappedOperatorTaskGroup const &task_group, TensorSlotName const &slot_name) { std::set slot_names = get_slot_names_for_task_group(task_group); @@ -106,11 +106,11 @@ OneToMany std::unordered_map m = map_values(task_group.get_shard_bindings().as_unordered_map(), - [&](OperatorAtomicTaskShardBinding const &b) -> ParallelTensorSpaceCoordinate { - return ptensor_space_coord_for_slot_name(b, slot_name); - }); + [&](OperatorAtomicTaskShardBinding const &b) -> ParallelTensorSpaceCoordinate { + return ptensor_space_coord_for_slot_name(b, slot_name); + }); - return invert_many_to_one(many_to_one_from_unstructured_relation(unordered_set_of(m))); + return bidict_from_unstructured_relation(unordered_set_of(m)).reversed(); } std::set get_slot_names_for_task_group(MappedOperatorTaskGroup const &g) { diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml index 1e13736fc7..5ece2ad8f1 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml @@ -14,7 +14,7 @@ includes = [ "op-attrs/parallel_tensor_shape.dtg.h", "op-attrs/parallel_tensor_space_coordinate.dtg.h", "pcg/machine_space_coordinate.dtg.h", - "utils/one_to_many/one_to_many.h", + "utils/bidict/bidict.h", "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_role.dtg.h", ] @@ -52,7 +52,7 @@ For a \ref DynamicOpenDataflowGraph originating from a \ref MappedParallelComput [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::OneToMany<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" +type = "std::optional<::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" docstring = ''' \brief The location (i.e., \ref MachineSpaceCoordinate) of each shard (i.e., \ref ParallelTensorSpaceCoordinate) of this (usually parallel) tensor. diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h index aa9fbd2874..ac8098e553 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h @@ -10,7 +10,7 @@ DynamicValueAttrs decide_dynamic_value_attrs_role(DynamicValueAttrs const &, DynamicValueAttrs decide_dynamic_value_attrs_mapping( DynamicValueAttrs const &, - OneToMany const &); + bidict const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml index 5a9234d815..63ad868e68 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml @@ -10,13 +10,11 @@ features = [ ] includes = [ - "utils/one_to_many/one_to_many.h", "op-attrs/parallel_tensor_space_coordinate.dtg.h", "pcg/machine_space_coordinate.dtg.h", ] -src_includes = [ -] +src_includes = [] [[fields]] name = "shard_coord" @@ -24,4 +22,4 @@ type = "::FlexFlow::ParallelTensorSpaceCoordinate" [[fields]] name = "mapping" -type = "::FlexFlow::OneToMany<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>" +type = "::FlexFlow::MachineSpaceCoordinate" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml index d05d8e011e..e20720334d 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml @@ -15,7 +15,7 @@ includes = [ "op-attrs/parallel_tensor_shape.dtg.h", "op-attrs/parallel_tensor_space_coordinate.dtg.h", "pcg/machine_space_coordinate.dtg.h", - "utils/one_to_many/one_to_many.h", + "utils/bidict/bidict.h", "task-spec/dynamic_graph/dynamic_tensor_role.dtg.h", ] @@ -38,7 +38,7 @@ type = "std::optional<::FlexFlow::ParallelTensorSpaceCoordinate>" [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::OneToMany<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" +type = "std::optional<::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" [[fields]] name = "role" diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index ebf84f2d81..b639c0072b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -19,6 +19,8 @@ #include "utils/containers/transform.h" #include "utils/optional.h" #include "task-spec/dynamic_graph/training_operation_attrs.h" +#include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" +#include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" namespace FlexFlow { @@ -92,10 +94,10 @@ static std::pair DynamicValueAttrs const &output) { std::unordered_set< std::pair> - input_mapping = unstructured_relation_from_one_to_many(assert_unwrap(input.mapping)); + input_mapping = unstructured_relation_from_bidict(assert_unwrap(input.mapping)); std::unordered_set< std::pair> - output_mapping = unstructured_relation_from_one_to_many(assert_unwrap(output.mapping)); + output_mapping = unstructured_relation_from_bidict(assert_unwrap(output.mapping)); // Exclude the point shared between the input and output mappings, because // those will not result in actual copies once shard expansion is performed @@ -105,11 +107,11 @@ static std::pair DynamicValueAttrs filtered_input = input; filtered_input.mapping = - one_to_many_from_unstructured_relation(set_difference(input_mapping, remove)); + bidict_from_unstructured_relation(set_difference(input_mapping, remove)); DynamicValueAttrs filtered_output = output; filtered_output.mapping = - one_to_many_from_unstructured_relation(set_difference(output_mapping, remove)); + bidict_from_unstructured_relation(set_difference(output_mapping, remove)); return std::pair{filtered_input, filtered_output}; } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc index 9a70c5cdd0..c6c9dcdf09 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc @@ -15,7 +15,7 @@ DynamicValueAttrs DynamicValueAttrs decide_dynamic_value_attrs_mapping( DynamicValueAttrs const &attrs, - OneToMany const &mapping) + bidict const &mapping) { ASSERT(!attrs.mapping.has_value()); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index f910c6c6c2..440668ea95 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -10,15 +10,14 @@ #include "utils/containers/binary_merge_disjoint_maps.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.h" #include "utils/containers/map_from_unordered.h" -#include "utils/one_to_many/one_to_many_filter_keys.h" #include "task-spec/dynamic_graph/training_operation_attrs.h" -#include "utils/one_to_many/require_one_to_many_is_bijection.h" #include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.h" #include "utils/bidict/algorithms/bidict_filter_values.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/map_keys.h" #include "utils/containers/require_only_key.h" +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -118,16 +117,6 @@ bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { slot_is_shard_expanded); } -static OneToMany - restrict_tensor_mapping_keys_to_coord( - OneToMany const - &mapping, - ParallelTensorSpaceCoordinate const ¶llel_tensor_coord) { - return one_to_many_filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { - return p == parallel_tensor_coord; - }); -} - static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( DynamicNodeInvocation const &i, MachineSpaceCoordinate const &machine_coord, @@ -140,7 +129,7 @@ static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( return DynamicValueAttrsShardingInfo{ /*shard_coord=*/parallel_tensor_coord, - /*mapping=*/restrict_tensor_mapping_keys_to_coord(v.mapping.value(), parallel_tensor_coord), + /*mapping=*/v.mapping.value().at_l(parallel_tensor_coord), }; }; @@ -171,14 +160,6 @@ static DynamicNodeInvocation shard_invocation_for_binding( DynamicValueAttrs result = v; result.shard_coord = parallel_tensor_coord; - result.mapping = transform( - v.mapping, - [&](OneToMany const &mapping) - -> OneToMany - { - return restrict_tensor_mapping_keys_to_coord(mapping, - parallel_tensor_coord); - }); return result; }; @@ -200,13 +181,13 @@ static std::set auto [input_slot, input] = get_only(i.inputs); auto [output_slot, output] = get_only(i.outputs); - OneToMany input_mapping = + bidict input_mapping = assert_unwrap(input.mapping); require_same(input_mapping.left_values(), assert_unwrap(output.mapping).left_values()); return transform( - input_mapping.left_values(), + set_of(input_mapping.left_values()), [&](ParallelTensorSpaceCoordinate const &p) -> DynamicNodeInvocationShardingInfo { // The machine coord for a copy is inherently nebulous because it // doesn't strictly run in any single location. Further, Realm has the @@ -215,7 +196,7 @@ static std::set // because we expect this to align with the most efficient way to issue // copies in Realm, although the current Realm backend uses a // centralized controller and thus issues copies all from a single node. - MachineSpaceCoordinate machine_coord = get_only(input_mapping.at_l(p)); + MachineSpaceCoordinate machine_coord = input_mapping.at_l(p); return invocation_sharding_info_for_binding(i, machine_coord, @@ -253,14 +234,12 @@ static std::set DynamicValueAttrs output = require_only_key(i.outputs, expected_output_slot); bidict - input_value_mapping = require_one_to_many_is_bijection( - assert_unwrap(input.mapping)); + input_value_mapping = assert_unwrap(input.mapping); std::set input_tensor_shards = set_of(input_value_mapping.left_values()); - + bidict - output_value_mapping = require_one_to_many_is_bijection( - assert_unwrap(output.mapping)); + output_value_mapping = assert_unwrap(output.mapping); auto get_task_shard_machine_coords_for_input_tensor_shard = [&](ParallelTensorSpaceCoordinate const &input_tensor_shard) @@ -291,12 +270,7 @@ static std::set return DynamicValueAttrsShardingInfo{ /*shard_coord=*/pc, - /*mapping=*/OneToMany{ - { - pc, - {mc}, - }, - }, + /*mapping=*/mc, }; }); @@ -318,15 +292,10 @@ static std::set DynamicValueAttrsShardingInfo input_sharding_info = DynamicValueAttrsShardingInfo{ /*shard_coord=*/c, - /*mapping=*/OneToMany{ - { - c, - {input_value_mapping.at_l(c)}, - }, - }, + /*mapping=*/input_value_mapping.at_l(c), }; - std::map sharding_infos = + std::map sharding_infos = binary_merge_disjoint_maps( keyed_output_sharding_infos, std::map{ @@ -368,12 +337,10 @@ static std::set DynamicValueAttrs input_grad = require_only_key(i.outputs, expected_input_grad_slot); bidict - output_grad_value_mapping = require_one_to_many_is_bijection( - assert_unwrap(output_grad.mapping)); + output_grad_value_mapping = assert_unwrap(output_grad.mapping); bidict - input_grad_value_mapping = require_one_to_many_is_bijection( - assert_unwrap(input_grad.mapping)); + input_grad_value_mapping = assert_unwrap(input_grad.mapping); std::set input_grad_tensor_shards = set_of(input_grad_value_mapping.left_values()); @@ -406,12 +373,7 @@ static std::set return DynamicValueAttrsShardingInfo{ /*shard_coord=*/pc, - /*mapping=*/OneToMany{ - { - pc, - {mc}, - }, - }, + /*mapping=*/mc, }; }); @@ -433,15 +395,10 @@ static std::set DynamicValueAttrsShardingInfo input_grad_sharding_info = DynamicValueAttrsShardingInfo{ /*shard_coord=*/c, - /*mapping=*/OneToMany{ - { - c, - {input_grad_value_mapping.at_l(c)}, - }, - }, + /*mapping=*/input_grad_value_mapping.at_l(c), }; - std::map sharding_infos = + std::map sharding_infos = binary_merge_disjoint_maps( keyed_output_grad_sharding_infos, std::map{ @@ -514,7 +471,17 @@ DynamicValueAttrs apply_dynamic_value_attrs_sharding_info( { DynamicValueAttrs result = value_attrs; result.shard_coord = value_sharding_info.shard_coord; - result.mapping = value_sharding_info.mapping; + + { + bidict value_mapping = + assert_unwrap(result.mapping); + + MachineSpaceCoordinate from_mapping = value_mapping.at_l(value_sharding_info.shard_coord); + MachineSpaceCoordinate from_sharding_info = value_sharding_info.mapping; + + ASSERT(from_mapping == from_sharding_info); + } + return result; } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index 05324c2195..9bb0e7a7c4 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -452,10 +452,6 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicValueAttrs graph_input_unmapped = mk_value(0, TensorSlotName::OUTPUT); - DynamicValueAttrs graph_input_use_mapped = - decide_dynamic_value_attrs_mapping( - graph_input_unmapped, - get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::INPUT)); DynamicValueAttrs invocation_output_unmapped = mk_value(invocation_id, TensorSlotName::OUTPUT); @@ -502,10 +498,10 @@ TEST_SUITE(FF_TEST_SUITE) { graph_input_unmapped, decide_dynamic_value_attrs_mapping( graph_input_unmapped, - OneToMany{ + bidict{ { mc_input_coord, - {mc3}, + mc3, }, }) }, @@ -521,118 +517,5 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result_j == correct_j); } - - // SUBCASE("reduction operator") { - - // auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, - // ParallelTensorSpaceCoordinate const &c2) - // -> OperatorAtomicTaskShardBinding { - // return OperatorAtomicTaskShardBinding{ - // /*tensor_coords=*/{ - // { - // TensorSlotName::INPUT, - // c1, - // }, - // { - // TensorSlotName::OUTPUT, - // c2, - // }, - // }, - // }; - // }; - - // ParallelTensorSpaceCoordinate mc1_input_coord = - // mk_pt_coord(0_n, 0_n, 0_n, 0_n); - // ParallelTensorSpaceCoordinate mc2_input_coord = - // mk_pt_coord(1_n, 0_n, 0_n, 0_n); - - // ParallelTensorSpaceCoordinate mc_output_coord = - // mk_pt_coord(0_n, 0_n, 0_n, 0_n); - - // MappedOperatorTaskGroup invocation_mapping = MappedOperatorTaskGroup{ - // bidict{ - // { - // mc3, - // mk_shard_binding(mc1_input_coord, - // mc_output_coord), - // }, - // }, - // }; - - // DynamicValueAttrs graph_input_unmapped = - // mk_value(0, TensorSlotName::OUTPUT); - // DynamicValueAttrs graph_input_use_mapped = - // decide_dynamic_value_attrs_mapping( - // graph_input_unmapped, - // get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::INPUT)); - - // DynamicValueAttrs invocation_output_unmapped = - // mk_value(invocation_id, TensorSlotName::OUTPUT); - // DynamicValueAttrs invocation_output_src_mapped = - // decide_dynamic_value_attrs_mapping( - // invocation_output_unmapped, - // get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT)); - - // DynamicNodeInvocation input = DynamicNodeInvocation{ - // /*inputs=*/{ - // { - // mk_slot(TensorSlotName::INPUT), - // graph_input_unmapped, - // }, - // }, - // /*node_attrs=*/DynamicNodeAttrs{ - // /*task_type=*/DynamicTaskType::FWD, - // /*device_coord=*/std::nullopt, - // /*mapping=*/invocation_mapping, - // /*op_attrs=*/TrainingOperationAttrs{ - // PCGOperatorAttrs{ - // ReductionAttrs{ - // /*reduction_degree=*/2_p, - // }, - // }, - // }, - // /*layer_guid=*/dynamic_layer_guid_t{ - // parallel_layer_guid_t{ - // Node{invocation_id}, - // }, - // }, - // /*per_device_op_state=*/std::nullopt, - // }, - // /*outputs=*/{ - // { - // mk_slot(TensorSlotName::OUTPUT), - // invocation_output_unmapped, - // }, - // }, - // }; - - // std::unordered_map unmapped_to_mapped_source_value = { - // { - // graph_input_unmapped, - // decide_dynamic_value_attrs_mapping( - // graph_input_unmapped, - // OneToMany{ - // { - // mc1_input_coord, - // {mc1}, - // }, - // { - // mc2_input_coord, - // {mc2}, - // }, - // }) - // }, - // }; - - // std::unordered_set result = copies_for_invocation_inputs( - // input, unmapped_to_mapped_source_value); - - // std::unordered_set correct = {}; - - // nlohmann::json result_j = transform(result, dynamic_node_invocation_to_serializable); - // nlohmann::json correct_j = transform(correct, dynamic_node_invocation_to_serializable); - - // CHECK(result_j == correct_j); - // } } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index 7f23053943..60dbfea9aa 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -7,8 +7,8 @@ #include #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "op-attrs/ops/element_unary.h" -#include "utils/one_to_many/one_to_many_filter_keys.h" -#include "utils/one_to_many/one_to_many_filter_values.h" +#include "utils/bidict/algorithms/bidict_filter_keys.h" +#include "utils/bidict/algorithms/bidict_filter_values.h" #include "utils/containers/map_from_pairs.h" #include "utils/containers/binary_merge_disjoint_maps.h" @@ -50,16 +50,16 @@ DynamicTensorSlot mk_slot(TensorSlotName const &slot_name, DynamicValueAttrs mk_value(size_t src_node_id, TensorSlotName src_slot_name, - OneToMany const &tensor_binding, + bidict const &tensor_binding, std::optional const &shard_coord, std::optional const &role = std::nullopt) { - OneToMany mapping = tensor_binding; + bidict mapping = tensor_binding; if (shard_coord.has_value()) { - mapping = one_to_many_filter_keys(mapping, - [&](ParallelTensorSpaceCoordinate const &p) { - return p == shard_coord.value(); - }); + mapping = bidict_filter_keys(mapping, + [&](ParallelTensorSpaceCoordinate const &p) { + return p == shard_coord.value(); + }); } return DynamicValueAttrs{ @@ -87,7 +87,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::optional const &shard_coord, std::optional const &role = std::nullopt) -> DynamicValueAttrs { - OneToMany + bidict tensor_binding = get_tensor_bindings_for_slot_name(mapped_task_group, use_slot_name); return mk_value(src_node_id, src_slot_name, tensor_binding, shard_coord, role); @@ -95,21 +95,17 @@ TEST_SUITE(FF_TEST_SUITE) { auto mk_sharding_info = [&](TensorSlotName slot_name, ParallelTensorSpaceCoordinate const &shard_coord, - MappedOperatorTaskGroup const &mapped_op_task_group, - MachineSpaceCoordinate const &device_coord) + MappedOperatorTaskGroup const &mapped_op_task_group) -> std::pair { - OneToMany + bidict tensor_binding = get_tensor_bindings_for_slot_name(mapped_op_task_group, slot_name); return std::pair{ mk_slot(slot_name), DynamicValueAttrsShardingInfo{ /*shard_coord=*/shard_coord, - /*mapping=*/one_to_many_filter_values(tensor_binding, - [&](MachineSpaceCoordinate const &c) -> bool { - return device_coord == c; - }), + /*mapping=*/tensor_binding.at_l(shard_coord), }, }; }; @@ -251,10 +247,10 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicNodeInvocationShardingInfo{ /*device_coord=*/nonempty_set{device_coord}, /*value_sharding=*/{ - mk_sharding_info(TensorSlotName::INPUT, input_shard_coord, mapped_task_group, device_coord), - mk_sharding_info(TensorSlotName::WEIGHT, weight_shard_coord, mapped_task_group, device_coord), - mk_sharding_info(TensorSlotName::OUTPUT_1, output_1_shard_coord, mapped_task_group, device_coord), - mk_sharding_info(TensorSlotName::OUTPUT_2, output_2_shard_coord, mapped_task_group, device_coord), + mk_sharding_info(TensorSlotName::INPUT, input_shard_coord, mapped_task_group), + mk_sharding_info(TensorSlotName::WEIGHT, weight_shard_coord, mapped_task_group), + mk_sharding_info(TensorSlotName::OUTPUT_1, output_1_shard_coord, mapped_task_group), + mk_sharding_info(TensorSlotName::OUTPUT_2, output_2_shard_coord, mapped_task_group), }, }; }; @@ -289,14 +285,14 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); - OneToMany src_binding{ - {pt1, {mc1}}, - {pt2, {mc2}}, + bidict src_binding{ + {pt1, mc1}, + {pt2, mc2}, }; - OneToMany dst_binding{ - {pt1, {mc3}}, - {pt2, {mc4}}, + bidict dst_binding{ + {pt1, mc3}, + {pt2, mc4}, }; DynamicNodeInvocation input = DynamicNodeInvocation{ @@ -339,20 +335,14 @@ TEST_SUITE(FF_TEST_SUITE) { mk_slot(TensorSlotName::INPUT), DynamicValueAttrsShardingInfo{ tensor_shard_coord, - one_to_many_filter_keys(src_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == tensor_shard_coord; - }), + src_binding.at_l(tensor_shard_coord), }, }, { mk_slot(TensorSlotName::OUTPUT), DynamicValueAttrsShardingInfo{ tensor_shard_coord, - one_to_many_filter_keys(dst_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == tensor_shard_coord; - }), + dst_binding.at_l(tensor_shard_coord), }, }, }, @@ -418,16 +408,16 @@ TEST_SUITE(FF_TEST_SUITE) { }; SUBCASE("fwd") { - OneToMany src_binding{ - {pt1, {mc1}}, - {pt2, {mc2}}, + bidict src_binding{ + {pt1, mc1}, + {pt2, mc2}, }; - OneToMany dst_binding{ - {pt1, {mc1}}, - {pt2, {mc2}}, - {pt3, {mc3}}, - {pt4, {mc4}}, + bidict dst_binding{ + {pt1, mc1}, + {pt2, mc2}, + {pt3, mc3}, + {pt4, mc4}, }; DynamicNodeInvocation input = DynamicNodeInvocation{ @@ -484,10 +474,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, DynamicValueAttrsShardingInfo{ dst_binding.at_r(mc), - one_to_many_filter_keys(dst_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == dst_binding.at_r(mc); - }), + mc, }, }; }; @@ -511,11 +498,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, DynamicValueAttrsShardingInfo{ input_shard_coord, - one_to_many_filter_keys( - src_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == input_shard_coord; - }), + src_binding.at_l(input_shard_coord), }, }, }, @@ -533,16 +516,16 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("bwd") { - OneToMany output_grad_binding{ - {pt1, {mc1}}, - {pt2, {mc2}}, - {pt3, {mc3}}, - {pt4, {mc4}}, + bidict output_grad_binding{ + {pt1, mc1}, + {pt2, mc2}, + {pt3, mc3}, + {pt4, mc4}, }; - OneToMany input_grad_binding{ - {pt1, {mc1}}, - {pt2, {mc2}}, + bidict input_grad_binding{ + {pt1, mc1}, + {pt2, mc2}, }; DynamicNodeInvocation input = DynamicNodeInvocation{ @@ -598,10 +581,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, DynamicValueAttrsShardingInfo{ output_grad_binding.at_r(mc), - one_to_many_filter_keys(output_grad_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == output_grad_binding.at_r(mc); - }), + mc, }, }; }; @@ -625,11 +605,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, DynamicValueAttrsShardingInfo{ input_grad_shard_coord, - one_to_many_filter_keys( - input_grad_binding, - [&](ParallelTensorSpaceCoordinate const &pt_coord) -> bool { - return pt_coord == input_grad_shard_coord; - }), + input_grad_binding.at_l(input_grad_shard_coord), }, }, }, From 768c1ebf9fcbca4828b4ef913e72c2c7f5bf379f Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 5 Jun 2026 01:18:09 -0700 Subject: [PATCH 21/35] Minor include fixes (#1655) --- .../src/compiler/machine_mapping/allowed_machine_views.cc | 2 +- lib/compiler/test/src/compiler/mcmc/generic_mcmc_algorithm.cc | 2 +- lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc | 2 +- .../test/src/compiler/unity_algorithm/graph_optimize_state.cc | 2 +- .../test/src/compiler/unity_algorithm/unity_algorithm.cc | 2 +- lib/op-attrs/include/op-attrs/ff_dim_t.h | 2 +- lib/op-attrs/src/op-attrs/relative_ff_dim_t.cc | 2 +- .../realm-execution/tasks/impl/controller_task_args.dtg.toml | 2 +- .../serializer/serializable_device_specific_ptr.dtg.toml | 4 ++-- lib/utils/include/utils/variant.h | 2 +- lib/utils/test/common/include/test/utils/rapidcheck/doctest.h | 2 +- lib/utils/test/common/include/test/utils/rapidcheck/gen.h | 2 +- .../graph/series_parallel/non_normal_sp_decomposition.cc | 2 +- 13 files changed, 14 insertions(+), 14 deletions(-) diff --git a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc index 2a0402a791..d280e929c5 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc @@ -1,5 +1,5 @@ #include "compiler/machine_mapping/allowed_machine_views.h" -#include "doctest/doctest.h" +#include #include "utils/containers/extend.h" #include "utils/containers/range.h" #include "utils/containers/transform.h" diff --git a/lib/compiler/test/src/compiler/mcmc/generic_mcmc_algorithm.cc b/lib/compiler/test/src/compiler/mcmc/generic_mcmc_algorithm.cc index b21ee4333f..d3c96dee92 100644 --- a/lib/compiler/test/src/compiler/mcmc/generic_mcmc_algorithm.cc +++ b/lib/compiler/test/src/compiler/mcmc/generic_mcmc_algorithm.cc @@ -1,5 +1,5 @@ #include "compiler/mcmc/generic_mcmc_algorithm.h" -#include "doctest/doctest.h" +#include using namespace FlexFlow; diff --git a/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc b/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc index 2584f6b3a6..4c03af5382 100644 --- a/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc +++ b/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc @@ -1,6 +1,6 @@ #include "compiler/mcmc/mcmc_over_mapped_pcg.h" #include "compiler/task_graph_simulator/task_simulator.h" -#include "doctest/doctest.h" +#include #include "internal/runtime_only_cost_estimator_for_test.h" #include "op-attrs/parallel_tensor_dims.h" #include "op-attrs/parallel_tensor_shape.dtg.h" diff --git a/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc b/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc index bf7884df29..cc7ca9425f 100644 --- a/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc +++ b/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc @@ -3,7 +3,7 @@ #include "compiler/machine_mapping/machine_mapping.h" #include "compiler/machine_mapping/machine_view.dtg.h" #include "compiler/machine_mapping/machine_view.h" -#include "doctest/doctest.h" +#include #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_builder.h" #include "test/utils/doctest/check_without_stringify.h" diff --git a/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc b/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc index f5278612aa..0d4123d381 100644 --- a/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc +++ b/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc @@ -1,6 +1,6 @@ #include "compiler/unity_algorithm/unity_algorithm.h" #include "compiler/cost_estimator/runtime_only_cost_estimator_from_cost_estimator.h" -#include "doctest/doctest.h" +#include #include "internal/cost_estimator_for_test.h" #include "op-attrs/parallel_tensor_dims.h" #include "op-attrs/parallel_tensor_shape.dtg.h" diff --git a/lib/op-attrs/include/op-attrs/ff_dim_t.h b/lib/op-attrs/include/op-attrs/ff_dim_t.h index 1411886eee..ae1f3a5cc0 100644 --- a/lib/op-attrs/include/op-attrs/ff_dim_t.h +++ b/lib/op-attrs/include/op-attrs/ff_dim_t.h @@ -3,7 +3,7 @@ #include "op-attrs/ff_dim_t.dtg.h" #include "op-attrs/relative_ff_dim_t.dtg.h" -#include "rapidcheck.h" +#include namespace FlexFlow { diff --git a/lib/op-attrs/src/op-attrs/relative_ff_dim_t.cc b/lib/op-attrs/src/op-attrs/relative_ff_dim_t.cc index 91caa03f36..1dcc8a845e 100644 --- a/lib/op-attrs/src/op-attrs/relative_ff_dim_t.cc +++ b/lib/op-attrs/src/op-attrs/relative_ff_dim_t.cc @@ -1,5 +1,5 @@ #include "op-attrs/relative_ff_dim_t.h" -#include "rapidcheck.h" +#include namespace FlexFlow { ff_dim_t ff_dim_t_from_relative_ff_dim_t(relative_ff_dim_t ff_dim, diff --git a/lib/realm-execution/include/realm-execution/tasks/impl/controller_task_args.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/impl/controller_task_args.dtg.toml index 0c0bd7b96c..9b63503aff 100644 --- a/lib/realm-execution/include/realm-execution/tasks/impl/controller_task_args.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/impl/controller_task_args.dtg.toml @@ -5,7 +5,7 @@ features = [] includes = [ "realm-execution/realm_context.h", - "functional", + "", ] [[fields]] diff --git a/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml index 07cf61f7e1..6a5675772e 100644 --- a/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml @@ -10,8 +10,8 @@ features = [ includes = [ "pcg/device_id_t.dtg.h", - "cstdint", - "optional", + "", + "", ] src_includes = [ diff --git a/lib/utils/include/utils/variant.h b/lib/utils/include/utils/variant.h index c5947413e1..ba689d57be 100644 --- a/lib/utils/include/utils/variant.h +++ b/lib/utils/include/utils/variant.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_UTILS_VARIANT_H #define _FLEXFLOW_UTILS_VARIANT_H -#include "rapidcheck.h" +#include #include "utils/type_traits.h" #include #include diff --git a/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h b/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h index 0121f03926..15ea9fc663 100644 --- a/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h +++ b/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_RAPIDCHECK_DOCTEST_H #define _FLEXFLOW_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_RAPIDCHECK_DOCTEST_H -#include "rapidcheck.h" +#include #include namespace FlexFlow { diff --git a/lib/utils/test/common/include/test/utils/rapidcheck/gen.h b/lib/utils/test/common/include/test/utils/rapidcheck/gen.h index 4fd7c0e069..ad9879d6e7 100644 --- a/lib/utils/test/common/include/test/utils/rapidcheck/gen.h +++ b/lib/utils/test/common/include/test/utils/rapidcheck/gen.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_UTILS_LIB_TEST_COMMON_INCLUDE_UTILS_TEST_RAPIDCHECK_GEN_H #define _FLEXFLOW_UTILS_LIB_TEST_COMMON_INCLUDE_UTILS_TEST_RAPIDCHECK_GEN_H -#include "rapidcheck.h" +#include #include namespace rc { diff --git a/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc b/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc index 5f41d77305..4d91d13c05 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/non_normal_sp_decomposition.h" -#include "doctest/doctest.h" +#include #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" using namespace ::FlexFlow; From bfbb318dc0414d4d0b3ca9d3b62a92b1cbdacc3a Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 29 May 2026 17:43:32 -0700 Subject: [PATCH 22/35] Remove DeviceType from MachineSpaceCoord, move task-spec to use device_id_t, flesh out copy_insertion logic --- .../cost_estimator/op_cost_estimate_key.h | 1 - .../machine_mapping/allowed_machine_views.h | 3 +- .../machine_mapping/machine_mapping.h | 1 - .../machine_mapping_mutation_set.h | 6 +- .../compiler/machine_mapping/machine_view.h | 8 - .../start_invariant_machine_view.dtg.toml | 6 - .../start_invariant_machine_view.h | 2 - .../unstructured_device_mapping.dtg.toml | 27 - .../mcmc/mcmc_over_mapped_pcg_config.dtg.toml | 5 - .../pcg_task_graph.dtg.toml | 4 +- .../cost_estimator/op_cost_estimate_key.cc | 1 - .../machine_mapping/allowed_machine_views.cc | 16 +- .../machine_mapping_mutation_set.cc | 10 +- .../machine_mapping/machine_resource_split.cc | 2 - .../compiler/machine_mapping/machine_view.cc | 19 +- .../start_invariant_machine_view.cc | 13 +- .../unstructured_device_mapping.cc | 26 - .../src/compiler/mcmc/mcmc_over_mapped_pcg.cc | 4 +- .../task_graph_simulator/pcg_task_graph.cc | 8 +- .../task_graph_simulator/task_simulator.cc | 11 +- .../unity_algorithm/unity_algorithm.cc | 2 +- .../machine_mapping/allowed_machine_views.cc | 9 +- .../get_optimal_machine_mapping.cc | 3 - .../get_tensor_set_movement_across_split.cc | 6 - .../machine_mapping/machine_mapping.cc | 4 - .../machine_mapping/machine_mapping_result.cc | 7 - .../compiler/machine_mapping/machine_view.cc | 148 +--- ...get_optimal_machine_mapping_with_memory.cc | 5 - .../machine_mapping_with_memory_result.cc | 8 - .../start_invariant_machine_view.cc | 47 +- .../src/compiler/mcmc/mcmc_over_mapped_pcg.cc | 6 +- .../task_graph_simulator/task_simulator.cc | 16 +- .../computation_graph_instance.h | 2 +- .../cost_estimator/local_cost_estimator.h | 2 +- .../local_task_argument_accessor.h | 2 +- .../per_device_op_state_initialization.h | 2 +- .../cost_estimator/local_cost_estimator.cc | 3 +- .../local_task_argument_accessor.cc | 4 +- .../computation_graph_instance.cc | 30 +- .../cost_estimator/local_cost_estimator.cc | 24 +- .../local_task_argument_accessor.cc | 10 +- lib/pcg/include/pcg/device_id.h | 22 - lib/pcg/include/pcg/device_id_t.dtg.toml | 23 - lib/pcg/include/pcg/device_id_t.h | 14 - .../machine_compute_resource_slice.dtg.toml | 2 + .../pcg/machine_compute_specification.h | 7 - .../pcg/machine_space_coordinate.dtg.toml | 7 +- .../include/pcg/machine_space_offset.dtg.toml | 8 +- lib/pcg/src/pcg/device_id.cc | 51 -- lib/pcg/src/pcg/device_id_t.cc | 17 - .../src/pcg/machine_compute_resource_slice.cc | 1 - .../src/pcg/machine_compute_specification.cc | 17 - lib/pcg/src/pcg/machine_space_offset.cc | 3 - .../src/pcg/machine_compute_specification.cc | 23 - .../mapped_operator_task_group.cc | 23 +- .../mapped_parallel_computation_graph.cc | 1 - ...ce_specific_managed_per_device_ff_handle.h | 2 +- .../realm-execution/device_specific_ptr.h | 3 +- .../realm-execution/distributed_ff_handle.h | 6 +- ...buted_per_device_op_state_initialization.h | 3 +- .../realm-execution/instance_allocation.h | 10 +- .../include/realm-execution/pcg_instance.h | 5 +- .../include/realm-execution/processor_kind.h | 14 + .../include/realm-execution/realm_context.h | 17 +- .../include/realm-execution/realm_manager.h | 2 +- ...uted_per_device_op_state_initialization.cc | 3 +- .../realm-execution/instance_allocation.cc | 16 +- .../src/realm-execution/pcg_instance.cc | 42 +- .../src/realm-execution/processor_kind.cc | 29 + .../src/realm-execution/realm_allocator.cc | 10 +- .../src/realm-execution/realm_context.cc | 98 +-- .../src/realm-execution/realm_manager.cc | 1 + .../tasks/impl/controller_task.cc | 6 +- .../test/src/realm-execution/test_e2e.cc | 529 +++++-------- .../include/task-spec/concrete_arg_spec.h | 1 - .../include/task-spec/device_id_t.dtg.toml | 23 + .../include/task-spec/device_specific.h | 16 +- .../dynamic_graph/dynamic_graph_edge.dtg.toml | 29 + .../dynamic_graph/dynamic_graph_edge.h | 15 + .../dynamic_graph/dynamic_node_attrs.dtg.toml | 13 +- .../dynamic_graph/dynamic_node_invocation.h | 22 + .../dynamic_node_mapping.dtg.toml | 26 + .../dynamic_graph/dynamic_node_mapping.h | 18 + .../dynamic_open_dataflow_graph.h | 27 + .../dynamic_graph/dynamic_slot_site.dtg.toml | 21 + .../dynamic_graph/dynamic_slot_site.h | 13 + .../dynamic_value_attrs.dtg.toml | 9 +- .../dynamic_graph/dynamic_value_attrs.h | 5 + .../external_dynamic_slot_site.dtg.toml | 19 + .../internal_dynamic_slot_site.dtg.toml | 29 + .../task-spec/dynamic_graph/loss_insertion.h | 2 +- .../task-spec/dynamic_graph/machine_slicing.h | 4 +- ...amic_open_dataflow_graph_from_mapped_pcg.h | 2 +- .../parallel_tensor_mapping.dtg.toml | 22 + .../serializable_dynamic_node_attrs.dtg.toml | 8 +- .../serializable_dynamic_value_attrs.dtg.toml | 3 +- .../training_only_op_type.dtg.toml | 15 + .../dynamic_graph/training_op_type.dtg.toml | 22 + .../dynamic_graph/training_operation_attrs.h | 13 + .../include/task-spec/serialization.h | 113 --- .../itask_argument_accessor.h | 2 +- .../task-spec/dynamic_graph/copy_insertion.cc | 263 ++++--- .../dynamic_graph/dynamic_graph_edge.cc | 20 + .../dynamic_graph/dynamic_node_invocation.cc | 62 ++ .../dynamic_graph/dynamic_node_mapping.cc | 31 + .../dynamic_open_dataflow_graph.cc | 181 ++++- .../dynamic_graph/dynamic_slot_site.cc | 26 + .../dynamic_graph/dynamic_value_attrs.cc | 9 + .../task-spec/dynamic_graph/loss_insertion.cc | 2 +- .../dynamic_graph/machine_slicing.cc | 4 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 8 +- .../dynamic_graph/shard_expansion.cc | 39 +- .../dynamic_graph/training_operation_attrs.cc | 24 + .../dynamic_graph/update_insertion.cc | 4 +- lib/task-spec/src/task-spec/serialization.cc | 1 - lib/task-spec/test/CMakeLists.txt | 1 - .../test/src/task-spec/device_specific.cc | 10 +- .../task-spec/dynamic_graph/copy_insertion.cc | 731 ++++++++++-------- .../dynamic_open_dataflow_graph.cc | 501 ++++++++++-- .../dynamic_graph/machine_slicing.cc | 21 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 1 + .../dynamic_graph/shard_expansion.cc | 508 ++++++------ .../dynamic_graph/update_insertion.cc | 157 ++++ 123 files changed, 2679 insertions(+), 1945 deletions(-) delete mode 100644 lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml delete mode 100644 lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc delete mode 100644 lib/pcg/include/pcg/device_id.h delete mode 100644 lib/pcg/include/pcg/device_id_t.dtg.toml delete mode 100644 lib/pcg/include/pcg/device_id_t.h delete mode 100644 lib/pcg/src/pcg/device_id.cc delete mode 100644 lib/pcg/src/pcg/device_id_t.cc create mode 100644 lib/realm-execution/include/realm-execution/processor_kind.h create mode 100644 lib/realm-execution/src/realm-execution/processor_kind.cc create mode 100644 lib/task-spec/include/task-spec/device_id_t.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.h create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/training_only_op_type.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h delete mode 100644 lib/task-spec/include/task-spec/serialization.h create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc delete mode 100644 lib/task-spec/src/task-spec/serialization.cc create mode 100644 lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc create mode 100644 lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc diff --git a/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.h b/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.h index d905abeb77..dae6989da8 100644 --- a/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.h +++ b/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.h @@ -3,7 +3,6 @@ #include "compiler/cost_estimator/op_cost_estimate_key.dtg.h" #include "compiler/cost_estimator/runtime_only_op_cost_estimate_key.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/machine_specification.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.dtg.h" #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" diff --git a/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h b/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h index 5201f7fa31..de899ab74d 100644 --- a/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h +++ b/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h @@ -14,8 +14,7 @@ bool is_valid_machine_view(MachineView const &mv, std::unordered_set get_allowed_machine_views(MachineComputeResourceSlice const &machine_spec, - OperatorTaskSpace const &task, - DeviceType device_type); + OperatorTaskSpace const &task); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping.h index 28d1fc943b..7ace9989fb 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping.h @@ -4,7 +4,6 @@ #include "compiler/machine_mapping/machine_mapping.dtg.h" #include "compiler/machine_mapping/machine_mapping_result.h" #include "compiler/series_parallel/pcg/pcg_binary_sp_decomposition.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/machine_specification.dtg.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.dtg.h" diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_mutation_set.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping_mutation_set.h index d05e3fab7c..005e2a1ee5 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_mutation_set.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_mutation_set.h @@ -7,13 +7,11 @@ namespace FlexFlow { std::optional get_random_mapping(ParallelComputationGraph const &pcg, - MachineComputeSpecification const &resources, - DeviceType const &device_type); + MachineComputeSpecification const &resources); std::optional get_random_mutation(SearchResult const &mapped_pcg, - MachineComputeSpecification const &resource, - DeviceType const &device_type); + MachineComputeSpecification const &resource); } // namespace FlexFlow #endif diff --git a/lib/compiler/include/compiler/machine_mapping/machine_view.h b/lib/compiler/include/compiler/machine_mapping/machine_view.h index 6888fa6b94..4a9e12215c 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_view.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_view.h @@ -6,7 +6,6 @@ #include "op-attrs/operator_task_space.dtg.h" #include "op-attrs/parallel_tensor_dim_degrees.dtg.h" #include "op-attrs/task_space_coordinate.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/machine_compute_specification.dtg.h" #include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" #include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.h" @@ -20,8 +19,6 @@ namespace FlexFlow { nonnegative_int mv_get_expected_task_space_num_dims(MachineView const &mv); -DeviceType get_device_type(MachineView const &mv); - std::vector get_strides(MachineView const &mv); std::vector @@ -50,11 +47,6 @@ std::unordered_set get_machine_space_coordinates(OperatorTaskSpace const &task, MachineView const &mv); -std::unordered_set - get_device_ids(OperatorTaskSpace const &task, - MachineView const &mv, - MachineComputeSpecification const &ms); - MachineView make_1d_machine_view(MachineSpaceCoordinate const &start, MachineSpecificationDimension const &dim, stride_t stride); diff --git a/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.dtg.toml b/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.dtg.toml index b271ed1959..7a0460a3b6 100644 --- a/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.dtg.toml @@ -12,7 +12,6 @@ features = [ includes = [ "compiler/machine_mapping/machine_view_dimension.dtg.h", - "pcg/device_type.dtg.h" ] src_includes = [ @@ -23,8 +22,3 @@ src_includes = [ [[fields]] name = "dimensions" type = "std::vector<::FlexFlow::MachineViewDimension>" - - -[[fields]] -name = "device_type" -type = "::FlexFlow::DeviceType" diff --git a/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h b/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h index 631d23b07c..7b5a4f57f0 100644 --- a/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h +++ b/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h @@ -19,8 +19,6 @@ StartInvariantMachineView nonnegative_int num_dims(StartInvariantMachineView const &mv); -DeviceType get_device_type(StartInvariantMachineView const &mv); - std::vector get_strides(StartInvariantMachineView const &mv); std::vector diff --git a/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml b/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml deleted file mode 100644 index 28391eddc0..0000000000 --- a/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml +++ /dev/null @@ -1,27 +0,0 @@ -namespace = "FlexFlow" -name = "UnstructuredDeviceMapping" -type = "struct" -features = [ - "eq", - # "ord", - "hash", - # "json", - # "rapidcheck", - "fmt", -] - -includes = [ - "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", - "pcg/device_id_t.dtg.h" -] - -src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h" -] - -[[fields]] -name = "raw_device_map" -type = "std::unordered_map<::FlexFlow::parallel_layer_guid_t, std::unordered_set<::FlexFlow::device_id_t>>" diff --git a/lib/compiler/include/compiler/mcmc/mcmc_over_mapped_pcg_config.dtg.toml b/lib/compiler/include/compiler/mcmc/mcmc_over_mapped_pcg_config.dtg.toml index 99320f735e..cf8476e882 100644 --- a/lib/compiler/include/compiler/mcmc/mcmc_over_mapped_pcg_config.dtg.toml +++ b/lib/compiler/include/compiler/mcmc/mcmc_over_mapped_pcg_config.dtg.toml @@ -8,7 +8,6 @@ features = [ ] includes = [ - "pcg/device_type.dtg.h", "utils/nonnegative_int/nonnegative_int.h" ] @@ -23,7 +22,3 @@ type = "::FlexFlow::nonnegative_int" [[fields]] name = "substitution_frequency" type = "float" - -[[fields]] -name = "device_type" -type = "::FlexFlow::DeviceType" \ No newline at end of file diff --git a/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml index 2c5b5f56fc..14de2d9455 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml @@ -9,8 +9,8 @@ includes = [ "utils/graph/digraph/digraph_view.h", "utils/bidict/bidict.h", "compiler/task_graph_simulator/pcg_task.dtg.h", - "pcg/device_id_t.dtg.h", "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", + "pcg/machine_space_coordinate.dtg.h", "", "" ] @@ -32,4 +32,4 @@ type = "::FlexFlow::bidict<::FlexFlow::Node, ::FlexFlow::PCGTask>" [[fields]] name = "node_to_devices" -type = "std::unordered_map<::FlexFlow::Node, std::unordered_set<::FlexFlow::device_id_t>>" +type = "std::unordered_map<::FlexFlow::Node, std::unordered_set<::FlexFlow::MachineSpaceCoordinate>>" diff --git a/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc b/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc index f3edd6a69a..62b645e3da 100644 --- a/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc +++ b/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc @@ -4,7 +4,6 @@ #include "compiler/machine_mapping/machine_view.dtg.h" #include "compiler/machine_mapping/machine_view.h" #include "op-attrs/parallel_tensor_shape.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/machine_specification.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.dtg.h" #include diff --git a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc index c3d9ae7bfb..9194c2e982 100644 --- a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc @@ -49,8 +49,7 @@ bool is_valid_machine_view(MachineView const &mv, */ static std::unordered_set get_candidate_machine_views(MachineComputeResourceSlice const &machine_spec, - OperatorTaskSpace const &task_space, - DeviceType const &device_type) { + OperatorTaskSpace const &task_space) { auto get_max_stride_upper_bound = [](std::vector const &tensor_dims, @@ -90,17 +89,15 @@ static std::unordered_set return strides; }; - auto get_candidate_starts = [](MachineComputeResourceSlice const &slice, - DeviceType const &device_type) + auto get_candidate_starts = [](MachineComputeResourceSlice const &slice) -> std::unordered_set { - ASSERT(device_type == DeviceType::GPU); std::unordered_set result; for (nonnegative_int node_idx : nonnegative_range(slice.num_nodes)) { for (nonnegative_int device_idx : nonnegative_range(slice.num_gpus_per_node)) { result.insert( - MachineSpaceCoordinate{node_idx, device_idx, device_type}); + MachineSpaceCoordinate{node_idx, device_idx}); } } return result; @@ -127,7 +124,7 @@ static std::unordered_set ASSERT(candidate_strides.size() > 0); std::unordered_set candidate_starts = - get_candidate_starts(machine_spec, device_type); + get_candidate_starts(machine_spec); ASSERT(candidate_starts.size() > 0); std::unordered_multiset> @@ -151,11 +148,10 @@ static std::unordered_set std::unordered_set get_allowed_machine_views(MachineComputeResourceSlice const &machine_spec, - OperatorTaskSpace const &task_space, - DeviceType device_type) { + OperatorTaskSpace const &task_space) { std::unordered_set views = - get_candidate_machine_views(machine_spec, task_space, device_type); + get_candidate_machine_views(machine_spec, task_space); return filter(views, [&](MachineView const &mv) { return is_valid_machine_view(mv, task_space, machine_spec); }); diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc index 47639ff88a..d6cdca97d1 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc @@ -11,15 +11,14 @@ namespace FlexFlow { std::optional get_random_mapping(ParallelComputationGraph const &pcg, - MachineComputeSpecification const &resources, - DeviceType const &device_type) { + MachineComputeSpecification const &resources) { std::vector layers = topological_ordering(pcg); std::unordered_map machine_views; for (parallel_layer_guid_t layer : layers) { OperatorTaskSpace task = get_operator_task_space(pcg, layer); std::unordered_set allowed_machine_views = get_allowed_machine_views( - compute_slice_from_specification(resources), task, DeviceType::GPU); + compute_slice_from_specification(resources), task); if (allowed_machine_views.empty()) { return std::nullopt; } @@ -31,8 +30,7 @@ std::optional std::optional get_random_mutation(SearchResult const &mapped_pcg, - MachineComputeSpecification const &resources, - DeviceType const &device_type) { + MachineComputeSpecification const &resources) { ParallelComputationGraph pcg = mapped_pcg.pcg; std::vector layers = topological_ordering(pcg); if (layers.size() == 0) { @@ -46,7 +44,7 @@ std::optional std::vector allowed_machine_views = vector_of(get_allowed_machine_views( - compute_slice_from_specification(resources), task, device_type)); + compute_slice_from_specification(resources), task)); MachineView random_new_machine_view = select_random(allowed_machine_views); machine_mapping.machine_views.at(random_layer) = random_new_machine_view; diff --git a/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc b/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc index 875f44a0c9..370b764a69 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc @@ -88,7 +88,6 @@ MachineSpaceCoordinate /*node_idx=*/(coord.node_idx + split.offset) .nonnegative_int_from_positive_int(), /*device_idx=*/coord.device_idx, - /*device_type=*/coord.device_type, }; } else { ASSERT(split.dimension == MachineSpecificationDimension::INTRA_NODE); @@ -97,7 +96,6 @@ MachineSpaceCoordinate /*node_idx=*/coord.node_idx, /*device_idx=*/ (coord.device_idx + split.offset).nonnegative_int_from_positive_int(), - /*device_type=*/coord.device_type, }; } } diff --git a/lib/compiler/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/src/compiler/machine_mapping/machine_view.cc index 090dec5845..5c38a66901 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_view.cc @@ -33,10 +33,6 @@ nonnegative_int mv_get_expected_task_space_num_dims(MachineView const &mv) { return num_elements(get_strides(mv)); } -DeviceType get_device_type(MachineView const &mv) { - return mv.start.device_type; -} - std::vector get_strides(MachineView const &mv) { return transform(mv.dimensions, [](MachineViewDimension const &dim) { return dim.stride; }); @@ -130,7 +126,7 @@ MachineSpaceCoordinate nonnegative_int device_idx = compute_index(machine_view.start.device_idx, intra_dimension_indices); MachineSpaceCoordinate ms_coord = MachineSpaceCoordinate{ - node_idx, device_idx, get_device_type(machine_view)}; + node_idx, device_idx}; return ms_coord; } @@ -177,19 +173,6 @@ std::unordered_set }); } -std::unordered_set - get_device_ids(OperatorTaskSpace const &task_space, - MachineView const &mv, - MachineComputeSpecification const &ms) { - ASSERT(op_task_space_num_dims(task_space) == - mv_get_expected_task_space_num_dims(mv)); - - return transform(get_machine_space_coordinates(task_space, mv), - [&](MachineSpaceCoordinate const &coord) { - return get_device_id(ms, coord); - }); -} - MachineView make_1d_machine_view(MachineSpaceCoordinate const &start, MachineSpecificationDimension const &dim, stride_t stride) { diff --git a/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc b/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc index cbb64d5bcf..4a2d66acc1 100644 --- a/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc @@ -18,17 +18,13 @@ MachineView machine_view_from_start_invariant( StartInvariantMachineView start_invariant_from_machine_view(MachineView const &mv) { - return StartInvariantMachineView{mv.dimensions, get_device_type(mv)}; + return StartInvariantMachineView{mv.dimensions}; } nonnegative_int num_dims(StartInvariantMachineView const &start_inv_mv) { return num_elements(start_inv_mv.dimensions); } -DeviceType get_device_type(StartInvariantMachineView const &start_inv_mv) { - return start_inv_mv.device_type; -} - std::vector get_strides(StartInvariantMachineView const &start_inv_mv) { return transform(start_inv_mv.dimensions, @@ -45,13 +41,12 @@ std::vector StartInvariantMachineView start_invariant_machine_view_from_strides_and_machine_spec_dimensions( std::vector const &strides, - std::vector const &dims, - DeviceType device_type) { + std::vector const &dims) { std::vector dimensions = transform(zip(strides, dims), [&](auto const &p) { return MachineViewDimension{p.first, p.second}; }); - return StartInvariantMachineView{dimensions, device_type}; + return StartInvariantMachineView{dimensions}; } MachineSpaceOffset get_machine_space_offset( @@ -60,7 +55,7 @@ MachineSpaceOffset get_machine_space_offset( TaskSpaceCoordinate const &coord) { MachineSpaceCoordinate dummy_start = - MachineSpaceCoordinate{0_n, 0_n, get_device_type(start_inv_machine_view)}; + MachineSpaceCoordinate{0_n, 0_n}; MachineView mv = machine_view_from_start_invariant(start_inv_machine_view, dummy_start); diff --git a/lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc b/lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc deleted file mode 100644 index 80c09d2dba..0000000000 --- a/lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc +++ /dev/null @@ -1,26 +0,0 @@ -#include "compiler/machine_mapping/unstructured_device_mapping.h" -#include "compiler/machine_mapping/machine_view.h" -#include "compiler/machine_mapping/unstructured_device_mapping.dtg.h" -#include "op-attrs/operator_task_space.dtg.h" -#include "op-attrs/operator_task_space.h" -#include "pcg/parallel_computation_graph/parallel_computation_graph.h" -#include "utils/containers/keys.h" -#include "utils/containers/map_values.h" - -namespace FlexFlow { - -UnstructuredDeviceMapping get_unstructured_device_mapping( - MachineMapping const &machine_mapping, - MachineComputeSpecification const &machine_spec, - ParallelComputationGraph const &pcg) { - std::unordered_map> - device_mapping; - for (auto const &[layer, machine_view] : machine_mapping.machine_views) { - OperatorTaskSpace op = get_operator_task_space(pcg, layer); - device_mapping.insert( - {layer, get_device_ids(op, machine_view, machine_spec)}); - } - return UnstructuredDeviceMapping{device_mapping}; -} - -} // namespace FlexFlow diff --git a/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc b/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc index 583a60b1ad..0d2c1e4455 100644 --- a/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc +++ b/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc @@ -22,7 +22,7 @@ SearchResult MachineComputeSpecification compute_spec = machine_spec.compute_specification; std::vector substitutions = get_substitution_set(compute_spec); MachineMapping random_mapping = assert_unwrap( - get_random_mapping(pcg, compute_spec, search_config.device_type)); + get_random_mapping(pcg, compute_spec)); SearchResult starting_state = SearchResult{pcg, random_mapping}; auto sampler = [&](SearchResult mapped_pcg) -> std::optional { @@ -43,7 +43,7 @@ SearchResult }); } else { MachineMapping new_machine_mapping = assert_unwrap(get_random_mutation( - mapped_pcg, compute_spec, search_config.device_type)); + mapped_pcg, compute_spec)); return SearchResult{mapped_pcg.pcg, new_machine_mapping}; } }; diff --git a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc index d4d5a78d6a..058f5b72bb 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc @@ -5,7 +5,6 @@ #include "compiler/machine_mapping/machine_view.dtg.h" #include "compiler/machine_mapping/machine_view.h" #include "op-attrs/operator_task_space.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/machine_specification.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.dtg.h" @@ -25,7 +24,7 @@ PCGTaskGraph DiGraph digraph = DiGraph::create(); bidict node_to_task; bidict node_to_layer; - std::unordered_map> node_to_devices; + std::unordered_map> node_to_devices; for (parallel_layer_guid_t const &layer : get_parallel_layers(pcg)) { MachineView mv = machine_mapping.machine_views.at(layer); @@ -35,9 +34,8 @@ PCGTaskGraph node_to_task.equate(node, PCGTask{op_key}); node_to_layer.equate(node, layer); node_to_devices[node] = - get_device_ids(get_operator_task_space(pcg, layer), - machine_mapping.machine_views.at(layer), - machine_spec); + get_machine_space_coordinates(get_operator_task_space(pcg, layer), + machine_mapping.machine_views.at(layer)); } for (ParallelComputationGraphEdge const &edge : get_edges(pcg)) { diff --git a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc index e514e6a753..28a6a3efae 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc @@ -1,8 +1,6 @@ #include "compiler/task_graph_simulator/task_simulator.h" #include "compiler/cost_estimator/cost_estimator.h" #include "compiler/cost_estimator/op_cost_estimate_key.h" -#include "compiler/machine_mapping/unstructured_device_mapping.dtg.h" -#include "compiler/machine_mapping/unstructured_device_mapping.h" #include "compiler/task_graph_simulator/pcg_task.dtg.h" #include "compiler/task_graph_simulator/pcg_task_graph.h" #include "compiler/task_graph_simulator/simulate_task_graph_execution.h" @@ -49,21 +47,18 @@ milliseconds_t task_simulator_estimate_forward_pass_time( std::unordered_set const &finished_tasks) -> bool { PCGTask current_task = task_graph.node_to_task.at_l(task); - UnstructuredDeviceMapping device_map = get_unstructured_device_mapping( - machine_mapping, machine_spec.compute_specification, pcg); - if (current_task.is_tensor_movement()) { return true; } assert(current_task.is_operator()); - auto get_devices = [&](Node const &n) { + auto get_devices = [&](Node const &n) -> std::unordered_set { return task_graph.node_to_devices.at(n); }; - std::unordered_set devices_occupied = + std::unordered_set devices_occupied = set_union(transform(in_progress_tasks, get_devices)); - std::unordered_set required_devices = get_devices(task); + std::unordered_set required_devices = get_devices(task); return set_intersection(devices_occupied, required_devices).empty(); }; diff --git a/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc b/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc index be8c7c4f98..240ad0b1c4 100644 --- a/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc +++ b/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc @@ -68,7 +68,7 @@ SearchResult graph_optimize(ParallelComputationGraph &pcg, get_operator_task_space_for_runtime_only_op_cost_estimate_key(key); return get_allowed_machine_views( - resources, op_task_space, DeviceType::GPU); + resources, op_task_space); }, }; diff --git a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc index d280e929c5..6a867b16f3 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc @@ -40,7 +40,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate{ start_node_idx, start_device_idx, - DeviceType::GPU, }, strides, }; @@ -65,7 +64,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; std::unordered_set result = - get_allowed_machine_views(ms, task, DeviceType::GPU); + get_allowed_machine_views(ms, task); CHECK(correct == result); } @@ -96,7 +95,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; std::unordered_set result = - get_allowed_machine_views(ms, task, DeviceType::GPU); + get_allowed_machine_views(ms, task); CHECK(correct == result); } @@ -110,7 +109,7 @@ TEST_SUITE(FF_TEST_SUITE) { OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{}}}; std::unordered_set result = - get_allowed_machine_views(full_machine_spec, task, DeviceType::GPU); + get_allowed_machine_views(full_machine_spec, task); std::unordered_set correct = { make_machine_view(0_n, 0_n), @@ -129,7 +128,7 @@ TEST_SUITE(FF_TEST_SUITE) { OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{2_ge2}}}; std::unordered_set result = - get_allowed_machine_views(full_machine_spec, task, DeviceType::GPU); + get_allowed_machine_views(full_machine_spec, task); std::unordered_set correct = { make_machine_view(0_n, 0_n, /*stride_1=*/1_p, intra), diff --git a/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc b/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc index 392e16bec5..59ed83d6fa 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc @@ -53,7 +53,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -68,7 +67,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -557,7 +555,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/2_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { diff --git a/lib/compiler/test/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc b/lib/compiler/test/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc index b3901e08ca..e2a87204f0 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc @@ -93,7 +93,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -108,7 +107,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/1_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -123,7 +121,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/2_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -138,7 +135,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/3_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -160,13 +156,11 @@ TEST_SUITE(FF_TEST_SUITE) { /*src=*/MachineSpaceCoordinate{ /*node_idx=*/src_mv.start.node_idx, /*device_idx=*/src_task_idx, - /*device_type=*/DeviceType::GPU, }, /*dst=*/ MachineSpaceCoordinate{ /*node_idx=*/dst_mv.start.node_idx, /*device_idx=*/dst_task_idx, - /*device_type=*/DeviceType::GPU, }, }; }; diff --git a/lib/compiler/test/src/compiler/machine_mapping/machine_mapping.cc b/lib/compiler/test/src/compiler/machine_mapping/machine_mapping.cc index 8af07a032c..6cae9a2e50 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/machine_mapping.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/machine_mapping.cc @@ -11,7 +11,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -26,7 +25,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -57,7 +55,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -72,7 +69,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { diff --git a/lib/compiler/test/src/compiler/machine_mapping/machine_mapping_result.cc b/lib/compiler/test/src/compiler/machine_mapping/machine_mapping_result.cc index 75dd63cccb..230602be98 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/machine_mapping_result.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/machine_mapping_result.cc @@ -10,7 +10,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -25,7 +24,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -191,7 +189,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -206,7 +203,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -285,7 +281,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/3_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -335,7 +330,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -350,7 +344,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { diff --git a/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc index 2ea8312991..29fdfff3f6 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc @@ -16,7 +16,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - DeviceType::GPU, }, { MachineViewDimension{ @@ -33,28 +32,6 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(mv_get_expected_task_space_num_dims(mv) == 2_n); } - TEST_CASE("get_device_type") { - MachineView mv = MachineView{ - MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - DeviceType::GPU, - }, - { - MachineViewDimension{ - stride_t{2_p}, - MachineSpecificationDimension::INTER_NODE, - }, - MachineViewDimension{ - stride_t{2_p}, - MachineSpecificationDimension::INTER_NODE, - }, - }, - }; - - CHECK(get_device_type(mv) == DeviceType::GPU); - } - TEST_CASE("get_machine_space_coordinate") { SUBCASE("1D case") { /** @@ -81,7 +58,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/1_n, - DeviceType::GPU, }, { MachineViewDimension{ @@ -100,7 +76,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate correct = MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/1_n, - DeviceType::GPU, }; CHECK(result == correct); @@ -115,7 +90,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate correct = MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/3_n, - DeviceType::GPU, }; CHECK(result == correct); @@ -130,7 +104,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate correct = MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/5_n, - DeviceType::GPU, }; CHECK(result == correct); @@ -174,7 +147,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate{ /*node_idx=*/1_n, /*device_idx=*/2_n, - DeviceType::GPU, }, { MachineViewDimension{ @@ -191,7 +163,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 0_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/2_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/2_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -200,7 +172,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/4_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/4_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -209,7 +181,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/2_n, /*device_idx=*/2_n, DeviceType::GPU}; + /*node_idx=*/2_n, /*device_idx=*/2_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -218,7 +190,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/2_n, /*device_idx=*/4_n, DeviceType::GPU}; + /*node_idx=*/2_n, /*device_idx=*/4_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -250,7 +222,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate{ /*node_idx=*/1_n, /*device_idx=*/0_n, - DeviceType::GPU, }, { MachineViewDimension{ @@ -267,7 +238,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 0_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/0_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/0_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -276,7 +247,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/4_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/4_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -285,7 +256,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/1_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/1_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -294,7 +265,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/5_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/5_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -332,7 +303,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; MachineView mv = MachineView{ MachineSpaceCoordinate{ - /*node_idx=*/0_n, /*device_idx=*/1_n, DeviceType::GPU}, + /*node_idx=*/0_n, /*device_idx=*/1_n}, {MachineViewDimension{stride_t{1_p}, MachineSpecificationDimension::INTER_NODE}, MachineViewDimension{stride_t{2_p}, @@ -343,7 +314,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n, 0_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/0_n, /*device_idx=*/3_n, DeviceType::GPU}; + /*node_idx=*/0_n, /*device_idx=*/3_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -352,7 +323,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n, 1_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/5_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/5_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -361,106 +332,11 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n, 1_n}); MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/7_n, DeviceType::GPU}; + /*node_idx=*/1_n, /*device_idx=*/7_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); } } } - - TEST_CASE("get_device_ids") { - SUBCASE("1D machine view") { - /** - * This operator has shape (3,), and thus 3 tasks. - * The (only) dimension is projected onto the INTRA (device) dimension - * with a stride of 2. The start of the projection defined by MachineView - * is at MachineSpaceCoordinate (0, 1), and the machine space has 1 node - * and 6 devices per node. - * - * The tasks will thus be distributed like this: - * +-------+-------+-------+-------+-------+-------+ - * | 0 | ((1)) | 2 | ((3)) | 4 | ((5)) | - * +-------+-------+-------+-------+-------+-------+ - * Where the integers are the device ids and ((x)) are the devices we - * select - */ - MachineComputeSpecification ms = MachineComputeSpecification{ - /*num_nodes=*/1_p, - /*num_cpus_per_node=*/6_p, - /*num_gpus_per_node=*/6_p, - }; - - OperatorTaskSpace task = OperatorTaskSpace{ - MinimalOrthotope{{ - 3_ge2, - }}, - }; - MachineView mv = MachineView{ - MachineSpaceCoordinate{ - /*node_idx=*/0_n, /*device_idx=*/1_n, DeviceType::GPU}, - {MachineViewDimension{stride_t{2_p}, - MachineSpecificationDimension::INTRA_NODE}}}; - - std::unordered_set correct = { - device_id_t{gpu_id_t{1_n}}, - device_id_t{gpu_id_t{3_n}}, - device_id_t{gpu_id_t{5_n}}, - }; - std::unordered_set result = get_device_ids(task, mv, ms); - CHECK(result == correct); - } - - SUBCASE("2D machine view") { - /** - * This operator has shape (2, 2), and thus 2 * 2 = 4 tasks. - * - The first dimension is projected onto the INTER (node) dimension with - * stride 1, - * - The second dimension is projected onto the INTRA (device) dimension - * with stride 2. The start of the projection defined by MachineView is at - * MachineSpaceCoordinate (1, 2), and the machine space has 3 nodes and 5 - * devices per node. - * - * The tasks will thus be distributed like this: - * +-------+-------+-------+-------+-------+ - * | 0 | 1 | 2 | 3 | 4 | - * +-------+-------+-------+-------+-------+ - * | 5 | 6 | ((7)) | 8 | ((9)) | - * +-------+-------+-------+-------+-------+ - * | 10 | 11 | ((12))| 13 | ((14))| - * +-------+-------+-------+-------+-------+ - * Where the integers are the device ids and ((x)) are the devices we - * select - */ - - MachineComputeSpecification ms = MachineComputeSpecification{ - /*num_nodes=*/3_p, - /*num_cpus_per_node=*/5_p, - /*num_gpus_per_node=*/5_p, - }; - - OperatorTaskSpace task = OperatorTaskSpace{ - MinimalOrthotope{{ - 2_ge2, - 2_ge2, - }}, - }; - MachineView mv = MachineView{ - MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/2_n, DeviceType::GPU}, - {MachineViewDimension{stride_t{1_p}, - MachineSpecificationDimension::INTER_NODE}, - MachineViewDimension{stride_t{2_p}, - MachineSpecificationDimension::INTRA_NODE}}}; - - std::unordered_set correct = { - device_id_t{gpu_id_t{7_n}}, - device_id_t{gpu_id_t{9_n}}, - device_id_t{gpu_id_t{12_n}}, - device_id_t{gpu_id_t{14_n}}, - }; - std::unordered_set result = get_device_ids(task, mv, ms); - CHECK(result == correct); - } - } } diff --git a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc index 54717d6699..8d21131c7d 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc @@ -52,7 +52,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/{}, }; @@ -61,7 +60,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -76,7 +74,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -91,7 +88,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/1_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -724,7 +720,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/2_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { diff --git a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc index 402dbe66d7..086c1ab984 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc @@ -104,7 +104,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -119,7 +118,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -307,7 +305,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -322,7 +319,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -410,7 +406,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/3_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -464,7 +459,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -479,7 +473,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { @@ -494,7 +487,6 @@ TEST_SUITE(FF_TEST_SUITE) { /*start=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }, /*dimensions=*/ { diff --git a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc index 3159f49118..e3b08e3805 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc @@ -12,8 +12,7 @@ TEST_SUITE(FF_TEST_SUITE) { {MachineViewDimension{stride_t{2_p}, MachineSpecificationDimension::INTER_NODE}, MachineViewDimension{stride_t{2_p}, - MachineSpecificationDimension::INTER_NODE}}, - DeviceType::GPU}; + MachineSpecificationDimension::INTER_NODE}}}; SUBCASE("num_dims") { nonnegative_int result = num_dims(simv); @@ -21,12 +20,6 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } - SUBCASE("get_device_type") { - DeviceType result = get_device_type(simv); - DeviceType correct = DeviceType::GPU; - CHECK(result == correct); - } - SUBCASE("get_strides") { std::vector result = get_strides(simv); std::vector correct = {stride_t{2_p}, stride_t{2_p}}; @@ -44,7 +37,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("StartInvariantMachineView - conversions") { MachineSpaceCoordinate start = - MachineSpaceCoordinate{1_n, 2_n, DeviceType::GPU}; + MachineSpaceCoordinate{1_n, 2_n}; std::vector dimensions = { MachineViewDimension{stride_t{2_p}, MachineSpecificationDimension::INTER_NODE}, @@ -53,7 +46,7 @@ TEST_SUITE(FF_TEST_SUITE) { MachineView mv = MachineView{start, dimensions}; StartInvariantMachineView simv = - StartInvariantMachineView{dimensions, DeviceType::GPU}; + StartInvariantMachineView{dimensions}; SUBCASE("start_invariant_from_machine_view") { StartInvariantMachineView result = start_invariant_from_machine_view(mv); @@ -102,8 +95,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; StartInvariantMachineView simv = StartInvariantMachineView{ {MachineViewDimension{stride_t{2_p}, - MachineSpecificationDimension::INTRA_NODE}}, - DeviceType::GPU}; + MachineSpecificationDimension::INTRA_NODE}}}; MachineComputeSpecification ms = MachineComputeSpecification{ /*num_nodes=*/1_p, /*num_cpus_per_node=*/6_p, @@ -114,7 +106,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n}); MachineSpaceOffset correct = - MachineSpaceOffset{0, 0, DeviceType::GPU}; + MachineSpaceOffset{0, 0}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -123,7 +115,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n}); MachineSpaceOffset correct = - MachineSpaceOffset{0, 2, DeviceType::GPU}; + MachineSpaceOffset{0, 2}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -132,7 +124,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (2,)") { TaskSpaceCoordinate coord = make_task_space_coordinate({2_n}); MachineSpaceOffset correct = - MachineSpaceOffset{0, 4, DeviceType::GPU}; + MachineSpaceOffset{0, 4}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -141,9 +133,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_machine_space_offsets") { std::unordered_set correct = { - MachineSpaceOffset{0, 0, DeviceType::GPU}, - MachineSpaceOffset{0, 2, DeviceType::GPU}, - MachineSpaceOffset{0, 4, DeviceType::GPU}}; + MachineSpaceOffset{0, 0}, + MachineSpaceOffset{0, 2}, + MachineSpaceOffset{0, 4}}; std::unordered_set result = get_machine_space_offsets(task, simv); CHECK(correct == result); @@ -176,8 +168,7 @@ TEST_SUITE(FF_TEST_SUITE) { {MachineViewDimension{stride_t{1_p}, MachineSpecificationDimension::INTER_NODE}, MachineViewDimension{stride_t{2_p}, - MachineSpecificationDimension::INTRA_NODE}}, - DeviceType::GPU}; + MachineSpecificationDimension::INTRA_NODE}}}; MachineComputeSpecification ms = MachineComputeSpecification{ /*num_nodes=*/2_p, /*num_cpus_per_node=*/4_p, @@ -188,7 +179,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 0_n}); MachineSpaceOffset correct = - MachineSpaceOffset{0, 0, DeviceType::GPU}; + MachineSpaceOffset{0, 0}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -197,7 +188,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n}); MachineSpaceOffset correct = - MachineSpaceOffset{0, 2, DeviceType::GPU}; + MachineSpaceOffset{0, 2}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -206,7 +197,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n}); MachineSpaceOffset correct = - MachineSpaceOffset{1, 0, DeviceType::GPU}; + MachineSpaceOffset{1, 0}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -215,7 +206,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n}); MachineSpaceOffset correct = - MachineSpaceOffset{1, 2, DeviceType::GPU}; + MachineSpaceOffset{1, 2}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -224,10 +215,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_machine_space_offsets") { std::unordered_set correct = { - MachineSpaceOffset{0, 0, DeviceType::GPU}, - MachineSpaceOffset{0, 2, DeviceType::GPU}, - MachineSpaceOffset{1, 0, DeviceType::GPU}, - MachineSpaceOffset{1, 2, DeviceType::GPU}}; + MachineSpaceOffset{0, 0}, + MachineSpaceOffset{0, 2}, + MachineSpaceOffset{1, 0}, + MachineSpaceOffset{1, 2}}; std::unordered_set result = get_machine_space_offsets(task, simv); CHECK(correct == result); diff --git a/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc b/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc index 4c03af5382..0fe4277ab6 100644 --- a/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc +++ b/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc @@ -65,8 +65,7 @@ TEST_SUITE(FF_TEST_SUITE) { MCMCOverMappedPCGConfig no_search = MCMCOverMappedPCGConfig{/*temperature=*/1.0, /*num_iterations=*/1_n, - /*substitution_frequency=*/0.2, - /*device_type=*/DeviceType::GPU}; + /*substitution_frequency=*/0.2}; SearchResult base_result = mcmc_over_mapped_pcg(pcg, cost_estimator, full_machine_spec, no_search); @@ -80,8 +79,7 @@ TEST_SUITE(FF_TEST_SUITE) { MCMCOverMappedPCGConfig search_config = MCMCOverMappedPCGConfig{/*temperature=*/1.0, /*num_iterations=*/100_n, - /*substitution_frequency=*/0.2, - /*device_type=*/DeviceType::GPU}; + /*substitution_frequency=*/0.2}; SearchResult result = mcmc_over_mapped_pcg( pcg, cost_estimator, full_machine_spec, search_config); diff --git a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc index 2846de6559..a7f8886846 100644 --- a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc @@ -13,8 +13,6 @@ #include "op-attrs/parallel_tensor_dims.dtg.h" #include "op-attrs/parallel_tensor_shape.dtg.h" #include "op-attrs/parallel_tensor_shape.h" -#include "pcg/device_id.h" -#include "pcg/device_type.dtg.h" #include "pcg/machine_space_coordinate.dtg.h" #include "pcg/machine_specification_dimension.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" @@ -68,9 +66,9 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector dims = {}; ParallelComputationGraph pcg = b.pcg; MachineView mv1 = - MachineView{MachineSpaceCoordinate{0_n, 0_n, DeviceType::GPU}, dims}; + MachineView{MachineSpaceCoordinate{0_n, 0_n}, dims}; MachineView mv2 = - MachineView{MachineSpaceCoordinate{0_n, 1_n, DeviceType::GPU}, dims}; + MachineView{MachineSpaceCoordinate{0_n, 1_n}, dims}; MachineMapping device_mapping = MachineMapping{{ {layer0, mv1}, @@ -149,13 +147,13 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("all different devices") { MachineView mv0 = MachineView{ - MachineSpaceCoordinate{0_n, 0_n, DeviceType::GPU}, dims}; + MachineSpaceCoordinate{0_n, 0_n}, dims}; MachineView mv1 = MachineView{ - MachineSpaceCoordinate{0_n, 1_n, DeviceType::GPU}, dims}; + MachineSpaceCoordinate{0_n, 1_n}, dims}; MachineView mv2 = MachineView{ - MachineSpaceCoordinate{1_n, 0_n, DeviceType::GPU}, dims}; + MachineSpaceCoordinate{1_n, 0_n}, dims}; MachineView mv3 = MachineView{ - MachineSpaceCoordinate{1_n, 1_n, DeviceType::GPU}, dims}; + MachineSpaceCoordinate{1_n, 1_n}, dims}; MachineMapping device_mapping = MachineMapping{{ {layer0, mv0}, @@ -209,7 +207,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("all the same device") { MachineView mv = MachineView{ - MachineSpaceCoordinate{0_n, 0_n, DeviceType::GPU}, dims}; + MachineSpaceCoordinate{0_n, 0_n}, dims}; MachineMapping device_mapping = MachineMapping{{ {layer0, mv}, {layer1, mv}, diff --git a/lib/local-execution/include/local-execution/computation_graph_instance.h b/lib/local-execution/include/local-execution/computation_graph_instance.h index a4ded5edaf..aa05cbd582 100644 --- a/lib/local-execution/include/local-execution/computation_graph_instance.h +++ b/lib/local-execution/include/local-execution/computation_graph_instance.h @@ -7,7 +7,7 @@ #include "kernels/profiling_settings.dtg.h" #include "local-execution/loss_config.dtg.h" #include "pcg/computation_graph.dtg.h" -#include "pcg/device_id_t.dtg.h" +#include "task-spec/device_id_t.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" diff --git a/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h b/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h index fdff1153ff..a978f3c996 100644 --- a/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h +++ b/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h @@ -5,7 +5,7 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" #include "kernels/profiling_settings.dtg.h" -#include "pcg/device_id_t.dtg.h" +#include "task-spec/device_id_t.dtg.h" #include "pcg/machine_interconnect_specification.dtg.h" namespace FlexFlow { diff --git a/lib/local-execution/include/local-execution/local_task_argument_accessor.h b/lib/local-execution/include/local-execution/local_task_argument_accessor.h index 12eab4a76d..a7df8814d0 100644 --- a/lib/local-execution/include/local-execution/local_task_argument_accessor.h +++ b/lib/local-execution/include/local-execution/local_task_argument_accessor.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_LOCAL_EXECUTION_INCLUDE_LOCAL_EXECUTION_LOCAL_TASK_ARGUMENT_ACCESSOR_H #include "kernels/accessor.h" -#include "pcg/device_id_t.dtg.h" +#include "task-spec/device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" #include "task-spec/task_argument_accessor/itask_argument_accessor.h" #include "task-spec/task_argument_accessor/task_tensor_parameter.dtg.h" diff --git a/lib/local-execution/include/local-execution/per_device_op_state_initialization.h b/lib/local-execution/include/local-execution/per_device_op_state_initialization.h index ff0ffd73f8..174e51f307 100644 --- a/lib/local-execution/include/local-execution/per_device_op_state_initialization.h +++ b/lib/local-execution/include/local-execution/per_device_op_state_initialization.h @@ -4,7 +4,7 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" #include "kernels/profiling_settings.dtg.h" -#include "pcg/device_id_t.dtg.h" +#include "task-spec/device_id_t.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" diff --git a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc index 0c6107c4af..27a1c39677 100644 --- a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -11,7 +11,6 @@ #include "op-attrs/tensor_slot_name.dtg.h" #include "pcg/computation_graph.h" #include "pcg/computation_graph/layer_added_result.dtg.h" -#include "pcg/device_id.h" #include "pcg/parallel_tensor_attrs.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/map_values.h" @@ -105,7 +104,7 @@ OpCostMetrics LocalCostEstimator::estimate_cost( // allocate memory std::shared_ptr tracked_allocator_ptr = std::make_shared(create_local_allocator_for_device_type( - get_device_type(this->device_idx))); + this->device_idx.device_type)); layer_guid_t layer_guid = layer_guid_t{Node{0}}; diff --git a/lib/local-execution/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/src/local-execution/local_task_argument_accessor.cc index b8feca720e..5fbe207e6c 100644 --- a/lib/local-execution/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/src/local-execution/local_task_argument_accessor.cc @@ -1,7 +1,5 @@ #include "local-execution/local_task_argument_accessor.h" #include "kernels/accessor.h" -#include "pcg/device_id.h" -#include "pcg/device_id_t.h" #include "utils/exception.h" #include "utils/optional.h" #include "utils/overload.h" @@ -84,7 +82,7 @@ device_handle_t LocalTaskArgumentAccessor::get_ff_handle() const { } DeviceType LocalTaskArgumentAccessor::get_kernel_device_type() const { - return get_device_type(this->device_idx); + return this->device_idx.device_type; } PCGOperatorAttrs LocalTaskArgumentAccessor::get_op_attrs() const { diff --git a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc index 2a4e204d59..a4049a609f 100644 --- a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc +++ b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc @@ -11,7 +11,6 @@ #include "op-attrs/ops/loss_functions/loss_attrs.dtg.h" #include "pcg/computation_graph.h" #include "pcg/computation_graph_builder.h" -#include "pcg/device_id_t.h" #include "pcg/device_type.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_guid_t.dtg.h" @@ -140,8 +139,14 @@ TEST_SUITE(FF_TEST_SUITE) { /*nesterov=*/false, /*weight_decay=*/0.001}}; device_handle_t ff_handle = cpu_make_device_handle_t(); - device_id_t device_idx = - make_device_id_t_from_idx(nonnegative_int{0}, DeviceType::CPU); + device_id_t device_idx = device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::CPU, + }; + std::unordered_map input_tensors; @@ -307,8 +312,13 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { /*weight_decay=*/0.001, }, }; - device_id_t device_idx = - make_device_id_t_from_idx(nonnegative_int{0}, DeviceType::GPU); + device_id_t device_idx = device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, + }; device_handle_t ff_handle = gpu_make_device_handle_t(managed_handle.raw_handle()); @@ -423,8 +433,14 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { }, }; - device_id_t device_idx = - make_device_id_t_from_idx(nonnegative_int{0}, DeviceType::GPU); + device_id_t device_idx = device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, + }; + device_handle_t ff_handle = gpu_make_device_handle_t(managed_handle.raw_handle()); diff --git a/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc index ac20c1ad75..34cfbb4e3d 100644 --- a/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -10,7 +10,6 @@ #include "op-attrs/parallel_tensor_shape.h" #include "op-attrs/tensor_slot_name.dtg.h" #include "pcg/computation_graph_builder.h" -#include "pcg/device_id_t.h" #include using namespace ::FlexFlow; @@ -19,8 +18,13 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("LocalCostEstimator") { Allocator allocator = create_local_cpu_memory_allocator(); device_handle_t ff_handle = cpu_make_device_handle_t(); - device_id_t device_idx = - make_device_id_t_from_idx(nonnegative_int{0}, DeviceType::CPU); + device_id_t device_idx = device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::CPU, + }; OptimizerAttrs optimizer_attrs = OptimizerAttrs{ SGDOptimizerAttrs{ @@ -66,7 +70,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*optimizer_attrs=*/optimizer_attrs, /*machine_view=*/ make_1d_machine_view( - MachineSpaceCoordinate{0_n, 0_n, DeviceType::CPU}, + MachineSpaceCoordinate{0_n, 0_n}, MachineSpecificationDimension::INTRA_NODE, stride_t{1_p}), }; @@ -89,8 +93,14 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { Allocator allocator = create_local_cuda_memory_allocator(); - device_id_t device_idx = - make_device_id_t_from_idx(nonnegative_int{0}, DeviceType::GPU); + device_id_t device_idx = device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, + }; + device_handle_t ff_handle = gpu_make_device_handle_t(managed_handle.raw_handle()); @@ -159,7 +169,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { /*optimizer_attrs=*/optimizer_attrs, /*machine_view=*/ make_1d_machine_view( - MachineSpaceCoordinate{0_n, 0_n, DeviceType::GPU}, + MachineSpaceCoordinate{0_n, 0_n}, MachineSpecificationDimension::INTRA_NODE, stride_t{1_p}), }; diff --git a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc index 07bb869d5f..9594e13336 100644 --- a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc @@ -3,7 +3,6 @@ #include "kernels/local_cpu_allocator.h" #include "kernels/profiling_settings.dtg.h" #include "op-attrs/ops/input_attrs.dtg.h" -#include "pcg/device_id_t.h" #include "task-spec/task_argument_accessor/task_tensor_parameter.h" #include "task-spec/task_impl_function.dtg.h" #include "utils/fmt/variant.h" @@ -52,8 +51,13 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - device_id_t device_idx = - make_device_id_t_from_idx(nonnegative_int{0}, DeviceType::CPU); + device_id_t device_idx = device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::CPU, + }; LocalTaskArgumentAccessor acc = LocalTaskArgumentAccessor{ /*allocator=*/allocator, diff --git a/lib/pcg/include/pcg/device_id.h b/lib/pcg/include/pcg/device_id.h deleted file mode 100644 index 36ea9de6b3..0000000000 --- a/lib/pcg/include/pcg/device_id.h +++ /dev/null @@ -1,22 +0,0 @@ -#ifndef _FLEXFLOW_PCG_INCLUDE_PCG_DEVICE_ID_H -#define _FLEXFLOW_PCG_INCLUDE_PCG_DEVICE_ID_H - -#include "pcg/cpu_id_t.dtg.h" -#include "pcg/device_id_t.dtg.h" -#include "pcg/device_type.dtg.h" -#include "pcg/gpu_id_t.dtg.h" - -namespace FlexFlow { - -device_id_t operator+(device_id_t, size_t); - -DeviceType get_device_type(device_id_t const &device_id); -gpu_id_t unwrap_gpu(device_id_t); -cpu_id_t unwrap_cpu(device_id_t); -nonnegative_int get_raw_id(device_id_t); - -device_id_t device_id_from_index(nonnegative_int, DeviceType); - -} // namespace FlexFlow - -#endif diff --git a/lib/pcg/include/pcg/device_id_t.dtg.toml b/lib/pcg/include/pcg/device_id_t.dtg.toml deleted file mode 100644 index 4efcb07975..0000000000 --- a/lib/pcg/include/pcg/device_id_t.dtg.toml +++ /dev/null @@ -1,23 +0,0 @@ -namespace = "FlexFlow" -name = "device_id_t" -type = "variant" -features = [ - "eq", - "ord", - "hash", - "json", - "fmt", -] - -includes = [ - "pcg/cpu_id_t.dtg.h", - "pcg/gpu_id_t.dtg.h", -] - -[[values]] -type = "::FlexFlow::gpu_id_t" -key = "gpu" - -[[values]] -type = "::FlexFlow::cpu_id_t" -key = "cpu" diff --git a/lib/pcg/include/pcg/device_id_t.h b/lib/pcg/include/pcg/device_id_t.h deleted file mode 100644 index e8e605b068..0000000000 --- a/lib/pcg/include/pcg/device_id_t.h +++ /dev/null @@ -1,14 +0,0 @@ -#ifndef _FLEXFLOW_LIB_PCG_INCLUDE_PCG_DEVICE_ID_T_H -#define _FLEXFLOW_LIB_PCG_INCLUDE_PCG_DEVICE_ID_T_H - -#include "pcg/device_id_t.dtg.h" -#include "pcg/device_type.dtg.h" - -namespace FlexFlow { - -device_id_t make_device_id_t_from_idx(nonnegative_int idx, - DeviceType device_type); - -} // namespace FlexFlow - -#endif diff --git a/lib/pcg/include/pcg/machine_compute_resource_slice.dtg.toml b/lib/pcg/include/pcg/machine_compute_resource_slice.dtg.toml index 77d6b7558c..41693ca0dc 100644 --- a/lib/pcg/include/pcg/machine_compute_resource_slice.dtg.toml +++ b/lib/pcg/include/pcg/machine_compute_resource_slice.dtg.toml @@ -6,6 +6,8 @@ features = [ "ord", "hash", "fmt", + "rapidcheck", + "json", ] includes = [ diff --git a/lib/pcg/include/pcg/machine_compute_specification.h b/lib/pcg/include/pcg/machine_compute_specification.h index 835e9040e0..2e8f910c08 100644 --- a/lib/pcg/include/pcg/machine_compute_specification.h +++ b/lib/pcg/include/pcg/machine_compute_specification.h @@ -1,7 +1,6 @@ #ifndef _FLEXFLOW_LIB_PCG_INCLUDE_PCG_MACHINE_COMPUTE_SPECIFICATION_H #define _FLEXFLOW_LIB_PCG_INCLUDE_PCG_MACHINE_COMPUTE_SPECIFICATION_H -#include "pcg/device_id_t.dtg.h" #include "pcg/device_type.dtg.h" #include "pcg/machine_compute_specification.dtg.h" #include "pcg/machine_space_coordinate.dtg.h" @@ -15,12 +14,6 @@ positive_int get_num_devices(MachineComputeSpecification const &ms, positive_int get_num_devices_per_node(MachineComputeSpecification const &ms, DeviceType const &device_type); -bool is_valid_machine_space_coordinate(MachineComputeSpecification const &ms, - MachineSpaceCoordinate const &coord); - -device_id_t get_device_id(MachineComputeSpecification const &ms, - MachineSpaceCoordinate const &coord); - } // namespace FlexFlow #endif diff --git a/lib/pcg/include/pcg/machine_space_coordinate.dtg.toml b/lib/pcg/include/pcg/machine_space_coordinate.dtg.toml index 41f4d563f3..8787b0d8f2 100644 --- a/lib/pcg/include/pcg/machine_space_coordinate.dtg.toml +++ b/lib/pcg/include/pcg/machine_space_coordinate.dtg.toml @@ -10,8 +10,7 @@ features = [ "fmt", ] -includes = [ - "pcg/device_type.dtg.h", +includes = [ "utils/nonnegative_int/nonnegative_int.h", ] @@ -22,7 +21,3 @@ type = "::FlexFlow::nonnegative_int" [[fields]] name = "device_idx" type = "::FlexFlow::nonnegative_int" - -[[fields]] -name = "device_type" -type = "::FlexFlow::DeviceType" diff --git a/lib/pcg/include/pcg/machine_space_offset.dtg.toml b/lib/pcg/include/pcg/machine_space_offset.dtg.toml index 57f884906b..54eda5cc76 100644 --- a/lib/pcg/include/pcg/machine_space_offset.dtg.toml +++ b/lib/pcg/include/pcg/machine_space_offset.dtg.toml @@ -10,9 +10,7 @@ features = [ "fmt", ] -includes = [ - "pcg/device_type.dtg.h", -] +includes = [] [[fields]] name = "node_offset" @@ -21,7 +19,3 @@ type = "int" [[fields]] name = "device_offset" type = "int" - -[[fields]] -name = "device_type" -type = "::FlexFlow::DeviceType" diff --git a/lib/pcg/src/pcg/device_id.cc b/lib/pcg/src/pcg/device_id.cc deleted file mode 100644 index 1a4f7b7d22..0000000000 --- a/lib/pcg/src/pcg/device_id.cc +++ /dev/null @@ -1,51 +0,0 @@ -#include "pcg/device_id.h" -#include "utils/exception.h" -#include - -namespace FlexFlow { - -device_id_t operator+(device_id_t, size_t) { - NOT_IMPLEMENTED(); -} - -DeviceType get_device_type(device_id_t const &device_id) { - if (device_id.has()) { - return DeviceType::GPU; - } else { - assert(device_id.has()); - return DeviceType::CPU; - } -} - -gpu_id_t unwrap_gpu(device_id_t device_id) { - return device_id.get(); -} - -cpu_id_t unwrap_cpu(device_id_t device_id) { - return device_id.get(); -} - -nonnegative_int get_raw_id(device_id_t device_id) { - switch (get_device_type(device_id)) { - case DeviceType::GPU: - return unwrap_gpu(device_id).gpu_index; - case DeviceType::CPU: - return unwrap_cpu(device_id).cpu_index; - default: - throw mk_runtime_error(fmt::format("Unsupported device {}", device_id)); - } -} - -device_id_t device_id_from_index(nonnegative_int idx, DeviceType device_type) { - switch (device_type) { - case DeviceType::GPU: - return device_id_t{gpu_id_t{idx}}; - case DeviceType::CPU: - return device_id_t{cpu_id_t{idx}}; - default: - throw mk_runtime_error( - fmt::format("Unsupported DeviceType {}", device_type)); - } -} - -} // namespace FlexFlow diff --git a/lib/pcg/src/pcg/device_id_t.cc b/lib/pcg/src/pcg/device_id_t.cc deleted file mode 100644 index eecaf1c81d..0000000000 --- a/lib/pcg/src/pcg/device_id_t.cc +++ /dev/null @@ -1,17 +0,0 @@ -#include "pcg/device_id_t.h" - -namespace FlexFlow { - -device_id_t make_device_id_t_from_idx(nonnegative_int idx, - DeviceType device_type) { - switch (device_type) { - case DeviceType::GPU: - return device_id_t{gpu_id_t{idx}}; - case DeviceType::CPU: - return device_id_t{cpu_id_t{idx}}; - default: - PANIC("Unhandled device_type", device_type); - } -} - -} // namespace FlexFlow diff --git a/lib/pcg/src/pcg/machine_compute_resource_slice.cc b/lib/pcg/src/pcg/machine_compute_resource_slice.cc index 2cb2386926..401ed3d850 100644 --- a/lib/pcg/src/pcg/machine_compute_resource_slice.cc +++ b/lib/pcg/src/pcg/machine_compute_resource_slice.cc @@ -20,7 +20,6 @@ positive_int bool is_valid_machine_space_coordinate_in_slice( MachineComputeResourceSlice const &slice, MachineSpaceCoordinate const &coord) { - ASSERT(coord.device_type == DeviceType::GPU); return (coord.node_idx < slice.num_nodes) && (coord.device_idx < slice.num_gpus_per_node); diff --git a/lib/pcg/src/pcg/machine_compute_specification.cc b/lib/pcg/src/pcg/machine_compute_specification.cc index a8dfb27524..0dbe546fff 100644 --- a/lib/pcg/src/pcg/machine_compute_specification.cc +++ b/lib/pcg/src/pcg/machine_compute_specification.cc @@ -1,5 +1,4 @@ #include "pcg/machine_compute_specification.h" -#include "pcg/device_id.h" #include "utils/containers/transform.h" #include @@ -37,20 +36,4 @@ positive_int get_num_devices_per_node(MachineComputeSpecification const &ms, } } -bool is_valid_machine_space_coordinate(MachineComputeSpecification const &ms, - MachineSpaceCoordinate const &coord) { - return (coord.node_idx < ms.num_nodes) && - (coord.device_idx < get_num_devices_per_node(ms, coord.device_type)); -} - -device_id_t get_device_id(MachineComputeSpecification const &ms, - MachineSpaceCoordinate const &coord) { - ASSERT(is_valid_machine_space_coordinate(ms, coord)); - - nonnegative_int raw_idx = - coord.node_idx * get_num_devices_per_node(ms, coord.device_type) + - coord.device_idx; - return device_id_from_index(raw_idx, coord.device_type); -} - } // namespace FlexFlow diff --git a/lib/pcg/src/pcg/machine_space_offset.cc b/lib/pcg/src/pcg/machine_space_offset.cc index 953dc38bc6..0240d3582d 100644 --- a/lib/pcg/src/pcg/machine_space_offset.cc +++ b/lib/pcg/src/pcg/machine_space_offset.cc @@ -13,14 +13,11 @@ MachineSpaceOffset get_machine_space_offset_from_coordinate( "The start node_idx is greater than one of the coord node_idx." "Are you sure you didn't swap them?"); - ASSERT(start.device_type == coord.device_type); - return MachineSpaceOffset{ /*node_offset=*/coord.node_idx.unwrap_nonnegative() - start.node_idx.unwrap_nonnegative(), /*device_offset=*/coord.device_idx.unwrap_nonnegative() - start.device_idx.unwrap_nonnegative(), - /*device_type=*/coord.device_type, }; } diff --git a/lib/pcg/test/src/pcg/machine_compute_specification.cc b/lib/pcg/test/src/pcg/machine_compute_specification.cc index c725da80ed..8a20d8d10d 100644 --- a/lib/pcg/test/src/pcg/machine_compute_specification.cc +++ b/lib/pcg/test/src/pcg/machine_compute_specification.cc @@ -1,5 +1,4 @@ #include "pcg/machine_compute_specification.h" -#include "pcg/device_id.h" #include using namespace FlexFlow; @@ -25,27 +24,5 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(get_num_devices(ms, DeviceType::GPU) == 4 * 8); CHECK(get_num_devices(ms, DeviceType::CPU) == 16 * 4); } - - SUBCASE("get_device_id") { - SUBCASE("valid MachineSpaceCoordinate") { - MachineSpaceCoordinate coord = MachineSpaceCoordinate{ - /*node_idx=*/2_n, - /*device_idx=*/12_n, - DeviceType::CPU, - }; - device_id_t correct = - device_id_from_index(nonnegative_int{2 * 16 + 12}, DeviceType::CPU); - device_id_t result = get_device_id(ms, coord); - CHECK(correct == result); - } - SUBCASE("MachineSpaceCoordinate out of bounds for given machine spec") { - MachineSpaceCoordinate coord = MachineSpaceCoordinate{ - /*node_idx=*/2_n, - /*device_idx=*/18_n, - DeviceType::CPU, - }; - CHECK_THROWS(get_device_id(ms, coord)); - } - } } } diff --git a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc index 1c3667afc7..3543ebe09f 100644 --- a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc +++ b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc @@ -13,14 +13,21 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("adl_serializer") { bidict shard_bindings{ - {MachineSpaceCoordinate{0_n, 0_n, DeviceType::CPU}, - OperatorAtomicTaskShardBinding{ - { - {TensorSlotName::INPUT, - ParallelTensorSpaceCoordinate{ - 0_n, 0_n, FFOrdered{1_n, 2_n, 3_n}}}, - }, - }}, + { + MachineSpaceCoordinate{0_n, 0_n}, + OperatorAtomicTaskShardBinding{ + { + { + TensorSlotName::INPUT, + ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_components=*/FFOrdered{1_n, 2_n, 3_n}, + }, + }, + }, + }, + }, }; MappedOperatorTaskGroup deserialized{shard_bindings}; nlohmann::json serialized = shard_bindings; diff --git a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 7856d89f27..5506afc63e 100644 --- a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -53,7 +53,6 @@ TEST_SUITE(FF_TEST_SUITE) { return MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/x, - /*device_type=*/DeviceType::GPU, }; }; diff --git a/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h b/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h index b698e613b5..e5ab8b597e 100644 --- a/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h +++ b/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h @@ -3,8 +3,8 @@ #include "kernels/device_handle_t.dtg.h" #include "kernels/managed_per_device_ff_handle.h" -#include "pcg/device_id_t.dtg.h" #include "realm-execution/device_specific_ptr.h" +#include "task-spec/device_id_t.dtg.h" #include namespace FlexFlow { diff --git a/lib/realm-execution/include/realm-execution/device_specific_ptr.h b/lib/realm-execution/include/realm-execution/device_specific_ptr.h index 32f8c13272..e8de4cfdad 100644 --- a/lib/realm-execution/include/realm-execution/device_specific_ptr.h +++ b/lib/realm-execution/include/realm-execution/device_specific_ptr.h @@ -1,7 +1,8 @@ #ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_DEVICE_SPECIFIC_PTR_H #define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_DEVICE_SPECIFIC_PTR_H -#include "pcg/device_id_t.dtg.h" +#include "task-spec/device_id_t.dtg.h" +#include #include namespace FlexFlow { diff --git a/lib/realm-execution/include/realm-execution/distributed_ff_handle.h b/lib/realm-execution/include/realm-execution/distributed_ff_handle.h index 8409a234a7..b2b61ee347 100644 --- a/lib/realm-execution/include/realm-execution/distributed_ff_handle.h +++ b/lib/realm-execution/include/realm-execution/distributed_ff_handle.h @@ -9,7 +9,7 @@ namespace FlexFlow { /** - * \brief Tracks the \ref device_handle_t (i.e., FFHandle) for each %GPU, both + * \brief Tracks the \ref device_handle_t (i.e., FFHandle) for each GPU, both * local and remote. A GPU here is represented by a Realm::Processor. */ struct DistributedFfHandle { @@ -31,8 +31,8 @@ struct DistributedFfHandle { /** * \brief Launches tasks (using \ref spawn_ff_handle_init_task) to create - * the \ref device_handle_t ""s for each %GPU and packages the results into a - * DistributedFfHandle. + * the \ref device_handle_t ""s for each GPU and packages the results into a + * \ref DistributedFfHandle. * * \relates DistributedFfHandle */ diff --git a/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h b/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h index 5d52f8caaf..f160146b96 100644 --- a/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h +++ b/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h @@ -25,7 +25,8 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( ProfilingSettings const &profiling_settings, DistributedFfHandle const &device_handle, OptimizerAttrs const &optimizer_attrs, - Realm::Event precondition); + Realm::Event precondition, + DeviceType device_type); } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/instance_allocation.h b/lib/realm-execution/include/realm-execution/instance_allocation.h index 66cc07af75..1f2a6d9134 100644 --- a/lib/realm-execution/include/realm-execution/instance_allocation.h +++ b/lib/realm-execution/include/realm-execution/instance_allocation.h @@ -12,10 +12,9 @@ namespace FlexFlow { * on the device represented by \p device_coord. */ std::pair - perform_instance_allocation_for_value( - MachineSpaceCoordinate const &device_coord, - DynamicValueAttrs const &value, - RealmContext &ctx); + perform_instance_allocation_for_value(device_id_t const &device_id, + DynamicValueAttrs const &value, + RealmContext &ctx); /** * @brief Allocates the (potentially remote) Realm instances for all of the @@ -28,7 +27,8 @@ TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, std::unordered_map const &preallocated, - RealmContext &ctx); + RealmContext &ctx, + DeviceType device_type); /** * @brief Destroys all of the instances held in \p instances. diff --git a/lib/realm-execution/include/realm-execution/pcg_instance.h b/lib/realm-execution/include/realm-execution/pcg_instance.h index 7b86d6d383..1e17856999 100644 --- a/lib/realm-execution/include/realm-execution/pcg_instance.h +++ b/lib/realm-execution/include/realm-execution/pcg_instance.h @@ -4,7 +4,6 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" #include "kernels/profiling_settings.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "realm-execution/distributed_ff_handle.h" @@ -12,6 +11,7 @@ #include "realm-execution/per_device_op_state_backing.dtg.h" #include "realm-execution/realm_context.h" #include "realm-execution/tensor_instance_backing.dtg.h" +#include "task-spec/device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" @@ -86,7 +86,8 @@ PCGInstance create_pcg_instance( std::unordered_map const &input_tensors, ProfilingSettings const &profiling_settings, - DistributedFfHandle const &ff_handle); + DistributedFfHandle const &ff_handle, + DeviceType device_type); /** * \brief Dispatch a training iteration for a \ref PCGInstance. diff --git a/lib/realm-execution/include/realm-execution/processor_kind.h b/lib/realm-execution/include/realm-execution/processor_kind.h new file mode 100644 index 0000000000..01bf6f8c35 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/processor_kind.h @@ -0,0 +1,14 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_PROCESSOR_KIND_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_PROCESSOR_KIND_H + +#include "pcg/device_type.dtg.h" +#include "realm-execution/realm.h" + +namespace FlexFlow { + +DeviceType device_type_from_processor_kind(Realm::Processor::Kind); +Realm::Processor::Kind processor_kind_from_device_type(DeviceType); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/realm_context.h b/lib/realm-execution/include/realm-execution/realm_context.h index ab89e916c0..0f2bed7de4 100644 --- a/lib/realm-execution/include/realm-execution/realm_context.h +++ b/lib/realm-execution/include/realm-execution/realm_context.h @@ -6,10 +6,10 @@ #include "kernels/managed_per_device_ff_handle.h" #include "op-attrs/parallel_tensor_shape.dtg.h" #include "op-attrs/tensor_shape.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/machine_space_coordinate.dtg.h" #include "realm-execution/realm.h" #include "realm-execution/tasks/task_id_t.dtg.h" +#include "task-spec/device_id_t.dtg.h" #include #include @@ -32,8 +32,8 @@ struct RealmContext { /** \name Device mapping */ ///\{ - Realm::Processor - map_device_coord_to_processor(MachineSpaceCoordinate const &); + Realm::Processor map_device_coord_to_processor(device_id_t const &) const; + device_id_t map_processor_to_device_coord(Realm::Processor) const; static Realm::Memory get_nearest_memory(Realm::Processor); ///\} @@ -97,8 +97,6 @@ struct RealmContext { */ [[nodiscard]] Realm::Event merge_outstanding_events(); - void discover_machine_topology(); - static std::optional make_device_handle_for_processor(Realm::Processor processor); @@ -111,14 +109,15 @@ struct RealmContext { */ Realm::Runtime get_runtime(); -private: + void discover_machine_topology(); + +public: Realm::Runtime runtime; Realm::Processor processor; Allocator allocator; std::vector outstanding_events; - std::unordered_map, - std::vector> - processors; + std::optional> processors = + std::nullopt; }; } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/realm_manager.h b/lib/realm-execution/include/realm-execution/realm_manager.h index 3984291641..ed01f4a022 100644 --- a/lib/realm-execution/include/realm-execution/realm_manager.h +++ b/lib/realm-execution/include/realm-execution/realm_manager.h @@ -3,10 +3,10 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "realm-execution/realm.h" #include "realm-execution/realm_context.h" #include "realm-execution/tasks/impl/controller_task.h" +#include "task-spec/device_id_t.dtg.h" namespace FlexFlow { diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index 8cf4c21b25..015af2b1ae 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -22,7 +22,8 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( ProfilingSettings const &profiling_settings, DistributedFfHandle const &device_handle, OptimizerAttrs const &optimizer_attrs, - Realm::Event precondition) { + Realm::Event precondition, + DeviceType device_type) { // Initialize all operators and save the per-device op state ASSERT(no_nodes_are_initialized(dg)); diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index 4ef2919b10..37bf0ca03d 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -22,15 +22,14 @@ namespace FlexFlow { std::pair - perform_instance_allocation_for_value( - MachineSpaceCoordinate const &device_coord, - DynamicValueAttrs const &value, - RealmContext &ctx) { + perform_instance_allocation_for_value(device_id_t const &device_id, + DynamicValueAttrs const &value, + RealmContext &ctx) { ASSERT(value.accessor == std::nullopt); TensorShape shape = get_piece_shape(value.parallel_tensor_shape.value()); - Realm::Processor proc = ctx.map_device_coord_to_processor(device_coord); + Realm::Processor proc = ctx.map_device_coord_to_processor(device_id); Realm::Memory memory = ctx.get_nearest_memory(proc); return ctx.create_instance(memory, shape, Realm::ProfilingRequestSet()); } @@ -39,7 +38,8 @@ TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, std::unordered_map const &preallocated, - RealmContext &ctx) { + RealmContext &ctx, + DeviceType device_type) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); for (DynamicValueAttrs const &v : keys(preallocated)) { @@ -53,9 +53,9 @@ TensorInstanceBacking perform_instance_allocation( NOT_IMPLEMENTED(); } else { if (!contains_key(result.backing, v)) { - MachineSpaceCoordinate device_coord = assert_unwrap(n.device_coord); + device_id_t device = assert_unwrap(n.device_coord); result.backing.insert(std::pair{ - v, perform_instance_allocation_for_value(device_coord, v, ctx)}); + v, perform_instance_allocation_for_value(device, v, ctx)}); } return result.backing.at(v); } diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 1ef0c4270b..62747fda7a 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -85,20 +85,31 @@ PCGInstance create_pcg_instance( std::unordered_map const &input_tensors, ProfilingSettings const &profiling_settings, - DistributedFfHandle const &device_handle) { + DistributedFfHandle const &device_handle, + DeviceType device_type) { DynamicOpenDataflowGraph dg = - make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg); + make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg, device_type); dg = perform_pass_expansion(dg); std::unordered_map inputs = input_tensors; std::optional logit_grad_value; if (loss.has_value()) { - auto [loss_attrs, label_tensor, logit_tensor, loss_mapping] = - assert_unwrap(loss); + ParallelLossConfig loss_config = assert_unwrap(loss); + + LossAttrs loss_attrs = loss_config.loss_attrs; + GenericTensorAccessorR label_tensor = loss_config.label_tensor; + parallel_tensor_guid_t logit_tensor = loss_config.logit_tensor; + MappedOperatorTaskGroup loss_op_task_group = loss_config.loss_mapping; + + DynamicNodeMapping mapping = DynamicNodeMapping{ + /*op_task_group=*/loss_op_task_group, + /*device_type=*/device_type, + }; + auto [dg2, label_v, logit_grad_v] = perform_loss_insertion( - dg, loss_attrs, dynamic_tensor_guid_t{logit_tensor}, loss_mapping); + dg, loss_attrs, dynamic_tensor_guid_t{logit_tensor}, mapping); dg = dg2; logit_grad_value = logit_grad_v; inputs.insert(std::pair{label_v, label_tensor}); @@ -106,9 +117,11 @@ PCGInstance create_pcg_instance( dg = perform_update_insertion(dg, optimizer_attrs); dg = perform_copy_insertion(dg); + debug_print_dynamic_open_dataflow_graph_as_dot(dg); dg = perform_shard_expansion(dg); + TensorInstanceBacking tensor_instance_backing = - perform_instance_allocation(dg, inputs, ctx); + perform_instance_allocation(dg, inputs, ctx, device_type); logit_grad_value = transform(logit_grad_value, [&](DynamicValueAttrs const &lgv) { @@ -141,7 +154,8 @@ PCGInstance create_pcg_instance( profiling_settings, device_handle, optimizer_attrs, - ctx.get_outstanding_events()); + ctx.get_outstanding_events(), + device_type); // Compute the topological ordering of the graph auto [kwarg_graph, node_map] = @@ -150,12 +164,14 @@ PCGInstance create_pcg_instance( std::vector invocation_topo_order = transform( node_topo_order, [&](Node node) { return node_map.at_l(node); }); - return PCGInstance{/*ctx=*/ctx, - /*execution_order=*/invocation_topo_order, - /*tensor_instance_backing=*/tensor_instance_backing, - /*device_state_backing=*/device_state_backing, - /*optimizer_attrs=*/optimizer_attrs, - /*logit_grad_tensor=*/logit_grad_tensor}; + return PCGInstance{ + /*ctx=*/ctx, + /*execution_order=*/invocation_topo_order, + /*tensor_instance_backing=*/tensor_instance_backing, + /*device_state_backing=*/device_state_backing, + /*optimizer_attrs=*/optimizer_attrs, + /*logit_grad_tensor=*/logit_grad_tensor, + }; } /** diff --git a/lib/realm-execution/src/realm-execution/processor_kind.cc b/lib/realm-execution/src/realm-execution/processor_kind.cc new file mode 100644 index 0000000000..3a40de0a42 --- /dev/null +++ b/lib/realm-execution/src/realm-execution/processor_kind.cc @@ -0,0 +1,29 @@ +#include "realm-execution/processor_kind.h" +#include + +namespace FlexFlow { + +DeviceType + device_type_from_processor_kind(Realm::Processor::Kind processor_kind) { + switch (processor_kind) { + case Realm::Processor::Kind::LOC_PROC: + return DeviceType::CPU; + case Realm::Processor::Kind::TOC_PROC: + return DeviceType::GPU; + default: + PANIC("Unhandled Realm::Processor::Kind", processor_kind); + } +} + +Realm::Processor::Kind processor_kind_from_device_type(DeviceType device_type) { + switch (device_type) { + case DeviceType::CPU: + return Realm::Processor::Kind::LOC_PROC; + case DeviceType::GPU: + return Realm::Processor::Kind::TOC_PROC; + default: + PANIC("Unhandled DeviceType", device_type); + } +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/realm_allocator.cc b/lib/realm-execution/src/realm-execution/realm_allocator.cc index 9d5f3ff6b4..842539f9c2 100644 --- a/lib/realm-execution/src/realm-execution/realm_allocator.cc +++ b/lib/realm-execution/src/realm-execution/realm_allocator.cc @@ -1,6 +1,7 @@ #include "realm-execution/realm_allocator.h" #include "kernels/device.h" #include "pcg/device_type.dtg.h" +#include "realm-execution/processor_kind.h" #include "utils/containers/contains_key.h" #include "utils/containers/values.h" @@ -45,14 +46,7 @@ void RealmAllocator::deallocate(void *ptr) { } DeviceType RealmAllocator::get_allocation_device_type() const { - switch (this->processor.kind()) { - case Realm::Processor::Kind::LOC_PROC: - return DeviceType::CPU; - case Realm::Processor::Kind::TOC_PROC: - return DeviceType::GPU; - default: - PANIC("Unhandled FwbTensorType", this->processor.kind()); - } + return device_type_from_processor_kind(this->processor.kind()); } Allocator get_realm_allocator(Realm::Processor processor, diff --git a/lib/realm-execution/src/realm-execution/realm_context.cc b/lib/realm-execution/src/realm-execution/realm_context.cc index 790c1bd613..4cd7bdc221 100644 --- a/lib/realm-execution/src/realm-execution/realm_context.cc +++ b/lib/realm-execution/src/realm-execution/realm_context.cc @@ -4,8 +4,7 @@ #include "op-attrs/datatype.h" #include "op-attrs/parallel_tensor_shape.h" #include "op-attrs/tensor_dims.dtg.h" -#include "pcg/device_id_t.h" -#include "pcg/device_type.dtg.h" +#include "realm-execution/processor_kind.h" #include "realm-execution/realm_allocator.h" #include "realm-execution/tasks/task_id_t.dtg.h" #include "realm-execution/tasks/task_id_t.h" @@ -14,6 +13,7 @@ #include "utils/exception.h" #include "utils/nonnegative_int/nonnegative_int.h" #include "utils/one_to_many/one_to_many.h" +#include "utils/optional.h" #include "utils/positive_int/positive_int.h" namespace FlexFlow { @@ -21,7 +21,11 @@ namespace FlexFlow { RealmContext::RealmContext(Realm::Processor processor) : processor(processor), allocator(get_realm_allocator( - processor, RealmContext::get_nearest_memory(processor))) {} + processor, RealmContext::get_nearest_memory(processor))) { + if (processor != Realm::Processor::NO_PROC) { + this->discover_machine_topology(); + } +} RealmContext::~RealmContext() { if (!this->outstanding_events.empty()) { @@ -31,31 +35,22 @@ RealmContext::~RealmContext() { } static std::tuple - convert_machine_space_coordinate( - MachineSpaceCoordinate const &device_coord) { + convert_machine_space_coordinate(MachineSpaceCoordinate const &device_coord, + DeviceType device_type) { Realm::AddressSpace as = int{device_coord.node_idx}; - Realm::Processor::Kind kind; - switch (device_coord.device_type) { - case DeviceType::CPU: - kind = Realm::Processor::Kind::LOC_PROC; - break; - case DeviceType::GPU: - kind = Realm::Processor::Kind::TOC_PROC; - break; - default: - PANIC("Unhandled DeviceType", fmt::to_string(device_coord.device_type)); - break; - } + Realm::Processor::Kind kind = processor_kind_from_device_type(device_type); nonnegative_int proc_in_node = device_coord.device_idx; return std::tuple{as, kind, proc_in_node}; } Realm::Processor RealmContext::map_device_coord_to_processor( - MachineSpaceCoordinate const &device_coord) { - this->discover_machine_topology(); - auto [as, kind, proc_in_node] = - convert_machine_space_coordinate(device_coord); - return this->processors.at(std::pair{as, kind}).at(int{proc_in_node}); + device_id_t const &device_id) const { + return assert_unwrap(this->processors).at_r(device_id); +} + +device_id_t + RealmContext::map_processor_to_device_coord(Realm::Processor p) const { + return assert_unwrap(this->processors).at_l(p); } Realm::Memory RealmContext::get_nearest_memory(Realm::Processor proc) { @@ -81,27 +76,7 @@ Allocator &RealmContext::get_current_device_allocator() { device_id_t RealmContext::get_current_device_idx() const { Realm::Processor proc = this->get_current_processor(); - - // FIXME: find a more efficient way to implement this than scanning the - // machine every time - Realm::Machine::ProcessorQuery pq(Realm::Machine::get_machine()); - pq.same_address_space_as(proc); - nonnegative_int idx{0}; - for (Realm::Processor p : pq) { - if (p == proc) { - break; - } - idx++; - } - - switch (proc.kind()) { - case Realm::Processor::LOC_PROC: - return make_device_id_t_from_idx(idx, DeviceType::CPU); - case Realm::Processor::TOC_PROC: - return make_device_id_t_from_idx(idx, DeviceType::GPU); - default: - PANIC("Unhandled Realm::ProcessorKind", fmt::to_string(int{proc.kind()})); - } + return this->map_processor_to_device_coord(proc); } Realm::Event @@ -316,16 +291,49 @@ Realm::Event RealmContext::merge_outstanding_events() { } void RealmContext::discover_machine_topology() { - if (!this->processors.empty()) { + if (this->processors.has_value()) { return; } + std::unordered_map, nonnegative_int> + next_device_idx; + + auto fresh_device_id = [&](nonnegative_int node_idx, + DeviceType device_type) -> device_id_t { + std::pair key = + std::pair{node_idx, device_type}; + if (!contains_key(next_device_idx, key)) { + next_device_idx.insert({key, 0_n}); + } + + nonnegative_int device_idx = next_device_idx.at(key); + next_device_idx.at(key)++; + + return device_id_t{ + MachineSpaceCoordinate{node_idx, device_idx}, + device_type, + }; + }; + + bidict procs; Realm::Machine::ProcessorQuery pq(Realm::Machine::get_machine()); for (Realm::Processor proc : pq) { Realm::AddressSpace as = proc.address_space(); Realm::Processor::Kind kind = proc.kind(); - this->processors[std::pair{as, kind}].push_back(proc); + + nonnegative_int node_idx = nonnegative_int{static_cast(as)}; + + if (kind != Realm::Processor::LOC_PROC && + kind != Realm::Processor::TOC_PROC) { + continue; + } + + DeviceType device_type = device_type_from_processor_kind(kind); + device_id_t coord = fresh_device_id(node_idx, device_type); + procs.equate_strict(proc, coord); } + + this->processors = procs; } Realm::Runtime RealmContext::get_runtime() { diff --git a/lib/realm-execution/src/realm-execution/realm_manager.cc b/lib/realm-execution/src/realm-execution/realm_manager.cc index e76be7054b..cbc6cc2367 100644 --- a/lib/realm-execution/src/realm-execution/realm_manager.cc +++ b/lib/realm-execution/src/realm-execution/realm_manager.cc @@ -22,6 +22,7 @@ RealmManager::~RealmManager() { ControllerTaskResult RealmManager::start_controller(std::function thunk, Realm::Event wait_on) { + Realm::Processor target_proc = Realm::Machine::ProcessorQuery(Realm::Machine::get_machine()) .only_kind(Realm::Processor::LOC_PROC) diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/controller_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/controller_task.cc index 925402bcb4..21253604b8 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/controller_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/controller_task.cc @@ -48,8 +48,10 @@ ControllerTaskResult sizeof(raw_ptr), precondition); - return ControllerTaskResult{std::unique_ptr(raw_ptr), - event}; + return ControllerTaskResult{ + std::unique_ptr(raw_ptr), + event, + }; } } // namespace FlexFlow diff --git a/lib/realm-execution/test/src/realm-execution/test_e2e.cc b/lib/realm-execution/test/src/realm-execution/test_e2e.cc index 39eae4e1cb..f698d2a07f 100644 --- a/lib/realm-execution/test/src/realm-execution/test_e2e.cc +++ b/lib/realm-execution/test/src/realm-execution/test_e2e.cc @@ -36,6 +36,184 @@ static bool did_loss_decrease(GenericTensorAccessorR const &first_epoch, compare_tensor_accessors_le(last_epoch, first_epoch, allocator)); } +struct E2ETrainingConfig { + MappedParallelComputationGraph mapped_pcg; + LossAttrs loss_attrs; + MappedOperatorTaskGroup loss_mapping; + OptimizerAttrs optimizer_attrs; + parallel_tensor_guid_t logit_tensor; + TensorShape input_shape; + TensorShape logit_shape; + TensorShape label_shape; + TensorShape loss_shape; +}; + +static E2ETrainingConfig create_e2e_test_case() { + positive_int batch_size = 10_p; + positive_int data_dim = 16_p; + positive_int hidden_dim = 32_p; + positive_int output_dim = 1_p; + + TensorShape input_tensor_shape = + TensorShape{TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + TensorShape label_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + TensorShape loss_tensor_shape = TensorShape{ + TensorDims{FFOrdered{output_dim, hidden_dim}}, DataType::FLOAT}; + + TensorShape weight_shape_1 = + TensorShape{TensorDims{FFOrdered{hidden_dim, data_dim}}, DataType::FLOAT}; + + TensorShape weight_shape_2 = TensorShape{ + TensorDims{FFOrdered{output_dim, hidden_dim}}, DataType::FLOAT}; + + ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + ParallelLayerAddedResult inputs_layer = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input = + require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult weights_layer_1 = add_parallel_layer( + pcg, + ParallelLayerAttrs{ + PCGOperatorAttrs{WeightAttrs{weight_shape_1, + InitializerAttrs{GlorotNormalAttrs{0}}}}, + std::nullopt}, + {}, + {}); + parallel_tensor_guid_t t_weights_1 = + require_only_key(weights_layer_1.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult weights_layer_2 = add_parallel_layer( + pcg, + ParallelLayerAttrs{ + PCGOperatorAttrs{WeightAttrs{weight_shape_2, + InitializerAttrs{GlorotNormalAttrs{0}}}}, + std::nullopt}, + {}, + {}); + parallel_tensor_guid_t t_weights_2 = + require_only_key(weights_layer_2.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult linear_operator_1 = add_parallel_layer( + pcg, + ParallelLayerAttrs{PCGOperatorAttrs{LinearAttrs{hidden_dim, + /*use_bias=*/false, + DataType::FLOAT, + Activation::RELU, + std::nullopt}}, + std::nullopt}, + { + { + TensorSlotName::INPUT, + t_input, + }, + }, + { + { + TensorSlotName::WEIGHT, + t_weights_1, + }, + }); + parallel_tensor_guid_t t_linear_1 = + require_only_key(linear_operator_1.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult linear_operator_2 = add_parallel_layer( + pcg, + ParallelLayerAttrs{PCGOperatorAttrs{LinearAttrs{output_dim, + /*use_bias=*/false, + DataType::FLOAT, + Activation::RELU, + std::nullopt}}, + std::nullopt}, + { + { + TensorSlotName::INPUT, + t_linear_1, + }, + }, + { + { + TensorSlotName::WEIGHT, + t_weights_2, + }, + }); + parallel_tensor_guid_t t_linear_2 = + require_only_key(linear_operator_2.outputs, TensorSlotName::OUTPUT); + + MachineSpaceCoordinate cpu0{0_n, 0_n}; + MachineSpaceCoordinate cpu1{0_n, 1_n}; + ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; + + std::unordered_map mapping = { + {inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, + {weights_layer_1.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu0, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, + {weights_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + {{cpu1, + OperatorAtomicTaskShardBinding{ + {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, + {linear_operator_1.parallel_layer, + MappedOperatorTaskGroup{{{cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::WEIGHT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}}}}, + {linear_operator_2.parallel_layer, + MappedOperatorTaskGroup{{{cpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::WEIGHT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}}}}}, + }; + + TensorShape output_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + MappedParallelComputationGraph mpcg = + mapped_pcg_from_pcg_and_mapped_op_task_groups(pcg, mapping); + + MappedOperatorTaskGroup loss_mapping{ + {{cpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::LOGIT, tensor_coord0}, + }}}}}; + + LossAttrs loss_attrs = LossAttrs{ + NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; + OptimizerAttrs optimizer_attrs = + OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001}}; + + return E2ETrainingConfig{ + /*mapped_pcg=*/mpcg, + /*loss_attrs=*/loss_attrs, + /*loss_mapping=*/loss_mapping, + /*optimizer_attrs=*/optimizer_attrs, + /*logit_tensor=*/t_linear_2, + /*input_shape=*/input_tensor_shape, + /*logit_shape=*/output_tensor_shape, + /*label_shape=*/label_tensor_shape, + /*loss_shape=*/loss_tensor_shape, + }; +} + TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("RealmBackend e2e Training (CPU Model Parallelism)") { std::vector fake_args = @@ -46,163 +224,16 @@ TEST_SUITE(FF_TEST_SUITE) { RealmManager manager = RealmManager{&fake_argc, &fake_argv}; (void)manager.start_controller([](RealmContext &ctx) { - Allocator allocator = ctx.get_current_device_allocator(); - - positive_int batch_size = 10_p; - positive_int data_dim = 16_p; - positive_int hidden_dim = 32_p; - positive_int output_dim = 1_p; + ASSERT(ctx.processors.has_value()); - TensorShape output_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + E2ETrainingConfig cfg = create_e2e_test_case(); - GenericTensorAccessorW label_tensor_backing = - allocator.allocate_tensor(output_tensor_shape); - - // construct computation graph - ParallelComputationGraph pcg = empty_parallel_computation_graph(); - - TensorShape input_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + Allocator allocator = ctx.get_current_device_allocator(); - TensorShape label_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + GenericTensorAccessorW output_tensor = + allocator.allocate_tensor(cfg.logit_shape); GenericTensorAccessorW label_tensor = - allocator.allocate_tensor(label_tensor_shape); - - TensorShape weight_shape_1 = TensorShape{ - TensorDims{FFOrdered{hidden_dim, data_dim}}, DataType::FLOAT}; - TensorShape weight_shape_2 = TensorShape{ - TensorDims{FFOrdered{output_dim, hidden_dim}}, DataType::FLOAT}; - - ParallelLayerAddedResult inputs_layer = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input = - require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult weights_layer_1 = add_parallel_layer( - pcg, - ParallelLayerAttrs{ - PCGOperatorAttrs{WeightAttrs{ - weight_shape_1, InitializerAttrs{GlorotNormalAttrs{0}}}}, - std::nullopt}, - {}, - {}); - parallel_tensor_guid_t t_weights_1 = - require_only_key(weights_layer_1.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult weights_layer_2 = add_parallel_layer( - pcg, - ParallelLayerAttrs{ - PCGOperatorAttrs{WeightAttrs{ - weight_shape_2, InitializerAttrs{GlorotNormalAttrs{0}}}}, - std::nullopt}, - {}, - {}); - parallel_tensor_guid_t t_weights_2 = - require_only_key(weights_layer_2.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult linear_operator_1 = add_parallel_layer( - pcg, - ParallelLayerAttrs{PCGOperatorAttrs{LinearAttrs{hidden_dim, - /*use_bias=*/false, - DataType::FLOAT, - Activation::RELU, - std::nullopt}}, - std::nullopt}, - { - { - TensorSlotName::INPUT, - t_input, - }, - }, - { - { - TensorSlotName::WEIGHT, - t_weights_1, - }, - }); - parallel_tensor_guid_t t_linear_1 = - require_only_key(linear_operator_1.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult linear_operator_2 = add_parallel_layer( - pcg, - ParallelLayerAttrs{PCGOperatorAttrs{LinearAttrs{output_dim, - /*use_bias=*/false, - DataType::FLOAT, - Activation::RELU, - std::nullopt}}, - std::nullopt}, - { - { - TensorSlotName::INPUT, - t_linear_1, - }, - }, - { - { - TensorSlotName::WEIGHT, - t_weights_2, - }, - }); - parallel_tensor_guid_t t_linear_2 = - require_only_key(linear_operator_2.outputs, TensorSlotName::OUTPUT); - - MachineSpaceCoordinate cpu0{0_n, 0_n, DeviceType::CPU}; - MachineSpaceCoordinate cpu1{0_n, 1_n, DeviceType::CPU}; - ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg = - mapped_pcg_from_pcg_and_mapped_op_task_groups( - /*pcg=*/pcg, - /*mapped_op_task_groups=*/{ - {inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {weights_layer_1.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {weights_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu1, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {linear_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::WEIGHT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, - {linear_operator_2.parallel_layer, - MappedOperatorTaskGroup{ - {{cpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::WEIGHT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, - }); - - MappedOperatorTaskGroup loss_mapping{ - {{cpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}}}}; - - // instantiate computation graph - LossAttrs loss_attrs = LossAttrs{ - NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; - OptimizerAttrs optimizer_attrs = - OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001}}; + allocator.allocate_tensor(cfg.label_shape); std::unordered_map input_tensors; @@ -214,18 +245,19 @@ TEST_SUITE(FF_TEST_SUITE) { PCGInstance pcg_instance = create_pcg_instance( /*ctx=*/ctx, - /*mpcg=*/mpcg, - /*optimizer=*/optimizer_attrs, + /*mpcg=*/cfg.mapped_pcg, + /*optimizer=*/cfg.optimizer_attrs, /*loss=*/ ParallelLossConfig{ - /*loss_attrs=*/loss_attrs, + /*loss_attrs=*/cfg.loss_attrs, /*label_tensor=*/label_tensor, - /*logit_tensor=*/t_linear_2, - /*loss_mapping=*/loss_mapping, + /*logit_tensor=*/cfg.logit_tensor, + /*loss_mapping=*/cfg.loss_mapping, }, /*input_tensors=*/input_tensors, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle); + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::CPU); // begin training loop int num_epochs = 5; @@ -240,9 +272,7 @@ TEST_SUITE(FF_TEST_SUITE) { dynamic_tensor_accessor_from_instance( pcg_instance.get_loss_tensor_instance().value(), Realm::Event::NO_EVENT, - lift_to_parallel( - TensorShape{TensorDims{FFOrdered{output_dim, hidden_dim}}, - DataType::FLOAT}), + lift_to_parallel(cfg.loss_shape), Permissions::RO, ctx.get_current_processor()) .require_read(), @@ -265,155 +295,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_SUITE(FF_CUDA_TEST_SUITE) { TEST_CASE("RealmBackend e2e Training (GPU Model Parallelism)") { - positive_int batch_size = 10_p; - positive_int data_dim = 16_p; - positive_int hidden_dim = 32_p; - positive_int output_dim = 1_p; - - TensorShape output_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - // construct computation graph - ParallelComputationGraph pcg = empty_parallel_computation_graph(); - - TensorShape input_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; - - TensorShape label_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - TensorShape weight_shape_1 = TensorShape{ - TensorDims{FFOrdered{hidden_dim, data_dim}}, DataType::FLOAT}; - TensorShape weight_shape_2 = TensorShape{ - TensorDims{FFOrdered{output_dim, hidden_dim}}, DataType::FLOAT}; - - ParallelLayerAddedResult inputs_layer = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input = - require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult weights_layer_1 = add_parallel_layer( - pcg, - ParallelLayerAttrs{ - PCGOperatorAttrs{WeightAttrs{ - weight_shape_1, InitializerAttrs{GlorotNormalAttrs{0}}}}, - std::nullopt}, - {}, - {}); - parallel_tensor_guid_t t_weights_1 = - require_only_key(weights_layer_1.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult weights_layer_2 = add_parallel_layer( - pcg, - ParallelLayerAttrs{ - PCGOperatorAttrs{WeightAttrs{ - weight_shape_2, InitializerAttrs{GlorotNormalAttrs{0}}}}, - std::nullopt}, - {}, - {}); - parallel_tensor_guid_t t_weights_2 = - require_only_key(weights_layer_2.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult linear_operator_1 = add_parallel_layer( - pcg, - ParallelLayerAttrs{PCGOperatorAttrs{LinearAttrs{hidden_dim, - /*use_bias=*/false, - DataType::FLOAT, - Activation::RELU, - std::nullopt}}, - std::nullopt}, - { - { - TensorSlotName::INPUT, - t_input, - }, - }, - { - { - TensorSlotName::WEIGHT, - t_weights_1, - }, - }); - parallel_tensor_guid_t t_linear_1 = - require_only_key(linear_operator_1.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult linear_operator_2 = add_parallel_layer( - pcg, - ParallelLayerAttrs{PCGOperatorAttrs{LinearAttrs{output_dim, - /*use_bias=*/false, - DataType::FLOAT, - Activation::RELU, - std::nullopt}}, - std::nullopt}, - { - { - TensorSlotName::INPUT, - t_linear_1, - }, - }, - { - { - TensorSlotName::WEIGHT, - t_weights_2, - }, - }); - parallel_tensor_guid_t t_linear_2 = - require_only_key(linear_operator_2.outputs, TensorSlotName::OUTPUT); - - MachineSpaceCoordinate gpu0{0_n, 0_n, DeviceType::GPU}; - ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; - MappedParallelComputationGraph mpcg = - mapped_pcg_from_pcg_and_mapped_op_task_groups( - /*pcg=*/pcg, - /*mapped_op_task_groups=*/{ - {inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {weights_layer_1.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {weights_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{ - {{TensorSlotName::OUTPUT, tensor_coord0}}}}}}}, - {linear_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::WEIGHT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, - {linear_operator_2.parallel_layer, - MappedOperatorTaskGroup{ - {{gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::WEIGHT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}}}}}, - }); - - MappedOperatorTaskGroup loss_mapping{ - {{gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::LOGIT, tensor_coord0}, - }}}}}; - - // instantiate computation graph - LossAttrs loss_attrs = LossAttrs{ - NonconfigurableLossAttrs{LossFunction::CATEGORICAL_CROSSENTROPY}}; - OptimizerAttrs optimizer_attrs = - OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001}}; + E2ETrainingConfig cfg = create_e2e_test_case(); //! [realm-execution example] std::vector fake_args = @@ -427,11 +309,10 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { manager.start_controller([&](RealmContext &ctx) { Allocator allocator = ctx.get_current_device_allocator(); - GenericTensorAccessorW label_tensor_backing = - allocator.allocate_tensor(output_tensor_shape); - + GenericTensorAccessorW logit_tensor = + allocator.allocate_tensor(cfg.logit_shape); GenericTensorAccessorW label_tensor = - allocator.allocate_tensor(label_tensor_shape); + allocator.allocate_tensor(cfg.label_shape); std::unordered_map input_tensors; @@ -443,18 +324,19 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { PCGInstance pcg_instance = create_pcg_instance( /*ctx=*/ctx, - /*mpcg=*/mpcg, - /*optimizer=*/optimizer_attrs, + /*mpcg=*/cfg.mapped_pcg, + /*optimizer=*/cfg.optimizer_attrs, /*loss=*/ ParallelLossConfig{ - /*loss_attrs=*/loss_attrs, + /*loss_attrs=*/cfg.loss_attrs, /*label_tensor=*/label_tensor, - /*logit_tensor=*/t_linear_2, - /*loss_mapping=*/loss_mapping, + /*logit_tensor=*/cfg.logit_tensor, + /*loss_mapping=*/cfg.loss_mapping, }, /*input_tensors=*/input_tensors, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle); + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::GPU); // begin training loop int num_epochs = 5; @@ -465,13 +347,12 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { /*instance=*/pcg_instance, /*profiling_settings=*/ProfilingSettings{0, 0}, /*device_handle=*/device_handle); + loss_values.push_back(copy_tensor_accessor_r( dynamic_tensor_accessor_from_instance( pcg_instance.get_loss_tensor_instance().value(), Realm::Event::NO_EVENT, - lift_to_parallel(TensorShape{ - TensorDims{FFOrdered{output_dim, hidden_dim}}, - DataType::FLOAT}), + lift_to_parallel(cfg.loss_shape), Permissions::RO, ctx.get_current_processor()) .require_read(), diff --git a/lib/task-spec/include/task-spec/concrete_arg_spec.h b/lib/task-spec/include/task-spec/concrete_arg_spec.h index 45bbd6ba6b..32df07d325 100644 --- a/lib/task-spec/include/task-spec/concrete_arg_spec.h +++ b/lib/task-spec/include/task-spec/concrete_arg_spec.h @@ -1,7 +1,6 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_CONCRETE_ARG_SPEC_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_CONCRETE_ARG_SPEC_H -#include "task-spec/serialization.h" #include "utils/hash-utils.h" #include "utils/type_index.h" #include diff --git a/lib/task-spec/include/task-spec/device_id_t.dtg.toml b/lib/task-spec/include/task-spec/device_id_t.dtg.toml new file mode 100644 index 0000000000..f6d29ee1a7 --- /dev/null +++ b/lib/task-spec/include/task-spec/device_id_t.dtg.toml @@ -0,0 +1,23 @@ +namespace = "FlexFlow" +name = "device_id_t" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/machine_space_coordinate.dtg.h", + "pcg/device_type.dtg.h", +] + +[[fields]] +name = "coord" +type = "::FlexFlow::MachineSpaceCoordinate" + +[[fields]] +name = "device_type" +type = "::FlexFlow::DeviceType" diff --git a/lib/task-spec/include/task-spec/device_specific.h b/lib/task-spec/include/task-spec/device_specific.h index 2055888b1b..49a2555411 100644 --- a/lib/task-spec/include/task-spec/device_specific.h +++ b/lib/task-spec/include/task-spec/device_specific.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_LOCAL_EXECUTION_DEVICE_SPECIFIC_H #define _FLEXFLOW_LOCAL_EXECUTION_DEVICE_SPECIFIC_H -#include "pcg/device_id_t.dtg.h" -#include "task-spec/serialization.h" +#include "task-spec/device_id_t.dtg.h" #include "utils/hash/tuple.h" +#include namespace FlexFlow { @@ -12,7 +12,8 @@ struct DeviceSpecific { DeviceSpecific() = delete; template - static DeviceSpecific create(device_id_t device_idx, Args &&...args) { + static DeviceSpecific create(device_id_t const &device_idx, + Args &&...args) { return DeviceSpecific(std::make_shared(std::forward(args)...), device_idx); } @@ -25,13 +26,13 @@ struct DeviceSpecific { return this->tie() != other.tie(); } - T const *get(device_id_t curr_device_idx) const { + T const *get(device_id_t const &curr_device_idx) const { ASSERT(curr_device_idx == this->device_idx); return (T const *)this->ptr.get(); } private: - DeviceSpecific(std::shared_ptr ptr, device_id_t device_idx) + DeviceSpecific(std::shared_ptr ptr, device_id_t const &device_idx) : ptr(ptr), device_idx(device_idx) {} private: @@ -57,11 +58,6 @@ std::ostream &operator<<(std::ostream &s, DeviceSpecific const &d) { return (s << fmt::to_string(d)); } -// manually force serialization to make DeviceSpecific trivially -// serializable -// template -// struct is_trivially_serializable> : std::true_type {}; - } // namespace FlexFlow namespace std { diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml new file mode 100644 index 0000000000..578bfe7585 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml @@ -0,0 +1,29 @@ +namespace = "FlexFlow" +name = "DynamicGraphEdge" +type = "struct" +features = [ + "eq", + "hash", + "fmt", +] + +includes = [ + "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", + "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h", + "task-spec/dynamic_graph/dynamic_slot_site.h", +] + +src_includes = [ +] + +[[fields]] +name = "src" +type = "::FlexFlow::DynamicSlotSite" + +[[fields]] +name = "dst_node" +type = "::FlexFlow::DynamicNodeInvocation" + +[[fields]] +name = "dst_slot" +type = "::FlexFlow::DynamicTensorSlot" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.h new file mode 100644 index 0000000000..b561dba8c4 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.h @@ -0,0 +1,15 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_GRAPH_EDGE_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_GRAPH_EDGE_H + +#include "task-spec/dynamic_graph/dynamic_graph_edge.dtg.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" + +namespace FlexFlow { + +DynamicGraphEdge + dynamic_graph_edge_from_slot_sites(DynamicSlotSite const &src, + InternalDynamicSlotSite const &dst); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml index 73c023fd40..7655bd24e1 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml @@ -10,11 +10,11 @@ features = [ includes = [ "", "task-spec/dynamic_graph/dynamic_task_type.dtg.h", - "pcg/machine_space_coordinate.dtg.h", - "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h", "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h", "task-spec/dynamic_graph/training_operation_attrs.dtg.h", "task-spec/device_specific_per_device_op_state.dtg.h", + "task-spec/device_id_t.dtg.h", + "task-spec/dynamic_graph/dynamic_node_mapping.dtg.h", ] src_includes = [ @@ -27,8 +27,13 @@ type = "std::optional<::FlexFlow::DynamicTaskType>" [[fields]] name = "device_coord" -type = "std::optional<::FlexFlow::MachineSpaceCoordinate>" +type = "std::optional<::FlexFlow::device_id_t>" docstring = ''' +\brief The device on which this task should execute. + +For a standard \ref ParallelComputationGraph, filled in by +\ref perform_shard_expansion. + \note Right now the \c device_coord for a copy node is sort of meaningless because we have one controller issuing all copies for the entire graph, no matter where they are. However the intention is this to be the "owner" or @@ -41,7 +46,7 @@ choices. [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::MappedOperatorTaskGroup>" +type = "std::optional<::FlexFlow::DynamicNodeMapping>" [[fields]] name = "op_attrs" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h new file mode 100644 index 0000000000..6f27b43ad3 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h @@ -0,0 +1,22 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_INVOCATION_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_INVOCATION_H + +#include "pcg/tensor_direction.dtg.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" +#include "task-spec/dynamic_graph/training_op_type.dtg.h" + +namespace FlexFlow { + +std::unordered_map + get_slot_map_for_direction(DynamicNodeInvocation const &, TensorDirection); + +TrainingOpType + dynamic_node_invocation_get_op_type(DynamicNodeInvocation const &); + +std::unordered_set + get_dynamic_slot_sites_for_invocation(DynamicNodeInvocation const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.dtg.toml new file mode 100644 index 0000000000..859902e2bf --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.dtg.toml @@ -0,0 +1,26 @@ +namespace = "FlexFlow" +name = "DynamicNodeMapping" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h", + "pcg/device_type.dtg.h", +] + +src_includes = [ +] + +[[fields]] +name = "op_task_group" +type = "::FlexFlow::MappedOperatorTaskGroup" + +[[fields]] +name = "device_type" +type = "::FlexFlow::DeviceType" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h new file mode 100644 index 0000000000..6a52773a71 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h @@ -0,0 +1,18 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_MAPPING_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_MAPPING_H + +#include "task-spec/device_id_t.dtg.h" +#include "task-spec/dynamic_graph/dynamic_node_mapping.dtg.h" + +namespace FlexFlow { + +bidict + dynamic_node_mapping_bindings_for_slot_name(DynamicNodeMapping const &, + TensorSlotName const &); + +std::unordered_set + target_devices_of_dynamic_node_mapping(DynamicNodeMapping const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h index 4ca62db5b1..6141c70e54 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h @@ -1,14 +1,19 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_OPEN_DATAFLOW_GRAPH_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_OPEN_DATAFLOW_GRAPH_H +#include "task-spec/dynamic_graph/dynamic_graph_edge.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h" namespace FlexFlow { DynamicOpenDataflowGraph make_empty_dynamic_open_dataflow_graph(); +void check_dynamic_open_dataflow_graph_is_valid( + DynamicOpenDataflowGraph const &); + nonnegative_int dynamic_graph_num_nodes(DynamicOpenDataflowGraph const &); bool full_dynamic_graph_satisfies( @@ -32,6 +37,28 @@ std::unordered_multiset std::unordered_set get_dynamic_invocation_set(DynamicOpenDataflowGraph const &); +std::unordered_set + get_dynamic_graph_edges(DynamicOpenDataflowGraph const &); +std::unordered_set + get_dynamic_graph_edges_incoming_to_invocation( + DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); +std::unordered_set + get_dynamic_graph_edges_outgoing_from_invocation( + DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); + +std::unordered_set + get_internal_dynamic_slot_sites(DynamicOpenDataflowGraph const &); + +std::unordered_set + get_dynamic_slot_sites(DynamicOpenDataflowGraph const &); + +DynamicSlotSite + dynamic_graph_find_source_of_value(DynamicOpenDataflowGraph const &, + DynamicValueAttrs const &); +std::unordered_set + dynamic_graph_find_sinks_of_value(DynamicOpenDataflowGraph const &, + DynamicValueAttrs const &); + std::optional find_output_value_attrs(DynamicOpenDataflowGraph const &, dynamic_tensor_guid_t, diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml new file mode 100644 index 0000000000..fa56b0f105 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml @@ -0,0 +1,21 @@ +namespace = "FlexFlow" +name = "DynamicSlotSite" +type = "variant" +features = [ + "eq", + "hash", + "fmt", +] + +includes = [ + "task-spec/dynamic_graph/external_dynamic_slot_site.dtg.h", + "task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.h", +] + +[[values]] +type = "::FlexFlow::InternalDynamicSlotSite" +key = "internal" + +[[values]] +type = "::FlexFlow::ExternalDynamicSlotSite" +key = "external" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h new file mode 100644 index 0000000000..4340cdcf39 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h @@ -0,0 +1,13 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_SLOT_SITE_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_SLOT_SITE_H + +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" +#include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" + +namespace FlexFlow { + +DynamicValueAttrs dynamic_value_attrs_for_slot_site(DynamicSlotSite const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml index 490a51f88d..926cc9893c 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml @@ -16,6 +16,7 @@ includes = [ "utils/bidict/bidict.h", "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_role.dtg.h", + "task-spec/dynamic_graph/parallel_tensor_mapping.dtg.h", ] src_includes = [ @@ -36,7 +37,13 @@ type = "std::optional<::FlexFlow::ParallelTensorSpaceCoordinate>" [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" +type = "std::optional<::FlexFlow::ParallelTensorMapping>" +docstring = ''' +\brief Which devices the shards of the parallel tensor are stored on. + +For computations originating from a \ref ParallelComputationGraph, +this field is usually resolved/assigned by \ref perform_copy_insertion. +''' [[fields]] name = "accessor" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h index 9cccc565cc..d6090e5f14 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h @@ -2,12 +2,17 @@ #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_VALUE_ATTRS_H #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" +#include "task-spec/dynamic_graph/parallel_tensor_mapping.dtg.h" namespace FlexFlow { DynamicValueAttrs decide_dynamic_value_attrs_role(DynamicValueAttrs const &, DynamicTensorRole); +DynamicValueAttrs + dynamic_value_attrs_with_mapping(DynamicValueAttrs const &, + ParallelTensorMapping const &); + } // namespace FlexFlow #endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml new file mode 100644 index 0000000000..bdcc93251c --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml @@ -0,0 +1,19 @@ +namespace = "FlexFlow" +name = "ExternalDynamicSlotSite" +type = "struct" +features = [ + "eq", + "hash", + "fmt", +] + +includes = [ + "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h", +] + +src_includes = [ +] + +[[fields]] +name = "value" +type = "::FlexFlow::DynamicValueAttrs" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml new file mode 100644 index 0000000000..bf29bf2b2f --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml @@ -0,0 +1,29 @@ +namespace = "FlexFlow" +name = "InternalDynamicSlotSite" +type = "struct" +features = [ + "eq", + "hash", + "fmt", +] + +includes = [ + "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", + "pcg/tensor_direction.dtg.h", + "task-spec/dynamic_graph/dynamic_tensor_slot.h", +] + +src_includes = [ +] + +[[fields]] +name = "invocation" +type = "::FlexFlow::DynamicNodeInvocation" + +[[fields]] +name = "direction" +type = "::FlexFlow::TensorDirection" + +[[fields]] +name = "slot_name" +type = "::FlexFlow::DynamicTensorSlot" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion.h index b3b2a465f8..8584b4bed8 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion.h @@ -14,7 +14,7 @@ LossInsertionResult perform_loss_insertion( DynamicOpenDataflowGraph const &dg, LossAttrs const &loss_attrs, dynamic_tensor_guid_t logit_tensor, - std::optional const &loss_mapping); + std::optional const &loss_mapping); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h b/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h index 823f962c25..b40fabe7bf 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h @@ -7,11 +7,11 @@ namespace FlexFlow { std::unordered_set perform_machine_slicing_for_invocation(DynamicNodeInvocation const &, - MachineSpaceCoordinate const &); + device_id_t const &); DynamicOpenDataflowGraph perform_machine_slicing(DynamicOpenDataflowGraph const &, - MachineSpaceCoordinate const &); + device_id_t const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h index 6a269ec3c9..c5db832a6a 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h @@ -7,7 +7,7 @@ namespace FlexFlow { DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( - MappedParallelComputationGraph const &); + MappedParallelComputationGraph const &, DeviceType device_type); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml new file mode 100644 index 0000000000..ba9785bc92 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml @@ -0,0 +1,22 @@ +namespace = "FlexFlow" +name = "ParallelTensorMapping" +type = "struct" +features = [ + "eq", + "hash", + "json", + "fmt", +] + +includes = [ + "utils/bidict/bidict.h", + "op-attrs/parallel_tensor_space_coordinate.dtg.h", + "task-spec/device_id_t.dtg.h", +] + +src_includes = [ +] + +[[fields]] +name = "raw" +type = "::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::device_id_t>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml index 3c43e1d637..43fb8490ec 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml @@ -11,8 +11,8 @@ features = [ includes = [ "", "task-spec/dynamic_graph/dynamic_task_type.dtg.h", - "pcg/machine_space_coordinate.dtg.h", - "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h", + "task-spec/device_id_t.dtg.h", + "task-spec/dynamic_graph/dynamic_node_mapping.dtg.h", "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h", "task-spec/dynamic_graph/training_operation_attrs.dtg.h", ] @@ -28,11 +28,11 @@ type = "std::optional<::FlexFlow::DynamicTaskType>" [[fields]] name = "device_coord" -type = "std::optional<::FlexFlow::MachineSpaceCoordinate>" +type = "std::optional<::FlexFlow::device_id_t>" [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::MappedOperatorTaskGroup>" +type = "std::optional<::FlexFlow::DynamicNodeMapping>" [[fields]] name = "op_attrs" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml index 454f1b7e8c..c884bd8abe 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml @@ -16,6 +16,7 @@ includes = [ "pcg/machine_space_coordinate.dtg.h", "utils/bidict/bidict.h", "task-spec/dynamic_graph/dynamic_tensor_role.dtg.h", + "task-spec/dynamic_graph/parallel_tensor_mapping.dtg.h", ] src_includes = [ @@ -37,7 +38,7 @@ type = "std::optional<::FlexFlow::ParallelTensorSpaceCoordinate>" [[fields]] name = "mapping" -type = "std::optional<::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::MachineSpaceCoordinate>>" +type = "std::optional<::FlexFlow::ParallelTensorMapping>" [[fields]] name = "role" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_only_op_type.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/training_only_op_type.dtg.toml new file mode 100644 index 0000000000..f1013771d1 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_only_op_type.dtg.toml @@ -0,0 +1,15 @@ +namespace = "FlexFlow" +name = "TrainingOnlyOpType" +type = "enum" +features = [ + "hash", + "json", + "rapidcheck", + "fmt", +] + +[[values]] +name = "COPY" + +[[values]] +name = "LOSS" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml new file mode 100644 index 0000000000..7744f3775f --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml @@ -0,0 +1,22 @@ +namespace = "FlexFlow" +name = "TrainingOpType" +type = "variant" +features = [ + "eq", + "hash", + "fmt", + "json", +] + +includes = [ + "task-spec/dynamic_graph/training_only_op_type.dtg.h", + "op-attrs/operator_type.dtg.h", +] + +[[values]] +type = "::FlexFlow::OperatorType" +key = "pcg_op" + +[[values]] +type = "::FlexFlow::TrainingOnlyOpType" +key = "training_only_op" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h new file mode 100644 index 0000000000..92446cb1b7 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_operation_attrs.h @@ -0,0 +1,13 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_TRAINING_OPERATION_ATTRS_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_TRAINING_OPERATION_ATTRS_H + +#include "task-spec/dynamic_graph/training_op_type.dtg.h" +#include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" + +namespace FlexFlow { + +TrainingOpType training_op_attrs_get_op_type(TrainingOperationAttrs const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/serialization.h b/lib/task-spec/include/task-spec/serialization.h deleted file mode 100644 index 29f9144a3b..0000000000 --- a/lib/task-spec/include/task-spec/serialization.h +++ /dev/null @@ -1,113 +0,0 @@ -#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_SERIALIZATION_H -#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_SERIALIZATION_H - -#include "kernels/device.h" -#include "kernels/nccl.h" -#include "op-attrs/ff_ordered/ff_ordered.h" -#include "utils/required.h" -#include "utils/type_traits.h" -#include "utils/variant.h" - -namespace FlexFlow { - -template -struct needs_serialization {}; - -template -struct visit_trivially_serializable; - -template -struct is_trivially_serializable : std::false_type {}; - -template -struct visit_trivially_serializable { - static constexpr bool value = is_trivially_serializable::value && - visit_trivially_serializable::value; -}; - -template -struct visit_trivially_serializable> { - static constexpr bool value = visit_trivially_serializable::value; -}; - -template <> -struct visit_trivially_serializable<> : std::true_type {}; - -template -struct is_trivially_serializable< - T, - typename std::enable_if::value>::type> - : std::true_type {}; - -template <> -struct is_trivially_serializable : std::true_type {}; -template <> -struct is_trivially_serializable : std::true_type {}; - -template -struct is_trivially_serializable< - T, - typename std::enable_if::value>::type> : std::true_type {}; - -template -struct is_trivially_serializable< - T, - typename std::enable_if::value>::type> - : std::true_type {}; - -template -struct is_trivially_serializable> - : is_trivially_serializable {}; - -template -struct is_trivially_serializable> : is_trivially_serializable { -}; - -template -struct is_trivially_serializable> - : elements_satisfy> {}; - -template -struct is_trivially_serializable> - : is_trivially_serializable {}; - -template -struct std_array_size_helper; - -template -struct std_array_size_helper> { - static const std::size_t value = N; -}; - -template -using std_array_size = std_array_size_helper; - -template -struct is_trivially_serializable< - T, - std::enable_if::value>>::value>> - : std::true_type {}; - -template -struct is_serializable : std::false_type {}; - -template -struct is_serializable< - T, - typename std::enable_if::value>::type> - : std::true_type {}; - -static_assert(is_trivially_serializable::value, ""); -static_assert(is_trivially_serializable::value, ""); -static_assert(is_trivially_serializable::value, ""); -static_assert(is_trivially_serializable::value, ""); -static_assert(is_trivially_serializable::value, ""); -static_assert(is_trivially_serializable::value, ""); -static_assert(is_trivially_serializable>::value, - ""); - -} // namespace FlexFlow - -#endif diff --git a/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h b/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h index b4c8dcdf36..165631889f 100644 --- a/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h +++ b/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h @@ -7,9 +7,9 @@ #include "op-attrs/ops/loss_functions/loss_attrs.dtg.h" #include "op-attrs/pcg_operator_attrs.dtg.h" #include "op-attrs/tensor_slot_name.dtg.h" -#include "pcg/device_id_t.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/concrete_arg_spec.h" +#include "task-spec/device_id_t.dtg.h" #include "task-spec/ops/arg_slot_id_t.dtg.h" #include "task-spec/per_device_op_state.dtg.h" #include "task-spec/privilege_tensor_accessor.h" diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 79bd091c0a..f611735e26 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -3,21 +3,35 @@ #include "op-attrs/tensor_slot_name.dtg.h" #include "pcg/machine_space_coordinate.dtg.h" #include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" +#include "task-spec/dynamic_graph/copy_insertion.h" #include "task-spec/dynamic_graph/dynamic_node_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" +#include "task-spec/dynamic_graph/dynamic_node_mapping.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.h" #include "task-spec/dynamic_graph/dynamic_task_type.h" #include "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" +#include "task-spec/dynamic_graph/dynamic_value_attrs.h" +#include "task-spec/dynamic_graph/parallel_tensor_mapping.dtg.h" #include "utils/bidict/algorithms/bidict_from_pairs.h" #include "utils/bidict/algorithms/unordered_set_of.h" #include "utils/containers/contains_key.h" +#include "utils/containers/count.h" +#include "utils/containers/filter_values.h" +#include "utils/containers/filtermap_keys.h" +#include "utils/containers/filtrans.h" #include "utils/containers/flatmap.h" #include "utils/containers/map_values2.h" +#include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/set_difference.h" #include "utils/containers/set_intersection.h" #include "utils/containers/transform.h" +#include "utils/containers/values.h" +#include "utils/containers/zip_values_strict_with.h" #include "utils/optional.h" +#include "utils/overload.h" namespace FlexFlow { @@ -44,107 +58,171 @@ bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { g, node_is_any, value_is_mapped, slot_is_mapped); } -static DynamicValueAttrs map_dynamic_value_attrs_for_task_group( - DynamicTensorSlot const &slot, - DynamicValueAttrs const &value, - MappedOperatorTaskGroup const &mapping) { - DynamicValueAttrs result = value; - result.mapping = get_tensor_bindings_for_slot_name(mapping, slot.slot_name); - return result; -} - static std::pair filter_mapping_to_avoid_degenerate_copies(DynamicValueAttrs const &input, DynamicValueAttrs const &output) { - std::unordered_set< - std::pair> - input_mapping = unordered_set_of(assert_unwrap(input.mapping)); - std::unordered_set< - std::pair> - output_mapping = unordered_set_of(assert_unwrap(output.mapping)); + std::unordered_set> + input_mapping = unordered_set_of(assert_unwrap(input.mapping).raw); + + std::unordered_set> + output_mapping = unordered_set_of(assert_unwrap(output.mapping).raw); // Exclude the point shared between the input and output mappings, because // those will not result in actual copies once shard expansion is performed std::unordered_set< - std::pair> + std::pair> remove = set_intersection(input_mapping, output_mapping); DynamicValueAttrs filtered_input = input; - filtered_input.mapping = - bidict_from_pairs(set_difference(input_mapping, remove)); + filtered_input.mapping = ParallelTensorMapping{ + bidict_from_pairs(set_difference(input_mapping, remove)), + }; DynamicValueAttrs filtered_output = output; - filtered_output.mapping = - bidict_from_pairs(set_difference(output_mapping, remove)); + filtered_output.mapping = ParallelTensorMapping{ + bidict_from_pairs(set_difference(output_mapping, remove)), + }; return std::pair{filtered_input, filtered_output}; } -std::unordered_set perform_copy_insertion_for_invocation( +std::unordered_map + get_mappings_for_invocation( + DynamicNodeInvocation const &i, + std::unordered_map const + &mappings) { + return filtermap_keys(mappings, + [&](InternalDynamicSlotSite const &s) + -> std::optional { + if (s.invocation == i) { + return s.slot_name; + } else { + return std::nullopt; + } + }); +} + +DynamicNodeInvocation apply_mappings_for_invocation( DynamicNodeInvocation const &i, - std::unordered_map const - &unmapped_value_to_mapped_source_value) { + std::unordered_map const + &all_mappings) { + std::unordered_map i_mappings = + get_mappings_for_invocation(i, all_mappings); + + std::unordered_map + i_input_mappings = restrict_keys(i_mappings, keys(i.inputs)); - MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); + std::unordered_map + i_output_mappings = restrict_keys(i_mappings, keys(i.outputs)); - auto map_tensor = [&](DynamicTensorSlot const &slot, - DynamicValueAttrs const &value) { - return map_dynamic_value_attrs_for_task_group(slot, value, mapping); + auto apply_mapping = + [&](DynamicValueAttrs const &v, + ParallelTensorMapping const &mapping) -> DynamicValueAttrs { + return dynamic_value_attrs_with_mapping(v, mapping); }; - std::unordered_map mapped_inputs = - map_values2(i.inputs, map_tensor); - std::unordered_map mapped_outputs = - map_values2(i.outputs, map_tensor); + return DynamicNodeInvocation{ + /*inputs=*/ + zip_values_strict_with(i.inputs, i_input_mappings, apply_mapping), + /*node_attrs=*/ + i.node_attrs, + /*outputs=*/ + zip_values_strict_with(i.outputs, i_output_mappings, apply_mapping), + }; +} - std::unordered_set result{DynamicNodeInvocation{ - /*inputs=*/mapped_inputs, - /*node_attrs=*/i.node_attrs, - /*outputs=*/mapped_outputs, - }}; +std::unordered_set copies_for_value( + DynamicOpenDataflowGraph const &g, + DynamicValueAttrs const &v, + std::unordered_map const + &mappings) { + InternalDynamicSlotSite src = ({ + DynamicSlotSite found = dynamic_graph_find_source_of_value(g, v); - for (auto const &[slot, input] : i.inputs) { - if (!contains_key(unmapped_value_to_mapped_source_value, input)) { - continue; + if (found.is_external()) { + return {}; } - DynamicValueAttrs source_value = - unmapped_value_to_mapped_source_value.at(input); - DynamicValueAttrs use_value = mapped_inputs.at(slot); - if (source_value != use_value) { - auto const &[filtered_source, filtered_use] = - filter_mapping_to_avoid_degenerate_copies(source_value, use_value); - DynamicNodeInvocation copy{ - /*inputs=*/{ - { - DynamicTensorSlot{TensorSlotName::INPUT, - slot.slot_tensor_role}, - filtered_source, - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/transform( - slot.slot_tensor_role, - dynamic_task_type_from_tensor_role_for_copy), - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - DynamicTensorSlot{TensorSlotName::OUTPUT, - slot.slot_tensor_role}, - filtered_use, - }, - }, - }; - result.insert(copy); - } - } + found.require_internal(); + }); + + std::unordered_set sinks = + dynamic_graph_find_sinks_of_value(g, DynamicValueAttrs{v}); + + ParallelTensorMapping src_mapping = mappings.at(src); + + std::unordered_map + mappings_for_sinks = generate_map( + sinks, + [&](InternalDynamicSlotSite const &s) -> ParallelTensorMapping { + return mappings.at(s); + }); + + std::unordered_set sink_mapping_set = + unordered_set_of(values(mappings_for_sinks)); + + std::unordered_set required_copies = + set_difference(sink_mapping_set, std::unordered_set{src_mapping}); + + auto make_copy_to = + [&](ParallelTensorMapping const &sink_mapping) -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + { + DynamicTensorSlot{ + TensorSlotName::INPUT, + src.slot_name.slot_tensor_role, + }, + dynamic_value_attrs_with_mapping(v, src_mapping), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/transform( + src.slot_name.slot_tensor_role, + dynamic_task_type_from_tensor_role_for_copy), + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + DynamicTensorSlot{ + TensorSlotName::OUTPUT, + src.slot_name.slot_tensor_role, + }, + dynamic_value_attrs_with_mapping(v, sink_mapping), + }, + }, + }; + }; + + return transform(required_copies, make_copy_to); +} + +std::unordered_map + resolve_tensor_mappings_from_node_mappings( + DynamicOpenDataflowGraph const &g) { + + auto get_mappings_for_invocation = [&](DynamicNodeInvocation const &i) + -> std::unordered_map { + return generate_map( + get_dynamic_slot_sites_for_invocation(i), + [&](InternalDynamicSlotSite const &s) -> ParallelTensorMapping { + return ParallelTensorMapping{ + dynamic_node_mapping_bindings_for_slot_name( + assert_unwrap(i.node_attrs.mapping), + s.slot_name.slot_name), + }; + }); + }; + + std::unordered_map result = + merge_disjoint_maps(transform(get_dynamic_invocation_set(g), + get_mappings_for_invocation)); return result; } @@ -154,25 +232,26 @@ DynamicOpenDataflowGraph ASSERT(no_part_of_graph_is_copy_inserted(g)); - std::unordered_map - unmapped_value_to_mapped_source_value; - for (DynamicNodeInvocation const &i : g.invocations) { - for (auto const &[slot, value] : i.outputs) { - unmapped_value_to_mapped_source_value.insert( - std::pair{value, - map_dynamic_value_attrs_for_task_group( - slot, value, assert_unwrap(i.node_attrs.mapping))}); - } - } + std::unordered_map + fully_resolved_tensor_mappings = + resolve_tensor_mappings_from_node_mappings(g); + + std::unordered_set all_copies = + flatmap(unordered_set_of(get_dynamic_values(g)), + [&](DynamicValueAttrs const &v) + -> std::unordered_set { + return copies_for_value(g, v, fully_resolved_tensor_mappings); + }); + + std::unordered_set mapped_invocations = transform( + get_dynamic_invocation_set(g), + [&](DynamicNodeInvocation const &i) -> DynamicNodeInvocation { + return apply_mappings_for_invocation(i, fully_resolved_tensor_mappings); + }); - // Use regular flatmap here to remove duplicates (we don't want to copy the - // same tensor to the same place multiple times) DynamicOpenDataflowGraph result = dynamic_open_dataflow_graph_from_invocation_set( - flatmap(g.invocations, [&](DynamicNodeInvocation const &i) { - return perform_copy_insertion_for_invocation( - i, unmapped_value_to_mapped_source_value); - })); + set_union(all_copies, mapped_invocations)); ASSERT(graph_is_fully_copy_inserted(result)); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc new file mode 100644 index 0000000000..788c010084 --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc @@ -0,0 +1,20 @@ +#include "task-spec/dynamic_graph/dynamic_graph_edge.h" + +namespace FlexFlow { + +DynamicGraphEdge + dynamic_graph_edge_from_slot_sites(DynamicSlotSite const &src, + InternalDynamicSlotSite const &dst) { + if (src.is_internal()) { + ASSERT(src.require_internal().direction == TensorDirection::OUTPUT); + } + ASSERT(dst.direction == TensorDirection::INCOMING); + + return DynamicGraphEdge{ + /*src=*/src, + /*dst_node=*/dst.invocation, + /*dst_slot=*/dst.slot_name, + }; +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc new file mode 100644 index 0000000000..93ad9269ce --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc @@ -0,0 +1,62 @@ +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" +#include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/training_operation_attrs.h" +#include "utils/containers/are_disjoint.h" +#include "utils/containers/set_union.h" +#include "utils/containers/unordered_set_of.h" +#include "utils/optional.h" + +namespace FlexFlow { + +std::unordered_map + get_slot_map_for_direction(DynamicNodeInvocation const &invocation, + TensorDirection direction) { + switch (direction) { + case TensorDirection::INCOMING: + return invocation.inputs; + case TensorDirection::OUTPUT: + return invocation.outputs; + default: + PANIC("Unexpected direction {}", direction); + } +} + +TrainingOpType + dynamic_node_invocation_get_op_type(DynamicNodeInvocation const &i) { + TrainingOperationAttrs training_op_attrs = + assert_unwrap(i.node_attrs.op_attrs); + + return training_op_attrs_get_op_type(training_op_attrs); +} + +std::unordered_set + get_dynamic_slot_sites_for_invocation(DynamicNodeInvocation const &i) { + + std::unordered_set input_slots = + transform(unordered_set_of(i.inputs), + [&](std::pair const &p) + -> InternalDynamicSlotSite { + return InternalDynamicSlotSite{ + /*invocation=*/i, + /*direction=*/TensorDirection::INCOMING, + /*slot_name=*/p.first, + }; + }); + + std::unordered_set output_slots = + transform(unordered_set_of(i.outputs), + [&](std::pair const &p) + -> InternalDynamicSlotSite { + return InternalDynamicSlotSite{ + /*invocation=*/i, + /*direction=*/TensorDirection::OUTPUT, + /*slot_name=*/p.first, + }; + }); + + ASSERT(are_disjoint(input_slots, output_slots)); + + return set_union(input_slots, output_slots); +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc new file mode 100644 index 0000000000..9e56dddcf9 --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc @@ -0,0 +1,31 @@ +#include "task-spec/dynamic_graph/dynamic_node_mapping.h" +#include "utils/bidict/algorithms/transform_values.h" +#include "utils/containers/transform.h" + +namespace FlexFlow { + +bidict + dynamic_node_mapping_bindings_for_slot_name( + DynamicNodeMapping const &mapping, TensorSlotName const &slot_name) { + bidict coord_bindings = + get_tensor_bindings_for_slot_name(mapping.op_task_group, slot_name); + + return transform_values( + coord_bindings, [&](MachineSpaceCoordinate const &coord) -> device_id_t { + return device_id_t{coord, mapping.device_type}; + }); +} + +std::unordered_set + target_devices_of_dynamic_node_mapping(DynamicNodeMapping const &mapping) { + + return transform(mapping.op_task_group.get_shard_bindings().left_values(), + [&](MachineSpaceCoordinate const &c) -> device_id_t { + return device_id_t{ + /*coord=*/c, + /*device_type=*/mapping.device_type, + }; + }); +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index d2a5b653e5..07c49664a5 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -1,19 +1,29 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "op-attrs/ff_ordered/filtrans.h" +#include "op-attrs/parallel_tensor_shape.h" +#include "op-attrs/pcg_operator_attrs.h" +#include "task-spec/dynamic_graph/dynamic_graph_edge.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.h" #include "task-spec/dynamic_graph/serializable_dynamic_node_attrs.h" #include "task-spec/dynamic_graph/serializable_dynamic_value_attrs.h" #include "utils/containers/all_of.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/contains_duplicates.h" +#include "utils/containers/contains_value.h" +#include "utils/containers/filter_values.h" #include "utils/containers/flatmap.h" +#include "utils/containers/get_only.h" #include "utils/containers/multiset_union.h" +#include "utils/containers/transform.h" #include "utils/containers/zip_strict.h" #include "utils/containers/zip_values_strict.h" #include "utils/graph/dataflow_graph/algorithms.h" #include "utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h" #include "utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h" +#include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/find_isomorphism.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.h" -#include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_dataflow_graph/algorithms/get_inputs.h" @@ -28,6 +38,30 @@ DynamicOpenDataflowGraph make_empty_dynamic_open_dataflow_graph() { }; } +void check_dynamic_open_dataflow_graph_is_valid( + DynamicOpenDataflowGraph const &g) { + std::unordered_map> + invocations_by_value_produced; + + for (DynamicNodeInvocation const &i : g.invocations) { + for (DynamicValueAttrs const &v : values(i.outputs)) { + invocations_by_value_produced[v].push_back(i); + } + } + + std::unordered_map> + values_produced_multiple_times = filter_values( + invocations_by_value_produced, + [](std::vector const &producers) -> bool { + return producers.size() > 1; + }); + + ASSERT(values_produced_multiple_times.size() == 0, + keys(values_produced_multiple_times)); + + labelled_open_kwarg_dataflow_graph_from_dynamic_open_dataflow_graph(g); +} + nonnegative_int dynamic_graph_num_nodes(DynamicOpenDataflowGraph const &g) { return num_elements(get_dynamic_nodes(g)); } @@ -88,6 +122,140 @@ std::unordered_set return g.invocations; } +std::unordered_set + get_dynamic_graph_edges(DynamicOpenDataflowGraph const &g) { + return flatmap(get_dynamic_invocation_set(g), + [&](DynamicNodeInvocation const &i) + -> std::unordered_set { + return get_dynamic_graph_edges_incoming_to_invocation(g, i); + }); +} + +std::unordered_set + get_dynamic_graph_edges_incoming_to_invocation( + DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { + return transform(unordered_set_of(i.inputs), + [&](std::pair const &p) + -> DynamicGraphEdge { + DynamicSlotSite src = + dynamic_graph_find_source_of_value(g, p.second); + + InternalDynamicSlotSite dst = InternalDynamicSlotSite{ + /*invocation=*/i, + /*direction=*/TensorDirection::INCOMING, + /*slot_name=*/p.first, + }; + + return dynamic_graph_edge_from_slot_sites(src, dst); + }); +} + +std::unordered_set + get_dynamic_graph_edges_outgoing_from_invocation( + DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { + return flatmap( + unordered_set_of(i.outputs), + [&](std::pair const &p) + -> std::unordered_set { + DynamicSlotSite src = DynamicSlotSite{ + InternalDynamicSlotSite{ + /*invocation=*/i, + /*direction=*/TensorDirection::OUTPUT, + /*slot_name=*/p.first, + }, + }; + + return transform( + dynamic_graph_find_sinks_of_value(g, p.second), + [&](InternalDynamicSlotSite const &sink) -> DynamicGraphEdge { + return dynamic_graph_edge_from_slot_sites(src, sink); + }); + }); +} + +std::unordered_set + get_internal_dynamic_slot_sites(DynamicOpenDataflowGraph const &g) { + return flatmap(get_dynamic_invocation_set(g), + [](DynamicNodeInvocation const &i) + -> std::unordered_set { + return get_dynamic_slot_sites_for_invocation(i); + }); +} + +std::unordered_set + get_dynamic_slot_sites(DynamicOpenDataflowGraph const &g) { + std::unordered_set internal_slot_sites = + get_internal_dynamic_slot_sites(g); + + std::unordered_set internal_values = + filtrans(internal_slot_sites, + [&](InternalDynamicSlotSite const &s) + -> std::optional { + if (s.direction == TensorDirection::OUTPUT) { + return dynamic_value_attrs_for_slot_site(DynamicSlotSite{s}); + } else { + return std::nullopt; + } + }); + + std::unordered_set all_values = + unordered_set_of(get_dynamic_values(g)); + + std::unordered_set external_values = + set_minus(all_values, internal_values); + + std::unordered_set external_slot_sites = transform( + external_values, + [](DynamicValueAttrs const &external_value) -> ExternalDynamicSlotSite { + return ExternalDynamicSlotSite{external_value}; + }); + + return set_union( + transform(internal_slot_sites, + [](InternalDynamicSlotSite const &s) -> DynamicSlotSite { + return DynamicSlotSite{s}; + }), + transform(external_slot_sites, + [](ExternalDynamicSlotSite const &s) -> DynamicSlotSite { + return DynamicSlotSite{s}; + })); +} + +std::unordered_set + dynamic_graph_find_sinks_of_value(DynamicOpenDataflowGraph const &g, + DynamicValueAttrs const &v) { + std::unordered_set found = filter( + get_internal_dynamic_slot_sites(g), + [&](InternalDynamicSlotSite const &s) -> bool { + return dynamic_value_attrs_for_slot_site(DynamicSlotSite{s}) == v && + s.direction == TensorDirection::INCOMING; + }); + + return found; +} + +DynamicSlotSite + dynamic_graph_find_source_of_value(DynamicOpenDataflowGraph const &g, + DynamicValueAttrs const &v) { + + auto is_source_of_value = [&](DynamicSlotSite const &s) -> bool { + return s.visit(overload{ + [&](InternalDynamicSlotSite const &internal_slot_site) -> bool { + return dynamic_value_attrs_for_slot_site(s) == v && + internal_slot_site.direction == TensorDirection::OUTPUT; + }, + [&](ExternalDynamicSlotSite const &external_slot_site) -> bool { + return external_slot_site.value == v; + }, + }); + }; + + std::unordered_set found = + filter(get_dynamic_slot_sites(g), is_source_of_value); + + return get_only(found); +} + std::optional find_output_value_attrs(DynamicOpenDataflowGraph const &dg, dynamic_tensor_guid_t tensor_guid, @@ -133,9 +301,13 @@ DynamicOpenDataflowGraph flatmap_dynamic_invocation_set( DynamicOpenDataflowGraph dynamic_open_dataflow_graph_from_invocation_set( std::unordered_set const &invocation_set) { - return DynamicOpenDataflowGraph{ + DynamicOpenDataflowGraph result = DynamicOpenDataflowGraph{ invocation_set, }; + + check_dynamic_open_dataflow_graph_is_valid(result); + + return result; } std::pair RecordFormatter { + return mk_record_for_map(mapping.raw.as_unordered_map()); + }; + std::function render_value_label = [&](DynamicValueAttrs const &a) -> nlohmann::json { nlohmann::json result = dynamic_value_attrs_to_serializable(a); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc new file mode 100644 index 0000000000..af13345abd --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc @@ -0,0 +1,26 @@ +#include "task-spec/dynamic_graph/dynamic_slot_site.h" +#include "utils/overload.h" + +namespace FlexFlow { + +DynamicValueAttrs + dynamic_value_attrs_for_slot_site(DynamicSlotSite const &slot) { + return slot.visit(overload{ + + [](ExternalDynamicSlotSite const &external_slot) -> DynamicValueAttrs { + return external_slot.value; + }, + + [](InternalDynamicSlotSite const &internal_slot) -> DynamicValueAttrs { + switch (internal_slot.direction) { + case TensorDirection::INCOMING: + return internal_slot.invocation.inputs.at(internal_slot.slot_name); + case TensorDirection::OUTPUT: + return internal_slot.invocation.outputs.at(internal_slot.slot_name); + default: + PANIC("Unexpected direction {}", internal_slot.direction); + } + }}); +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc index 282279edbe..a05330b1b8 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc @@ -13,4 +13,13 @@ DynamicValueAttrs return result; } +DynamicValueAttrs + dynamic_value_attrs_with_mapping(DynamicValueAttrs const &v, + ParallelTensorMapping const &m) { + ASSERT(v.mapping == std::nullopt); + DynamicValueAttrs result = v; + result.mapping = m; + return result; +} + } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc index 8066926262..9bcdc351fc 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc @@ -17,7 +17,7 @@ LossInsertionResult perform_loss_insertion( DynamicOpenDataflowGraph const &dg, LossAttrs const &loss_attrs, dynamic_tensor_guid_t logit_tensor, - std::optional const &loss_mapping) { + std::optional const &loss_mapping) { DynamicValueAttrs logit_value = assert_unwrap( find_output_value_attrs(dg, logit_tensor, mk_dynamic_tensor_role_fwd())); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc index 0a22015ddf..e46f8b1510 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc @@ -6,7 +6,7 @@ namespace FlexFlow { std::unordered_set perform_machine_slicing_for_invocation( DynamicNodeInvocation const &invocation, - MachineSpaceCoordinate const &device_coord) { + device_id_t const &device_coord) { ASSERT(invocation.node_attrs.device_coord.has_value()); @@ -19,7 +19,7 @@ std::unordered_set DynamicOpenDataflowGraph perform_machine_slicing(DynamicOpenDataflowGraph const &g, - MachineSpaceCoordinate const &device_coord) { + device_id_t const &device_coord) { DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [&](DynamicNodeInvocation const &invocation) diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 380c2d17a1..57efdd150b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -15,7 +15,7 @@ namespace FlexFlow { DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( - MappedParallelComputationGraph const &mpcg) { + MappedParallelComputationGraph const &mpcg, DeviceType device_type) { DynamicOpenDataflowGraph result = make_empty_dynamic_open_dataflow_graph(); ParallelComputationGraph pcg = pcg_from_mpcg(mpcg); @@ -24,7 +24,11 @@ DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( DynamicNodeAttrs result_attrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, - /*mapping=*/mpcg_get_mapping_for_layer(mpcg, layer), + /*mapping=*/ + DynamicNodeMapping{ + /*op_task_group=*/mpcg_get_mapping_for_layer(mpcg, layer), + /*device_type=*/device_type, + }, /*op_attrs=*/TrainingOperationAttrs{attrs.op_attrs}, /*pcg_layer_guid=*/dynamic_layer_guid_t{layer}, /*per_device_op_state=*/std::nullopt, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index fb6efb96d0..117819a639 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -1,6 +1,8 @@ #include "task-spec/dynamic_graph/shard_expansion.h" +#include "task-spec/dynamic_graph/dynamic_node_mapping.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" +#include "task-spec/dynamic_graph/shard_expansion.h" #include "utils/bidict/algorithms/filter_keys.h" #include "utils/containers/get_only.h" #include "utils/containers/map_values2.h" @@ -40,10 +42,9 @@ bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { slot_is_shard_expanded); } -static bidict +static bidict restrict_tensor_mapping_keys_to_coord( - bidict const - &mapping, + bidict const &mapping, ParallelTensorSpaceCoordinate const ¶llel_tensor_coord) { return filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { return p == parallel_tensor_coord; @@ -52,7 +53,7 @@ static bidict static DynamicNodeInvocation shard_invocation_for_binding( DynamicNodeInvocation const &i, - MachineSpaceCoordinate const &machine_coord, + device_id_t const &device_coord, OperatorAtomicTaskShardBinding const &binding) { auto shard_expand_value_attrs = [&](DynamicTensorSlot const &s, @@ -64,17 +65,18 @@ static DynamicNodeInvocation shard_invocation_for_binding( result.shard_coord = parallel_tensor_coord; result.mapping = transform( v.mapping, - [&](bidict const - &mapping) { - return restrict_tensor_mapping_keys_to_coord(mapping, - parallel_tensor_coord); + [&](ParallelTensorMapping const &mapping) -> ParallelTensorMapping { + return ParallelTensorMapping{ + restrict_tensor_mapping_keys_to_coord(mapping.raw, + parallel_tensor_coord), + }; }); return result; }; DynamicNodeAttrs expanded_node_attrs = [&]() { DynamicNodeAttrs result = i.node_attrs; - result.device_coord = machine_coord; + result.device_coord = device_coord; return result; }(); @@ -89,10 +91,10 @@ static std::unordered_set perform_shard_expansion_for_copy(DynamicNodeInvocation const &i) { auto [input_slot, input] = get_only(i.inputs); auto [output_slot, output] = get_only(i.outputs); - bidict input_mapping = - assert_unwrap(input.mapping); + bidict input_mapping = + assert_unwrap(input.mapping).raw; require_same(input_mapping.left_values(), - assert_unwrap(output.mapping).left_values()); + assert_unwrap(output.mapping).raw.left_values()); return transform( input_mapping.left_values(), [&](ParallelTensorSpaceCoordinate const &p) { @@ -103,7 +105,7 @@ static std::unordered_set // because we expect this to align with the most efficient way to issue // copies in Realm, although the current Realm backend uses a // centralized controller and thus issues copies all from a single node. - MachineSpaceCoordinate machine_coord = input_mapping.at_l(p); + device_id_t machine_coord = input_mapping.at_l(p); return shard_invocation_for_binding(i, machine_coord, @@ -121,16 +123,15 @@ std::unordered_set return perform_shard_expansion_for_copy(i); } - MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); + DynamicNodeMapping mapping = assert_unwrap(i.node_attrs.mapping); - std::unordered_set shard_machine_coords = - mapping.get_shard_bindings().left_values(); + std::unordered_set shard_machine_coords = + target_devices_of_dynamic_node_mapping(mapping); return transform( - shard_machine_coords, - [&](MachineSpaceCoordinate const &c) -> DynamicNodeInvocation { + shard_machine_coords, [&](device_id_t const &c) -> DynamicNodeInvocation { OperatorAtomicTaskShardBinding slot_bindings = - mapping.get_shard_bindings().at_l(c); + mapping.op_task_group.get_shard_bindings().at_l(c.coord); return shard_invocation_for_binding(i, c, slot_bindings); }); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc new file mode 100644 index 0000000000..90db73e617 --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/training_operation_attrs.cc @@ -0,0 +1,24 @@ +#include "task-spec/dynamic_graph/training_operation_attrs.h" +#include "op-attrs/pcg_operator_attrs.h" +#include "utils/overload.h" + +namespace FlexFlow { + +TrainingOpType training_op_attrs_get_op_type( + TrainingOperationAttrs const &training_op_attrs) { + return training_op_attrs.visit(overload{ + [](PCGOperatorAttrs const &a) -> TrainingOpType { + return TrainingOpType{ + pcg_op_attrs_get_op_type(a), + }; + }, + [](LossAttrs const &) -> TrainingOpType { + return TrainingOpType{TrainingOnlyOpType::LOSS}; + }, + [](CopyAttrs const &) -> TrainingOpType { + return TrainingOpType{TrainingOnlyOpType::COPY}; + }, + }); +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc index 58a32db6c1..0c3769512b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc @@ -58,8 +58,8 @@ static DynamicNodeInvocation get_update_invocation_for_invocation( auto create_binding_for_role = [&](DynamicTensorRole const &role) -> std::pair { DynamicTensorSlot binding_slot = tensor_slot_with_role(slot, role); - DynamicValueAttrs binding_attrs = dynamic_value_attrs_with_role( - value_attrs, mk_dynamic_tensor_role_fwd()); + DynamicValueAttrs binding_attrs = + dynamic_value_attrs_with_role(value_attrs, role); return std::pair{ binding_slot, diff --git a/lib/task-spec/src/task-spec/serialization.cc b/lib/task-spec/src/task-spec/serialization.cc deleted file mode 100644 index a2ad6eabfa..0000000000 --- a/lib/task-spec/src/task-spec/serialization.cc +++ /dev/null @@ -1 +0,0 @@ -#include "task-spec/serialization.h" diff --git a/lib/task-spec/test/CMakeLists.txt b/lib/task-spec/test/CMakeLists.txt index 9665dba88e..d417280c26 100644 --- a/lib/task-spec/test/CMakeLists.txt +++ b/lib/task-spec/test/CMakeLists.txt @@ -8,7 +8,6 @@ ff_add_test_executable( DEPS doctest utils-test-common - # local-execution kernels task-spec op-attrs diff --git a/lib/task-spec/test/src/task-spec/device_specific.cc b/lib/task-spec/test/src/task-spec/device_specific.cc index 34ef9b2bef..6a42a9b570 100644 --- a/lib/task-spec/test/src/task-spec/device_specific.cc +++ b/lib/task-spec/test/src/task-spec/device_specific.cc @@ -6,12 +6,14 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("DeviceSpecific") { DeviceSpecific device_specific1 = - DeviceSpecific::create(device_id_t{gpu_id_t{0_n}}, - "hello world"); + DeviceSpecific::create( + device_id_t{MachineSpaceCoordinate{0_n, 0_n}, DeviceType::GPU}, + "hello world"); DeviceSpecific device_specific2 = - DeviceSpecific::create(device_id_t{gpu_id_t{1_n}}, - "hello world"); + DeviceSpecific::create( + device_id_t{MachineSpaceCoordinate{0_n, 1_n}, DeviceType::GPU}, + "hello world"); std::string result1 = fmt::to_string(device_specific1); std::string result2 = fmt::to_string(device_specific2); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index 2160f6bf82..58a8a36fcd 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -1,433 +1,478 @@ #include "task-spec/dynamic_graph/copy_insertion.h" #include "op-attrs/tensor_slot_name.dtg.h" #include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" +#include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_task_type.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" +#include "test/utils/doctest/check_kv.h" #include "test/utils/doctest/fmt/unordered_set.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("perform_copy_insertion_for_invocation") { + TEST_CASE("perform_copy_insertion") { + + auto mk_slot = [](TensorSlotName tensor_slot_name) -> DynamicTensorSlot { + return DynamicTensorSlot{ + tensor_slot_name, + std::nullopt, + }; + }; + + auto mk_value_attrs = + [](size_t src_layer_guid, + TensorSlotName src_slot, + std::optional const &mapping) + -> DynamicValueAttrs { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{ + src_layer_guid, + }, + src_slot, + }, + }, + }, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/mapping, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + }; + + auto mk_ptensor_coord = + [](nonnegative_int shard_idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_components=*/ + FFOrdered{ + shard_idx, + }, + }; + }; + auto mk_machine_coord = - [](nonnegative_int node_idx, - nonnegative_int device_idx) -> MachineSpaceCoordinate { + [](nonnegative_int device_idx) -> MachineSpaceCoordinate { return MachineSpaceCoordinate{ - /*node_idx=*/node_idx, + /*node_idx=*/0_n, /*device_idx=*/device_idx, - /*device_type=*/DeviceType::GPU, }; }; - auto mk_pt_coord = - [](nonnegative_int idx1, - nonnegative_int idx2, - nonnegative_int idx3, - nonnegative_int idx4) -> ParallelTensorSpaceCoordinate { - return ParallelTensorSpaceCoordinate{ - /*sum_component=*/idx1, - /*discard_copy_component=*/idx2, - /*shard_components=*/ - FFOrdered{ - idx3, - idx4, + auto mk_device_id = [&](nonnegative_int device_idx) -> device_id_t { + return device_id_t{ + mk_machine_coord(device_idx), + DeviceType::GPU, + }; + }; + + auto mk_pcg_layer_guid = [](size_t pcg_layer_guid) -> dynamic_layer_guid_t { + return dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{pcg_layer_guid}, }, }; }; - auto mk_input_shard_binding = [&](ParallelTensorSpaceCoordinate const &c) + auto mk_node_attrs = + [](dynamic_layer_guid_t layer_guid, + std::optional const &mapping, + std::optional const &op_attrs) + -> DynamicNodeAttrs { + return DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }; + }; + + auto mk_binding = [&](nonnegative_int input_shard_idx, + nonnegative_int output_shard_idx) -> OperatorAtomicTaskShardBinding { return OperatorAtomicTaskShardBinding{ /*tensor_coords=*/{ + { + TensorSlotName::INPUT, + mk_ptensor_coord(input_shard_idx), + }, { TensorSlotName::OUTPUT, - c, + mk_ptensor_coord(output_shard_idx), }, }, }; }; - auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, - ParallelTensorSpaceCoordinate const &c2, - ParallelTensorSpaceCoordinate const &c3, - ParallelTensorSpaceCoordinate const &c4) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ + DynamicValueAttrs v1 = mk_value_attrs( + /*src_layer_guid=*/0, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/std::nullopt); + + DynamicValueAttrs v2 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/std::nullopt); + + DynamicValueAttrs v3 = mk_value_attrs( + /*src_layer_guid=*/2, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/std::nullopt); + + SUBCASE("inserts copy when necessary") { + DynamicNodeMapping mapping1 = DynamicNodeMapping{ + MappedOperatorTaskGroup{ + bidict{ + { + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(1_n), + mk_binding(1_n, 1_n), + }, + }, + }, + DeviceType::GPU, + }; + + DynamicNodeMapping mapping2 = DynamicNodeMapping{ + MappedOperatorTaskGroup{ + bidict{ + { + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(2_n), + mk_binding(1_n, 1_n), + }, + }, + }, + DeviceType::GPU, + }; + + DynamicNodeInvocation inv1 = DynamicNodeInvocation{ + /*inputs=*/{ { - TensorSlotName::INPUT, - c1, + mk_slot(TensorSlotName::INPUT), + v1, }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), + /*outputs=*/ + { { - TensorSlotName::WEIGHT, - c2, + mk_slot(TensorSlotName::OUTPUT), + v2, }, + }, + }; + + DynamicNodeInvocation inv2 = DynamicNodeInvocation{ + /*inputs=*/{ { - TensorSlotName::OUTPUT_1, - c3, + mk_slot(TensorSlotName::INPUT), + v2, }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), + /*outputs=*/ + { { - TensorSlotName::OUTPUT_2, - c4, + mk_slot(TensorSlotName::OUTPUT), + v3, }, }, }; - }; - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); - MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); - MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); - - ParallelTensorSpaceCoordinate mc1_input_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc1_weight_coord = - mk_pt_coord(0_n, 1_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc1_output_1_coord = - mk_pt_coord(1_n, 0_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc1_output_2_coord = - mk_pt_coord(3_n, 0_n, 0_n, 0_n); - - ParallelTensorSpaceCoordinate mc2_input_coord = - mk_pt_coord(0_n, 1_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc2_weight_coord = - mk_pt_coord(0_n, 4_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc2_output_1_coord = - mk_pt_coord(1_n, 2_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc2_output_2_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - - MappedOperatorTaskGroup input_mapping_same = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_input_coord), - }, - { - mc2, - mk_input_shard_binding(mc2_input_coord), - }, - }, - }; + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set({inv1, inv2}); - MappedOperatorTaskGroup weight_mapping_same = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_weight_coord), - }, - { - mc2, - mk_input_shard_binding(mc2_weight_coord), - }, - }, - }; + DynamicOpenDataflowGraph result = perform_copy_insertion(g); - MappedOperatorTaskGroup invocation_mapping = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_shard_binding(mc1_input_coord, - mc1_weight_coord, - mc1_output_1_coord, - mc1_output_2_coord), + DynamicOpenDataflowGraph correct = [&] { + DynamicValueAttrs mapped_v1 = mk_value_attrs( + /*src_layer_guid=*/0, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); + + DynamicValueAttrs mapped_v2_placement1 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); + + DynamicValueAttrs mapped_v2_placement2 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(2_n)}, + }, + }); + + DynamicValueAttrs mapped_v3 = mk_value_attrs( + /*src_layer_guid=*/2, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(2_n)}, + }, + }); + + DynamicNodeInvocation mapped_inv1 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v1, + }, }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), + /*outputs=*/ { - mc2, - mk_shard_binding(mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), - }, - }, - }; - - MappedOperatorTaskGroup invocation_mapping_diff_vs_copy1 = - MappedOperatorTaskGroup{ - bidict{ { - mc2, - mk_shard_binding(mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), + mk_slot(TensorSlotName::OUTPUT), + mapped_v2_placement1, }, }, }; - auto mk_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { - return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - }; - }; - auto mk_value = [&](size_t src_node_id, - TensorSlotName src_slot_name, - MappedOperatorTaskGroup const &mapping, - std::optional const &use_slot_name) - -> DynamicValueAttrs { - return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{src_node_id}, - src_slot_name, - }, - }}, - /*parallel_tensor_shape=*/std::nullopt, - /*shard_coord=*/std::nullopt, - /*mapping=*/ - transform(use_slot_name, - [&](TensorSlotName s) { - return get_tensor_bindings_for_slot_name(mapping, s); - }), - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }; - }; - - size_t invocation1_id = 20; - - DynamicValueAttrs graph_input1 = - mk_value(0, TensorSlotName::OUTPUT, invocation_mapping, std::nullopt); - DynamicValueAttrs graph_input1_use = mk_value( - 0, TensorSlotName::OUTPUT, invocation_mapping, TensorSlotName::INPUT); - DynamicValueAttrs graph_input1_use_diff_vs_copy1 = - mk_value(0, - TensorSlotName::OUTPUT, - invocation_mapping_diff_vs_copy1, - TensorSlotName::INPUT); - DynamicValueAttrs graph_input2 = - mk_value(1, TensorSlotName::OUTPUT, invocation_mapping, std::nullopt); - DynamicValueAttrs graph_input2_use = mk_value( - 1, TensorSlotName::OUTPUT, invocation_mapping, TensorSlotName::WEIGHT); - DynamicValueAttrs invocation1_output1 = mk_value(invocation1_id, - TensorSlotName::OUTPUT_1, - invocation_mapping, - std::nullopt); - DynamicValueAttrs invocation1_output1_src = - mk_value(invocation1_id, - TensorSlotName::OUTPUT_1, - invocation_mapping, - TensorSlotName::OUTPUT_1); - DynamicValueAttrs invocation1_output2 = mk_value(invocation1_id, - TensorSlotName::OUTPUT_2, - invocation_mapping, - std::nullopt); - DynamicValueAttrs invocation1_output2_src = - mk_value(invocation1_id, - TensorSlotName::OUTPUT_2, - invocation_mapping, - TensorSlotName::OUTPUT_2); - - DynamicValueAttrs graph_input1_src_same = mk_value( - 0, TensorSlotName::OUTPUT, input_mapping_same, TensorSlotName::OUTPUT); - DynamicValueAttrs graph_input2_src_same = mk_value( - 1, TensorSlotName::OUTPUT, weight_mapping_same, TensorSlotName::OUTPUT); - - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - graph_input1, - }, - { - mk_slot(TensorSlotName::WEIGHT), - graph_input2, - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/invocation_mapping, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/ - dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT_1), - invocation1_output1, + DynamicNodeInvocation inserted_copy = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v2_placement1, + }, }, + /*node_attrs=*/ + mk_node_attrs(dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT_2), - invocation1_output2, + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v2_placement2, + }, }, - }, - }; - DynamicNodeInvocation mapped = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - graph_input1_use, - }, - { - mk_slot(TensorSlotName::WEIGHT), - graph_input2_use, - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/invocation_mapping, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/ - dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT_1), - invocation1_output1_src, + }; + + DynamicNodeInvocation mapped_inv2 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v2_placement2, + }, }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT_2), - invocation1_output2_src, + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v3, + }, }, - }, - }; - - auto mk_copy = [&](DynamicValueAttrs const &src, - DynamicValueAttrs const &dst) { - return DynamicNodeInvocation{ - /*inputs=*/{{mk_slot(TensorSlotName::INPUT), src}}, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{{mk_slot(TensorSlotName::OUTPUT), dst}}, - }; - }; - - SUBCASE("same mapping, no copies") { - std::unordered_map sources_same{ - {graph_input1, graph_input1_src_same}, - {graph_input2, graph_input2_src_same}}; - - std::unordered_set result = - perform_copy_insertion_for_invocation(input, sources_same); + }; - std::unordered_set correct = {mapped}; + return dynamic_open_dataflow_graph_from_invocation_set( + {mapped_inv1, mapped_inv2, inserted_copy}); + }(); - CHECK(result.size() == correct.size()); - CHECK(result == correct); + CHECK_MESSAGE( + result == correct, + check_kv("result\n", dynamic_open_dataflow_graph_as_dot(result)), + check_kv("correct\n", dynamic_open_dataflow_graph_as_dot(correct))); } - SUBCASE("copy one tensor, one point") { - MappedOperatorTaskGroup input_mapping_copy1 = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_input_coord), - }, - { - mc3, - mk_input_shard_binding(mc2_input_coord), + SUBCASE("does not insert a copy when not necessary") { + DynamicNodeMapping mapping1 = DynamicNodeMapping{ + MappedOperatorTaskGroup{ + bidict{ + { + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(1_n), + mk_binding(1_n, 1_n), + }, }, }, + DeviceType::GPU, }; - MappedOperatorTaskGroup input_mapping_copy1_diff_vs_use = + + DynamicNodeMapping mapping2 = DynamicNodeMapping{ MappedOperatorTaskGroup{ bidict{ { - mc3, - mk_input_shard_binding(mc2_input_coord), + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(1_n), + mk_binding(1_n, 1_n), }, }, - }; - - DynamicValueAttrs graph_input1_src_copy1 = - mk_value(0, - TensorSlotName::OUTPUT, - input_mapping_copy1, - TensorSlotName::OUTPUT); - DynamicValueAttrs graph_input1_src_copy1_diff_vs_use = - mk_value(0, - TensorSlotName::OUTPUT, - input_mapping_copy1_diff_vs_use, - TensorSlotName::OUTPUT); - - std::unordered_map sources_copy1{ - {graph_input1, graph_input1_src_copy1}, - {graph_input2, graph_input2_src_same}}; - - std::unordered_set result = - perform_copy_insertion_for_invocation(input, sources_copy1); - - std::unordered_set correct = { - mapped, - mk_copy(graph_input1_src_copy1_diff_vs_use, - graph_input1_use_diff_vs_copy1), + }, + DeviceType::GPU, }; - CHECK(result.size() == correct.size()); - CHECK(result == correct); - } - - SUBCASE("copy two tensors, two points") { - MappedOperatorTaskGroup input_mapping_copy2 = MappedOperatorTaskGroup{ - bidict{ + DynamicNodeInvocation inv1 = DynamicNodeInvocation{ + /*inputs=*/{ { - mc3, - mk_input_shard_binding(mc1_input_coord), + mk_slot(TensorSlotName::INPUT), + v1, }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), + /*outputs=*/ + { { - mc4, - mk_input_shard_binding(mc2_input_coord), + mk_slot(TensorSlotName::OUTPUT), + v2, }, }, }; - MappedOperatorTaskGroup weight_mapping_copy2 = MappedOperatorTaskGroup{ - bidict{ + + DynamicNodeInvocation inv2 = DynamicNodeInvocation{ + /*inputs=*/{ { - mc4, - mk_input_shard_binding(mc1_weight_coord), + mk_slot(TensorSlotName::INPUT), + v2, }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), + /*outputs=*/ + { { - mc3, - mk_input_shard_binding(mc2_weight_coord), + mk_slot(TensorSlotName::OUTPUT), + v3, }, }, }; - DynamicValueAttrs graph_input1_src_copy2 = - mk_value(0, - TensorSlotName::OUTPUT, - input_mapping_copy2, - TensorSlotName::OUTPUT); - DynamicValueAttrs graph_input2_src_copy2 = - mk_value(1, - TensorSlotName::OUTPUT, - weight_mapping_copy2, - TensorSlotName::OUTPUT); - - std::unordered_map sources_copy2{ - {graph_input1, graph_input1_src_copy2}, - {graph_input2, graph_input2_src_copy2}}; - - std::unordered_set result = - perform_copy_insertion_for_invocation(input, sources_copy2); - - std::unordered_set correct = { - mapped, - mk_copy(graph_input1_src_copy2, graph_input1_use), - mk_copy(graph_input2_src_copy2, graph_input2_use), - }; + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set({inv1, inv2}); + + DynamicOpenDataflowGraph result = perform_copy_insertion(g); + + DynamicOpenDataflowGraph correct = [&] { + DynamicValueAttrs mapped_v1 = mk_value_attrs( + /*src_layer_guid=*/0, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); + + DynamicValueAttrs mapped_v2 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); + + DynamicValueAttrs mapped_v3 = mk_value_attrs( + /*src_layer_guid=*/2, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); + + DynamicNodeInvocation mapped_inv1 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v1, + }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), + /*outputs=*/ + { + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v2, + }, + }, + }; + + DynamicNodeInvocation mapped_inv2 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v2, + }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), + /*outputs=*/ + { + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v3, + }, + }, + }; + + return dynamic_open_dataflow_graph_from_invocation_set( + {mapped_inv1, mapped_inv2}); + }(); - CHECK(result.size() == correct.size()); - CHECK(result == correct); + CHECK_MESSAGE( + result == correct, + check_kv("result\n", dynamic_open_dataflow_graph_as_dot(result)), + check_kv("correct\n", dynamic_open_dataflow_graph_as_dot(correct))); } } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index 49b8d4a77a..9e529bba04 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -1,4 +1,8 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "op-attrs/initializer_attrs.h" +#include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include "task-spec/dynamic_graph/serializable_dynamic_value_attrs.h" +#include "utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h" #include "utils/graph/node/algorithms.h" #include @@ -6,48 +10,212 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("dynamic_op_dataflow_graph_from_invocation_set") { - DynamicValueAttrs value_1 = DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{1}, - TensorSlotName::OUTPUT, - }, - }}, - /*parallel_tensor_shape=*/std::nullopt, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*tensor_type=*/std::nullopt, + + auto mk_dynamic_value = [](size_t node_id, + TensorSlotName slot_name) -> DynamicValueAttrs { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{node_id}, + slot_name, + }, + }}, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*tensor_type=*/std::nullopt, + }; }; - DynamicValueAttrs value_2 = DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{2}, - TensorSlotName::OUTPUT, - }, - }}, - /*parallel_tensor_shape=*/std::nullopt, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*tensor_type=*/std::nullopt, + auto mk_slot = [](TensorSlotName slot_name) { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + }; }; - DynamicValueAttrs value_3 = DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{3}, - TensorSlotName::OUTPUT, - }, - }}, - /*parallel_tensor_shape=*/std::nullopt, - /*shard_coord=*/std::nullopt, + DynamicValueAttrs value_1 = mk_dynamic_value(1, TensorSlotName::OUTPUT); + DynamicValueAttrs value_2 = mk_dynamic_value(2, TensorSlotName::OUTPUT); + DynamicValueAttrs value_3 = mk_dynamic_value(3, TensorSlotName::OUTPUT); + + DynamicNodeAttrs node_attrs = DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*tensor_type=*/std::nullopt, + /*op_attrs=*/std::nullopt, + /*layer_guid=*/dynamic_layer_guid_t{parallel_layer_guid_t{Node{4}}}, + /*per_device_op_state=*/std::nullopt, }; + SUBCASE("correct usage") { + DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ + /*inputs=*/std::unordered_map{ + { + mk_slot(TensorSlotName::INPUT), + value_1, + }, + }, + /*node_attrs=*/node_attrs, + /*outputs=*/ + std::unordered_map{ + { + mk_slot(TensorSlotName::OUTPUT), + value_2, + }, + }, + }; + + DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ + /*inputs=*/std::unordered_map{}, + /*node_attrs=*/node_attrs, + /*outputs=*/ + std::unordered_map{ + { + mk_slot(TensorSlotName::OUTPUT), + value_3, + }, + }, + }; + + DynamicNodeInvocation invocation_3 = DynamicNodeInvocation{ + /*inputs=*/std::unordered_map{ + { + mk_slot(TensorSlotName::INPUT), + value_1, + }, + { + mk_slot(TensorSlotName::WEIGHT), + value_2, + }, + { + mk_slot(TensorSlotName::BIAS), + value_1, + }, + }, + /*node_attrs=*/node_attrs, + /*outputs=*/ + std::unordered_map{}, + }; + + std::unordered_set invocation_set = { + invocation_1, + invocation_2, + invocation_3, + }; + + DynamicOpenDataflowGraph result = + dynamic_open_dataflow_graph_from_invocation_set(invocation_set); + + CHECK(dynamic_graph_num_nodes(result) == 3); + } + + SUBCASE("throws if multiple invocations produce the same value") { + DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ + /*inputs=*/std::unordered_map{ + { + mk_slot(TensorSlotName::INPUT), + value_1, + }, + }, + /*node_attrs=*/node_attrs, + /*outputs=*/ + std::unordered_map{ + { + mk_slot(TensorSlotName::OUTPUT), + value_2, + }, + }, + }; + + DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ + /*inputs=*/std::unordered_map{}, + /*node_attrs=*/node_attrs, + /*outputs=*/ + std::unordered_map{ + { + mk_slot(TensorSlotName::OUTPUT), + value_2, + }, + }, + }; + + std::unordered_set invocation_set = { + invocation_1, + invocation_2, + }; + + CHECK_THROWS( + dynamic_open_dataflow_graph_from_invocation_set(invocation_set)); + } + + SUBCASE("throws if invocations contain/create cycle") { + DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ + /*inputs=*/std::unordered_map{ + { + mk_slot(TensorSlotName::INPUT), + value_1, + }, + }, + /*node_attrs=*/node_attrs, + /*outputs=*/ + std::unordered_map{ + { + mk_slot(TensorSlotName::OUTPUT), + value_2, + }, + }, + }; + + DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ + /*inputs=*/std::unordered_map{ + { + mk_slot(TensorSlotName::INPUT), + value_2, + }, + }, + /*node_attrs=*/node_attrs, + /*outputs=*/ + std::unordered_map{ + { + mk_slot(TensorSlotName::OUTPUT), + value_1, + }, + }, + }; + + std::unordered_set invocation_set = { + invocation_1, + invocation_2, + }; + + CHECK_THROWS( + dynamic_open_dataflow_graph_from_invocation_set(invocation_set)); + } + } + + TEST_CASE("get_dynamic_slot_sites") { + auto mk_dynamic_value = [](int node_id, + TensorSlotName slot_name) -> DynamicValueAttrs { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{static_cast(node_id)}, + slot_name, + }, + }}, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*tensor_type=*/std::nullopt, + }; + }; + + DynamicValueAttrs value_1 = mk_dynamic_value(1, TensorSlotName::OUTPUT); + DynamicValueAttrs value_2 = mk_dynamic_value(2, TensorSlotName::OUTPUT); + DynamicValueAttrs value_3 = mk_dynamic_value(3, TensorSlotName::OUTPUT); + DynamicNodeAttrs node_attrs = DynamicNodeAttrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, @@ -59,25 +227,14 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ /*inputs=*/std::unordered_map{ - {DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/std::nullopt, - }, - value_1}, - }, - /*node_attrs=*/node_attrs, - /*outputs=*/ - std::unordered_map{ - {DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - }, - value_2}, + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/std::nullopt, + }, + value_1, + }, }, - }; - - DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{}, /*node_attrs=*/node_attrs, /*outputs=*/ std::unordered_map{ @@ -86,48 +243,250 @@ TEST_SUITE(FF_TEST_SUITE) { /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/std::nullopt, }, - value_3, + value_2, }, }, }; - DynamicNodeInvocation invocation_3 = DynamicNodeInvocation{ + DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ /*inputs=*/std::unordered_map{ { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::INPUT, /*slot_tensor_role=*/std::nullopt, }, - value_1, + value_2, }, { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::WEIGHT, /*slot_tensor_role=*/std::nullopt, }, - value_2, + value_3, }, - { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::BIAS, - /*slot_tensor_role=*/std::nullopt, - }, + }, + /*node_attrs=*/node_attrs, + /*outputs=*/{}, + }; + + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set( + std::unordered_set{invocation_1, invocation_2}); + + std::unordered_set result = get_dynamic_slot_sites(g); + + auto mk_internal_slot_site = [](DynamicNodeInvocation const &invocation, + TensorDirection direction, + TensorSlotName slot_name) { + return DynamicSlotSite{ + InternalDynamicSlotSite{ + /*invocation=*/invocation, + /*direction=*/direction, + /*slot_name=*/ + DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + }, + }, + }; + }; + + std::unordered_set correct = { + DynamicSlotSite{ + ExternalDynamicSlotSite{ value_1, }, }, - /*node_attrs=*/node_attrs, - /*outputs=*/std::unordered_map{}, + DynamicSlotSite{ + ExternalDynamicSlotSite{ + value_3, + }, + }, + mk_internal_slot_site( + invocation_1, TensorDirection::INCOMING, TensorSlotName::INPUT), + mk_internal_slot_site( + invocation_1, TensorDirection::OUTPUT, TensorSlotName::OUTPUT), + mk_internal_slot_site( + invocation_2, TensorDirection::INCOMING, TensorSlotName::INPUT), + mk_internal_slot_site( + invocation_2, TensorDirection::INCOMING, TensorSlotName::WEIGHT), + }; + + CHECK(result == correct); + } + + TEST_CASE( + "labelled_open_kwarg_dataflow_graph_from_dynamic_open_dataflow_graph") { + dynamic_layer_guid_t layer_guid = dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{0}, + }, + }; + + dynamic_tensor_guid_t tensor_guid = dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + /*node=*/Node{0}, + /*slot_name=*/TensorSlotName::OUTPUT, + }, + }, + }; + + TrainingOperationAttrs weight_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + WeightAttrs{ + /*tensor_shape=*/TensorShape{ + /*dims=*/TensorDims{ + FFOrdered{ + 4_p, + 3_p, + }, + }, + /*data_type=*/DataType::FLOAT, + }, + /*initializer=*/make_zero_initializer(), + }, + }, + }; + + DynamicNodeAttrs fwd_weight_node_attrs = DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/weight_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }; + + DynamicTensorSlot fwd_weight_output_slot1 = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + }; + + DynamicValueAttrs fwd_weight_output_attrs1 = DynamicValueAttrs{ + /*tensor_guid=*/tensor_guid, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/mk_dynamic_tensor_role_fwd(), + }; + + DynamicNodeInvocation weight_invocation = DynamicNodeInvocation{ + /*inputs=*/{}, + /*node_attrs=*/fwd_weight_node_attrs, + /*outputs=*/ + std::unordered_map{ + { + fwd_weight_output_slot1, + fwd_weight_output_attrs1, + }, + }, + }; + + DynamicNodeAttrs upd_weight_node_attrs = DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::UPD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/weight_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }; + + DynamicTensorSlot upd_weight_input_slot2 = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + }; + + DynamicValueAttrs upd_weight_input_attrs2 = DynamicValueAttrs{ + /*tensor_guid=*/tensor_guid, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/mk_dynamic_tensor_role_bwd(), }; - std::unordered_set invocation_set = { - invocation_1, - invocation_2, - invocation_3, + DynamicTensorSlot upd_weight_input_slot3 = DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/ + mk_dynamic_tensor_role_opt(OptimizerSlotName::SGD_V), }; - DynamicOpenDataflowGraph result = - dynamic_open_dataflow_graph_from_invocation_set(invocation_set); + DynamicValueAttrs upd_weight_input_attrs3 = DynamicValueAttrs{ + /*tensor_guid=*/tensor_guid, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/mk_dynamic_tensor_role_opt(OptimizerSlotName::SGD_V), + }; - ASSERT(dynamic_graph_num_nodes(result) == 3); + DynamicOpenDataflowGraph input = + dynamic_open_dataflow_graph_from_invocation_set( + /*invocations=*/{ + weight_invocation, + DynamicNodeInvocation{/*inputs=*/{ + { + fwd_weight_output_slot1, + fwd_weight_output_attrs1, + }, + { + upd_weight_input_slot2, + upd_weight_input_attrs2, + }, + { + upd_weight_input_slot3, + upd_weight_input_attrs3, + }, + }, + /*node_attrs=*/upd_weight_node_attrs, + /*outputs=*/{}}, + }); + + std::pair, + bidict> + result = + labelled_open_kwarg_dataflow_graph_from_dynamic_open_dataflow_graph( + input); + + LabelledOpenKwargDataflowGraph + correct = LabelledOpenKwargDataflowGraph:: + create>(); + + KwargNodeAddedResult fwd_weight_added = correct.add_node( + /*node_label=*/fwd_weight_node_attrs, + /*inputs=*/{}, + /*output_labels=*/ + { + { + fwd_weight_output_slot1, + fwd_weight_output_attrs1, + }, + }); + + KwargNodeAddedResult upd_weight_added = correct.add_node( + /*node_label=*/fwd_weight_node_attrs, + /*inputs=*/{}, + /*output_labels=*/ + { + { + fwd_weight_output_slot1, + fwd_weight_output_attrs1, + }, + }); } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc index 40b3460ee5..6bc31999a6 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc @@ -6,13 +6,14 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("perform_machine_slicing_for_invocation") { - auto mk_machine_coord = - [](nonnegative_int node_idx, - nonnegative_int device_idx) -> MachineSpaceCoordinate { - return MachineSpaceCoordinate{ - /*node_idx=*/node_idx, - /*device_idx=*/device_idx, - /*device_type=*/DeviceType::GPU, + auto mk_device_id = [](nonnegative_int node_idx, + nonnegative_int device_idx) -> device_id_t { + return device_id_t{ + MachineSpaceCoordinate{ + /*node_idx=*/node_idx, + /*device_idx=*/device_idx, + }, + DeviceType::GPU, }; }; @@ -32,9 +33,9 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(2_n, 0_n); - MachineSpaceCoordinate mc3 = mk_machine_coord(4_n, 0_n); + device_id_t mc1 = mk_device_id(0_n, 0_n); + device_id_t mc2 = mk_device_id(2_n, 0_n); + device_id_t mc3 = mk_device_id(4_n, 0_n); ParallelTensorSpaceCoordinate mc1_input_coord = mk_pt_coord(0_n, 0_n, 0_n, 0_n); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc new file mode 100644 index 0000000000..8b13789179 --- /dev/null +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -0,0 +1 @@ + diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index efe21146db..65870192d7 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -2,6 +2,7 @@ #include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" #include "task-spec/dynamic_graph/copy_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_copy_layer_guid_t.dtg.h" +#include "task-spec/dynamic_graph/dynamic_node_mapping.h" #include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" #include "test/utils/doctest/fmt/unordered_set.h" #include "utils/bidict/algorithms/filter_keys.h" @@ -14,7 +15,6 @@ static MachineSpaceCoordinate mk_machine_coord(nonnegative_int node_idx, return MachineSpaceCoordinate{ /*node_idx=*/node_idx, /*device_idx=*/device_idx, - /*device_type=*/DeviceType::GPU, }; }; @@ -33,200 +33,157 @@ static ParallelTensorSpaceCoordinate mk_pt_coord(nonnegative_int idx1, }; }; -DynamicTensorSlot mk_slot(TensorSlotName const &slot_name) { - return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - }; -}; - -DynamicValueAttrs - mk_value(size_t src_node_id, - TensorSlotName src_slot_name, - bidict - tensor_binding, - std::optional const &shard_coord) { - if (shard_coord.has_value()) { - tensor_binding = filter_keys(tensor_binding, - [&](ParallelTensorSpaceCoordinate const &p) { - return p == shard_coord.value(); - }); - } - return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{src_node_id}, - src_slot_name, - }, - }}, - /*parallel_tensor_shape=*/std::nullopt, - /*shard_coord=*/shard_coord, - /*mapping=*/ - tensor_binding, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }; -}; - TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("perform_shard_expansion_for_invocation") { - auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, - ParallelTensorSpaceCoordinate const &c2, - ParallelTensorSpaceCoordinate const &c3, - ParallelTensorSpaceCoordinate const &c4) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ - { - TensorSlotName::INPUT, - c1, - }, - { - TensorSlotName::WEIGHT, - c2, - }, - { - TensorSlotName::OUTPUT_1, - c3, - }, - { - TensorSlotName::OUTPUT_2, - c4, - }, - }, + + auto mk_slot = [](TensorSlotName const &slot_name) { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, }; }; - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(2_n, 0_n); - - ParallelTensorSpaceCoordinate mc1_input_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc1_weight_coord = - mk_pt_coord(0_n, 1_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc1_output_1_coord = - mk_pt_coord(1_n, 0_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc1_output_2_coord = - mk_pt_coord(3_n, 0_n, 0_n, 0_n); - - ParallelTensorSpaceCoordinate mc2_input_coord = - mk_pt_coord(0_n, 1_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc2_weight_coord = - mk_pt_coord(0_n, 4_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc2_output_1_coord = - mk_pt_coord(1_n, 2_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc2_output_2_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - - MappedOperatorTaskGroup mapped_task_group = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_shard_binding(mc1_input_coord, - mc1_weight_coord, - mc1_output_1_coord, - mc1_output_2_coord), - }, - { - mc2, - mk_shard_binding(mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), - }, - }, + DeviceType device_type = DeviceType::GPU; + auto mk_device_id = [&](MachineSpaceCoordinate const &c) -> device_id_t { + return device_id_t{c, device_type}; }; - auto mk_op_value = + auto mk_value = [&](size_t src_node_id, TensorSlotName src_slot_name, - TensorSlotName use_slot_name, + bidict tensor_binding, std::optional const &shard_coord) -> DynamicValueAttrs { - bidict - tensor_binding = get_tensor_bindings_for_slot_name(mapped_task_group, - use_slot_name); - return mk_value(src_node_id, src_slot_name, tensor_binding, shard_coord); + if (shard_coord.has_value()) { + tensor_binding = + filter_keys(tensor_binding, + [&](ParallelTensorSpaceCoordinate const &p) -> bool { + return p == shard_coord.value(); + }); + } + + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{src_node_id}, + src_slot_name, + }, + }}, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/shard_coord, + /*mapping=*/ + ParallelTensorMapping{tensor_binding}, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; }; - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_op_value(0, - TensorSlotName::OUTPUT, - TensorSlotName::INPUT, - std::nullopt), - }, - { - mk_slot(TensorSlotName::WEIGHT), - mk_op_value(1, - TensorSlotName::OUTPUT, - TensorSlotName::WEIGHT, - std::nullopt), - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/mapped_task_group, - /*op_attrs=*/std::nullopt, - /*layer_guid=*/ - dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT_1), - mk_op_value(20, - TensorSlotName::OUTPUT_1, - TensorSlotName::OUTPUT_1, - std::nullopt), + SUBCASE("standard operators") { + auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, + ParallelTensorSpaceCoordinate const &c2, + ParallelTensorSpaceCoordinate const &c3, + ParallelTensorSpaceCoordinate const &c4) + -> OperatorAtomicTaskShardBinding { + return OperatorAtomicTaskShardBinding{ + /*tensor_coords=*/{ + { + TensorSlotName::INPUT, + c1, + }, + { + TensorSlotName::WEIGHT, + c2, + }, + { + TensorSlotName::OUTPUT_1, + c3, + }, + { + TensorSlotName::OUTPUT_2, + c4, + }, }, - { - mk_slot(TensorSlotName::OUTPUT_2), - mk_op_value(20, - TensorSlotName::OUTPUT_2, - TensorSlotName::OUTPUT_2, - std::nullopt), - }, - }, - }; + }; + }; + + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(2_n, 0_n); + + ParallelTensorSpaceCoordinate mc1_input_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate mc1_weight_coord = + mk_pt_coord(0_n, 1_n, 2_n, 0_n); + ParallelTensorSpaceCoordinate mc1_output_1_coord = + mk_pt_coord(1_n, 0_n, 0_n, 1_n); + ParallelTensorSpaceCoordinate mc1_output_2_coord = + mk_pt_coord(3_n, 0_n, 0_n, 0_n); + + ParallelTensorSpaceCoordinate mc2_input_coord = + mk_pt_coord(0_n, 1_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate mc2_weight_coord = + mk_pt_coord(0_n, 4_n, 2_n, 0_n); + ParallelTensorSpaceCoordinate mc2_output_1_coord = + mk_pt_coord(1_n, 2_n, 0_n, 1_n); + ParallelTensorSpaceCoordinate mc2_output_2_coord = + mk_pt_coord(0_n, 0_n, 0_n, 0_n); + + DynamicNodeMapping node_mapping = DynamicNodeMapping{ + MappedOperatorTaskGroup{ + bidict{ + { + mc1, + mk_shard_binding(mc1_input_coord, + mc1_weight_coord, + mc1_output_1_coord, + mc1_output_2_coord), + }, + { + mc2, + mk_shard_binding(mc2_input_coord, + mc2_weight_coord, + mc2_output_1_coord, + mc2_output_2_coord), + }, + }, + }, + device_type, + }; - std::unordered_set result = - perform_shard_expansion_for_invocation(input); + auto mk_op_value = + [&](size_t src_node_id, + TensorSlotName src_slot_name, + TensorSlotName use_slot_name, + std::optional const &shard_coord) + -> DynamicValueAttrs { + bidict tensor_binding = + dynamic_node_mapping_bindings_for_slot_name(node_mapping, + use_slot_name); + return mk_value( + src_node_id, src_slot_name, tensor_binding, shard_coord); + }; - auto mk_invocation_shard = - [&](MachineSpaceCoordinate const &device_coord, - ParallelTensorSpaceCoordinate const &input_shard_coord, - ParallelTensorSpaceCoordinate const &weight_shard_coord, - ParallelTensorSpaceCoordinate const &output_1_shard_coord, - ParallelTensorSpaceCoordinate const &output_2_shard_coord) - -> DynamicNodeInvocation { - return DynamicNodeInvocation{ + DynamicNodeInvocation input = DynamicNodeInvocation{ /*inputs=*/{ { mk_slot(TensorSlotName::INPUT), mk_op_value(0, TensorSlotName::OUTPUT, TensorSlotName::INPUT, - input_shard_coord), + std::nullopt), }, { mk_slot(TensorSlotName::WEIGHT), mk_op_value(1, TensorSlotName::OUTPUT, TensorSlotName::WEIGHT, - weight_shard_coord), + std::nullopt), }, }, /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/std::nullopt, - /*device_coord=*/device_coord, - /*mapping=*/mapped_task_group, + /*device_coord=*/std::nullopt, + /*mapping=*/node_mapping, /*op_attrs=*/std::nullopt, /*layer_guid=*/ dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, @@ -239,112 +196,173 @@ TEST_SUITE(FF_TEST_SUITE) { mk_op_value(20, TensorSlotName::OUTPUT_1, TensorSlotName::OUTPUT_1, - output_1_shard_coord), + std::nullopt), }, { mk_slot(TensorSlotName::OUTPUT_2), mk_op_value(20, TensorSlotName::OUTPUT_2, TensorSlotName::OUTPUT_2, - output_2_shard_coord), + std::nullopt), }, }, }; - }; - std::unordered_set correct = { - mk_invocation_shard(mc1, - mc1_input_coord, - mc1_weight_coord, - mc1_output_1_coord, - mc1_output_2_coord), - mk_invocation_shard(mc2, - mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), - }; + std::unordered_set result = + perform_shard_expansion_for_invocation(input); - CHECK(result.size() == correct.size()); - CHECK(result == correct); - } + auto mk_invocation_shard = + [&](device_id_t const &device_coord, + ParallelTensorSpaceCoordinate const &input_shard_coord, + ParallelTensorSpaceCoordinate const &weight_shard_coord, + ParallelTensorSpaceCoordinate const &output_1_shard_coord, + ParallelTensorSpaceCoordinate const &output_2_shard_coord) + -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_op_value(0, + TensorSlotName::OUTPUT, + TensorSlotName::INPUT, + input_shard_coord), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_op_value(1, + TensorSlotName::OUTPUT, + TensorSlotName::WEIGHT, + weight_shard_coord), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/device_coord, + /*mapping=*/node_mapping, + /*op_attrs=*/std::nullopt, + /*layer_guid=*/ + dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + mk_slot(TensorSlotName::OUTPUT_1), + mk_op_value(20, + TensorSlotName::OUTPUT_1, + TensorSlotName::OUTPUT_1, + output_1_shard_coord), + }, + { + mk_slot(TensorSlotName::OUTPUT_2), + mk_op_value(20, + TensorSlotName::OUTPUT_2, + TensorSlotName::OUTPUT_2, + output_2_shard_coord), + }, + }, + }; + }; - TEST_CASE("perform_shard_expansion_for_invocation (copy)") { - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); - MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); - MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); + std::unordered_set correct = { + mk_invocation_shard(mk_device_id(mc1), + mc1_input_coord, + mc1_weight_coord, + mc1_output_1_coord, + mc1_output_2_coord), + mk_invocation_shard(mk_device_id(mc2), + mc2_input_coord, + mc2_weight_coord, + mc2_output_1_coord, + mc2_output_2_coord), + }; - ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } - bidict src_binding{ - {pt1, mc1}, - {pt2, mc2}, - }; - bidict dst_binding{ - {pt1, mc3}, - {pt2, mc4}, - }; + SUBCASE("for copy operator") { + device_id_t dev1 = mk_device_id(mk_machine_coord(0_n, 0_n)); + device_id_t dev2 = mk_device_id(mk_machine_coord(1_n, 0_n)); + device_id_t dev3 = mk_device_id(mk_machine_coord(2_n, 0_n)); + device_id_t dev4 = mk_device_id(mk_machine_coord(3_n, 0_n)); - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ + ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); + ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); + + bidict src_binding{ + {pt1, dev1}, + {pt2, dev2}, + }; + bidict dst_binding{ + {pt1, dev3}, + {pt2, dev4}, + }; + + DynamicNodeInvocation input = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value( + 0, TensorSlotName::OUTPUT, src_binding, std::nullopt), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + mk_slot(TensorSlotName::OUTPUT), + mk_value( + 20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), + }, + }, + }; + + std::unordered_set result = + perform_shard_expansion_for_invocation(input); + + auto mk_invocation_shard = + [&](device_id_t const &device_coord, + ParallelTensorSpaceCoordinate const &tensor_shard_coord) + -> DynamicNodeInvocation { + DynamicNodeInvocation result = input; + result.inputs = { { mk_slot(TensorSlotName::INPUT), - mk_value(0, TensorSlotName::OUTPUT, src_binding, std::nullopt), + mk_value( + 0, TensorSlotName::OUTPUT, src_binding, tensor_shard_coord), }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { + }; + // See perform_shard_expansion_for_copy in shard_expansion.cc for explanation of the choice of device placement. + result.node_attrs.device_coord = device_coord; + result.outputs = { { mk_slot(TensorSlotName::OUTPUT), - mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), + mk_value(20, + TensorSlotName::OUTPUT, + dst_binding, + tensor_shard_coord), }, - }, - }; - - std::unordered_set result = - perform_shard_expansion_for_invocation(input); - - auto mk_invocation_shard = - [&](MachineSpaceCoordinate const &device_coord, - ParallelTensorSpaceCoordinate const &tensor_shard_coord) - -> DynamicNodeInvocation { - DynamicNodeInvocation result = input; - result.inputs = { - { - mk_slot(TensorSlotName::INPUT), - mk_value( - 0, TensorSlotName::OUTPUT, src_binding, tensor_shard_coord), - }, + }; + return result; }; - // See perform_shard_expansion_for_copy in shard_expansion.cc for explanation of the choice of device placement. - result.node_attrs.device_coord = device_coord; - result.outputs = { - { - mk_slot(TensorSlotName::OUTPUT), - mk_value( - 20, TensorSlotName::OUTPUT, dst_binding, tensor_shard_coord), - }, - }; - return result; - }; - std::unordered_set correct = { - mk_invocation_shard(mc1, pt1), - mk_invocation_shard(mc2, pt2), - }; + std::unordered_set correct = { + mk_invocation_shard(dev1, pt1), + mk_invocation_shard(dev2, pt2), + }; - CHECK(result.size() == correct.size()); - CHECK(result == correct); + CHECK(result.size() == correct.size()); + CHECK(result == correct); + } } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc new file mode 100644 index 0000000000..3fb94bb98a --- /dev/null +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc @@ -0,0 +1,157 @@ +#include "task-spec/dynamic_graph/update_insertion.h" +#include "op-attrs/initializer_attrs.h" +#include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("perform_update_insertion") { + dynamic_layer_guid_t layer_guid = dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{0}, + }, + }; + + dynamic_tensor_guid_t tensor_guid = dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + /*node=*/Node{0}, + /*slot_name=*/TensorSlotName::OUTPUT, + }, + }, + }; + + TrainingOperationAttrs weight_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + WeightAttrs{ + /*tensor_shape=*/TensorShape{ + /*dims=*/TensorDims{ + FFOrdered{ + 4_p, + 3_p, + }, + }, + /*data_type=*/DataType::FLOAT, + }, + /*initializer=*/make_zero_initializer(), + }, + }, + }; + + DynamicNodeInvocation weight_invocation = DynamicNodeInvocation{ + /*inputs=*/{}, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/weight_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + std::unordered_map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + }, + DynamicValueAttrs{ + /*tensor_guid=*/tensor_guid, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/mk_dynamic_tensor_role_fwd(), + }, + }, + }, + }; + + DynamicOpenDataflowGraph input = + dynamic_open_dataflow_graph_from_invocation_set({weight_invocation}); + + OptimizerAttrs optimizer_attrs = OptimizerAttrs{ + SGDOptimizerAttrs{ + /*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001, + }, + }; + + DynamicOpenDataflowGraph result = + perform_update_insertion(input, optimizer_attrs); + + DynamicOpenDataflowGraph correct = + dynamic_open_dataflow_graph_from_invocation_set( + /*invocations=*/{ + weight_invocation, + DynamicNodeInvocation{ + /*inputs=*/{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/ + mk_dynamic_tensor_role_fwd(), + }, + DynamicValueAttrs{ + /*tensor_guid=*/tensor_guid, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/mk_dynamic_tensor_role_fwd(), + }, + }, + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/ + mk_dynamic_tensor_role_bwd(), + }, + DynamicValueAttrs{ + /*tensor_guid=*/tensor_guid, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/mk_dynamic_tensor_role_bwd(), + }, + }, + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/ + mk_dynamic_tensor_role_opt( + OptimizerSlotName::SGD_V), + }, + DynamicValueAttrs{ + /*tensor_guid=*/tensor_guid, + /*parallel_tensor_shape=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/ + mk_dynamic_tensor_role_opt( + OptimizerSlotName::SGD_V), + }, + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::UPD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/weight_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{}}, + }); + + CHECK(result == correct); + } +} From aa57301e8d2ea390934afa9a6a42941636b9a663 Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 4 Jun 2026 16:08:52 -0700 Subject: [PATCH 23/35] Fix usage of node device_coords. --- ...istributed_per_device_op_state_initialization.cc | 13 +++++++++++-- .../src/realm-execution/instance_allocation.cc | 10 +++++----- .../src/realm-execution/pcg_instance.cc | 2 +- 3 files changed, 17 insertions(+), 8 deletions(-) diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index 1d517a8fe4..f7fd0949ab 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -7,6 +7,7 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" #include "utils/containers/map_values.h" +#include "utils/containers/maybe_get_only.h" #include "utils/containers/values.h" #include "utils/optional.h" #include @@ -32,8 +33,16 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( DeviceSpecificPtr *> device_state_map; for (DynamicNodeInvocation const &invocation : dg.invocations) { - Realm::Processor target_proc = ctx.map_device_coord_to_processor( - assert_unwrap(invocation.node_attrs.device_coord)); + // Nodes mapped to multiple devices are always parallel operators and don't + // have any initialization to perform anyway + std::optional device_coord = + maybe_get_only(assert_unwrap(invocation.node_attrs.device_coords)); + if (!device_coord.has_value()) { + continue; + } + + Realm::Processor target_proc = + ctx.map_device_coord_to_processor(assert_unwrap(device_coord)); TensorInstanceBacking tensor_backing = subset_tensor_instance_backing_for_invocation(tensor_instance_backing, diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index 79eef476df..a305b2f20c 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -47,15 +47,15 @@ TensorInstanceBacking perform_instance_allocation( } TensorInstanceBacking result = make_empty_tensor_instance_backing(); - auto allocate = [&](DynamicNodeAttrs const &n, DynamicValueAttrs const &v) { + auto allocate = [&](DynamicValueAttrs const &v) { if (contains_key(preallocated, v)) { // FIXME: Attach external instance to existing allocation and use that NOT_IMPLEMENTED(); } else { if (!contains_key(result.backing, v)) { - MachineSpaceCoordinate device_coord = assert_unwrap(n.device_coord); + MachineSpaceCoordinate device_coord = v.mapping.at(assert_unwrap(v.shard_coord)); result.backing.insert(std::pair{ - v, perform_instance_allocation_for_value(device_coord, v, ctx)}); + v, perform_instance_allocation_for_value(assert_unwrap(device_coord), v, ctx)}); } return result.backing.at(v); } @@ -63,10 +63,10 @@ TensorInstanceBacking perform_instance_allocation( for (DynamicNodeInvocation const &invocation : g.invocations) { for (DynamicValueAttrs const &input : values(invocation.inputs)) { - allocate(invocation.node_attrs, input); + allocate(input); } for (DynamicValueAttrs const &output : values(invocation.outputs)) { - allocate(invocation.node_attrs, output); + allocate(output); } } diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 4b068d70be..1cbb4b9e10 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -188,7 +188,7 @@ static Realm::Event spawn_dynamic_node_invocation( auto spawn_task = [&]() { Realm::Processor target_proc = ctx.map_device_coord_to_processor( - assert_unwrap(invocation.node_attrs.device_coord)); + get_only(assert_unwrap(invocation.node_attrs.device_coords))); return spawn_op_task(ctx, target_proc, invocation, From b2c6b3d06963f087cb502934d333177076f9c29c Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Thu, 4 Jun 2026 16:17:11 -0700 Subject: [PATCH 24/35] Fix more compile errors. --- .../src/realm-execution/instance_allocation.cc | 5 +++-- lib/realm-execution/src/realm-execution/pcg_instance.cc | 4 ++-- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index a305b2f20c..f1d88f672c 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -53,9 +53,10 @@ TensorInstanceBacking perform_instance_allocation( NOT_IMPLEMENTED(); } else { if (!contains_key(result.backing, v)) { - MachineSpaceCoordinate device_coord = v.mapping.at(assert_unwrap(v.shard_coord)); + MachineSpaceCoordinate device_coord = + assert_unwrap(v.mapping).at_l(assert_unwrap(v.shard_coord)); result.backing.insert(std::pair{ - v, perform_instance_allocation_for_value(assert_unwrap(device_coord), v, ctx)}); + v, perform_instance_allocation_for_value(device_coord, v, ctx)}); } return result.backing.at(v); } diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 1cbb4b9e10..cfefdae619 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -233,10 +233,10 @@ static Realm::Event spawn_dynamic_node_invocation( // chain reductions sequentially to avoid write races on dst Realm::Event result = precondition; - for (auto const &[p, m] : unstructured_relation_from_one_to_many(assert_unwrap(output_grad.mapping))) { + for (auto const &[p, m] : assert_unwrap(output_grad.mapping)) { DynamicValueAttrs replica_key = output_grad; replica_key.mapping = - OneToMany{{p, {m}}}; + bidict{{p, m}}; replica_key.shard_coord = p; Realm::RegionInstance src_inst = From 43db2f37daa5d7a86a797bf27c89b64879a3d203 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 12 Jun 2026 01:50:02 -0700 Subject: [PATCH 25/35] Move unordered_set to set, unordered_map to map, and other minor fixes --- bin/run-model/src/run-model/main.cc | 2 +- .../sp-ization-benchmarking/distributions.h | 16 +- .../sp-ization-benchmarking/sample_graphs.cc | 6 +- flake.nix | 1 + .../op_cost_estimate_key.dtg.toml | 11 +- .../cost_estimator/op_cost_metrics.dtg.toml | 1 + .../compiler/cost_estimator/op_cost_metrics.h | 2 +- ...runtime_only_op_cost_estimate_key.dtg.toml | 11 +- .../single_tensor_movement.dtg.toml | 8 +- .../tensor_set_movement.dtg.toml | 9 +- .../abstracted_device.h | 2 +- ...tracted_single_tensor_communication_edge.h | 2 +- ...abstracted_single_tensor_movement.dtg.toml | 9 +- .../abstracted_single_tensor_movement.h | 10 +- .../abstracted_tensor_set_movement.dtg.toml | 9 +- .../abstracted_tensor_set_movement.h | 8 +- .../machine_mapping/allowed_machine_views.h | 2 +- .../feasible_machine_mapping_result.dtg.toml | 1 + .../machine_mapping/machine_mapping.dtg.toml | 13 +- .../machine_mapping_cache.dtg.toml | 8 +- .../machine_mapping_constraints.dtg.toml | 13 +- .../machine_mapping_constraints.h | 8 +- .../machine_mapping_context.dtg.toml | 2 +- .../machine_mapping_problem_tree.dtg.toml | 1 + .../machine_mapping_problem_tree.h | 6 +- .../mm_problem_tree_parallel_split.dtg.toml | 1 + .../mm_problem_tree_series_split.dtg.toml | 1 + .../unmapped_op_cost_estimate_key.dtg.toml | 12 +- ...runtime_only_op_cost_estimate_key.dtg.toml | 11 +- .../machine_mapping_result.dtg.toml | 1 + .../machine_mapping/machine_mapping_result.h | 2 +- .../machine_mapping_state.dtg.toml | 1 + .../machine_mapping/machine_resource_split.h | 2 +- .../compiler/machine_mapping/machine_view.h | 8 +- ...machine_mapping_with_memory_cache.dtg.toml | 8 +- ...chine_mapping_with_memory_context.dtg.toml | 2 +- .../machine_mapping_with_memory_result.h | 17 +- .../pareto_optimal_machine_mapping.dtg.toml | 1 + .../pareto_optimal_machine_mapping.h | 2 +- ...er_guid_oblivious_machine_mapping.dtg.toml | 7 +- ...lel_layer_guid_oblivious_machine_mapping.h | 6 +- .../pcg_split_boundary_layers.dtg.toml | 8 +- .../start_invariant_machine_view.h | 2 +- .../machine_mapping/transitive_reduced_pcg.h | 4 +- .../unstructured_device_mapping.dtg.toml | 10 +- ...omputation_graph_binary_sp_decomposition.h | 2 +- .../pcg/pcg_binary_sp_decomposition.h | 10 +- .../in_progress_task.dtg.toml | 2 - .../task_graph_simulator/pcg_task.dtg.toml | 1 + .../task_execution_constraint.dtg.toml | 4 +- .../task_graph_execution_state.dtg.toml | 8 +- .../task_graph_execution_trace.dtg.toml | 8 +- .../cost_estimator/op_cost_estimate_key.cc | 2 +- .../cost_estimator/op_cost_metrics.cc | 2 +- .../cost_estimator/tensor_set_movement.cc | 5 +- .../abstracted_device.cc | 2 +- ...racted_single_tensor_communication_edge.cc | 2 +- .../abstracted_single_tensor_movement.cc | 24 +-- .../abstracted_tensor_set_movement.cc | 16 +- ...racted_tensor_set_movement_across_split.cc | 16 +- .../machine_mapping/allowed_machine_views.cc | 32 +-- ...substitution_and_update_machine_mapping.cc | 16 +- .../get_optimal_machine_mapping.cc | 30 +-- .../get_tensor_set_movement_across_split.cc | 6 +- .../machine_mapping/machine_mapping.cc | 22 +- .../machine_mapping_constraints.cc | 24 +-- .../machine_mapping_mutation_set.cc | 4 +- .../get_machine_mapping_problem_tree.cc | 2 +- .../machine_mapping_problem_tree.cc | 8 +- .../machine_mapping/machine_mapping_result.cc | 2 +- .../machine_mapping/machine_resource_split.cc | 4 +- .../compiler/machine_mapping/machine_view.cc | 16 +- ...get_optimal_machine_mapping_with_memory.cc | 24 +-- .../machine_mapping_with_memory_result.cc | 36 +++- .../pareto_optimal_machine_mapping.cc | 2 +- ...el_layer_guid_oblivious_machine_mapping.cc | 28 +-- .../start_invariant_machine_view.cc | 2 +- .../machine_mapping/transitive_reduced_pcg.cc | 16 +- .../unstructured_device_mapping.cc | 2 +- ...mputation_graph_binary_sp_decomposition.cc | 2 +- ...ion_graph_series_parallel_decomposition.cc | 8 +- .../get_pcg_series_parallel_decomposition.cc | 10 +- .../pcg/pcg_binary_sp_decomposition.cc | 12 +- .../task_graph_simulator/pcg_task_graph.cc | 6 +- .../simulate_task_graph_execution.cc | 10 +- .../task_graph_execution_trace.cc | 2 +- .../task_graph_simulator/task_simulator.cc | 12 +- .../unity_algorithm/graph_optimize_state.cc | 22 +- .../unity_algorithm/unity_algorithm.cc | 4 +- .../machine_mapping/allowed_machine_views.cc | 20 +- .../get_optimal_machine_mapping.cc | 44 ++-- .../machine_mapping/machine_resource_split.cc | 18 +- .../compiler/machine_mapping/machine_view.cc | 10 +- ...get_optimal_machine_mapping_with_memory.cc | 40 ++-- .../machine_mapping_with_memory_result.cc | 10 +- .../start_invariant_machine_view.cc | 10 +- .../pcg/pcg_binary_sp_decomposition.cc | 2 +- .../simulate_task_graph_execution.cc | 24 +-- .../task_graph_simulator/task_simulator.cc | 4 +- .../src/internal/cost_estimator_for_test.cc | 4 +- .../src/internal/cost_estimator_for_test.h | 4 +- .../runtime_only_cost_estimator_for_test.cc | 4 +- .../runtime_only_cost_estimator_for_test.h | 4 +- .../include/kernels/local_cpu_allocator.h | 4 +- .../include/kernels/local_cuda_allocator.h | 4 +- .../include/kernels/reduce_tensor_accessor.h | 6 +- .../src/kernels/reduce_tensor_accessor.cc | 2 +- .../computation_graph_instance.h | 10 +- .../cost_estimator/tracked_allocator.h | 2 +- .../local_task_argument_accessor.h | 6 +- .../local-execution/tensor_allocation.h | 2 +- .../computation_graph_instance.cc | 18 +- .../cost_estimator/local_cost_estimator.cc | 24 +-- .../local_task_argument_accessor.cc | 2 +- .../src/local-execution/task_execution.cc | 7 +- .../src/local-execution/tensor_allocation.cc | 10 +- .../computation_graph_instance.cc | 6 +- .../local_task_argument_accessor.cc | 2 +- .../op-attrs/ff_ordered/ff_ordered_from_map.h | 4 +- .../op-attrs/ff_ordered/map_from_ff_ordered.h | 4 +- .../op-attrs/get_incoming_tensor_roles.h | 4 +- ..._space_to_parallel_tensor_space_mappings.h | 28 +-- .../op-attrs/get_operator_task_space.h | 2 +- .../include/op-attrs/operator_task_space.h | 6 +- lib/op-attrs/include/op-attrs/ops/attention.h | 8 +- .../include/op-attrs/ops/batch_norm.h | 10 +- lib/op-attrs/include/op-attrs/ops/conv_2d.h | 8 +- lib/op-attrs/include/op-attrs/ops/embedding.h | 2 +- lib/op-attrs/include/op-attrs/ops/index.dox | 2 +- .../include/op-attrs/ops/layer_norm.h | 8 +- lib/op-attrs/include/op-attrs/ops/linear.h | 8 +- .../op-attrs/parallel_tensor_dim_degrees.h | 6 +- .../op-attrs/parallel_tensor_dims.dtg.toml | 4 +- .../include/op-attrs/parallel_tensor_dims.h | 2 +- .../include/op-attrs/parallel_tensor_shape.h | 4 +- .../parallel_tensor_space_coordinate.h | 4 +- .../op-attrs/replica_parallel_dim_set.h | 2 +- .../include/op-attrs/shape_inference.h | 16 +- lib/op-attrs/include/op-attrs/tensor_dims.h | 6 +- lib/op-attrs/src/op-attrs/datatype.cc | 4 +- .../ff_ordered/ff_ordered_from_map.cc | 5 +- .../src/op-attrs/ff_ordered/ff_ordered_of.cc | 2 +- .../ff_ordered/map_from_ff_ordered.cc | 2 +- .../src/op-attrs/get_incoming_tensor_roles.cc | 54 ++--- ...space_to_parallel_tensor_space_mappings.cc | 70 +++---- .../src/op-attrs/get_operator_task_space.cc | 2 +- ..._space_to_parallel_tensor_space_mapping.cc | 4 +- .../src/op-attrs/operator_task_space.cc | 16 +- lib/op-attrs/src/op-attrs/ops/attention.cc | 18 +- lib/op-attrs/src/op-attrs/ops/batch_norm.cc | 22 +- lib/op-attrs/src/op-attrs/ops/conv_2d.cc | 14 +- lib/op-attrs/src/op-attrs/ops/embedding.cc | 2 +- lib/op-attrs/src/op-attrs/ops/layer_norm.cc | 14 +- lib/op-attrs/src/op-attrs/ops/linear.cc | 28 +-- .../op-attrs/parallel_tensor_dim_degrees.cc | 39 ++-- .../src/op-attrs/parallel_tensor_dims.cc | 2 +- .../src/op-attrs/parallel_tensor_shape.cc | 6 +- .../parallel_tensor_space_coordinate.cc | 15 +- ..._space_to_parallel_tensor_space_mapping.cc | 12 +- .../src/op-attrs/replica_parallel_dim_set.cc | 4 +- lib/op-attrs/src/op-attrs/shape_inference.cc | 196 +++++++++--------- .../src/op-attrs/task_space_coordinate.cc | 11 +- .../src/op-attrs/tensor_dim_permutation.cc | 4 +- lib/op-attrs/src/op-attrs/tensor_dims.cc | 14 +- .../ff_ordered/ff_ordered_from_map.cc | 2 +- .../src/op-attrs/get_incoming_tensor_roles.cc | 4 +- .../test/src/op-attrs/operator_task_space.cc | 14 +- .../test/src/op-attrs/ops/attention.cc | 12 +- .../test/src/op-attrs/ops/batch_norm.cc | 8 +- lib/op-attrs/test/src/op-attrs/ops/conv_2d.cc | 8 +- .../test/src/op-attrs/ops/layer_norm.cc | 8 +- lib/op-attrs/test/src/op-attrs/ops/linear.cc | 8 +- .../op-attrs/parallel_tensor_dim_degrees.cc | 12 +- .../src/op-attrs/parallel_tensor_dim_idx_t.cc | 2 +- lib/op-attrs/test/src/op-attrs/tensor_dims.cc | 12 +- lib/pcg/include/pcg/computation_graph.h | 36 ++-- .../layer_added_result.dtg.toml | 4 +- .../include/pcg/computation_graph_builder.h | 8 +- .../graphs/v1_kwarg_dataflow_graph.dtg.toml | 10 +- .../v1/graphs/v1_kwarg_dataflow_graph.h | 28 +-- .../v1_labelled_kwarg_dataflow_graph.dtg.toml | 10 +- .../graphs/v1_labelled_kwarg_dataflow_graph.h | 16 +- ...mapped_parallel_computation_graph.dtg.toml | 2 +- .../mapped_parallel_computation_graph.h | 12 +- ...perator_atomic_task_shard_binding.dtg.toml | 7 +- lib/pcg/include/pcg/metric_attrs.h | 4 +- ...or_space_to_machine_space_mapping.dtg.toml | 1 + lib/pcg/include/pcg/optimizer_attrs.h | 2 +- .../generate_weight_transform.h | 2 +- .../parallel_computation_graph.h | 48 ++--- .../parallel_computation_graph_builder.h | 6 +- .../parallel_layer_added_result.dtg.toml | 4 +- lib/pcg/src/pcg/computation_graph.cc | 92 ++++---- lib/pcg/src/pcg/computation_graph_builder.cc | 52 ++--- .../v1/graphs/v1_kwarg_dataflow_graph.cc | 4 +- .../v1_labelled_kwarg_dataflow_graph.cc | 4 +- .../mapped_operator_task_group.cc | 15 +- .../mapped_parallel_computation_graph.cc | 20 +- lib/pcg/src/pcg/metric_attrs.cc | 2 +- lib/pcg/src/pcg/optimizer_attrs.cc | 8 +- .../generate_weight_transform.cc | 4 +- .../parallel_computation_graph.cc | 121 ++++++----- .../parallel_computation_graph_builder.cc | 44 ++-- lib/pcg/test/src/pcg/computation_graph.cc | 24 +-- .../mapped_parallel_computation_graph.cc | 2 +- .../parallel_computation_graph.cc | 28 +-- .../parallel_computation_graph_builder.cc | 84 ++++---- .../include/realm-execution/dependency_set.h | 4 +- .../realm-execution/distributed_ff_handle.h | 6 +- .../realm-execution/instance_allocation.h | 2 +- .../include/realm-execution/pcg_instance.h | 10 +- .../per_device_op_state_backing.dtg.toml | 4 +- .../include/realm-execution/realm_allocator.h | 2 +- .../include/realm-execution/realm_context.h | 4 +- ...ializable_tensor_instance_backing.dtg.toml | 8 +- .../tensor_instance_backing.dtg.toml | 8 +- .../realm-execution/distributed_ff_handle.cc | 4 +- ...uted_per_device_op_state_initialization.cc | 6 +- .../realm-execution/instance_allocation.cc | 6 +- .../src/realm-execution/pcg_instance.cc | 20 +- .../test/src/realm-execution/test_e2e.cc | 4 +- .../src/realm-execution/test_op_replicate.cc | 16 +- .../perform_shape_inference.h | 2 +- .../operator_pattern/get_attribute_map.h | 2 +- .../operator_attribute_pattern.dtg.toml | 8 +- .../materialize_operator_from_attrs_map.h | 2 +- .../output_graph/output_graph_expr.h | 6 +- .../output_operator_attribute_expr.h | 2 +- .../output_operator_attrs_assignment.dtg.toml | 8 +- .../output_operator_attrs_assignment.h | 2 +- .../include/substitutions/pcg_pattern.h | 6 +- .../substitutions/pcg_pattern_match.dtg.toml | 9 +- .../sub_parallel_computation_graph.h | 20 +- ...b_parallel_computation_graph_data.dtg.toml | 20 +- .../substitutions/substitution_builder.h | 12 +- .../tensor_attribute_pattern.dtg.toml | 8 +- .../substitutions/unlabelled/pattern_edge.h | 4 +- .../unlabelled/pattern_split.dtg.toml | 10 +- .../unlabelled/unlabelled_graph_pattern.h | 14 +- ...warg_dataflow_graph_pattern_match.dtg.toml | 8 +- ...elled_kwarg_dataflow_graph_pattern_match.h | 6 +- .../apply_substitution/apply_substitution.cc | 44 ++-- .../evaluate_substitution_output.cc | 6 +- .../output_expr_to_result_sub_pcg_mapping.cc | 4 +- .../perform_shape_inference.cc | 30 +-- .../operator_pattern/get_attribute_map.cc | 4 +- .../materialize_operator_from_attrs_map.cc | 8 +- .../output_graph/output_graph_expr.cc | 12 +- .../output_operator_attribute_expr.cc | 2 +- .../output_operator_attrs_assignment.cc | 14 +- .../src/substitutions/pcg_pattern.cc | 12 +- .../src/substitutions/pcg_pattern_match.cc | 22 +- .../sub_parallel_computation_graph.cc | 28 +-- .../src/substitutions/substitution.cc | 10 +- .../src/substitutions/substitution_builder.cc | 14 +- .../substitutions/unity_substitution_set.cc | 32 +-- .../unlabelled/find_pattern_matches.cc | 16 +- .../substitutions/unlabelled/pattern_edge.cc | 8 +- .../unlabelled/pattern_matching.cc | 22 +- .../substitutions/unlabelled/pattern_split.cc | 4 +- .../unlabelled/unlabelled_graph_pattern.cc | 14 +- ...lled_kwarg_dataflow_graph_pattern_match.cc | 12 +- .../apply_substitution/apply_substitution.cc | 2 +- .../evaluate_substitution_output.cc | 10 +- .../perform_shape_inference.cc | 2 +- .../test/src/substitutions/pcg_pattern.cc | 10 +- .../src/substitutions/substitution_builder.cc | 2 +- .../substitutions/unity_substitution_set.cc | 32 +-- .../unlabelled/find_pattern_matches.cc | 28 +-- .../unlabelled/pattern_matching.cc | 28 +-- .../substitutions/unlabelled/pattern_split.cc | 24 +-- .../task-spec/dynamic_graph/copy_insertion.h | 8 +- .../dynamic_node_invocation.dtg.toml | 2 +- .../dynamic_open_dataflow_graph.dtg.toml | 6 +- .../dynamic_open_dataflow_graph.h | 12 +- .../task-spec/dynamic_graph/machine_slicing.h | 2 +- ...ializable_dynamic_node_invocation.dtg.toml | 2 +- .../task-spec/dynamic_graph/shard_expansion.h | 4 +- .../dynamic_graph/update_insertion.h | 2 +- .../task-spec/dynamic_graph/copy_insertion.cc | 22 +- .../dynamic_open_dataflow_graph.cc | 55 +++-- .../dynamic_graph/machine_slicing.cc | 4 +- ...ake_dynamic_open_dataflow_graph_from_cg.cc | 10 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 5 +- .../task-spec/dynamic_graph/pass_expansion.cc | 4 +- .../dynamic_graph/shard_expansion.cc | 14 +- .../dynamic_graph/update_insertion.cc | 10 +- .../task-spec/dynamic_graph/copy_insertion.cc | 26 +-- .../dynamic_open_dataflow_graph.cc | 2 +- .../dynamic_graph/machine_slicing.cc | 2 +- .../task-spec/dynamic_graph/pass_expansion.cc | 4 +- .../dynamic_graph/shard_expansion.cc | 22 +- .../benchmark/src/internal/random_dag.cc | 2 +- .../utils/bidict/algorithms/bidict_from_map.h | 4 +- .../bidict_from_unstructured_relation.h | 2 +- .../bidict_transform_keys_and_values.h | 24 +++ .../utils/bidict/algorithms/left_entries.h | 6 +- .../utils/bidict/algorithms/right_entries.h | 6 +- .../unstructured_relation_from_bidict.h | 12 +- lib/utils/include/utils/bidict/bidict.h | 93 +++++---- .../utils/cli/cli_argument_key.dtg.toml | 1 + .../include/utils/cli/cli_flag_spec.dtg.toml | 1 + .../utils/cli/cli_parse_result.dtg.toml | 11 +- .../cli/cli_positional_argument_key.dtg.toml | 1 + .../cli/cli_positional_argument_spec.dtg.toml | 3 +- lib/utils/include/utils/cli/cli_spec.dtg.toml | 7 +- lib/utils/include/utils/cli/cli_spec.h | 2 +- lib/utils/include/utils/commutative_pair.h | 4 +- .../utils/containers/are_all_distinct.h | 6 +- .../include/utils/containers/are_disjoint.h | 6 + .../containers/binary_cartesian_product.h | 11 +- .../binary_merge_disjoint_unordered_maps.h | 4 +- .../utils/containers/binary_merge_maps_with.h | 2 +- .../binary_merge_unordered_maps_with.h | 4 +- .../utils/containers/cartesian_product.h | 6 +- .../include/utils/containers/enumerate.h | 4 +- .../include/utils/containers/filter_values.h | 6 +- lib/utils/include/utils/containers/find.h | 6 +- lib/utils/include/utils/containers/flatmap.h | 20 +- .../utils/containers/get_all_assignments.h | 41 +++- .../get_all_permutations_with_repetition.h | 7 +- .../utils/containers/get_element_counts.h | 8 +- .../include/utils/containers/get_one_of.h | 6 +- lib/utils/include/utils/containers/get_only.h | 7 + lib/utils/include/utils/containers/group_by.h | 4 +- lib/utils/include/utils/containers/index.dox | 2 +- .../include/utils/containers/invert_map.h | 9 +- .../utils/containers/invert_unordered_map.h | 22 ++ .../include/utils/containers/is_submapeq_of.h | 10 +- .../utils/containers/is_superseteq_of.h | 6 +- lib/utils/include/utils/containers/keys.h | 10 - .../containers/lift_optional_through_map.h | 8 +- .../include/utils/containers/lookup_in_map.h | 6 +- .../containers/map_from_keys_and_values.h | 6 +- .../include/utils/containers/map_keys2.h | 6 +- .../containers/map_keys_with_value_merging.h | 27 +++ .../utils/containers/merge_maps_with.h | 2 +- lib/utils/include/utils/containers/minimum.h | 10 +- .../include/utils/containers/multiset_union.h | 4 +- .../utils/containers/require_two_keys.h | 4 +- .../include/utils/containers/set_difference.h | 6 +- lib/utils/include/utils/containers/set_of.h | 11 + .../include/utils/containers/set_union.h | 4 +- .../containers/try_merge_nondisjoint_maps.h | 40 ++++ .../unordered_map_from_keys_and_values.h | 26 +++ .../unstructured_exhaustive_relational_join.h | 22 +- .../include/utils/containers/value_all.h | 2 +- lib/utils/include/utils/containers/values.h | 6 +- .../utils/containers/vector_from_idx_map.h | 17 +- .../utils/containers/without_nullopts.h | 8 +- .../utils/containers/zip_values_strict_with.h | 16 +- .../utils/deduplicated_priority_queue.h | 13 +- lib/utils/include/utils/disjoint_set.h | 4 +- lib/utils/include/utils/dot/dot_file.h | 10 +- lib/utils/include/utils/fmt.h | 2 +- lib/utils/include/utils/fmt/unordered_map.h | 10 +- .../full_binary_tree/find_paths_to_leaf.h | 10 +- .../full_binary_tree/get_all_leaf_paths.h | 10 +- .../utils/full_binary_tree/get_leaves.h | 10 +- .../full_binary_tree/get_path_to_leaf_map.h | 20 +- lib/utils/include/utils/graph/algorithms.h | 40 ++-- .../utils/graph/dataflow_graph/algorithms.h | 4 +- .../algorithms/dataflow_graph_data.dtg.toml | 13 +- .../dataflow_graph_isomorphism.dtg.toml | 1 + .../algorithms/find_isomorphisms.h | 2 +- .../get_dataflow_edges_from_node_to_node.h | 2 +- .../algorithms/get_incoming_edges.h | 4 +- .../algorithms/get_outgoing_edges.h | 6 +- .../algorithms/get_subgraph_incoming_edges.h | 4 +- .../algorithms/get_subgraph_outgoing_edges.h | 4 +- ...et_transitive_reduced_edges_across_split.h | 2 +- ..._transitive_reduced_outputs_across_split.h | 2 +- .../split_boundary_nodes.dtg.toml | 10 +- .../algorithms/view_as_open_dataflow_graph.h | 8 +- .../view_from_dataflow_graph_data.h | 6 +- .../graph/dataflow_graph/dataflow_graph.h | 6 +- .../dataflow_graph/dataflow_graph_view.h | 6 +- .../dataflow_output_query.dtg.toml | 2 +- .../dataflow_graph/dataflow_output_query.h | 4 +- .../dataflow_graph/i_dataflow_graph_view.h | 6 +- .../digraph/algorithms/apply_contraction.h | 2 +- .../digraph/algorithms/calculate_topo_rank.h | 2 +- .../bipartite_component.dtg.toml | 11 +- ...bipartite_composite_decomposition.dtg.toml | 6 +- ...mplete_bipartite_composite_decomposition.h | 4 +- .../is_complete_bipartite_digraph.h | 2 +- .../graph/digraph/algorithms/contract_node.h | 4 +- .../utils/graph/digraph/algorithms/flipped.h | 4 +- .../graph/digraph/algorithms/get_ancestors.h | 2 +- .../digraph/algorithms/get_descendants.h | 2 +- .../graph/digraph/algorithms/get_dominators.h | 6 +- .../digraph/algorithms/get_dominators_map.h | 2 +- .../graph/digraph/algorithms/get_edges.h | 2 +- .../get_edges_from_subgraph_to_subgraph.h | 6 +- .../algorithms/get_imm_dominators_map.h | 2 +- .../algorithms/get_imm_post_dominator.h | 2 +- .../algorithms/get_imm_post_dominators_map.h | 2 +- .../digraph/algorithms/get_incoming_edges.h | 6 +- .../digraph/algorithms/get_initial_nodes.h | 2 +- .../get_longest_path_lengths_from_root.h | 12 +- .../algorithms/get_lowest_common_ancestors.h | 4 +- .../get_node_with_greatest_topo_rank.h | 2 +- .../digraph/algorithms/get_outgoing_edges.h | 6 +- .../digraph/algorithms/get_post_dominators.h | 2 +- .../algorithms/get_post_dominators_map.h | 2 +- .../digraph/algorithms/get_predecessors.h | 8 +- .../algorithms/get_strict_dominators.h | 2 +- .../algorithms/get_strict_dominators_map.h | 2 +- .../algorithms/get_subgraph_outgoing_edges.h | 4 +- .../algorithms/get_subgraph_successors.h | 4 +- .../graph/digraph/algorithms/get_successors.h | 8 +- .../digraph/algorithms/get_terminal_nodes.h | 2 +- .../get_weakly_connected_components.h | 2 +- .../digraph/algorithms/transitive_reduction.h | 8 +- .../include/utils/graph/digraph/digraph.h | 4 +- .../utils/graph/digraph/digraph_view.h | 4 +- .../utils/graph/digraph/i_digraph_view.h | 2 +- .../include/utils/graph/graph_split.dtg.toml | 10 +- lib/utils/include/utils/graph/index.dox | 4 +- .../utils/graph/instances/adjacency_digraph.h | 12 +- .../graph/instances/adjacency_multidigraph.h | 16 +- .../instances/hashmap_undirected_graph.h | 6 +- .../instances/unordered_set_dataflow_graph.h | 24 +-- .../unordered_set_kwarg_dataflow_graph.h | 34 +-- ...ordered_set_labelled_open_dataflow_graph.h | 56 ++--- ...d_set_labelled_open_kwarg_dataflow_graph.h | 72 +++---- .../unordered_set_open_kwarg_dataflow_graph.h | 34 +-- .../unordered_set_undirected_graph.h | 12 +- ...raph_data_from_kwarg_dataflow_graph_data.h | 24 +-- ...dataflow_graph_from_kwarg_dataflow_graph.h | 2 +- ...somorphism_between_kwarg_dataflow_graphs.h | 2 +- .../algorithms/get_all_kwarg_dataflow_edges.h | 2 +- .../get_all_kwarg_dataflow_inputs.h | 2 +- .../get_all_kwarg_dataflow_outputs.h | 2 +- ...t_incoming_kwarg_dataflow_edges_for_node.h | 6 +- ...incoming_kwarg_dataflow_outputs_for_node.h | 2 +- .../algorithms/get_incoming_slots_for_node.h | 6 +- ...t_kwarg_dataflow_edges_from_node_to_node.h | 2 +- .../get_kwarg_dataflow_graph_data.h | 6 +- .../get_kwarg_dataflow_graph_subgraph.h | 10 +- ...t_kwarg_dataflow_subgraph_incoming_edges.h | 6 +- ...t_kwarg_dataflow_subgraph_outgoing_edges.h | 6 +- .../get_kwarg_dataflow_value_uses.h | 4 +- ...outgoing_kwarg_dataflow_outputs_for_node.h | 4 +- .../algorithms/get_outgoing_slots_for_node.h | 6 +- .../algorithms/kwarg_dataflow_graph_as_dot.h | 7 +- .../kwarg_dataflow_graph_data.dtg.toml | 13 +- .../algorithms/kwarg_dataflow_graph_data.h | 10 +- ...ary_nodes_for_kwarg_dataflow_graph_split.h | 6 +- ...educed_kwarg_dataflow_edges_across_split.h | 12 +- ...uced_kwarg_dataflow_outputs_across_split.h | 2 +- .../view_as_open_kwarg_dataflow_graph.h | 8 +- .../view_from_kwarg_dataflow_graph_data.h | 12 +- .../i_kwarg_dataflow_graph.h | 8 +- .../i_kwarg_dataflow_graph_view.h | 8 +- .../kwarg_dataflow_graph.h | 14 +- .../kwarg_dataflow_graph_view.h | 6 +- .../kwarg_dataflow_output_query.dtg.toml | 2 +- .../kwarg_node_added_result.dtg.toml | 6 +- ...azy_copy_of_labelled_dataflow_graph_view.h | 6 +- .../view_as_labelled_open_dataflow_graph.h | 8 +- ...lled_kwarg_dataflow_graph_node_label_map.h | 4 +- ...ed_kwarg_dataflow_graph_output_label_map.h | 4 +- ...t_labelled_kwarg_dataflow_graph_subgraph.h | 6 +- ...kwarg_dataflow_graph_view_with_labelling.h | 18 +- ...abelled_kwarg_dataflow_graph_data.dtg.toml | 18 +- ...abelled_kwarg_dataflow_graph_view_as_dot.h | 2 +- ...ew_as_labelled_open_kwarg_dataflow_graph.h | 8 +- .../i_labelled_kwarg_dataflow_graph.h | 4 +- .../labelled_kwarg_dataflow_graph.h | 4 +- .../algorithms/find_isomorphism.h | 2 +- .../from_labelled_open_dataflow_graph_data.h | 4 +- .../algorithms/get_graph_data.h | 12 +- ...labelled_open_dataflow_graph_data.dtg.toml | 20 +- .../algorithms/permute_input_ids.h | 8 +- .../algorithms/permute_node_ids.h | 10 +- .../algorithms/rewrite_labels.h | 10 +- .../algorithms/with_labelling.h | 20 +- ...ween_labelled_open_kwarg_dataflow_graphs.h | 2 +- ..._labelled_open_kwarg_dataflow_graph_data.h | 11 +- ...ed_open_kwarg_dataflow_graph_data.dtg.toml | 22 +- .../labelled_open_kwarg_dataflow_graph_data.h | 4 +- ...ed_open_kwarg_dataflow_graph_view_as_dot.h | 2 +- ...kwarg_dataflow_graph_view_with_labelling.h | 20 +- ...lled_open_kwarg_dataflow_graph_input_ids.h | 10 +- ...elled_open_kwarg_dataflow_graph_node_ids.h | 10 +- ...abelled_open_kwarg_dataflow_graph_labels.h | 10 +- .../i_labelled_open_kwarg_dataflow_graph.h | 4 +- .../labelled_open_kwarg_dataflow_graph.h | 4 +- .../multidigraph/algorithms/get_edge_counts.h | 2 +- .../graph/multidigraph/algorithms/get_edges.h | 2 +- .../algorithms/get_incoming_edges.h | 6 +- .../get_multidiedge_to_diedge_map.h | 2 +- .../algorithms/get_outgoing_edges.h | 8 +- .../graph/multidigraph/i_multidigraph_view.h | 4 +- .../utils/graph/multidigraph/multidigraph.h | 4 +- .../graph/multidigraph/multidigraph_view.h | 4 +- .../include/utils/graph/node/algorithms.h | 2 +- lib/utils/include/utils/graph/node/graph.h | 2 +- .../include/utils/graph/node/graph_view.h | 2 +- .../include/utils/graph/node/i_graph_view.h | 2 +- .../include/utils/graph/node/node_query.h | 4 +- .../algorithms/find_isomorphisms.h | 2 +- .../from_open_dataflow_graph_data.h | 8 +- .../algorithms/get_edges.h | 2 +- .../algorithms/get_incoming_edges.h | 6 +- .../get_open_dataflow_graph_inputs.h | 2 +- .../algorithms/get_open_dataflow_value_uses.h | 2 +- .../algorithms/get_open_dataflow_values.h | 2 +- .../algorithms/get_source_nodes.h | 2 +- .../algorithms/get_subgraph.h | 6 +- .../algorithms/get_subgraph_incoming_edges.h | 4 +- .../algorithms/get_subgraph_inputs.h | 4 +- .../get_unused_open_dataflow_graph_inputs.h | 2 +- .../open_dataflow_graph_data.dtg.toml | 14 +- .../open_dataflow_graph_isomorphism.dtg.toml | 1 + .../i_open_dataflow_graph_view.h | 6 +- .../open_dataflow_edge_query.h | 4 +- .../open_dataflow_graph_view.h | 4 +- .../unordered_set_open_dataflow_graph.h | 28 +-- ...hisms_between_open_kwarg_dataflow_graphs.h | 32 +-- ...warg_dataflow_graph_input_id_permutation.h | 2 +- .../get_all_kwarg_dataflow_graph_inputs.h | 2 +- .../get_all_open_kwarg_dataflow_edges.h | 2 +- .../get_all_open_kwarg_dataflow_values.h | 6 +- ...oming_open_kwarg_dataflow_edges_for_node.h | 6 +- ...ming_open_kwarg_dataflow_values_for_node.h | 2 +- .../get_open_kwarg_dataflow_graph_data.h | 9 +- .../get_open_kwarg_dataflow_graph_subgraph.h | 30 +-- ...n_kwarg_dataflow_subgraph_incoming_edges.h | 6 +- .../get_open_kwarg_dataflow_subgraph_inputs.h | 4 +- .../get_open_kwarg_dataflow_value_uses.h | 4 +- ..._unused_open_kwarg_dataflow_graph_inputs.h | 2 +- .../open_kwarg_dataflow_graph_as_dot.h | 13 +- .../open_kwarg_dataflow_graph_data.dtg.toml | 15 +- .../open_kwarg_dataflow_graph_data.h | 2 +- ..._kwarg_dataflow_graph_isomorphism.dtg.toml | 1 + ...mute_open_kwarg_dataflow_graph_input_ids.h | 2 +- ...phism_between_open_kwarg_dataflow_graphs.h | 2 +- ...g_dataflow_graph_by_materializing_inputs.h | 20 +- ...view_from_open_kwarg_dataflow_graph_data.h | 17 +- .../i_open_kwarg_dataflow_graph.h | 4 +- .../i_open_kwarg_dataflow_graph_view.h | 8 +- .../open_kwarg_dataflow_graph.h | 4 +- .../open_kwarg_dataflow_graph_view.h | 4 +- lib/utils/include/utils/graph/query_set.h | 20 +- lib/utils/include/utils/graph/render_dot.h | 6 +- .../binary_parallel_split.dtg.toml | 1 + .../binary_series_split.dtg.toml | 1 + .../binary_sp_decomposition_tree.dtg.toml | 1 + .../binary_sp_decomposition_tree.h | 4 +- .../find_paths_to_leaf.h | 2 +- .../get_all_leaf_paths.h | 2 +- .../get_leaves.h | 2 +- .../get_path_to_leaf_map.h | 2 +- .../series_parallel/digraph_generation.h | 4 +- .../extended_parallel_reduction.dtg.toml | 11 +- .../extended_series_reduction.dtg.toml | 3 +- .../graph/series_parallel/get_ancestors.h | 4 +- .../non_normal_sp_decomposition.h | 4 +- .../series_parallel/parallel_reduction.h | 4 +- .../series_parallel_decomposition.h | 10 +- .../series_parallel/series_parallel_metrics.h | 16 +- .../graph/series_parallel/series_reduction.h | 2 +- .../sp_ization/dependencies_are_maintained.h | 2 +- .../sp_ization/escribano_algo.h | 10 +- .../sp_ization/flexible_algo.h | 6 +- .../series_parallel/sp_ization/node_role.h | 6 +- ...ization_combined_benchmark_result.dtg.toml | 4 +- .../sp_ization/up_down_partition.dtg.toml | 11 +- .../sp_ization/up_down_partition.h | 6 +- .../sp_ization/work_duplicating_sp_ization.h | 2 +- lib/utils/include/utils/graph/traversal.h | 36 ++-- .../algorithms/get_connected_components.h | 2 +- .../graph/undirected/algorithms/get_edges.h | 2 +- .../algorithms/get_neighboring_nodes.h | 2 +- .../graph/undirected/i_undirected_graph.h | 2 +- .../undirected/i_undirected_graph_view.h | 2 +- .../utils/graph/undirected/undirected_edge.h | 2 +- .../utils/graph/undirected/undirected_graph.h | 4 +- .../graph/undirected/undirected_graph_view.h | 4 +- lib/utils/include/utils/graph/views/views.h | 38 ++-- .../include/utils/many_to_one/many_to_one.h | 58 +++--- .../utils/many_to_one/many_to_one_from_map.h | 4 +- .../include/utils/nonempty_set/nonempty_set.h | 3 +- .../nonempty_unordered_set.h | 2 +- .../include/utils/one_to_many/one_to_many.h | 20 +- .../one_to_many_from_l_to_r_mapping.h | 2 +- .../one_to_many_transform_values.h | 2 +- lib/utils/include/utils/ord/unordered_map.h | 2 +- .../utils/orthotope/dim_coord.dtg.toml | 8 +- lib/utils/include/utils/orthotope/dim_coord.h | 37 ++-- .../utils/orthotope/dim_domain.dtg.toml | 8 +- .../include/utils/orthotope/dim_domain.h | 14 +- .../include/utils/orthotope/dim_projection.h | 22 +- .../include/utils/orthotope/down_projection.h | 20 +- .../include/utils/orthotope/eq_projection.h | 4 +- .../orthotope/minimal_dim_domain.dtg.toml | 8 +- .../utils/orthotope/minimal_dim_domain.h | 22 +- .../orthotope/minimal_dim_domain_mapping.h | 55 ++--- lib/utils/include/utils/orthotope/orthotope.h | 2 +- .../include/utils/orthotope/up_projection.h | 26 +-- lib/utils/include/utils/record_formatter.h | 2 +- .../bidict/algorithms/bidict_filter_keys.cc | 6 +- .../bidict/algorithms/bidict_filter_values.cc | 6 +- .../bidict/algorithms/bidict_filtrans_keys.cc | 8 +- .../algorithms/bidict_filtrans_values.cc | 8 +- .../algorithms/bidict_from_enumerating.cc | 4 +- .../algorithms/bidict_from_keys_and_values.cc | 13 ++ .../bidict/algorithms/bidict_from_map.cc | 9 +- .../bidict/algorithms/bidict_from_pairs.cc | 11 + .../bidict_from_unstructured_relation.cc | 8 +- .../bidict_transform_keys_and_values.cc | 15 ++ .../algorithms/bidict_unordered_set_of.cc | 6 +- .../binary_merge_disjoint_bidicts.cc | 6 +- .../algorithms/exhaustive_relational_join.cc | 8 +- .../utils/bidict/algorithms/filter_bidict.cc | 6 +- .../utils/bidict/algorithms/left_entries.cc | 10 + .../algorithms/merge_disjoint_bidicts.cc | 6 +- .../utils/bidict/algorithms/right_entries.cc | 10 + .../src/utils/bidict/algorithms/transform.cc | 10 +- .../utils/bidict/algorithms/transform_keys.cc | 8 +- .../bidict/algorithms/transform_values.cc | 8 +- .../unstructured_relation_from_bidict.cc | 8 +- lib/utils/src/utils/bidict/bidict.cc | 22 +- lib/utils/src/utils/cli/cli_parse.cc | 4 +- .../src/utils/containers/are_disjoint.cc | 14 ++ lib/utils/src/utils/containers/argmax.cc | 6 +- lib/utils/src/utils/containers/argmin.cc | 6 +- .../containers/binary_cartesian_product.cc | 15 +- ...rge_unordered_maps_with_left_dominating.cc | 3 +- ...ge_unordered_maps_with_right_dominating.cc | 9 +- .../utils/containers/contains_duplicates.cc | 4 +- .../src/utils/containers/contains_value.cc | 6 +- lib/utils/src/utils/containers/enumerate.cc | 2 +- lib/utils/src/utils/containers/extend.cc | 4 +- .../src/utils/containers/extend_vector.cc | 4 +- lib/utils/src/utils/containers/flatmap.cc | 9 + .../containers/generate_unordered_map.cc | 4 +- .../utils/containers/get_all_assignments.cc | 8 +- .../get_all_permutations_with_repetition.cc | 6 +- .../utils/containers/get_element_counts.cc | 2 +- lib/utils/src/utils/containers/get_only.cc | 8 +- lib/utils/src/utils/containers/group_by.cc | 9 +- lib/utils/src/utils/containers/invert_map.cc | 11 +- .../utils/containers/invert_unordered_map.cc | 13 ++ .../src/utils/containers/is_submapeq_of.cc | 2 +- .../src/utils/containers/is_subseteq_of.cc | 4 +- lib/utils/src/utils/containers/items.cc | 4 +- lib/utils/src/utils/containers/keys.cc | 1 - .../containers/lift_optional_through_map.cc | 10 +- .../src/utils/containers/lookup_in_map.cc | 5 +- .../containers/map_from_keys_and_values.cc | 5 +- .../src/utils/containers/map_from_pairs.cc | 2 +- lib/utils/src/utils/containers/map_keys.cc | 18 +- lib/utils/src/utils/containers/map_keys2.cc | 9 +- .../containers/map_keys_with_value_merging.cc | 9 + lib/utils/src/utils/containers/map_values2.cc | 20 +- .../merge_disjoint_unordered_maps.cc | 5 +- .../containers/merge_unordered_maps_with.cc | 6 +- .../src/utils/containers/multiset_union.cc | 23 ++ .../src/utils/containers/require_all_of.cc | 8 +- .../src/utils/containers/require_all_same.cc | 4 +- .../src/utils/containers/require_all_same1.cc | 6 +- .../src/utils/containers/require_two_keys.cc | 5 +- .../src/utils/containers/restrict_keys.cc | 5 +- lib/utils/src/utils/containers/set_of.cc | 18 ++ lib/utils/src/utils/containers/set_union.cc | 11 +- .../src/utils/containers/transform_pairs.cc | 10 +- .../src/utils/containers/try_get_one_of.cc | 2 +- .../containers/try_merge_nondisjoint_maps.cc | 15 ++ .../try_merge_nondisjoint_unordered_maps.cc | 13 ++ .../src/utils/containers/unordered_items.cc | 9 +- .../src/utils/containers/unordered_keys.cc | 6 +- .../unordered_map_from_keys_and_values.cc | 14 ++ .../containers/unordered_map_from_pairs.cc | 2 +- .../utils/containers/unordered_multiset_of.cc | 2 +- ...unstructured_exhaustive_relational_join.cc | 14 +- .../utils/containers/vector_from_idx_map.cc | 3 + lib/utils/src/utils/containers/vector_of.cc | 6 +- .../src/utils/containers/without_nullopts.cc | 4 +- .../containers/zip_values_strict_with.cc | 7 +- lib/utils/src/utils/disjoint_set.cc | 4 +- lib/utils/src/utils/fmt/unordered_map.cc | 2 +- lib/utils/src/utils/fmt/unordered_multiset.cc | 2 +- lib/utils/src/utils/fmt/unordered_set.cc | 2 +- .../full_binary_tree/find_paths_to_leaf.cc | 2 +- .../full_binary_tree/get_all_leaf_paths.cc | 2 +- .../src/utils/full_binary_tree/get_leaves.cc | 5 +- .../full_binary_tree/get_path_to_leaf_map.cc | 2 +- lib/utils/src/utils/graph/algorithms.cc | 32 +-- .../utils/graph/dataflow_graph/algorithms.cc | 4 +- .../algorithms/find_isomorphism.cc | 2 +- .../algorithms/find_isomorphisms.cc | 4 +- .../get_dataflow_edges_from_node_to_node.cc | 2 +- .../algorithms/get_incoming_edges.cc | 4 +- .../algorithms/get_outgoing_edges.cc | 6 +- .../algorithms/get_subgraph_incoming_edges.cc | 6 +- .../algorithms/get_subgraph_outgoing_edges.cc | 6 +- ...sitive_reduced_boundary_nodes_for_split.cc | 6 +- ...t_transitive_reduced_edges_across_split.cc | 12 +- ...transitive_reduced_outputs_across_split.cc | 2 +- .../algorithms/view_as_open_dataflow_graph.cc | 10 +- .../view_from_dataflow_graph_data.cc | 12 +- .../graph/dataflow_graph/dataflow_graph.cc | 6 +- .../dataflow_graph/dataflow_graph_view.cc | 6 +- .../dataflow_graph/dataflow_output_query.cc | 4 +- .../dataflow_graph/i_dataflow_graph_view.cc | 4 +- .../digraph/algorithms/apply_contraction.cc | 2 +- .../digraph/algorithms/calculate_topo_rank.cc | 4 +- ...plete_bipartite_composite_decomposition.cc | 10 +- .../get_cbc_decomposition.cc | 12 +- .../is_complete_bipartite_digraph.cc | 8 +- .../graph/digraph/algorithms/contract_node.cc | 4 +- .../utils/graph/digraph/algorithms/flipped.cc | 6 +- .../graph/digraph/algorithms/get_ancestors.cc | 2 +- .../digraph/algorithms/get_descendants.cc | 6 +- .../digraph/algorithms/get_dominators.cc | 10 +- .../digraph/algorithms/get_dominators_map.cc | 16 +- .../graph/digraph/algorithms/get_edges.cc | 2 +- .../get_edges_from_subgraph_to_subgraph.cc | 6 +- .../algorithms/get_imm_dominators_map.cc | 14 +- .../algorithms/get_imm_post_dominator.cc | 8 +- .../algorithms/get_imm_post_dominators_map.cc | 2 +- .../digraph/algorithms/get_incoming_edges.cc | 19 +- .../digraph/algorithms/get_initial_nodes.cc | 6 +- .../get_longest_path_lengths_from_root.cc | 16 +- .../algorithms/get_lowest_common_ancestors.cc | 14 +- .../get_node_with_greatest_topo_rank.cc | 4 +- .../digraph/algorithms/get_outgoing_edges.cc | 17 +- .../digraph/algorithms/get_post_dominators.cc | 2 +- .../algorithms/get_post_dominators_map.cc | 2 +- .../digraph/algorithms/get_predecessors.cc | 12 +- .../algorithms/get_strict_dominators.cc | 4 +- .../algorithms/get_strict_dominators_map.cc | 6 +- .../algorithms/get_subgraph_outgoing_edges.cc | 6 +- .../algorithms/get_subgraph_successors.cc | 6 +- .../digraph/algorithms/get_successors.cc | 8 +- .../digraph/algorithms/get_terminal_nodes.cc | 2 +- .../algorithms/get_topological_ordering.cc | 4 +- ...topological_ordering_from_starting_node.cc | 2 +- .../get_weakly_connected_components.cc | 2 +- .../get_inverse_line_graph.cc | 4 +- .../graph/digraph/algorithms/is_acyclic.cc | 8 +- .../digraph/algorithms/transitive_closure.cc | 2 +- .../algorithms/transitive_reduction.cc | 8 +- lib/utils/src/utils/graph/digraph/digraph.cc | 4 +- .../src/utils/graph/digraph/digraph_view.cc | 4 +- .../graph/digraph/directed_edge_query.cc | 4 +- .../graph/instances/adjacency_digraph.cc | 8 +- .../graph/instances/adjacency_multidigraph.cc | 46 ++-- .../instances/hashmap_undirected_graph.cc | 10 +- .../instances/unordered_set_dataflow_graph.cc | 26 +-- .../unordered_set_undirected_graph.cc | 8 +- ...aph_data_from_kwarg_dataflow_graph_data.cc | 2 +- ...ataflow_graph_from_kwarg_dataflow_graph.cc | 2 +- .../get_all_kwarg_dataflow_edges.cc | 2 +- .../get_all_kwarg_dataflow_inputs.cc | 2 +- .../get_all_kwarg_dataflow_outputs.cc | 2 +- ..._incoming_kwarg_dataflow_edges_for_node.cc | 2 +- ...ncoming_kwarg_dataflow_outputs_for_node.cc | 2 +- .../algorithms/get_incoming_slots_for_node.cc | 2 +- ..._kwarg_dataflow_edges_from_node_to_node.cc | 2 +- .../get_kwarg_dataflow_graph_subgraph.cc | 2 +- ..._kwarg_dataflow_subgraph_incoming_edges.cc | 4 +- ..._kwarg_dataflow_subgraph_outgoing_edges.cc | 4 +- .../get_kwarg_dataflow_value_uses.cc | 2 +- ...utgoing_kwarg_dataflow_outputs_for_node.cc | 2 +- .../algorithms/get_outgoing_slots_for_node.cc | 2 +- .../algorithms/kwarg_dataflow_graph_as_dot.cc | 2 +- ...duced_kwarg_dataflow_edges_across_split.cc | 2 +- ...ced_kwarg_dataflow_outputs_across_split.cc | 2 +- ...led_kwarg_dataflow_graph_node_label_map.cc | 2 +- ...d_kwarg_dataflow_graph_output_label_map.cc | 2 +- ..._labelled_kwarg_dataflow_graph_subgraph.cc | 2 +- ...warg_dataflow_graph_view_with_labelling.cc | 4 +- ...belled_kwarg_dataflow_graph_view_as_dot.cc | 2 +- ...d_open_kwarg_dataflow_graph_view_as_dot.cc | 2 +- ...warg_dataflow_graph_view_with_labelling.cc | 4 +- .../algorithms/get_edge_counts.cc | 2 +- .../multidigraph/algorithms/get_edges.cc | 2 +- .../algorithms/get_incoming_edges.cc | 15 +- .../get_multidiedge_to_diedge_map.cc | 6 +- .../algorithms/get_outgoing_edges.cc | 17 +- .../graph/multidigraph/i_multidigraph_view.cc | 2 +- .../utils/graph/multidigraph/multidigraph.cc | 4 +- .../graph/multidigraph/multidigraph_view.cc | 4 +- lib/utils/src/utils/graph/node/algorithms.cc | 2 +- lib/utils/src/utils/graph/node/graph.cc | 2 +- lib/utils/src/utils/graph/node/graph_view.cc | 2 +- lib/utils/src/utils/graph/node/node_query.cc | 6 +- .../algorithms/find_isomorphism.cc | 2 +- .../algorithms/find_isomorphisms.cc | 42 ++-- .../from_open_dataflow_graph_data.cc | 8 +- .../algorithms/get_edges.cc | 2 +- .../algorithms/get_incoming_edge.cc | 2 +- .../algorithms/get_incoming_edges.cc | 12 +- .../get_open_dataflow_graph_inputs.cc | 2 +- .../get_open_dataflow_value_uses.cc | 4 +- .../algorithms/get_open_dataflow_values.cc | 4 +- .../algorithms/get_source_nodes.cc | 2 +- .../algorithms/get_subgraph.cc | 18 +- .../algorithms/get_subgraph_incoming_edges.cc | 6 +- .../algorithms/get_subgraph_inputs.cc | 6 +- .../get_unused_open_dataflow_graph_inputs.cc | 2 +- .../i_open_dataflow_graph_view.cc | 4 +- .../open_dataflow_edge_query.cc | 4 +- .../open_dataflow_graph_view.cc | 4 +- .../unordered_set_open_dataflow_graph.cc | 22 +- ...isms_between_open_kwarg_dataflow_graphs.cc | 2 +- .../get_all_kwarg_dataflow_graph_inputs.cc | 2 +- .../get_all_open_kwarg_dataflow_edges.cc | 2 +- .../get_all_open_kwarg_dataflow_values.cc | 2 +- ...ming_open_kwarg_dataflow_edges_for_node.cc | 2 +- ...ing_open_kwarg_dataflow_values_for_node.cc | 2 +- .../get_open_kwarg_dataflow_graph_subgraph.cc | 6 +- ..._kwarg_dataflow_subgraph_incoming_edges.cc | 4 +- ...get_open_kwarg_dataflow_subgraph_inputs.cc | 4 +- .../get_open_kwarg_dataflow_value_uses.cc | 2 +- ...unused_open_kwarg_dataflow_graph_inputs.cc | 2 +- .../open_kwarg_dataflow_graph_as_dot.cc | 2 +- lib/utils/src/utils/graph/render_dot.cc | 8 +- .../balanced_binary_sp_tree_from_nary.cc | 2 +- .../binary_sp_decomposition_tree.cc | 2 +- .../find_paths_to_leaf.cc | 2 +- .../get_all_leaf_paths.cc | 2 +- .../get_leaves.cc | 9 +- .../get_path_to_leaf_map.cc | 2 +- .../series_parallel/digraph_generation.cc | 12 +- .../graph/series_parallel/get_ancestors.cc | 18 +- .../get_series_parallel_decomposition.cc | 16 +- .../non_normal_sp_decomposition.cc | 10 +- .../normalize_sp_decomposition.cc | 4 +- .../series_parallel/parallel_reduction.cc | 22 +- .../series_parallel_decomposition.cc | 24 +-- .../series_parallel_metrics.cc | 34 +-- .../graph/series_parallel/series_reduction.cc | 14 +- .../sp_ization/dependencies_are_maintained.cc | 8 +- .../sp_ization/escribano_algo.cc | 78 +++---- .../sp_ization/flexible_algo.cc | 116 +++++------ .../sp_ization/naive_stratum_sync.cc | 20 +- .../series_parallel/sp_ization/node_role.cc | 8 +- .../sp_ization/up_down_partition.cc | 4 +- .../sp_ization/work_duplicating_sp_ization.cc | 24 +-- lib/utils/src/utils/graph/traversal.cc | 28 +-- .../algorithms/get_connected_components.cc | 12 +- .../graph/undirected/algorithms/get_edges.cc | 2 +- .../algorithms/get_neighboring_nodes.cc | 8 +- .../utils/graph/undirected/undirected_edge.cc | 2 +- .../graph/undirected/undirected_graph.cc | 4 +- .../graph/undirected/undirected_graph_view.cc | 4 +- lib/utils/src/utils/graph/views/views.cc | 40 ++-- lib/utils/src/utils/hash/unordered_map.cc | 2 +- .../src/utils/hash/unordered_multiset.cc | 2 +- lib/utils/src/utils/hash/unordered_set.cc | 2 +- .../src/utils/many_to_one/many_to_one.cc | 16 +- .../many_to_one/many_to_one_from_bidict.cc | 6 +- .../utils/many_to_one/many_to_one_from_map.cc | 13 +- .../src/utils/one_to_many/one_to_many.cc | 2 +- .../one_to_many_from_l_to_r_mapping.cc | 2 +- lib/utils/src/utils/orthotope/dim_coord.cc | 16 +- lib/utils/src/utils/orthotope/dim_domain.cc | 10 +- .../src/utils/orthotope/dim_projection.cc | 4 +- .../src/utils/orthotope/down_projection.cc | 6 +- .../src/utils/orthotope/eq_projection.cc | 4 +- .../src/utils/orthotope/minimal_dim_domain.cc | 12 +- .../orthotope/minimal_dim_domain_mapping.cc | 4 +- lib/utils/src/utils/orthotope/orthotope.cc | 8 +- .../src/utils/orthotope/up_projection.cc | 6 +- lib/utils/src/utils/record_formatter.cc | 2 +- .../utils/doctest/check_without_stringify.h | 4 +- .../test/utils/doctest/fmt/unordered_map.h | 6 +- .../utils/doctest/fmt/unordered_multiset.h | 6 +- .../test/utils/doctest/fmt/unordered_set.h | 6 +- .../include/test/utils/rapidcheck/gen.h | 6 +- .../test/utils/doctest/fmt/unordered_map.cc | 2 +- .../utils/doctest/fmt/unordered_multiset.cc | 2 +- .../test/utils/doctest/fmt/unordered_set.cc | 2 +- .../algorithms/bidict_from_enumerating.cc | 14 +- .../bidict/algorithms/bidict_from_map.cc | 6 +- .../bidict_from_unstructured_relation.cc | 8 +- .../algorithms/bidict_unordered_set_of.cc | 4 +- .../unstructured_relation_from_bidict.cc | 6 +- lib/utils/test/src/utils/bidict/bidict.cc | 14 +- lib/utils/test/src/utils/commutative_pair.cc | 4 +- .../test/src/utils/containers/are_disjoint.cc | 18 +- lib/utils/test/src/utils/containers/argmax.cc | 6 +- lib/utils/test/src/utils/containers/argmin.cc | 6 +- .../containers/binary_cartesian_product.cc | 26 +-- .../binary_merge_disjoint_unordered_maps.cc | 18 +- .../binary_merge_unordered_maps_with.cc | 54 ++--- ...rge_unordered_maps_with_left_dominating.cc | 16 +- ...ge_unordered_maps_with_right_dominating.cc | 16 +- .../src/utils/containers/cartesian_product.cc | 32 +-- .../test/src/utils/containers/contains.cc | 6 +- .../utils/containers/contains_duplicates.cc | 2 +- .../test/src/utils/containers/contains_key.cc | 6 +- .../src/utils/containers/contains_value.cc | 4 +- .../test/src/utils/containers/enumerate.cc | 20 +- lib/utils/test/src/utils/containers/extend.cc | 8 +- lib/utils/test/src/utils/containers/filter.cc | 34 +-- .../test/src/utils/containers/filter_keys.cc | 10 +- .../src/utils/containers/filtermap_keys.cc | 10 +- .../src/utils/containers/filtermap_values.cc | 10 +- .../test/src/utils/containers/filtrans.cc | 10 +- lib/utils/test/src/utils/containers/find.cc | 6 +- .../test/src/utils/containers/flatmap.cc | 40 ++-- .../utils/containers/get_all_assignments.cc | 22 +- .../utils/containers/get_all_permutations.cc | 22 +- .../get_all_permutations_with_repetition.cc | 22 +- .../utils/containers/get_element_counts.cc | 6 +- .../test/src/utils/containers/get_one_of.cc | 6 +- .../test/src/utils/containers/get_only.cc | 2 +- .../test/src/utils/containers/group_by.cc | 12 +- .../src/utils/containers/inplace_filter.cc | 20 +- .../src/utils/containers/is_submapeq_of.cc | 14 +- .../src/utils/containers/is_subseteq_of.cc | 6 +- .../src/utils/containers/is_superseteq_of.cc | 8 +- lib/utils/test/src/utils/containers/keys.cc | 4 +- .../containers/lift_optional_through_map.cc | 20 +- .../src/utils/containers/lookup_in_map.cc | 4 +- .../test/src/utils/containers/map_keys.cc | 12 +- .../test/src/utils/containers/map_keys2.cc | 10 +- .../utils/containers/map_keys_and_values.cc | 10 +- .../test/src/utils/containers/map_values.cc | 10 +- .../test/src/utils/containers/map_values2.cc | 8 +- .../merge_disjoint_unordered_maps.cc | 32 +-- .../containers/merge_unordered_maps_with.cc | 56 ++--- .../src/utils/containers/multiset_union.cc | 14 +- .../src/utils/containers/permute_with_key.cc | 12 +- .../test/src/utils/containers/product.cc | 4 +- lib/utils/test/src/utils/containers/range.cc | 4 +- .../src/utils/containers/repeat_element.cc | 12 +- .../src/utils/containers/require_all_same1.cc | 10 +- .../utils/containers/require_no_duplicates.cc | 20 +- .../src/utils/containers/require_only_key.cc | 8 +- .../src/utils/containers/require_two_keys.cc | 10 +- .../src/utils/containers/restrict_keys.cc | 10 +- .../src/utils/containers/set_intersection.cc | 6 +- .../test/src/utils/containers/set_union.cc | 12 +- .../test/src/utils/containers/sorted_by.cc | 6 +- .../test/src/utils/containers/sum_where.cc | 2 +- .../test/src/utils/containers/transform.cc | 10 +- lib/utils/test/src/utils/containers/try_at.cc | 2 +- .../src/utils/containers/try_get_one_of.cc | 6 +- .../try_merge_nondisjoint_unordered_maps.cc | 48 ++--- .../containers/unordered_map_from_pairs.cc | 24 +-- .../utils/containers/unordered_multiset_of.cc | 10 +- .../src/utils/containers/unordered_set_of.cc | 10 +- ...unstructured_exhaustive_relational_join.cc | 18 +- lib/utils/test/src/utils/containers/values.cc | 10 +- .../src/utils/containers/zip_values_strict.cc | 18 +- lib/utils/test/src/utils/dot/dot_file.cc | 2 +- lib/utils/test/src/utils/fmt/unordered_map.cc | 12 +- lib/utils/test/src/utils/fmt/unordered_set.cc | 6 +- lib/utils/test/src/utils/graph/algorithms.cc | 8 +- lib/utils/test/src/utils/graph/cow_ptr_t.cc | 2 +- .../utils/graph/dataflow_graph/algorithms.cc | 2 +- .../algorithms/dataflow_graph_as_dot.cc | 2 +- .../get_dataflow_edges_from_node_to_node.cc | 20 +- .../algorithms/get_outgoing_edges.cc | 26 +-- .../algorithms/get_subgraph_incoming_edges.cc | 8 +- .../algorithms/get_subgraph_outgoing_edges.cc | 8 +- ...t_transitive_reduced_edges_across_split.cc | 12 +- ...transitive_reduced_outputs_across_split.cc | 4 +- .../digraph/algorithms/apply_contraction.cc | 8 +- ...plete_bipartite_composite_decomposition.cc | 12 +- .../get_cbc_decomposition.cc | 2 +- .../is_complete_bipartite_digraph.cc | 14 +- .../graph/digraph/algorithms/contract_node.cc | 8 +- .../utils/graph/digraph/algorithms/flipped.cc | 8 +- .../digraph/algorithms/get_descendants.cc | 56 ++--- .../digraph/algorithms/get_dominators.cc | 20 +- .../digraph/algorithms/get_dominators_map.cc | 4 +- .../graph/digraph/algorithms/get_edges.cc | 14 +- .../get_edges_from_subgraph_to_subgraph.cc | 48 ++--- .../algorithms/get_imm_dominators_map.cc | 4 +- .../algorithms/get_imm_post_dominators_map.cc | 16 +- .../digraph/algorithms/get_incoming_edges.cc | 4 +- .../digraph/algorithms/get_initial_nodes.cc | 18 +- .../get_longest_path_lengths_from_root.cc | 4 +- .../algorithms/get_lowest_common_ancestors.cc | 92 ++++---- .../digraph/algorithms/get_outgoing_edges.cc | 4 +- .../algorithms/get_post_dominators_map.cc | 12 +- .../digraph/algorithms/get_predecessors.cc | 4 +- .../digraph/algorithms/get_successors.cc | 4 +- .../digraph/algorithms/get_terminal_nodes.cc | 14 +- .../get_weakly_connected_components.cc | 24 +-- .../get_inverse_line_graph.cc | 24 +-- .../digraph/algorithms/transitive_closure.cc | 8 +- .../algorithms/transitive_reduction.cc | 40 ++-- .../graph/instances/adjacency_digraph.cc | 32 +-- .../graph/instances/adjacency_multidigraph.cc | 44 ++-- .../instances/unordered_set_dataflow_graph.cc | 40 ++-- .../unordered_set_kwarg_dataflow_graph.cc | 40 ++-- ..._set_labelled_open_kwarg_dataflow_graph.cc | 40 ++-- ...unordered_set_open_kwarg_dataflow_graph.cc | 40 ++-- ...aph_data_from_kwarg_dataflow_graph_data.cc | 4 +- ...ataflow_graph_from_kwarg_dataflow_graph.cc | 16 +- .../get_kwarg_dataflow_graph_subgraph.cc | 12 +- .../view_from_kwarg_dataflow_graph_data.cc | 19 +- .../get_labelled_kwarg_dataflow_graph_data.cc | 22 +- ...led_kwarg_dataflow_graph_node_label_map.cc | 18 +- ...d_kwarg_dataflow_graph_output_label_map.cc | 18 +- ..._labelled_kwarg_dataflow_graph_subgraph.cc | 26 +-- ...warg_dataflow_graph_view_with_labelling.cc | 4 +- .../multidigraph/algorithms/add_edges.cc | 6 +- .../multidigraph/algorithms/add_nodes.cc | 4 +- .../multidigraph/algorithms/get_edges.cc | 4 +- .../algorithms/get_incoming_edges.cc | 18 +- .../algorithms/get_outgoing_edges.cc | 18 +- .../utils/graph/multidigraph/multidigraph.cc | 70 +++---- .../get_open_dataflow_graph_inputs.cc | 4 +- .../get_open_dataflow_value_uses.cc | 8 +- .../algorithms/get_subgraph.cc | 29 +-- .../get_unused_open_dataflow_graph_inputs.cc | 8 +- .../algorithms/permute_node_ids.cc | 24 +-- ...rg_dataflow_graphs_are_isomorphic_under.cc | 4 +- ..._dataflow_graph_by_materializing_inputs.cc | 20 +- .../balanced_binary_sp_tree_from_nary.cc | 8 +- .../get_leaves.cc | 26 +-- ...ft_associative_binary_sp_tree_from_nary.cc | 10 +- ...ht_associative_binary_sp_tree_from_nary.cc | 10 +- .../graph/series_parallel/get_ancestors.cc | 26 +-- .../series_parallel/parallel_reduction.cc | 32 +-- .../series_parallel_decomposition.cc | 18 +- .../graph/series_parallel/series_reduction.cc | 56 ++--- .../sp_ization/escribano_algo.cc | 56 ++--- .../sp_ization/flexible_algo.cc | 32 +-- .../sp_ization/naive_stratum_sync.cc | 8 +- .../series_parallel/sp_ization/node_role.cc | 10 +- .../sp_ization/work_duplicating_sp_ization.cc | 8 +- .../algorithms/get_connected_components.cc | 22 +- .../graph/undirected/undirected_graph.cc | 32 +-- lib/utils/test/src/utils/graph/views/views.cc | 42 ++-- .../test/src/utils/many_to_one/many_to_one.cc | 20 +- .../utils/nonnegative_int/nonnegative_int.cc | 2 +- .../test/src/utils/one_to_many/one_to_many.cc | 14 +- .../test/src/utils/orthotope/dim_coord.cc | 26 +-- .../test/src/utils/orthotope/dim_domain.cc | 20 +- .../src/utils/orthotope/dim_projection.cc | 12 +- .../src/utils/positive_int/positive_int.cc | 4 +- 1042 files changed, 5710 insertions(+), 5203 deletions(-) create mode 100644 lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h create mode 100644 lib/utils/include/utils/containers/invert_unordered_map.h create mode 100644 lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h create mode 100644 lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h create mode 100644 lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc create mode 100644 lib/utils/src/utils/containers/invert_unordered_map.cc create mode 100644 lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc create mode 100644 lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc diff --git a/bin/run-model/src/run-model/main.cc b/bin/run-model/src/run-model/main.cc index 46e1186b6b..e794356699 100644 --- a/bin/run-model/src/run-model/main.cc +++ b/bin/run-model/src/run-model/main.cc @@ -86,7 +86,7 @@ int main(int argc, char **argv) { /*nesterov=*/false, /*weight_decay=*/0.001}}; - std::unordered_map input_tensors; + std::map input_tensors; DistributedFfHandle device_handle = create_distributed_ff_handle(ctx, diff --git a/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h b/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h index 08f0534b07..572c9c5f5c 100644 --- a/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h +++ b/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h @@ -3,8 +3,8 @@ #include "utils/graph/node/node.dtg.h" #include -#include -#include +#include +#include namespace FlexFlow { @@ -55,10 +55,10 @@ struct GaussianNoise { }; template -std::unordered_map - make_cost_map(std::unordered_set const &nodes, +std::map + make_cost_map(std::set const &nodes, Dist const &distribution) { - std::unordered_map cost_map; + std::map cost_map; for (Node const &node : nodes) { cost_map[node] = distribution(); } @@ -66,10 +66,10 @@ std::unordered_map } template -std::unordered_map - add_noise_to_cost_map(std::unordered_map cost_map, +std::map + add_noise_to_cost_map(std::map cost_map, Noise const &noise) { - std::unordered_map noisy_cost_map; + std::map noisy_cost_map; for (auto const &[node, cost] : cost_map) { noisy_cost_map[node] = noise() * cost; } diff --git a/bin/sp-ization-benchmarking/src/sp-ization-benchmarking/sample_graphs.cc b/bin/sp-ization-benchmarking/src/sp-ization-benchmarking/sample_graphs.cc index d88d0d451d..72c257382d 100644 --- a/bin/sp-ization-benchmarking/src/sp-ization-benchmarking/sample_graphs.cc +++ b/bin/sp-ization-benchmarking/src/sp-ization-benchmarking/sample_graphs.cc @@ -108,7 +108,7 @@ DiGraph make_full_taso_nasnet(size_t num_reduction_cells, size_t N) { (i % (N + 1) == N) ? make_reduction_taso_nasnet_cell() : make_normal_taso_nasnet_cell(); Node cell_output = get_only(get_terminal_nodes(s)); - std::unordered_map node_map = parallel_extend(g, s); + std::map node_map = parallel_extend(g, s); later_input = node_map.at(later_input); earlier_input = node_map.at(earlier_input); cell_output = node_map.at(cell_output); @@ -331,8 +331,8 @@ DiGraph make_2_terminal_random_dag(size_t num_nodes, float p, size_t step) { } } } - std::unordered_set sinks = get_terminal_nodes(g); - std::unordered_set sources = get_initial_nodes(g); + std::set sinks = get_terminal_nodes(g); + std::set sources = get_initial_nodes(g); Node sink = g.add_node(); Node source = g.add_node(); for (Node s : sources) { diff --git a/flake.nix b/flake.nix index 3c06347214..ad71cbefb4 100644 --- a/flake.nix +++ b/flake.nix @@ -155,6 +155,7 @@ expect universal-ctags ninja + tig ]) (with pkgs.python3Packages; [ gitpython diff --git a/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.dtg.toml b/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.dtg.toml index 42435312c3..6fb5f9304e 100644 --- a/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.dtg.toml +++ b/lib/compiler/include/compiler/cost_estimator/op_cost_estimate_key.dtg.toml @@ -24,9 +24,8 @@ includes = [ ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", - "utils/ord/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] @@ -35,15 +34,15 @@ type = "::FlexFlow::PCGOperatorAttrs" [[fields]] name = "input_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "weight_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "output_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "optimizer_attrs" diff --git a/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.dtg.toml b/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.dtg.toml index 7a673c83b2..d90541c41b 100644 --- a/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.dtg.toml +++ b/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.dtg.toml @@ -3,6 +3,7 @@ name = "OpCostMetrics" type = "struct" features = [ "eq", + "ord", "fmt", "hash", ] diff --git a/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.h b/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.h index aa638f7287..62345892c7 100644 --- a/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.h +++ b/lib/compiler/include/compiler/cost_estimator/op_cost_metrics.h @@ -7,7 +7,7 @@ namespace FlexFlow { bool is_pareto_optimal_in(OpCostMetrics const &, - std::unordered_set const &); + std::set const &); OpCostMetrics make_op_cost_metrics_from_runtime_only( RuntimeOnlyOpCostMetrics const &runtime_only, diff --git a/lib/compiler/include/compiler/cost_estimator/runtime_only_op_cost_estimate_key.dtg.toml b/lib/compiler/include/compiler/cost_estimator/runtime_only_op_cost_estimate_key.dtg.toml index 99501645ce..c9d4d15063 100644 --- a/lib/compiler/include/compiler/cost_estimator/runtime_only_op_cost_estimate_key.dtg.toml +++ b/lib/compiler/include/compiler/cost_estimator/runtime_only_op_cost_estimate_key.dtg.toml @@ -17,9 +17,8 @@ includes = [ ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/ord/unordered_map.h", - "utils/fmt/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] @@ -28,15 +27,15 @@ type = "::FlexFlow::PCGOperatorAttrs" [[fields]] name = "input_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "weight_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "output_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "machine_view" diff --git a/lib/compiler/include/compiler/cost_estimator/single_tensor_movement.dtg.toml b/lib/compiler/include/compiler/cost_estimator/single_tensor_movement.dtg.toml index 9dd77f4fe2..881331f937 100644 --- a/lib/compiler/include/compiler/cost_estimator/single_tensor_movement.dtg.toml +++ b/lib/compiler/include/compiler/cost_estimator/single_tensor_movement.dtg.toml @@ -10,14 +10,14 @@ features = [ includes = [ "compiler/cost_estimator/communication_edge.h", "utils/units/num_bytes_t.h", - "", + "", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "edge_to_size" -type = "std::unordered_map<::FlexFlow::CommunicationEdge, ::FlexFlow::num_bytes_t>" +type = "std::map<::FlexFlow::CommunicationEdge, ::FlexFlow::num_bytes_t>" diff --git a/lib/compiler/include/compiler/cost_estimator/tensor_set_movement.dtg.toml b/lib/compiler/include/compiler/cost_estimator/tensor_set_movement.dtg.toml index 2660b0a3c3..164bea7cff 100644 --- a/lib/compiler/include/compiler/cost_estimator/tensor_set_movement.dtg.toml +++ b/lib/compiler/include/compiler/cost_estimator/tensor_set_movement.dtg.toml @@ -3,6 +3,7 @@ name = "TensorSetMovement" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -10,14 +11,14 @@ features = [ includes = [ "compiler/cost_estimator/communication_edge.h", "utils/units/num_bytes_t.h", - "", + "", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "edge_to_size" -type = "std::unordered_map<::FlexFlow::CommunicationEdge, ::FlexFlow::num_bytes_t>" +type = "std::map<::FlexFlow::CommunicationEdge, ::FlexFlow::num_bytes_t>" diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.h b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.h index b0a17309e9..39ebd7b1fe 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.h +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.h @@ -13,7 +13,7 @@ namespace FlexFlow { MachineSpaceCoordinate concretize_abstracted_device( AbstractedDevice const &abstracted_device, - std::unordered_map const + std::map const &machine_space_stencils); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.h b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.h index 0b4f7a3a43..937ec30233 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.h +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.h @@ -14,7 +14,7 @@ std::optional concretize_abstracted_single_tensor_communication_edge( AbstractedSingleTensorCommunicationEdge const &edge, MachineSpaceStencil const &src_machine_stencil, - std::unordered_map const + std::map const &dst_machine_stencils); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.dtg.toml b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.dtg.toml index f8a8280e9f..ae6d63eb70 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.dtg.toml @@ -3,6 +3,7 @@ name = "AbstractedSingleTensorMovement" type = "struct" features = [ "eq", + "ord", "hash", "fmt", "json", @@ -11,12 +12,12 @@ features = [ includes = [ "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.dtg.h", "utils/units/num_bytes_t.h", - "", + "", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] @@ -25,4 +26,4 @@ type = "::FlexFlow::BinaryTreePath" [[fields]] name = "edge_to_size" -type = "std::unordered_map<::FlexFlow::AbstractedSingleTensorCommunicationEdge, ::FlexFlow::num_bytes_t>" +type = "std::map<::FlexFlow::AbstractedSingleTensorCommunicationEdge, ::FlexFlow::num_bytes_t>" diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h index 1a4062cc4c..556a0a83e0 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h @@ -8,24 +8,24 @@ namespace FlexFlow { -std::unordered_set +std::set abstracted_single_tensor_movement_get_dst_layers( AbstractedSingleTensorMovement const &); AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( - std::unordered_multiset const &); + std::multiset const &); AbstractedSingleTensorMovement abstracted_single_tensor_movement_from_communications( BinaryTreePath const &src_op_tree_path, - std::unordered_set const + std::set const &communications); TensorSetMovement concretize_abstracted_single_tensor_movement( AbstractedSingleTensorMovement const &, - std::unordered_map const + std::map const &pre_machine_stencils, - std::unordered_map const + std::map const &post_machine_stencils); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.dtg.toml b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.dtg.toml index 7b1468f4c9..5a8268acfa 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.dtg.toml @@ -3,21 +3,22 @@ name = "AbstractedTensorSetMovement" type = "struct" features = [ "eq", + "ord", "hash", "fmt", "json", ] includes = [ - "", + "", "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.dtg.h", ] src_includes = [ - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h", + "utils/fmt/set.h", + "utils/hash/set.h", ] [[fields]] name = "single_tensor_movements" -type = "std::unordered_set<::FlexFlow::AbstractedSingleTensorMovement>" +type = "std::set<::FlexFlow::AbstractedSingleTensorMovement>" diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h index d925df2762..15f6654918 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h @@ -19,16 +19,16 @@ AbstractedTensorSetMovement abstracted_tensor_set_movement_from_single_tensor_movement( AbstractedSingleTensorMovement const &); -std::unordered_set +std::set get_src_layers(AbstractedTensorSetMovement const &); -std::unordered_set +std::set get_dst_layers(AbstractedTensorSetMovement const &); TensorSetMovement concretize_abstracted_tensor_set_movement( AbstractedTensorSetMovement const &, - std::unordered_map const + std::map const &pre_machine_stencils, - std::unordered_map const + std::map const &post_machine_stencils); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h b/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h index 5201f7fa31..89da978ffd 100644 --- a/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h +++ b/lib/compiler/include/compiler/machine_mapping/allowed_machine_views.h @@ -12,7 +12,7 @@ bool is_valid_machine_view(MachineView const &mv, OperatorTaskSpace const &task, MachineComputeResourceSlice const &ms); -std::unordered_set +std::set get_allowed_machine_views(MachineComputeResourceSlice const &machine_spec, OperatorTaskSpace const &task, DeviceType device_type); diff --git a/lib/compiler/include/compiler/machine_mapping/feasible_machine_mapping_result.dtg.toml b/lib/compiler/include/compiler/machine_mapping/feasible_machine_mapping_result.dtg.toml index ef2d7899ed..35f61c2838 100644 --- a/lib/compiler/include/compiler/machine_mapping/feasible_machine_mapping_result.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/feasible_machine_mapping_result.dtg.toml @@ -3,6 +3,7 @@ name = "FeasibleMachineMappingResult" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping.dtg.toml index 18b61840fd..fed1dd8c39 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping.dtg.toml @@ -13,14 +13,13 @@ features = [ includes = [ "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", "compiler/machine_mapping/machine_view.dtg.h", -] - -src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", - "utils/ord/unordered_map.h", +] + +src_includes = [ + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] name = "machine_views" -type = "std::unordered_map<::FlexFlow::parallel_layer_guid_t, ::FlexFlow::MachineView>" +type = "std::map<::FlexFlow::parallel_layer_guid_t, ::FlexFlow::MachineView>" diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_cache.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_cache.dtg.toml index 5683206177..df82b09d6d 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_cache.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_cache.dtg.toml @@ -8,16 +8,16 @@ features = [ ] includes = [ - "", + "", "compiler/machine_mapping/machine_mapping_state.dtg.h", "compiler/machine_mapping/machine_mapping_result.dtg.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "raw_map" -type = "std::unordered_map<::FlexFlow::MachineMappingState, ::FlexFlow::MachineMappingResult>" +type = "std::map<::FlexFlow::MachineMappingState, ::FlexFlow::MachineMappingResult>" diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.dtg.toml index a83a7caa02..18f0ea20a2 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.dtg.toml @@ -3,6 +3,7 @@ name = "MachineMappingConstraints" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -11,14 +12,14 @@ includes = [ "compiler/machine_mapping/machine_view.dtg.h", "utils/full_binary_tree/binary_tree_path.dtg.h", "", -] - -src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", +] + +src_includes = [ + "utils/hash/map.h", + "utils/fmt/map.h", "utils/fmt/optional.h", ] [[fields]] name = "machine_views" -type = "std::unordered_map<::FlexFlow::BinaryTreePath, std::optional<::FlexFlow::MachineView>>" +type = "std::map<::FlexFlow::BinaryTreePath, std::optional<::FlexFlow::MachineView>>" diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h index 15d02dd64a..5562c39f38 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h @@ -11,15 +11,15 @@ namespace FlexFlow { MachineMappingConstraints get_unconstrained_solution_for_layers( - std::unordered_set const &); + std::set const &); -std::unordered_set +std::set get_unconstrained_layers(MachineMappingConstraints const &); -std::unordered_set +std::set get_constrained_layers(MachineMappingConstraints const &); -std::unordered_set +std::set get_all_layers(MachineMappingConstraints const &); std::optional diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_context.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_context.dtg.toml index 04d2eb1378..3bc07377c7 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_context.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_context.dtg.toml @@ -16,4 +16,4 @@ type = "::FlexFlow::RuntimeOnlyCostEstimator" [[fields]] name = "allowed_machine_views" -type = "std::function(::FlexFlow::UnmappedRuntimeOnlyOpCostEstimateKey const &, ::FlexFlow::MachineComputeResourceSlice const &)>" +type = "std::function(::FlexFlow::UnmappedRuntimeOnlyOpCostEstimateKey const &, ::FlexFlow::MachineComputeResourceSlice const &)>" diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.dtg.toml index 2456dca145..43bedc9624 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.dtg.toml @@ -3,6 +3,7 @@ name = "MachineMappingProblemTree" type = "variant" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h index 62a4206d54..8bb78a9086 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h @@ -19,16 +19,16 @@ GenericBinarySPDecompositionTreeImplementation< SPDecompositionTreeNodeType get_node_type(MachineMappingProblemTree const &); -std::unordered_multiset +std::multiset get_leaves(MachineMappingProblemTree const &); -std::unordered_set +std::set get_all_leaf_paths(MachineMappingProblemTree const &); std::optional mm_problem_tree_get_subtree_at_path(MachineMappingProblemTree const &, BinaryTreePath const &); -std::unordered_map +std::map mm_problem_tree_get_path_to_leaf_map(MachineMappingProblemTree const &); std::string as_dot(MachineMappingProblemTree const &); diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_parallel_split.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_parallel_split.dtg.toml index b0dad05430..5183f1895b 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_parallel_split.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_parallel_split.dtg.toml @@ -3,6 +3,7 @@ name = "MMProblemTreeParallelSplit" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_series_split.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_series_split.dtg.toml index 64b05b0101..60e0023b40 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_series_split.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/mm_problem_tree_series_split.dtg.toml @@ -3,6 +3,7 @@ name = "MMProblemTreeSeriesSplit" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_op_cost_estimate_key.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_op_cost_estimate_key.dtg.toml index 4bad66f7ee..a879b319c7 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_op_cost_estimate_key.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_op_cost_estimate_key.dtg.toml @@ -10,14 +10,14 @@ features = [ includes = [ "op-attrs/pcg_operator_attrs.dtg.h", "op-attrs/parallel_tensor_shape.dtg.h", - "", + "", "pcg/optimizer_attrs.dtg.h", "op-attrs/tensor_slot_name.dtg.h", ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] @@ -26,15 +26,15 @@ type = "::FlexFlow::PCGOperatorAttrs" [[fields]] name = "input_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "weight_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "output_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "optimizer_attrs" diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_runtime_only_op_cost_estimate_key.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_runtime_only_op_cost_estimate_key.dtg.toml index 8db92162a1..f75dd8efe6 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_runtime_only_op_cost_estimate_key.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/unmapped_runtime_only_op_cost_estimate_key.dtg.toml @@ -3,6 +3,7 @@ name = "UnmappedRuntimeOnlyOpCostEstimateKey" type = "struct" features = [ "eq", + "ord", "fmt", "hash", ] @@ -15,8 +16,8 @@ includes = [ ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] @@ -25,12 +26,12 @@ type = "::FlexFlow::PCGOperatorAttrs" [[fields]] name = "input_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "weight_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" [[fields]] name = "output_shapes" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorShape>" diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.dtg.toml index 1c9b664246..cc51c3ff0c 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.dtg.toml @@ -3,6 +3,7 @@ name = "MachineMappingResult" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h index fd48f1b02c..8d5d741187 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h @@ -13,7 +13,7 @@ namespace FlexFlow { FeasibleMachineMappingResult require_feasible(MachineMappingResult const &); [[nodiscard]] MachineMappingResult get_mapping_with_minimal_runtime( - std::unordered_set const &); + std::set const &); [[nodiscard]] MachineMappingResult series_combine(milliseconds_t comm_cost, diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_state.dtg.toml b/lib/compiler/include/compiler/machine_mapping/machine_mapping_state.dtg.toml index fece560df6..e87ebc388c 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_state.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_state.dtg.toml @@ -3,6 +3,7 @@ name = "MachineMappingState" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/machine_mapping/machine_resource_split.h b/lib/compiler/include/compiler/machine_mapping/machine_resource_split.h index ce8c2029c6..2a9c23b2a3 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_resource_split.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_resource_split.h @@ -12,7 +12,7 @@ std::pair apply_resource_split(MachineResourceSplit const &split, MachineComputeResourceSlice const &resources); -std::unordered_set +std::set get_machine_resource_splits(MachineComputeResourceSlice const &); MachineSpaceCoordinate diff --git a/lib/compiler/include/compiler/machine_mapping/machine_view.h b/lib/compiler/include/compiler/machine_mapping/machine_view.h index 6888fa6b94..d22319fd1d 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_view.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_view.h @@ -14,7 +14,7 @@ #include "utils/bidict/bidict.h" #include #include -#include +#include namespace FlexFlow { @@ -46,11 +46,11 @@ OperatorSpaceToMachineSpaceMapping get_coordinate_mapping_for_machine_view( OperatorTaskSpace const &operator_task_space, MachineView const &machine_view); -std::unordered_set +std::set get_machine_space_coordinates(OperatorTaskSpace const &task, MachineView const &mv); -std::unordered_set +std::set get_device_ids(OperatorTaskSpace const &task, MachineView const &mv, MachineComputeSpecification const &ms); @@ -70,7 +70,7 @@ OperatorAtomicTaskShardBinding MappedOperatorTaskGroup mapped_operator_task_group_from_machine_view( ComputationGraphOpAttrs const &, - std::unordered_map const &, + std::map const &, MachineView const &); bidict diff --git a/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_cache.dtg.toml b/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_cache.dtg.toml index bfe5981466..d6b8c66a0e 100644 --- a/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_cache.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_cache.dtg.toml @@ -8,16 +8,16 @@ features = [ ] includes = [ - "", + "", "compiler/machine_mapping/machine_mapping_state.dtg.h", "compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "raw_map" -type = "std::unordered_map<::FlexFlow::MachineMappingState, ::FlexFlow::MachineMappingWithMemoryResult>" +type = "std::map<::FlexFlow::MachineMappingState, ::FlexFlow::MachineMappingWithMemoryResult>" diff --git a/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_context.dtg.toml b/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_context.dtg.toml index 7c31c0d16b..6253c2fbad 100644 --- a/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_context.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_context.dtg.toml @@ -21,4 +21,4 @@ type = "::FlexFlow::OptimizerAttrs" [[fields]] name = "allowed_machine_views" -type = "std::function(::FlexFlow::UnmappedRuntimeOnlyOpCostEstimateKey const &, ::FlexFlow::MachineComputeResourceSlice const &)>" +type = "std::function(::FlexFlow::UnmappedRuntimeOnlyOpCostEstimateKey const &, ::FlexFlow::MachineComputeResourceSlice const &)>" diff --git a/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.h b/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.h index ab648f48f3..325585a31c 100644 --- a/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.h +++ b/lib/compiler/include/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.h @@ -12,16 +12,21 @@ struct MachineMappingWithMemoryResult { MachineMappingWithMemoryResult() = delete; explicit MachineMappingWithMemoryResult( - std::unordered_set const &); + std::set const &); - bool operator==(MachineMappingWithMemoryResult const &) const; - bool operator!=(MachineMappingWithMemoryResult const &) const; + [[nodiscard]] bool operator==(MachineMappingWithMemoryResult const &) const; + [[nodiscard]] bool operator!=(MachineMappingWithMemoryResult const &) const; - std::unordered_set const & + [[nodiscard]] bool operator<(MachineMappingWithMemoryResult const &) const; + [[nodiscard]] bool operator>(MachineMappingWithMemoryResult const &) const; + [[nodiscard]] bool operator<=(MachineMappingWithMemoryResult const &) const; + [[nodiscard]] bool operator>=(MachineMappingWithMemoryResult const &) const; + + [[nodiscard]] std::set const & get_pareto_frontier() const; private: - std::unordered_set m_pareto_frontier; + std::set m_pareto_frontier; private: std::tuple tie() const; @@ -38,7 +43,7 @@ std::ostream &operator<<(std::ostream &, [[nodiscard]] bool is_empty(MachineMappingWithMemoryResult const &); [[nodiscard]] MachineMappingWithMemoryResult get_mapping_with_minimal_runtime( - std::unordered_set const &); + std::set const &); [[nodiscard]] MachineMappingWithMemoryResult series_combine(milliseconds_t comm_cost, diff --git a/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.dtg.toml b/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.dtg.toml index fc33be6aae..e75d676c07 100644 --- a/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.dtg.toml @@ -3,6 +3,7 @@ name = "ParetoOptimalMachineMapping" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h b/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h index 6e263fc412..dcb909d59f 100644 --- a/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h +++ b/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h @@ -7,7 +7,7 @@ namespace FlexFlow { bool is_pareto_optimal_in( ParetoOptimalMachineMapping const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.dtg.toml b/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.dtg.toml index 344817bffc..3ba7018a50 100644 --- a/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.dtg.toml @@ -3,6 +3,7 @@ name = "ParallelLayerGuidObliviousMachineMapping" type = "struct" features = [ "eq", + "ord", "hash", "fmt", "rapidcheck", @@ -14,10 +15,10 @@ includes = [ ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "raw_mapping" -type = "std::unordered_map<::FlexFlow::BinaryTreePath, ::FlexFlow::MachineView>" +type = "std::map<::FlexFlow::BinaryTreePath, ::FlexFlow::MachineView>" diff --git a/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.h b/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.h index 9f2871239d..ceaa754697 100644 --- a/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.h +++ b/lib/compiler/include/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.h @@ -23,18 +23,18 @@ std::optional get_machine_view_for_path(ParallelLayerGuidObliviousMachineMapping const &, BinaryTreePath const &); -std::unordered_map +std::map get_machine_stencils_for_decomposition( ParallelComputationGraph const &pcg, PCGBinarySPDecomposition const &decomposition, ParallelLayerGuidObliviousMachineMapping const &mapping); -std::unordered_map> +std::map> get_machine_stencils_for_mm_problem_tree( MachineMappingProblemTree const &, ParallelLayerGuidObliviousMachineMapping const &mapping); -std::unordered_map +std::map get_machine_stencils_for_partially_mapped_mm_problem_tree( MachineMappingProblemTree const &, ParallelLayerGuidObliviousMachineMapping const &); diff --git a/lib/compiler/include/compiler/machine_mapping/pcg_split_boundary_layers.dtg.toml b/lib/compiler/include/compiler/machine_mapping/pcg_split_boundary_layers.dtg.toml index cbdedd47e7..7a7b399669 100644 --- a/lib/compiler/include/compiler/machine_mapping/pcg_split_boundary_layers.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/pcg_split_boundary_layers.dtg.toml @@ -9,17 +9,17 @@ features = [ includes = [ "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", - "", + "", ] src_includes = [ - "utils/hash/unordered_set.h", "utils/fmt/unordered_set.h", + "utils/hash/set.h", "utils/fmt/set.h", ] [[fields]] name = "pre_split_boundary" -type = "std::unordered_set<::FlexFlow::parallel_layer_guid_t>" +type = "std::set<::FlexFlow::parallel_layer_guid_t>" [[fields]] name = "post_split_boundary" -type = "std::unordered_set<::FlexFlow::parallel_layer_guid_t>" +type = "std::set<::FlexFlow::parallel_layer_guid_t>" diff --git a/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h b/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h index 631d23b07c..a59a78570e 100644 --- a/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h +++ b/lib/compiler/include/compiler/machine_mapping/start_invariant_machine_view.h @@ -36,7 +36,7 @@ MachineSpaceOffset StartInvariantMachineView const &mv, TaskSpaceCoordinate const &coordinates); -std::unordered_set +std::set get_machine_space_offsets(OperatorTaskSpace const &task, StartInvariantMachineView const &mv); diff --git a/lib/compiler/include/compiler/machine_mapping/transitive_reduced_pcg.h b/lib/compiler/include/compiler/machine_mapping/transitive_reduced_pcg.h index 8055d15b4e..7c84615dfb 100644 --- a/lib/compiler/include/compiler/machine_mapping/transitive_reduced_pcg.h +++ b/lib/compiler/include/compiler/machine_mapping/transitive_reduced_pcg.h @@ -18,11 +18,11 @@ TransitiveReducedKwargDataflowGraphView TransitiveReducedPCG pcg_get_transitive_reduction(ParallelComputationGraph const &); -std::unordered_set +std::set pcg_get_transitive_reduced_edges_across_split(TransitiveReducedPCG const &, PCGBinarySeriesSplit const &); -std::unordered_set +std::set pcg_get_transitive_reduced_tensors_across_split( TransitiveReducedPCG const &, PCGBinarySeriesSplit const &); diff --git a/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml b/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml index 28391eddc0..5e9555ae5e 100644 --- a/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml +++ b/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.dtg.toml @@ -16,12 +16,12 @@ includes = [ ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h" + "utils/hash/map.h", + "utils/fmt/map.h", + "utils/hash/set.h", + "utils/fmt/set.h" ] [[fields]] name = "raw_device_map" -type = "std::unordered_map<::FlexFlow::parallel_layer_guid_t, std::unordered_set<::FlexFlow::device_id_t>>" +type = "std::map<::FlexFlow::parallel_layer_guid_t, std::set<::FlexFlow::device_id_t>>" diff --git a/lib/compiler/include/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.h b/lib/compiler/include/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.h index 8a7c467303..e08715259d 100644 --- a/lib/compiler/include/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.h +++ b/lib/compiler/include/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.h @@ -33,7 +33,7 @@ std::optional ComputationGraph const &); bool is_left_associative(ComputationGraphBinarySPDecomposition const &); bool is_right_associative(ComputationGraphBinarySPDecomposition const &); -std::unordered_multiset +std::multiset get_layers(ComputationGraphBinarySPDecomposition const &); V1BinarySPDecomposition diff --git a/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h b/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h index c680644f30..21ffc11af3 100644 --- a/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h +++ b/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h @@ -22,8 +22,8 @@ GenericBinarySPDecompositionTreeImplementation - get_parallel_layers(PCGBinarySPDecomposition const &); +std::multiset + pcg_sp_tree_get_parallel_layers(PCGBinarySPDecomposition const &); PCGBinarySPDecomposition pcg_binary_sp_decomposition_from_binary_sp_decomposition_tree( @@ -31,14 +31,14 @@ PCGBinarySPDecomposition SPDecompositionTreeNodeType get_node_type(PCGBinarySPDecomposition const &); -std::unordered_set +std::set pcg_sp_tree_get_all_leaf_paths(PCGBinarySPDecomposition const &); -std::unordered_set +std::set find_paths_to_leaf(PCGBinarySPDecomposition const &, parallel_layer_guid_t const &); -std::unordered_map +std::map pcg_sp_tree_get_path_to_leaf_map(PCGBinarySPDecomposition const &); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/task_graph_simulator/in_progress_task.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/in_progress_task.dtg.toml index 0788cb196e..d89ea1b9ee 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/in_progress_task.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/in_progress_task.dtg.toml @@ -1,7 +1,6 @@ namespace = "FlexFlow" name = "InProgressTask" type = "struct" - features = [ "eq", "hash", @@ -13,7 +12,6 @@ includes = [ "utils/graph/node/node.dtg.h" ] - [[fields]] name = "start_time" type = "float" diff --git a/lib/compiler/include/compiler/task_graph_simulator/pcg_task.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/pcg_task.dtg.toml index 48eb99e9c6..d1b0cd270b 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/pcg_task.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/pcg_task.dtg.toml @@ -3,6 +3,7 @@ name = "PCGTask" type = "variant" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/compiler/include/compiler/task_graph_simulator/task_execution_constraint.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/task_execution_constraint.dtg.toml index a39b072fb3..10ebfcaeac 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/task_execution_constraint.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/task_execution_constraint.dtg.toml @@ -7,10 +7,10 @@ features = [ includes = [ "utils/graph/node/node.dtg.h", "", - "" + "" ] [[fields]] name = "is_satisfied" -type = "std::function const &, std::unordered_set const &)>" +type = "std::function const &, std::set const &)>" diff --git a/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_state.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_state.dtg.toml index bc93b7b8bc..672a00fa6a 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_state.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_state.dtg.toml @@ -10,14 +10,14 @@ includes = [ "utils/graph/node/node.dtg.h", "compiler/task_graph_simulator/in_progress_task.dtg.h", "compiler/task_graph_simulator/in_progress_task_comparator.h", - "", + "", "", "" ] src_includes = [ - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", "utils/hash/set.h", "utils/fmt/set.h", "utils/fmt/vector.h", @@ -34,7 +34,7 @@ type = "::FlexFlow::DeduplicatedPriorityQueue<::FlexFlow::InProgressTask, std::v [[fields]] name = "finished_tasks" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "current_time" diff --git a/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_trace.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_trace.dtg.toml index 629e222920..7c336ad656 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_trace.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/task_graph_execution_trace.dtg.toml @@ -10,15 +10,15 @@ features = [ includes = [ "compiler/task_graph_simulator/task_profile.dtg.h", - "" + "" ] src_includes = [ - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h" + "utils/fmt/set.h", + "utils/hash/set.h" ] [[fields]] name = "task_profiles" -type = "std::unordered_set<::FlexFlow::TaskProfile>" +type = "std::set<::FlexFlow::TaskProfile>" diff --git a/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc b/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc index f3edd6a69a..fb4b9de77e 100644 --- a/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc +++ b/lib/compiler/src/compiler/cost_estimator/op_cost_estimate_key.cc @@ -7,7 +7,7 @@ #include "pcg/device_id_t.dtg.h" #include "pcg/machine_specification.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.dtg.h" -#include +#include namespace FlexFlow { diff --git a/lib/compiler/src/compiler/cost_estimator/op_cost_metrics.cc b/lib/compiler/src/compiler/cost_estimator/op_cost_metrics.cc index 7eab1f0a2a..e961fb28db 100644 --- a/lib/compiler/src/compiler/cost_estimator/op_cost_metrics.cc +++ b/lib/compiler/src/compiler/cost_estimator/op_cost_metrics.cc @@ -4,7 +4,7 @@ namespace FlexFlow { bool is_pareto_optimal_in(OpCostMetrics const &m, - std::unordered_set const &others) { + std::set const &others) { return all_of(others, [&](OpCostMetrics const &other) { return m.forward_runtime <= other.forward_runtime || m.backward_runtime <= other.backward_runtime || diff --git a/lib/compiler/src/compiler/cost_estimator/tensor_set_movement.cc b/lib/compiler/src/compiler/cost_estimator/tensor_set_movement.cc index d5b9d8a7f5..e3131698ac 100644 --- a/lib/compiler/src/compiler/cost_estimator/tensor_set_movement.cc +++ b/lib/compiler/src/compiler/cost_estimator/tensor_set_movement.cc @@ -3,7 +3,6 @@ #include "compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" -#include "utils/containers/unordered_multiset_of.h" #include "utils/full_binary_tree/binary_tree_path.dtg.h" namespace FlexFlow { @@ -49,11 +48,11 @@ TensorSetMovement get_tensor_set_movement_from_pcg_edge( return concretize_abstracted_tensor_set_movement( abstracted_tensor_set_movement, /*pre_machine_stencils=*/ - std::unordered_map{ + std::map{ {src_path, src_machine_stencil}, }, /*post_machine_stencils=*/ - std::unordered_map{ + std::map{ {dst_path, dst_machine_stencil}, }); } diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.cc index b5bdb42ece..34ed687ba9 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_device.cc @@ -15,7 +15,7 @@ namespace FlexFlow { MachineSpaceCoordinate concretize_abstracted_device( AbstractedDevice const &abstracted_device, - std::unordered_map const &stencils) { + std::map const &stencils) { return machine_space_stencil_compute_machine_coord( stencils.at(abstracted_device.operator_tree_path), diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.cc index 2a35b76849..ad37be27ee 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.cc @@ -9,7 +9,7 @@ std::optional concretize_abstracted_single_tensor_communication_edge( AbstractedSingleTensorCommunicationEdge const &edge, MachineSpaceStencil const &src_machine_stencil, - std::unordered_map const + std::map const &dst_machine_stencils) { MachineSpaceCoordinate src = machine_space_stencil_compute_machine_coord( diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc index e8bd602289..a789e76789 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc @@ -6,25 +6,25 @@ #include "utils/containers/require_same.h" #include "utils/containers/transform.h" #include "utils/containers/values.h" -#include "utils/containers/merge_unordered_maps_with.h" -#include "utils/containers/unordered_map_from_pairs.h" +#include "utils/containers/merge_maps_with.h" +#include "utils/containers/map_from_pairs.h" namespace FlexFlow { -std::unordered_set +std::set abstracted_single_tensor_movement_get_dst_layers( AbstractedSingleTensorMovement const &m) { return transform( - unordered_keys(m.edge_to_size), + keys(m.edge_to_size), [](AbstractedSingleTensorCommunicationEdge const &e) -> BinaryTreePath { return e.dst.operator_tree_path; }); } AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( - std::unordered_multiset const &movements) { + std::multiset const &movements) { - std::unordered_multiset src_paths = + std::multiset src_paths = transform(movements, [](AbstractedSingleTensorMovement const &m) { return m.src_op_tree_path; }); @@ -34,7 +34,7 @@ AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( return AbstractedSingleTensorMovement{ /*src_op_tree_path=*/require_all_same1(src_paths), /*edge_to_size=*/ - merge_unordered_maps_with(transform(vector_of(movements), + merge_maps_with(transform(vector_of(movements), [](AbstractedSingleTensorMovement const &m) { return m.edge_to_size; }), @@ -45,13 +45,13 @@ AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( AbstractedSingleTensorMovement abstracted_single_tensor_movement_from_communications( BinaryTreePath const &src_op_tree_path, - std::unordered_set const + std::set const &communications) { return AbstractedSingleTensorMovement{ /*src_op_tree_path=*/src_op_tree_path, /*edge_to_size=*/ - unordered_map_from_pairs( + map_from_pairs( transform(communications, [](AbstractedSingleTensorCommunication const &c) { return std::pair{c.edge, c.size}; @@ -61,16 +61,16 @@ AbstractedSingleTensorMovement TensorSetMovement concretize_abstracted_single_tensor_movement( AbstractedSingleTensorMovement const &abstracted, - std::unordered_map const + std::map const &pre_machine_stencils, - std::unordered_map const + std::map const &post_machine_stencils) { ASSERT(contains_key(pre_machine_stencils, abstracted.src_op_tree_path)); MachineSpaceStencil pre_machine_stencil = pre_machine_stencils.at(abstracted.src_op_tree_path); - std::unordered_map, num_bytes_t> + std::map, num_bytes_t> communication_edges = map_keys_with_value_merging( abstracted.edge_to_size, /*key_func=*/ diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc index 37bf62029f..b75eed3fdf 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc @@ -7,9 +7,7 @@ #include "utils/containers/map_keys_with_value_merging.h" #include "utils/containers/merge_maps_with.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" -#include "utils/hash/unordered_map.h" -#include "utils/containers/binary_merge_unordered_maps_with.h" +#include "utils/containers/binary_merge_maps_with.h" namespace FlexFlow { @@ -25,7 +23,7 @@ AbstractedTensorSetMovement }; } -std::unordered_set +std::set get_src_layers(AbstractedTensorSetMovement const &m) { return transform( m.single_tensor_movements, @@ -34,20 +32,20 @@ std::unordered_set }); } -std::unordered_set +std::set get_dst_layers(AbstractedTensorSetMovement const &m) { return flatmap(m.single_tensor_movements, [](AbstractedSingleTensorMovement const &m) - -> std::unordered_set { + -> std::set { return abstracted_single_tensor_movement_get_dst_layers(m); }); } TensorSetMovement concretize_abstracted_tensor_set_movement( AbstractedTensorSetMovement const &abstracted, - std::unordered_map const + std::map const &pre_machine_stencils, - std::unordered_map const + std::map const &post_machine_stencils) { std::vector single_tensor_movements = @@ -63,7 +61,7 @@ TensorSetMovement concretize_abstracted_tensor_set_movement( [](TensorSetMovement const &lhs, TensorSetMovement const &rhs) -> TensorSetMovement { return TensorSetMovement{ - binary_merge_unordered_maps_with( + binary_merge_maps_with( lhs.edge_to_size, rhs.edge_to_size, [](num_bytes_t l, num_bytes_t r) { return l + r; }), diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc index 192ada3fb6..9fbbdf0cfe 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc @@ -10,7 +10,6 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" -#include "utils/bidict/algorithms/bidict_unordered_set_of.h" #include "utils/containers/binary_cartesian_product.h" #include "utils/containers/flatmap.h" #include "utils/containers/get_only.h" @@ -18,9 +17,10 @@ #include "utils/containers/map_from_pairs.h" #include "utils/containers/merge_maps_with.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/values.h" #include "utils/containers/vector_of.h" +#include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" namespace FlexFlow { @@ -43,9 +43,9 @@ AbstractedSingleTensorMovement get_abstracted_single_tensor_movement_along_edge( bidict coord_mapping = op_to_op_get_coord_mapping(mapping); - std::unordered_map - single_comms = unordered_map_from_pairs(transform( - bidict_unordered_set_of(coord_mapping), + std::map + single_comms = map_from_pairs(transform( + unstructured_relation_from_bidict(coord_mapping), [&](std::pair const & src_dst) -> std::pair { @@ -69,7 +69,7 @@ AbstractedSingleTensorMovement get_abstracted_single_tensor_movement_along_edge( AbstractedTensorSetMovement get_abstracted_tensor_set_movement_across_split( TransitiveReducedPCG const &tr_pcg, PCGBinarySeriesSplit const &split) { - std::unordered_set edges_across_split = + std::set edges_across_split = pcg_get_transitive_reduced_edges_across_split(tr_pcg, split); OneToMany @@ -100,11 +100,11 @@ AbstractedTensorSetMovement get_abstracted_tensor_set_movement_across_split( }; return AbstractedTensorSetMovement{ - transform(unordered_set_of(edges_by_tensor.right_groups()), + transform(edges_by_tensor.right_groups(), [&](nonempty_set const &edges) { return merge_abstracted_single_tensor_movements(transform( - unordered_multiset_of(edges.unwrap_as_unordered_set()), + multiset_of(edges.unwrap_as_set()), to_abstracted_single_tensor_movement)); }), }; diff --git a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc index c3d9ae7bfb..7d7a2ddcc8 100644 --- a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc @@ -15,8 +15,8 @@ #include "utils/containers/repeat_element.h" #include "utils/containers/sorted.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_multiset_of.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/multiset_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/zip.h" #include "utils/nonnegative_int/nonnegative_range.h" #include "utils/nonnegative_int/num_elements.h" @@ -47,7 +47,7 @@ bool is_valid_machine_view(MachineView const &mv, * returned set contains a valid machine view (i.e. it's possible for all * the returned `MachineView`s to be invalid) */ -static std::unordered_set +static std::set get_candidate_machine_views(MachineComputeResourceSlice const &machine_spec, OperatorTaskSpace const &task_space, DeviceType const &device_type) { @@ -67,7 +67,7 @@ static std::unordered_set auto get_candidate_strides = [&](std::vector const &tensor_dims, positive_int total_devices) - -> std::unordered_multiset { + -> std::multiset { positive_int max_stride_upper_bound = get_max_stride_upper_bound(tensor_dims, total_devices); @@ -77,12 +77,12 @@ static std::unordered_set max_stride_upper_bound.nonnegative_int_from_positive_int() + 1_n), [](nonnegative_int stride) { return stride_t{positive_int{stride}}; }); - std::unordered_multiset> raw_stride_vectors = + std::multiset> raw_stride_vectors = cartesian_product( repeat_element(/*num_times=*/num_elements(tensor_dims), /*element=*/single_stride_range)); - std::unordered_multiset strides = + std::multiset strides = transform(raw_stride_vectors, [](auto const &stride_vec) { return MultiDimensionalStride{stride_vec}; }); @@ -92,10 +92,10 @@ static std::unordered_set auto get_candidate_starts = [](MachineComputeResourceSlice const &slice, DeviceType const &device_type) - -> std::unordered_set { + -> std::set { ASSERT(device_type == DeviceType::GPU); - std::unordered_set result; + std::set result; for (nonnegative_int node_idx : nonnegative_range(slice.num_nodes)) { for (nonnegative_int device_idx : nonnegative_range(slice.num_gpus_per_node)) { @@ -107,8 +107,8 @@ static std::unordered_set }; auto get_candidate_dimensions = [](OperatorTaskSpace const &task_space) - -> std::unordered_multiset> { - std::unordered_set options = { + -> std::multiset> { + std::set options = { MachineSpecificationDimension::INTER_NODE, MachineSpecificationDimension::INTRA_NODE}; return get_all_permutations_with_repetition( @@ -122,19 +122,19 @@ static std::unordered_set positive_int total_devices = get_total_num_devices_in_slice(machine_spec); - std::unordered_multiset candidate_strides = + std::multiset candidate_strides = get_candidate_strides(tensor_dims, total_devices); ASSERT(candidate_strides.size() > 0); - std::unordered_set candidate_starts = + std::set candidate_starts = get_candidate_starts(machine_spec, device_type); ASSERT(candidate_starts.size() > 0); - std::unordered_multiset> + std::multiset> candidate_dimensions = get_candidate_dimensions(task_space); ASSERT(candidate_dimensions.size() > 0); - std::unordered_set machine_views; + std::set machine_views; for (MultiDimensionalStride const &strides : candidate_strides) { for (MachineSpaceCoordinate start : candidate_starts) { @@ -149,12 +149,12 @@ static std::unordered_set return machine_views; } -std::unordered_set +std::set get_allowed_machine_views(MachineComputeResourceSlice const &machine_spec, OperatorTaskSpace const &task_space, DeviceType device_type) { - std::unordered_set views = + std::set views = get_candidate_machine_views(machine_spec, task_space, device_type); return filter(views, [&](MachineView const &mv) { return is_valid_machine_view(mv, task_space, machine_spec); diff --git a/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc index 4e38750de3..198fa29326 100644 --- a/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc @@ -36,18 +36,18 @@ SearchResult apply_substitution_and_update_machine_mapping( apply_substitution_from_output_result( substitution_output_result, spcg, sub, match); - std::unordered_map post_node_data = + std::map post_node_data = get_sub_pcg_data(post_substitution_graph).node_data; - std::unordered_set + std::set substitution_output_parallel_layers = - get_parallel_layers(substitution_output_result.first); + spcg_get_parallel_layers(substitution_output_result.first); - std::unordered_map machine_views = + std::map machine_views = mapped_pcg.machine_mapping.machine_views; - std::unordered_set matched_nodes = - unordered_set_of(values(match.node_assignment)); + std::set matched_nodes = + set_of(values(match.node_assignment)); std::vector substituted_machine_views = vector_of( transform(matched_nodes, [&](parallel_layer_guid_t const &node) { @@ -59,9 +59,9 @@ SearchResult apply_substitution_and_update_machine_mapping( select_random(substituted_machine_views)); } - ASSERT(is_subseteq_of(unordered_keys(post_node_data), unordered_keys(machine_views))); + ASSERT(is_subseteq_of(keys(post_node_data), keys(machine_views))); - std::unordered_map + std::map post_node_machine_views = filter(machine_views, [&](std::pair const &p) { diff --git a/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc index 48f3bd9eed..521217e599 100644 --- a/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc @@ -23,11 +23,11 @@ #include "utils/containers/contains.h" #include "utils/containers/contains_key.h" #include "utils/containers/flatmap.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/keys.h" #include "utils/containers/set_minus.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/exception.h" #include "utils/overload.h" @@ -90,22 +90,22 @@ MachineMappingResult ¶llel_split_transformation) { auto get_boundary_machine_view_assignments = - [&](std::unordered_set const &boundary_layers, + [&](std::set const &boundary_layers, MachineMappingProblemTree const &root, BinaryTreePathEntry const &prefix) - -> std::unordered_set { + -> std::set { MachineMappingConstraints sub_constraints = restrict_to_child(constraints, prefix); ASSERT(get_all_layers(sub_constraints) == get_all_leaf_paths(root)); - std::unordered_set unconstrained_boundary_layers = - set_minus(boundary_layers, get_constrained_layers(sub_constraints)); + std::set unconstrained_boundary_layers = + set_minus(boundary_layers, set_of(get_constrained_layers(sub_constraints))); - std::unordered_map> - allowed = generate_unordered_map( + std::map> + allowed = generate_map( unconstrained_boundary_layers, - [&](BinaryTreePath const &l) -> std::unordered_set { + [&](BinaryTreePath const &l) -> std::set { UnmappedRuntimeOnlyOpCostEstimateKey leaf = mm_problem_tree_get_subtree_at_path(root, l) .value() @@ -113,12 +113,12 @@ MachineMappingResult return context.allowed_machine_views(leaf, resources); }); - std::unordered_set> + std::set> assignments = get_all_assignments(allowed); return transform( assignments, - [](std::unordered_map const &m) { + [](std::map const &m) { return ParallelLayerGuidObliviousMachineMapping{m}; }); }; @@ -255,7 +255,7 @@ MachineMappingResult get_optimal_machine_mapping( return parallel_combine(resource_split, left_result, right_result); }; - std::unordered_set parallel_results = transform( + std::set parallel_results = transform( get_machine_resource_splits(resources), evaluate_resource_split); return minimize_runtime(series_result, @@ -269,10 +269,10 @@ MachineMappingResult get_optimal_machine_mapping( MachineComputeResourceSlice const &resource, MachineMappingConstraints const &constraints) { - std::unordered_set candidates = [&] { + std::set candidates = [&] { std::optional machine_view = require_only_root(constraints); if (machine_view.has_value()) { - return std::unordered_set{machine_view.value()}; + return std::set{machine_view.value()}; } else { return context.allowed_machine_views(leaf, resource); } @@ -287,7 +287,7 @@ MachineMappingResult get_optimal_machine_mapping( return make_singleton_machine_mapping_result(cost, machine_view); }; - std::unordered_set candidate_results = + std::set candidate_results = transform(candidates, get_mapping_result); return get_mapping_with_minimal_runtime(candidate_results); diff --git a/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc b/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc index f7dbdd1d05..aa8bd7642f 100644 --- a/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/get_tensor_set_movement_across_split.cc @@ -24,7 +24,7 @@ TensorSetMovement get_tensor_set_movement_across_split( get_abstracted_tensor_set_movement_across_split(tr_pcg, split); auto get_task_spaces = [&](PCGBinarySPDecomposition const &t) - -> std::unordered_map { + -> std::map { return map_values(pcg_sp_tree_get_path_to_leaf_map(t), [&](parallel_layer_guid_t parallel_layer_guid) { return get_operator_task_space(tr_pcg.full_pcg, @@ -32,11 +32,11 @@ TensorSetMovement get_tensor_set_movement_across_split( }); }; - std::unordered_map pre_stencils = + std::map pre_stencils = get_machine_stencils_for_decomposition( tr_pcg.full_pcg, split.get_left_child(), pre_mapping); - std::unordered_map post_stencils = + std::map post_stencils = get_machine_stencils_for_decomposition( tr_pcg.full_pcg, split.get_right_child(), post_mapping); diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc index c7b068d121..13bb389efb 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc @@ -7,8 +7,8 @@ #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/containers/are_disjoint.h" -#include "utils/containers/unordered_keys.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/keys.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -16,11 +16,11 @@ MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapping(ParallelComputationGraph const &pcg, MachineMapping const &mapping) { - std::unordered_set pcg_layers = - get_parallel_layers(pcg); + std::set pcg_layers = + pcg_get_parallel_layers(pcg); - std::unordered_set mapped_layers = - unordered_keys(mapping.machine_views); + std::set mapped_layers = + keys(mapping.machine_views); ASSERT(mapped_layers == pcg_layers); @@ -29,7 +29,7 @@ MappedParallelComputationGraph ComputationGraphOpAttrs op_attrs = assert_unwrap( compgraph_op_attrs_from_pcg_op_attrs(pcg_get_op_attrs(pcg, l))); - std::unordered_map + std::map inputs_dim_degrees = get_incoming_input_degrees(pcg, l); ASSERT(contains_key(mapping.machine_views, l)); @@ -39,8 +39,8 @@ MappedParallelComputationGraph op_attrs, inputs_dim_degrees, machine_view); }; - std::unordered_map - mapped_op_task_groups = generate_unordered_map(mapped_layers, mapping_for_layer); + std::map + mapped_op_task_groups = generate_map(mapped_layers, mapping_for_layer); return mapped_pcg_from_pcg_and_mapped_op_task_groups(pcg, mapped_op_task_groups); @@ -49,12 +49,12 @@ MappedParallelComputationGraph MachineMapping combine_disjoint_mappings(MachineMapping const &m1, MachineMapping const &m2) { return MachineMapping{ - binary_merge_disjoint_unordered_maps(m1.machine_views, m2.machine_views), + binary_merge_disjoint_maps(m1.machine_views, m2.machine_views), }; } bool nodes_are_disjoint(MachineMapping const &m1, MachineMapping const &m2) { - return are_disjoint(unordered_keys(m1.machine_views), unordered_keys(m2.machine_views)); + return are_disjoint(keys(m1.machine_views), keys(m2.machine_views)); } std::optional get_machine_mapping_from_machine_mapping_result( diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc index f77d424795..3b73b98a47 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc @@ -3,8 +3,8 @@ #include "utils/containers/filter_values.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/flatmap.h" -#include "utils/containers/generate_unordered_map.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/generate_map.h" +#include "utils/containers/keys.h" #include "utils/containers/map_values.h" #include "utils/containers/restrict_keys.h" #include "utils/full_binary_tree/binary_tree_path.h" @@ -12,34 +12,34 @@ namespace FlexFlow { MachineMappingConstraints get_unconstrained_solution_for_layers( - std::unordered_set const &layers) { + std::set const &layers) { return MachineMappingConstraints{ - generate_unordered_map(layers, + generate_map(layers, [](BinaryTreePath const &) -> std::optional { return std::nullopt; }), }; } -std::unordered_set +std::set get_unconstrained_layers(MachineMappingConstraints const &constraints) { - return unordered_keys(filter_values( + return keys(filter_values( constraints.machine_views, [](std::optional const &mv) { return !mv.has_value(); })); } -std::unordered_set +std::set get_constrained_layers(MachineMappingConstraints const &constraints) { - return unordered_keys(filter_values( + return keys(filter_values( constraints.machine_views, [](std::optional const &mv) { return mv.has_value(); })); } -std::unordered_set +std::set get_all_layers(MachineMappingConstraints const &partial_solution) { - return unordered_keys(partial_solution.machine_views); + return keys(partial_solution.machine_views); } std::optional get_machine_view_for_layer( @@ -103,8 +103,8 @@ MachineMappingConstraints with_additional_constraints( std::optional require_only_root(MachineMappingConstraints const &constraints) { - ASSERT(unordered_keys(constraints.machine_views) == - std::unordered_set{binary_tree_root_path()}, + ASSERT(keys(constraints.machine_views) == + std::set{binary_tree_root_path()}, fmt::format("require_only_root expected constraints to have only a " "single key (the root path), but received {}", constraints)); diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc index 47639ff88a..32ba9ad966 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc @@ -14,10 +14,10 @@ std::optional MachineComputeSpecification const &resources, DeviceType const &device_type) { std::vector layers = topological_ordering(pcg); - std::unordered_map machine_views; + std::map machine_views; for (parallel_layer_guid_t layer : layers) { OperatorTaskSpace task = get_operator_task_space(pcg, layer); - std::unordered_set allowed_machine_views = + std::set allowed_machine_views = get_allowed_machine_views( compute_slice_from_specification(resources), task, DeviceType::GPU); if (allowed_machine_views.empty()) { diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc index 88d64f6359..f632c99190 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc @@ -20,7 +20,7 @@ bool is_valid_machine_mapping_problem_tree( auto contains_paths = [](MachineMappingProblemTree const &t, - std::unordered_set const &paths) { + std::set const &paths) { return all_of(paths, [&](BinaryTreePath const &p) { return mm_problem_tree_get_subtree_at_path(t, p).has_value(); }); diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.cc index 8fc6d52c0f..67f02ad3fa 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.cc @@ -75,12 +75,12 @@ SPDecompositionTreeNodeType }); } -std::unordered_multiset +std::multiset get_leaves(MachineMappingProblemTree const &tree) { return get_leaves(tree, generic_binary_sp_impl_for_mm_problem_tree()); } -std::unordered_set +std::set get_all_leaf_paths(MachineMappingProblemTree const &tree) { return get_all_leaf_paths(tree, generic_binary_sp_impl_for_mm_problem_tree()); } @@ -92,7 +92,7 @@ std::optional tree, generic_binary_sp_impl_for_mm_problem_tree(), path); } -std::unordered_map +std::map mm_problem_tree_get_path_to_leaf_map( MachineMappingProblemTree const &tree) { return get_path_to_leaf_map(tree, @@ -119,7 +119,7 @@ std::string as_dot(MachineMappingProblemTree const &tree) { }; auto path_set_as_dot = - [&](std::unordered_set const &path_set) -> std::string { + [&](std::set const &path_set) -> std::string { return "(" + join_strings(path_set, ", ", path_as_dot) + ")"; }; diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_result.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_result.cc index 33d550474c..3d380491d2 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_result.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_result.cc @@ -21,7 +21,7 @@ FeasibleMachineMappingResult } [[nodiscard]] MachineMappingResult get_mapping_with_minimal_runtime( - std::unordered_set const &candidates) { + std::set const &candidates) { MachineMappingResult result = infeasible_machine_mapping_result(); for (MachineMappingResult const &candidate : candidates) { diff --git a/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc b/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc index 875f44a0c9..97c99f3014 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_resource_split.cc @@ -44,10 +44,10 @@ std::pair } } -std::unordered_set +std::set get_machine_resource_splits(MachineComputeResourceSlice const &resources) { - std::unordered_set result; + std::set result; for (positive_int i = 1_p; i < resources.num_nodes; i *= 2_p) { result.insert(MachineResourceSplit{ diff --git a/lib/compiler/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/src/compiler/machine_mapping/machine_view.cc index 7d00707b0d..06c58bd2d5 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_view.cc @@ -163,7 +163,7 @@ OperatorSpaceToMachineSpaceMapping get_coordinate_mapping_for_machine_view( }; } -std::unordered_set +std::set get_machine_space_coordinates(OperatorTaskSpace const &task_space, MachineView const &machine_view) { @@ -177,7 +177,7 @@ std::unordered_set }); } -std::unordered_set +std::set get_device_ids(OperatorTaskSpace const &task_space, MachineView const &mv, MachineComputeSpecification const &ms) { @@ -206,7 +206,7 @@ MachineView static OperatorAtomicTaskShardBinding operator_atomic_task_shard_binding_from_machine_view( ComputationGraphOpAttrs const &op_attrs, - std::unordered_map const + std::map const &inputs_dim_degrees, MachineView const &machine_view, MachineSpaceCoordinate const &machine_space_coord) { @@ -217,12 +217,12 @@ static OperatorAtomicTaskShardBinding mv_task_space_coord_for_machine_space_coord( machine_view, op_task_space, machine_space_coord); - std::unordered_map + std::map mappings = get_operator_to_ptensor_mappings(op_attrs, inputs_dim_degrees); - std::unordered_map - ptensor_coords = generate_unordered_map( - unordered_keys(inputs_dim_degrees), + std::map + ptensor_coords = generate_map( + keys(inputs_dim_degrees), [&](TensorSlotName const &slot_name) -> ParallelTensorSpaceCoordinate { num_ptensor_shard_dims_t num_shard_dims = @@ -240,7 +240,7 @@ static OperatorAtomicTaskShardBinding MappedOperatorTaskGroup mapped_operator_task_group_from_machine_view( ComputationGraphOpAttrs const &op_attrs, - std::unordered_map const + std::map const &inputs_dim_degrees, MachineView const &machine_view) { diff --git a/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc b/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc index dad91dd317..ab54c034a1 100644 --- a/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc +++ b/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc @@ -17,9 +17,9 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "utils/containers/contains.h" #include "utils/containers/flatmap.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_all_assignments.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/exception.h" #include "utils/overload.h" @@ -81,12 +81,12 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( auto get_boundary_machine_view_assignments = [&](MachineMappingProblemTree const &root, - std::unordered_set const &boundary_layers) - -> std::unordered_set { - std::unordered_map> - allowed = generate_unordered_map( + std::set const &boundary_layers) + -> std::set { + std::map> + allowed = generate_map( boundary_layers, - [&](BinaryTreePath const &l) -> std::unordered_set { + [&](BinaryTreePath const &l) -> std::set { UnmappedRuntimeOnlyOpCostEstimateKey leaf = mm_problem_tree_get_subtree_at_path(root, l) .value() @@ -96,7 +96,7 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( return transform( get_all_assignments(allowed), - [](std::unordered_map const &m) { + [](std::map const &m) { return ParallelLayerGuidObliviousMachineMapping{m}; }); }; @@ -226,7 +226,7 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( return parallel_combine(resource_split, left_result, right_result); }; - std::unordered_set parallel_results = + std::set parallel_results = transform(get_machine_resource_splits(resources), evaluate_resource_split); @@ -241,10 +241,10 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( MachineComputeResourceSlice const &resource, MachineMappingConstraints const &constraints) { - std::unordered_set candidates = [&] { + std::set candidates = [&] { std::optional machine_view = require_only_root(constraints); if (machine_view.has_value()) { - return std::unordered_set{machine_view.value()}; + return std::set{machine_view.value()}; } else { return context.allowed_machine_views(leaf, resource); } @@ -261,7 +261,7 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( machine_view); }; - std::unordered_set candidate_results = + std::set candidate_results = transform(candidates, get_mapping_result); return get_mapping_with_minimal_runtime(candidate_results); diff --git a/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc b/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc index 9021e0d382..4f09c569d2 100644 --- a/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc +++ b/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc @@ -7,12 +7,12 @@ #include "utils/containers/transform.h" #include "utils/full_binary_tree/binary_tree_path.h" #include "utils/hash/tuple.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" namespace FlexFlow { MachineMappingWithMemoryResult::MachineMappingWithMemoryResult( - std::unordered_set const &pareto_frontier) + std::set const &pareto_frontier) : m_pareto_frontier(pareto_frontier) { ASSERT(all_of(pareto_frontier, [&](ParetoOptimalMachineMapping const &m) { return is_pareto_optimal_in(m, pareto_frontier); @@ -29,7 +29,27 @@ bool MachineMappingWithMemoryResult::operator!=( return this->tie() != other.tie(); } -std::unordered_set const & +bool MachineMappingWithMemoryResult::operator<( + MachineMappingWithMemoryResult const &other) const { + return this->tie() < other.tie(); +} + +bool MachineMappingWithMemoryResult::operator>( + MachineMappingWithMemoryResult const &other) const { + return this->tie() > other.tie(); +} + +bool MachineMappingWithMemoryResult::operator<=( + MachineMappingWithMemoryResult const &other) const { + return this->tie() <= other.tie(); +} + +bool MachineMappingWithMemoryResult::operator>=( + MachineMappingWithMemoryResult const &other) const { + return this->tie() >= other.tie(); +} + +std::set const & MachineMappingWithMemoryResult::get_pareto_frontier() const { return this->m_pareto_frontier; } @@ -44,7 +64,7 @@ std::ostream &operator<<(std::ostream &s, return (s << fmt::to_string(r)); } -std::tuple const &> +std::tuple const &> MachineMappingWithMemoryResult::tie() const { return std::tie(this->m_pareto_frontier); } @@ -56,7 +76,7 @@ MachineMappingWithMemoryResult empty_machine_mapping_with_memory_result() { } MachineMappingWithMemoryResult get_mapping_with_minimal_runtime( - std::unordered_set const &candidates) { + std::set const &candidates) { MachineMappingWithMemoryResult result = empty_machine_mapping_with_memory_result(); @@ -100,7 +120,7 @@ MachineMappingWithMemoryResult return ParetoOptimalMachineMapping{cost, mapping}; }; - std::unordered_set result; + std::set result; for (ParetoOptimalMachineMapping const &pre_mm : pre_result.get_pareto_frontier()) { @@ -143,7 +163,7 @@ MachineMappingWithMemoryResult return ParetoOptimalMachineMapping{cost, mapping}; }; - std::unordered_set result; + std::set result; for (ParetoOptimalMachineMapping const &lhs_mm : lhs_result.get_pareto_frontier()) { @@ -165,7 +185,7 @@ MachineMappingWithMemoryResult MachineMappingWithMemoryResult minimize_runtime(MachineMappingWithMemoryResult const &m1, MachineMappingWithMemoryResult const &m2) { - std::unordered_set result = + std::set result = set_union(m1.get_pareto_frontier(), m2.get_pareto_frontier()); return MachineMappingWithMemoryResult{ diff --git a/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc index ca6d762eed..96ea3c6b30 100644 --- a/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc @@ -6,7 +6,7 @@ namespace FlexFlow { bool is_pareto_optimal_in( ParetoOptimalMachineMapping const &m, - std::unordered_set const &others) { + std::set const &others) { return is_pareto_optimal_in( m.cost, transform(others, [](ParetoOptimalMachineMapping const &m) { return m.cost; diff --git a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc index 97814288ab..1c61ba4fd9 100644 --- a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc @@ -9,7 +9,7 @@ #include "utils/containers/require_same.h" #include "utils/containers/try_at.h" #include "utils/full_binary_tree/binary_tree_path.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -17,7 +17,7 @@ ParallelLayerGuidObliviousMachineMapping binary_combine_mappings( ParallelLayerGuidObliviousMachineMapping const &lhs, ParallelLayerGuidObliviousMachineMapping const &rhs) { return ParallelLayerGuidObliviousMachineMapping{ - binary_merge_disjoint_unordered_maps( + binary_merge_disjoint_maps( map_keys(lhs.raw_mapping, nest_inside_left_child), map_keys(rhs.raw_mapping, nest_inside_right_child)), }; @@ -39,22 +39,22 @@ std::optional get_machine_view_for_path( return try_at(mapping.raw_mapping, path); } -std::unordered_map +std::map get_machine_stencils_for_decomposition( ParallelComputationGraph const &pcg, PCGBinarySPDecomposition const &decomposition, ParallelLayerGuidObliviousMachineMapping const &mapping) { - std::unordered_set leaf_paths = require_same( - pcg_sp_tree_get_all_leaf_paths(decomposition), unordered_keys(mapping.raw_mapping)); + std::set leaf_paths = require_same( + pcg_sp_tree_get_all_leaf_paths(decomposition), keys(mapping.raw_mapping)); - std::unordered_map + std::map path_to_op_task_space_map = map_values(pcg_sp_tree_get_path_to_leaf_map(decomposition), [&](parallel_layer_guid_t l) -> OperatorTaskSpace { return get_operator_task_space(pcg, l); }); - return generate_unordered_map( + return generate_map( leaf_paths, [&](BinaryTreePath const &p) -> MachineSpaceStencil { return MachineSpaceStencil{ /*operator_task_space=*/path_to_op_task_space_map.at(p), @@ -63,20 +63,20 @@ std::unordered_map }); } -std::unordered_map> +std::map> get_machine_stencils_for_mm_problem_tree( MachineMappingProblemTree const &tree, ParallelLayerGuidObliviousMachineMapping const &mapping) { - std::unordered_map + std::map tree_leaf_map = mm_problem_tree_get_path_to_leaf_map(tree); - std::unordered_set mapping_paths = unordered_keys(mapping.raw_mapping); - std::unordered_set tree_paths = unordered_keys(tree_leaf_map); + std::set mapping_paths = keys(mapping.raw_mapping); + std::set tree_paths = keys(tree_leaf_map); ASSERT(is_subseteq_of(mapping_paths, tree_paths)); - return generate_unordered_map( + return generate_map( tree_paths, [&](BinaryTreePath const &p) -> std::optional { if (!contains_key(mapping.raw_mapping, p)) { @@ -88,7 +88,7 @@ std::unordered_map> ComputationGraphOpAttrs leaf_op_attrs = compgraph_op_attrs_from_pcg_op_attrs(leaf.op_attrs).value(); - std::unordered_map + std::map leaf_input_degrees = map_values(leaf.input_shapes, [](ParallelTensorShape const &s) { return get_parallel_degrees(s); @@ -102,7 +102,7 @@ std::unordered_map> }); } -std::unordered_map +std::map get_machine_stencils_for_partially_mapped_mm_problem_tree( MachineMappingProblemTree const &tree, ParallelLayerGuidObliviousMachineMapping const &mappings) { diff --git a/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc b/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc index cbb64d5bcf..19038f77f8 100644 --- a/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc @@ -71,7 +71,7 @@ MachineSpaceOffset get_machine_space_offset( return get_machine_space_offset_from_coordinate(dummy_start, ms_coord); } -std::unordered_set get_machine_space_offsets( +std::set get_machine_space_offsets( OperatorTaskSpace const &task, StartInvariantMachineView const &start_inv_machine_view) { return transform( diff --git a/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc b/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc index 5779edc382..16f418cec2 100644 --- a/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc +++ b/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc @@ -34,7 +34,7 @@ TransitiveReducedPCG }; } -std::unordered_set +std::set pcg_get_transitive_reduced_edges_across_split( TransitiveReducedPCG const &tr_pcg, PCGBinarySeriesSplit const &split) { @@ -44,16 +44,16 @@ std::unordered_set BinarySeriesSplit raw_split = binary_series_split_from_pcg_series_split(split); - std::unordered_set> raw_edges = - get_transitive_reduced_kwarg_dataflow_edges_across_split(raw_tr_g, - raw_split); + std::set> raw_edges = + set_of(get_transitive_reduced_kwarg_dataflow_edges_across_split(raw_tr_g, + raw_split)); return transform(raw_edges, [](KwargDataflowEdge const &e) { return ParallelComputationGraphEdge{e}; }); } -std::unordered_set +std::set pcg_get_transitive_reduced_tensors_across_split( TransitiveReducedPCG const &tr_pcg, PCGBinarySeriesSplit const &split) { TransitiveReducedKwargDataflowGraphView raw_tr_g = @@ -62,9 +62,9 @@ std::unordered_set BinarySeriesSplit raw_split = binary_series_split_from_pcg_series_split(split); - std::unordered_set> raw_outputs = - get_transitive_reduced_kwarg_dataflow_outputs_across_split(raw_tr_g, - raw_split); + std::set> raw_outputs = + set_of(get_transitive_reduced_kwarg_dataflow_outputs_across_split(raw_tr_g, + raw_split)); return transform(raw_outputs, [](KwargDataflowOutput const &o) { diff --git a/lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc b/lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc index 80c09d2dba..afd64c2a97 100644 --- a/lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/unstructured_device_mapping.cc @@ -13,7 +13,7 @@ UnstructuredDeviceMapping get_unstructured_device_mapping( MachineMapping const &machine_mapping, MachineComputeSpecification const &machine_spec, ParallelComputationGraph const &pcg) { - std::unordered_map> + std::map> device_mapping; for (auto const &[layer, machine_view] : machine_mapping.machine_views) { OperatorTaskSpace op = get_operator_task_space(pcg, layer); diff --git a/lib/compiler/src/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.cc b/lib/compiler/src/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.cc index 9886468386..c5b668cd65 100644 --- a/lib/compiler/src/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.cc +++ b/lib/compiler/src/compiler/series_parallel/computation_graph/computation_graph_binary_sp_decomposition.cc @@ -157,7 +157,7 @@ bool is_right_associative(ComputationGraphBinarySPDecomposition const &tree) { tree, generic_impl_for_computation_graph_sp_tree()); } -std::unordered_multiset +std::multiset get_layers(ComputationGraphBinarySPDecomposition const &tree) { return get_leaves(tree, generic_impl_for_computation_graph_sp_tree()); } diff --git a/lib/compiler/src/compiler/series_parallel/computation_graph/get_computation_graph_series_parallel_decomposition.cc b/lib/compiler/src/compiler/series_parallel/computation_graph/get_computation_graph_series_parallel_decomposition.cc index 50b95f3c3e..922751bf3c 100644 --- a/lib/compiler/src/compiler/series_parallel/computation_graph/get_computation_graph_series_parallel_decomposition.cc +++ b/lib/compiler/src/compiler/series_parallel/computation_graph/get_computation_graph_series_parallel_decomposition.cc @@ -13,13 +13,13 @@ namespace FlexFlow { std::string render_preprocessed_computation_graph_for_sp_decomposition( ComputationGraph const &cg) { - std::unordered_set weight_and_input_layers = + std::set weight_and_input_layers = filter(get_layers(cg), [&](layer_guid_t const &l) { ComputationGraphOpAttrs op_attrs = get_layer_attrs(cg, l).op_attrs; return op_attrs.has() || op_attrs.has(); }); - std::unordered_set weight_and_input_layer_successors = + std::set weight_and_input_layer_successors = get_subgraph_successors(cg, weight_and_input_layers); // dot has is incapable of rendering the number of edges in the all-to-all @@ -65,13 +65,13 @@ std::optional } DiGraphView preprocessed_digraph = [&] { - std::unordered_set weight_and_input_layers = + std::set weight_and_input_layers = filter(get_layers(cg), [&](layer_guid_t const &l) { ComputationGraphOpAttrs op_attrs = get_layer_attrs(cg, l).op_attrs; return op_attrs.has() || op_attrs.has(); }); - std::unordered_set weight_and_input_layer_successors = + std::set weight_and_input_layer_successors = get_subgraph_successors(cg, weight_and_input_layers); DiGraph digraph = materialize_digraph_view(cg.raw_graph); diff --git a/lib/compiler/src/compiler/series_parallel/pcg/get_pcg_series_parallel_decomposition.cc b/lib/compiler/src/compiler/series_parallel/pcg/get_pcg_series_parallel_decomposition.cc index 30a0655b2d..5ca11c4cd9 100644 --- a/lib/compiler/src/compiler/series_parallel/pcg/get_pcg_series_parallel_decomposition.cc +++ b/lib/compiler/src/compiler/series_parallel/pcg/get_pcg_series_parallel_decomposition.cc @@ -36,7 +36,7 @@ std::optional assert(layer_is_weight_or_input(starting_point) || layer_is_parallel_op(starting_point)); - std::unordered_set successors = + std::set successors = get_successors(pcg, starting_point); if (successors.size() != 1) { @@ -55,13 +55,13 @@ std::optional }; DiGraphView preprocessed_digraph = [&] { - std::unordered_set weight_and_input_layers = - filter(get_parallel_layers(pcg), layer_is_weight_or_input); + std::set weight_and_input_layers = + filter(pcg_get_parallel_layers(pcg), layer_is_weight_or_input); - std::unordered_set par_chain_endpoints = + std::set par_chain_endpoints = transform(weight_and_input_layers, follow_to_last_parallel_op); - std::unordered_set par_chain_endpoint_successors = + std::set par_chain_endpoint_successors = get_subgraph_successors(pcg, par_chain_endpoints); DiGraph digraph = materialize_digraph_view(pcg.raw_graph); diff --git a/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc b/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc index 4d1c88d9eb..8becc755e2 100644 --- a/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc +++ b/lib/compiler/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc @@ -130,8 +130,8 @@ PCGBinarySPDecomposition }); } -std::unordered_multiset - get_parallel_layers(PCGBinarySPDecomposition const &tree) { +std::multiset + pcg_sp_tree_get_parallel_layers(PCGBinarySPDecomposition const &tree) { return get_leaves(tree, generic_impl_for_pcg_sp_tree()); } @@ -150,18 +150,18 @@ SPDecompositionTreeNodeType }); } -std::unordered_set +std::set pcg_sp_tree_get_all_leaf_paths(PCGBinarySPDecomposition const &tree) { - return unordered_keys(pcg_sp_tree_get_path_to_leaf_map(tree)); + return keys(pcg_sp_tree_get_path_to_leaf_map(tree)); } -std::unordered_set +std::set find_paths_to_leaf(PCGBinarySPDecomposition const &tree, parallel_layer_guid_t const &leaf) { return find_paths_to_leaf(tree, generic_impl_for_pcg_sp_tree(), leaf); } -std::unordered_map +std::map pcg_sp_tree_get_path_to_leaf_map(PCGBinarySPDecomposition const &tree) { return get_path_to_leaf_map(tree, generic_impl_for_pcg_sp_tree()); } diff --git a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc index b016b106e9..917f63949e 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc @@ -13,8 +13,8 @@ #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" #include "utils/bidict/bidict.h" #include "utils/graph/instances/adjacency_digraph.h" -#include -#include +#include +#include #include "utils/containers/set_of.h" namespace FlexFlow { @@ -28,7 +28,7 @@ PCGTaskGraph bidict node_to_layer; std::map> node_to_devices; - for (parallel_layer_guid_t const &layer : get_parallel_layers(pcg)) { + for (parallel_layer_guid_t const &layer : pcg_get_parallel_layers(pcg)) { MachineView mv = machine_mapping.machine_views.at(layer); RuntimeOnlyOpCostEstimateKey op_key = get_mapped_runtime_only_op_cost_estimate_key_for_layer(pcg, layer, mv); diff --git a/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc b/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc index 30e345243c..e708d4b0f6 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc @@ -16,7 +16,7 @@ #include "utils/graph/node/algorithms.h" #include "utils/overload.h" #include -#include +#include namespace FlexFlow { @@ -35,7 +35,7 @@ TaskGraphExecutionTrace simulate_task_graph_execution( /*finished_tasks=*/{}, /*current_time=*/0.0}; - std::unordered_set task_profiles; + std::set task_profiles; auto start_task_processing = [&](Node const &task) { float cost = cost_function(task); @@ -47,7 +47,7 @@ TaskGraphExecutionTrace simulate_task_graph_execution( }; auto dependencies_are_satisfied = [&](Node const &task) { - std::unordered_set incoming_dependencies = + std::set incoming_dependencies = get_predecessors(task_graph, task); return is_subseteq_of(incoming_dependencies, execution_state.finished_tasks); @@ -80,8 +80,8 @@ TaskGraphExecutionTrace simulate_task_graph_execution( while (!is_processing_done()) { auto ready_tasks_copy = execution_state.ready_tasks; for (Node const &task : ready_tasks_copy) { - std::unordered_set raw_in_progress_tasks = transform( - unordered_set_of(execution_state.in_progress_tasks.contents()), + std::set raw_in_progress_tasks = transform( + set_of(execution_state.in_progress_tasks.contents()), [](InProgressTask const &t) { return t.node; }); if (constraint.is_satisfied( diff --git a/lib/compiler/src/compiler/task_graph_simulator/task_graph_execution_trace.cc b/lib/compiler/src/compiler/task_graph_simulator/task_graph_execution_trace.cc index 1e15931174..5458792f35 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/task_graph_execution_trace.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/task_graph_execution_trace.cc @@ -3,7 +3,7 @@ #include "utils/containers/minimum.h" #include "utils/containers/transform.h" #include "utils/exception.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" namespace FlexFlow { diff --git a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc index c3b29c85fb..078d190f0f 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc @@ -15,8 +15,8 @@ #include "utils/containers/set_union.h" #include "utils/containers/transform.h" #include "utils/graph/digraph/digraph.h" -#include "utils/hash/unordered_set.h" -#include +#include "utils/hash/set.h" +#include namespace FlexFlow { @@ -45,8 +45,8 @@ milliseconds_t task_simulator_estimate_forward_pass_time( auto is_allowed_to_run = [&](Node const &task, - std::unordered_set const &in_progress_tasks, - std::unordered_set const &finished_tasks) -> bool { + std::set const &in_progress_tasks, + std::set const &finished_tasks) -> bool { PCGTask current_task = task_graph.node_to_task.at_l(task); UnstructuredDeviceMapping device_map = get_unstructured_device_mapping( @@ -61,9 +61,9 @@ milliseconds_t task_simulator_estimate_forward_pass_time( return task_graph.node_to_devices.at(n); }; - std::unordered_set devices_occupied = + std::set devices_occupied = set_union(transform(in_progress_tasks, get_devices)); - std::unordered_set required_devices = unordered_set_of(get_devices(task)); + std::set required_devices = set_of(get_devices(task)); return set_intersection(devices_occupied, required_devices).empty(); }; diff --git a/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc b/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc index 6883098dab..8a81e97255 100644 --- a/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc +++ b/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc @@ -9,8 +9,10 @@ #include "utils/containers/zip_values_strict.h" #include "utils/containers/zip_values_strict_with.h" #include "utils/hash/tuple.h" -#include "utils/hash/unordered_map.h" -#include "utils/hash/unordered_multiset.h" +#include "utils/hash/map.h" +#include "utils/hash/multiset.h" +#include "utils/containers/transform.h" +#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -18,24 +20,24 @@ GraphOptimizeState::GraphOptimizeState(ParallelComputationGraph const &pcg, milliseconds_t runtime) : pcg(pcg), runtime(runtime) {} -static std::unordered_multiset>, - std::unordered_map>> + std::map>> get_layer_signature_set(ParallelComputationGraph const &pcg) { auto get_layer_signature = [&](parallel_layer_guid_t l) -> std::tuple>, - std::unordered_map> { + std::map> { ParallelLayerAttrs layer_attrs = get_parallel_layer_attrs(pcg, l); - std::unordered_map< + std::map< TensorSlotName, std::tuple> inputs = map_values( @@ -52,7 +54,7 @@ static std::unordered_multiset outputs = + std::map outputs = map_values(get_outgoing_tensors(pcg, l), [&](parallel_tensor_guid_t const &o) { return get_parallel_tensor_attrs(pcg, o); @@ -65,7 +67,7 @@ static std::unordered_multiset std::unordered_set { + -> std::set { OperatorTaskSpace op_task_space = get_operator_task_space_for_runtime_only_op_cost_estimate_key(key); diff --git a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc index 3bb224cd8d..d5acf667e6 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc @@ -2,9 +2,9 @@ #include "utils/containers/extend.h" #include "utils/containers/range.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/zip.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include #include @@ -57,14 +57,14 @@ TEST_SUITE(FF_TEST_SUITE) { OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{3_ge2}}}; - std::unordered_set correct = { + std::set correct = { make_machine_view(0_n, 0_n, 1_p, intra), make_machine_view(0_n, 1_n, 1_p, intra), make_machine_view(0_n, 2_n, 1_p, intra), make_machine_view(0_n, 0_n, 2_p, intra), }; - std::unordered_set result = + std::set result = get_allowed_machine_views(ms, task, DeviceType::GPU); CHECK(correct == result); @@ -79,7 +79,7 @@ TEST_SUITE(FF_TEST_SUITE) { OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{2_ge2, 3_ge2}}}; - std::unordered_set correct = { + std::set correct = { make_machine_view( 0_n, 0_n, /*stride_1=*/1_p, inter, /*stride_2=*/1_p, intra), make_machine_view( @@ -95,7 +95,7 @@ TEST_SUITE(FF_TEST_SUITE) { 0_n, 0_n, /*stride_1=*/2_p, intra, /*stride_2=*/1_p, inter), }; - std::unordered_set result = + std::set result = get_allowed_machine_views(ms, task, DeviceType::GPU); CHECK(correct == result); @@ -109,10 +109,10 @@ TEST_SUITE(FF_TEST_SUITE) { }; OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{}}}; - std::unordered_set result = + std::set result = get_allowed_machine_views(full_machine_spec, task, DeviceType::GPU); - std::unordered_set correct = { + std::set correct = { make_machine_view(0_n, 0_n), make_machine_view(1_n, 0_n), }; @@ -128,10 +128,10 @@ TEST_SUITE(FF_TEST_SUITE) { }; OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{2_ge2}}}; - std::unordered_set result = + std::set result = get_allowed_machine_views(full_machine_spec, task, DeviceType::GPU); - std::unordered_set correct = { + std::set correct = { make_machine_view(0_n, 0_n, /*stride_1=*/1_p, intra), make_machine_view(0_n, 0_n, /*stride_1=*/1_p, inter), make_machine_view(1_n, 0_n, /*stride_1=*/1_p, intra), diff --git a/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc b/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc index 392e16bec5..6d9ab13592 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc @@ -219,7 +219,7 @@ TEST_SUITE(FF_TEST_SUITE) { ASSERT(k == k1); ASSERT(resources == four_nodes_resources); - return std::unordered_set{ + return std::set{ mv_stride_1, mv_stride_2, }; @@ -231,7 +231,7 @@ TEST_SUITE(FF_TEST_SUITE) { mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 2), }, - std::unordered_map{{}}); + std::map{{}}); MachineMappingContext context = MachineMappingContext{ /*cost_estimator=*/runtime_only_cost_estimator, @@ -316,17 +316,17 @@ TEST_SUITE(FF_TEST_SUITE) { auto allowed_machine_views = [&](UnmappedRuntimeOnlyOpCostEstimateKey const &k, MachineComputeResourceSlice const &resources) - -> std::unordered_set { - std::unordered_set result; + -> std::set { + std::set result; if (resources == four_nodes_resources) { - result = std::unordered_set{mv_stride_1, mv_stride_2}; + result = std::set{mv_stride_1, mv_stride_2}; } else if (resources == three_nodes_resources) { - result = std::unordered_set{mv_stride_1, mv_stride_2}; + result = std::set{mv_stride_1, mv_stride_2}; } else if (resources == two_nodes_resources) { - result = std::unordered_set{mv_stride_1}; + result = std::set{mv_stride_1}; } else { - result = std::unordered_set{}; + result = std::set{}; } for (MachineView const &mv : result) { @@ -345,14 +345,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("solution requires taking comm cost into account") { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( - std::unordered_map{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 4), mk_cost_entry(k2, mv_stride_2, 1), }}, - std::unordered_map{{ + std::map{{ { TensorSetMovement{{}}, 0.0_ms, @@ -402,14 +402,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("solution places operators on different machine views") { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( - std::unordered_map{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 4), mk_cost_entry(k2, mv_stride_2, 1), }}, - std::unordered_map{{ + std::map{{ { TensorSetMovement{{}}, 0.0_ms, @@ -469,27 +469,27 @@ TEST_SUITE(FF_TEST_SUITE) { [&](UnmappedRuntimeOnlyOpCostEstimateKey const &k, MachineComputeResourceSlice const &resources) { if (resources == four_nodes_resources) { - return std::unordered_set{mv_stride_1, mv_stride_2}; + return std::set{mv_stride_1, mv_stride_2}; } else if (resources == three_nodes_resources) { - return std::unordered_set{mv_stride_1, mv_stride_2}; + return std::set{mv_stride_1, mv_stride_2}; } else if (resources == two_nodes_resources) { - return std::unordered_set{mv_stride_1}; + return std::set{mv_stride_1}; } else { - return std::unordered_set{}; + return std::set{}; } }; SUBCASE("cannot use overlapping machine views in parallel") { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( - std::unordered_map{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 4), mk_cost_entry(k2, mv_stride_2, 1), }}, - std::unordered_map{{ + std::map{{ { TensorSetMovement{{}}, 0.0_ms, @@ -531,14 +531,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("solution is running operators in parallel") { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( - std::unordered_map{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 3), mk_cost_entry(k2, mv_stride_2, 4), }}, - std::unordered_map{{ + std::map{{ { TensorSetMovement{{}}, 0.0_ms, @@ -595,14 +595,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("solution is running operators in series") { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( - std::unordered_map{{ mk_cost_entry(k1, mv_stride_1, 3), mk_cost_entry(k1, mv_stride_2, 1), mk_cost_entry(k2, mv_stride_1, 4), mk_cost_entry(k2, mv_stride_2, 1), }}, - std::unordered_map{{ + std::map{{ { TensorSetMovement{{}}, 0.0_ms, diff --git a/lib/compiler/test/src/compiler/machine_mapping/machine_resource_split.cc b/lib/compiler/test/src/compiler/machine_mapping/machine_resource_split.cc index 3b47e63143..e34573dd2f 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/machine_resource_split.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/machine_resource_split.cc @@ -1,7 +1,7 @@ #include "compiler/machine_mapping/machine_resource_split.h" #include "pcg/machine_compute_specification.dtg.h" #include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/hash/pair.h" #include @@ -15,9 +15,9 @@ TEST_SUITE(FF_TEST_SUITE) { /*num_gpus_per_node=*/1_p, }; - std::unordered_set result = + std::set result = get_machine_resource_splits(input); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } @@ -28,10 +28,10 @@ TEST_SUITE(FF_TEST_SUITE) { /*num_gpus_per_node=*/2_p, }; - std::unordered_set result = + std::set result = get_machine_resource_splits(input); - std::unordered_set correct = { + std::set correct = { MachineResourceSplit{ /*offset=*/1_p, /*dimension=*/MachineSpecificationDimension::INTRA_NODE, @@ -50,10 +50,10 @@ TEST_SUITE(FF_TEST_SUITE) { /*num_gpus_per_node=*/1_p, }; - std::unordered_set result = + std::set result = get_machine_resource_splits(input); - std::unordered_set correct = { + std::set correct = { MachineResourceSplit{ /*offset=*/1_p, /*dimension=*/MachineSpecificationDimension::INTER_NODE, @@ -85,10 +85,10 @@ TEST_SUITE(FF_TEST_SUITE) { /*num_gpus_per_node=*/8_p, }; - std::unordered_set result = + std::set result = get_machine_resource_splits(input); - std::unordered_set correct = { + std::set correct = { MachineResourceSplit{ /*offset=*/1_p, /*dimension=*/MachineSpecificationDimension::INTRA_NODE, diff --git a/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc index 2ea8312991..cac7b23386 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc @@ -4,7 +4,7 @@ #include "pcg/gpu_id_t.dtg.h" #include "test/utils/doctest/fmt/optional.h" #include "utils/containers/transform.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/fmt/vector.h" #include @@ -402,12 +402,12 @@ TEST_SUITE(FF_TEST_SUITE) { {MachineViewDimension{stride_t{2_p}, MachineSpecificationDimension::INTRA_NODE}}}; - std::unordered_set correct = { + std::set correct = { device_id_t{gpu_id_t{1_n}}, device_id_t{gpu_id_t{3_n}}, device_id_t{gpu_id_t{5_n}}, }; - std::unordered_set result = get_device_ids(task, mv, ms); + std::set result = get_device_ids(task, mv, ms); CHECK(result == correct); } @@ -453,13 +453,13 @@ TEST_SUITE(FF_TEST_SUITE) { MachineViewDimension{stride_t{2_p}, MachineSpecificationDimension::INTRA_NODE}}}; - std::unordered_set correct = { + std::set correct = { device_id_t{gpu_id_t{7_n}}, device_id_t{gpu_id_t{9_n}}, device_id_t{gpu_id_t{12_n}}, device_id_t{gpu_id_t{14_n}}, }; - std::unordered_set result = get_device_ids(task, mv, ms); + std::set result = get_device_ids(task, mv, ms); CHECK(result == correct); } } diff --git a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc index 54717d6699..05a7a4cc4d 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc @@ -266,13 +266,13 @@ TEST_SUITE(FF_TEST_SUITE) { }; CostEstimator cost_estimator = make_fake_cost_estimator( - std::unordered_map{{ + std::map{{ { map_unmapped_op_cost_estimate_key(k1, mv1), k1_on_mv1_cost, }, }}, - std::unordered_map{ + std::map{ { empty_tensor_set_movement(), 0_ms, @@ -290,7 +290,7 @@ TEST_SUITE(FF_TEST_SUITE) { MachineComputeResourceSlice const &resources) { ASSERT(k == runtime_only_from_unmapped_op_cost_estimate_key(k1)); ASSERT(resources == four_nodes_resources); - return std::unordered_set{mv1}; + return std::set{mv1}; }; MachineMappingWithMemoryContext context = MachineMappingWithMemoryContext{ @@ -325,7 +325,7 @@ TEST_SUITE(FF_TEST_SUITE) { MachineComputeResourceSlice const &resources) { ASSERT(k == runtime_only_from_unmapped_op_cost_estimate_key(k3)); ASSERT(resources == four_nodes_resources); - return std::unordered_set{mv2, mv3, mv4}; + return std::set{mv2, mv3, mv4}; }; OpCostMetrics k3_on_mv2_cost = OpCostMetrics{ @@ -347,7 +347,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; CostEstimator cost_estimator = make_fake_cost_estimator( - std::unordered_map{{ + std::map{{ { map_unmapped_op_cost_estimate_key(k3, mv2), k3_on_mv2_cost, @@ -361,7 +361,7 @@ TEST_SUITE(FF_TEST_SUITE) { k3_on_mv4_cost, }, }}, - std::unordered_map{ + std::map{ { empty_tensor_set_movement(), 0_ms, @@ -474,7 +474,7 @@ TEST_SUITE(FF_TEST_SUITE) { milliseconds_t mv3_to_mv2_cost, milliseconds_t mv3_to_mv3_cost) { return make_fake_cost_estimator( - std::unordered_map{{ + std::map{{ { map_unmapped_op_cost_estimate_key(k2, mv2), OpCostMetrics{ @@ -508,7 +508,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, }}, - std::unordered_map{{ + std::map{{ { empty_tensor_set_movement(), 0_ms, @@ -561,18 +561,18 @@ TEST_SUITE(FF_TEST_SUITE) { [&](UnmappedRuntimeOnlyOpCostEstimateKey const &k, MachineComputeResourceSlice const &resources) { if (k == runtime_only_from_unmapped_op_cost_estimate_key(k1)) { - return std::unordered_set{ + return std::set{ mv1, }; } else { if (resources == four_nodes_resources) { - return std::unordered_set{mv2, mv3}; + return std::set{mv2, mv3}; } else if (resources == three_nodes_resources) { - return std::unordered_set{mv2, mv3}; + return std::set{mv2, mv3}; } else if (resources == two_nodes_resources) { - return std::unordered_set{mv2}; + return std::set{mv2}; } else { - return std::unordered_set{}; + return std::set{}; } }; }; @@ -628,7 +628,7 @@ TEST_SUITE(FF_TEST_SUITE) { milliseconds_t k3_on_mv3_cost, num_bytes_t k3_on_mv3_mem_usage) { return make_fake_cost_estimator( - std::unordered_map{{ + std::map{{ { map_unmapped_op_cost_estimate_key(k2, mv2), OpCostMetrics{ @@ -662,7 +662,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, }}, - std::unordered_map{ + std::map{ { empty_tensor_set_movement(), 0_ms, @@ -684,18 +684,18 @@ TEST_SUITE(FF_TEST_SUITE) { [&](UnmappedRuntimeOnlyOpCostEstimateKey const &k, MachineComputeResourceSlice const &resources) { if (k == runtime_only_from_unmapped_op_cost_estimate_key(k1)) { - return std::unordered_set{ + return std::set{ mv1, }; } else { if (resources == four_nodes_resources) { - return std::unordered_set{mv2, mv3}; + return std::set{mv2, mv3}; } else if (resources == three_nodes_resources) { - return std::unordered_set{mv2, mv3}; + return std::set{mv2, mv3}; } else if (resources == two_nodes_resources) { - return std::unordered_set{mv2}; + return std::set{mv2}; } else { - return std::unordered_set{}; + return std::set{}; } }; }; diff --git a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc index 402dbe66d7..9ef6aa90d0 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc @@ -1,6 +1,6 @@ #include "compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.h" #include "compiler/machine_mapping/machine_view.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/rapidcheck/some.h" #include "utils/nonnegative_int/nonnegative_int.h" #include @@ -71,10 +71,10 @@ TEST_SUITE(FF_TEST_SUITE) { mapping3, }}; - std::unordered_set result = + std::set result = mapping_result.get_pareto_frontier(); - std::unordered_set correct = { + std::set correct = { mapping1, mapping2, mapping3, @@ -87,10 +87,10 @@ TEST_SUITE(FF_TEST_SUITE) { MachineMappingWithMemoryResult mapping_result = MachineMappingWithMemoryResult{{}}; - std::unordered_set result = + std::set result = mapping_result.get_pareto_frontier(); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } diff --git a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc index 3159f49118..84f08fdd52 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc @@ -1,6 +1,6 @@ #include "compiler/machine_mapping/start_invariant_machine_view.h" #include "op-attrs/task_space_coordinate.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/fmt/vector.h" #include @@ -140,11 +140,11 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_machine_space_offsets") { - std::unordered_set correct = { + std::set correct = { MachineSpaceOffset{0, 0, DeviceType::GPU}, MachineSpaceOffset{0, 2, DeviceType::GPU}, MachineSpaceOffset{0, 4, DeviceType::GPU}}; - std::unordered_set result = + std::set result = get_machine_space_offsets(task, simv); CHECK(correct == result); } @@ -223,12 +223,12 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_machine_space_offsets") { - std::unordered_set correct = { + std::set correct = { MachineSpaceOffset{0, 0, DeviceType::GPU}, MachineSpaceOffset{0, 2, DeviceType::GPU}, MachineSpaceOffset{1, 0, DeviceType::GPU}, MachineSpaceOffset{1, 2, DeviceType::GPU}}; - std::unordered_set result = + std::set result = get_machine_space_offsets(task, simv); CHECK(correct == result); } diff --git a/lib/compiler/test/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc b/lib/compiler/test/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc index ced2634000..5e02dbba48 100644 --- a/lib/compiler/test/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc +++ b/lib/compiler/test/src/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.cc @@ -1,5 +1,5 @@ #include "compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "test/utils/rapidcheck.h" #include diff --git a/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc b/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc index e88f2b7840..3fd3bdcb56 100644 --- a/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc +++ b/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc @@ -27,8 +27,8 @@ TEST_SUITE(FF_TEST_SUITE) { auto is_allowed_to_run = [&](Node const &n, - std::unordered_set const &in_progress_tasks, - std::unordered_set const &finished_tasks) { return true; }; + std::set const &in_progress_tasks, + std::set const &finished_tasks) { return true; }; TaskExecutionConstraint constraint = TaskExecutionConstraint{is_allowed_to_run}; @@ -59,8 +59,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("no processing constraints") { auto is_allowed_to_run = [&](Node const &n, - std::unordered_set const &in_progress_tasks, - std::unordered_set const &finished_tasks) { + std::set const &in_progress_tasks, + std::set const &finished_tasks) { return true; }; @@ -80,8 +80,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("one node at a time") { auto is_allowed_to_run = [&](Node const &n, - std::unordered_set const &in_progress_tasks, - std::unordered_set const &finished_tasks) { + std::set const &in_progress_tasks, + std::set const &finished_tasks) { return in_progress_tasks.size() == 0; }; @@ -123,8 +123,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("no processing constraints") { auto is_allowed_to_run = [&](Node const &n, - std::unordered_set const &in_progress_tasks, - std::unordered_set const &finished_tasks) { + std::set const &in_progress_tasks, + std::set const &finished_tasks) { return true; }; @@ -146,8 +146,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("one node at a time") { auto is_allowed_to_run = [&](Node const &n, - std::unordered_set const &in_progress_tasks, - std::unordered_set const &finished_tasks) { + std::set const &in_progress_tasks, + std::set const &finished_tasks) { return in_progress_tasks.size() == 0; }; @@ -187,8 +187,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("at most two nodes at a time") { auto is_allowed_to_run = [&](Node const &n, - std::unordered_set const &in_progress_tasks, - std::unordered_set const &finished_tasks) { + std::set const &in_progress_tasks, + std::set const &finished_tasks) { return in_progress_tasks.size() < 2; }; diff --git a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc index 2846de6559..07f59ac8d5 100644 --- a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc @@ -29,8 +29,8 @@ #include "utils/nonnegative_int/nonnegative_int.h" #include #include -#include -#include +#include +#include namespace FlexFlow { diff --git a/lib/compiler/test/src/internal/cost_estimator_for_test.cc b/lib/compiler/test/src/internal/cost_estimator_for_test.cc index 7092a3848b..dd161c411d 100644 --- a/lib/compiler/test/src/internal/cost_estimator_for_test.cc +++ b/lib/compiler/test/src/internal/cost_estimator_for_test.cc @@ -35,8 +35,8 @@ CostEstimator make_fake_cost_estimator( } CostEstimator make_fake_cost_estimator( - std::unordered_map const &op_cost_map, - std::unordered_map const + std::map const &op_cost_map, + std::map const &comm_cost_map) { return make_fake_cost_estimator( [op_cost_map](OpCostEstimateKey const &k) { diff --git a/lib/compiler/test/src/internal/cost_estimator_for_test.h b/lib/compiler/test/src/internal/cost_estimator_for_test.h index 12708210f3..7af8f9c8ad 100644 --- a/lib/compiler/test/src/internal/cost_estimator_for_test.h +++ b/lib/compiler/test/src/internal/cost_estimator_for_test.h @@ -34,8 +34,8 @@ CostEstimator make_fake_cost_estimator( &get_communication_cost); CostEstimator make_fake_cost_estimator( - std::unordered_map const &op_cost_map, - std::unordered_map const &comm_cost_map); + std::map const &op_cost_map, + std::map const &comm_cost_map); CostEstimator make_fake_constant_cost_estimator(milliseconds_t forward_op_cost, milliseconds_t backward_op_cost, diff --git a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc index 59bf08a399..5a27e78d15 100644 --- a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc +++ b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc @@ -26,9 +26,9 @@ RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( } RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( - std::unordered_map const &op_cost_map, - std::unordered_map const + std::map const &comm_cost_map) { return make_fake_runtime_only_cost_estimator( [op_cost_map](RuntimeOnlyOpCostEstimateKey const &k) { diff --git a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h index 2b5824263d..a09b75c967 100644 --- a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h +++ b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h @@ -12,9 +12,9 @@ RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( &get_communication_cost); RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( - std::unordered_map const &op_cost_map, - std::unordered_map const &comm_cost_map); + std::map const &comm_cost_map); RuntimeOnlyCostEstimator make_fake_constant_runtime_only_cost_estimator( milliseconds_t forward_op_cost, diff --git a/lib/kernels/include/kernels/local_cpu_allocator.h b/lib/kernels/include/kernels/local_cpu_allocator.h index 9653dcf00e..9cee590a08 100644 --- a/lib/kernels/include/kernels/local_cpu_allocator.h +++ b/lib/kernels/include/kernels/local_cpu_allocator.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_KERNELS_INCLUDE_KERNELS_LOCAL_CPU_ALLOCATOR_H #include "kernels/allocation.h" -#include +#include namespace FlexFlow { @@ -18,7 +18,7 @@ struct LocalCPUAllocator : public IAllocator { DeviceType get_allocation_device_type() const override; private: - std::unordered_map> ptrs; + std::map> ptrs; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(LocalCPUAllocator); diff --git a/lib/kernels/include/kernels/local_cuda_allocator.h b/lib/kernels/include/kernels/local_cuda_allocator.h index b8e0540974..f95db7a9ef 100644 --- a/lib/kernels/include/kernels/local_cuda_allocator.h +++ b/lib/kernels/include/kernels/local_cuda_allocator.h @@ -1,5 +1,5 @@ #include "kernels/allocation.h" -#include +#include namespace FlexFlow { @@ -15,7 +15,7 @@ struct LocalCudaAllocator : public IAllocator { DeviceType get_allocation_device_type() const override; private: - std::unordered_set ptrs; + std::set ptrs; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(LocalCudaAllocator); diff --git a/lib/kernels/include/kernels/reduce_tensor_accessor.h b/lib/kernels/include/kernels/reduce_tensor_accessor.h index 02ff63544f..12adf3bc10 100644 --- a/lib/kernels/include/kernels/reduce_tensor_accessor.h +++ b/lib/kernels/include/kernels/reduce_tensor_accessor.h @@ -22,7 +22,7 @@ struct CPUReduceTensorAccessorInDims { template void operator()(GenericTensorAccessorR const &input, GenericTensorAccessorW &output, - std::unordered_set const &dims_to_reduce, + std::set const &dims_to_reduce, F &&f) { using T = real_type_t
; @@ -60,7 +60,7 @@ struct CPUReduceTensorAccessorInDims { template GenericTensorAccessorW reduce_tensor_accessor_in_dims(GenericTensorAccessorR const &input, - std::unordered_set const &dims, + std::set const &dims, Allocator &output_allocator, F &&f) { @@ -89,7 +89,7 @@ real_type_t
F &&f) { Allocator cpu_allocator = create_local_cpu_memory_allocator(); - std::unordered_set input_dims = get_ff_dim_t_set(input.shape.dims); + std::set input_dims = get_ff_dim_t_set(input.shape.dims); GenericTensorAccessorW reduced = reduce_tensor_accessor_in_dims(input, input_dims, cpu_allocator, f); diff --git a/lib/kernels/src/kernels/reduce_tensor_accessor.cc b/lib/kernels/src/kernels/reduce_tensor_accessor.cc index b51306d0e8..5cdea23210 100644 --- a/lib/kernels/src/kernels/reduce_tensor_accessor.cc +++ b/lib/kernels/src/kernels/reduce_tensor_accessor.cc @@ -6,7 +6,7 @@ using F = std::function; template GenericTensorAccessorW reduce_tensor_accessor_in_dims(GenericTensorAccessorR const &, - std::unordered_set const &, + std::set const &, Allocator &, F &&); diff --git a/lib/local-execution/include/local-execution/computation_graph_instance.h b/lib/local-execution/include/local-execution/computation_graph_instance.h index a4ded5edaf..4d9b5617fe 100644 --- a/lib/local-execution/include/local-execution/computation_graph_instance.h +++ b/lib/local-execution/include/local-execution/computation_graph_instance.h @@ -15,7 +15,7 @@ #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" #include "utils/units/milliseconds_t.h" #include -#include +#include namespace FlexFlow { @@ -44,26 +44,26 @@ ComputationGraphInstance create_computation_graph_instance( ComputationGraph const &cg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::unordered_map const + std::map const &input_tensors, Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, device_id_t device_idx); -std::unordered_map> +std::map> perform_all_passes_for_computation_graph_instance( ComputationGraphInstance &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, device_id_t device_idx); -std::unordered_map> +std::map> perform_forward_pass_for_computation_graph_instance( ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, device_id_t device_idx); -std::unordered_map> +std::map> perform_backward_pass_for_computation_graph_instance( ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, diff --git a/lib/local-execution/include/local-execution/cost_estimator/tracked_allocator.h b/lib/local-execution/include/local-execution/cost_estimator/tracked_allocator.h index 0b531f9b3d..7be631f87d 100644 --- a/lib/local-execution/include/local-execution/cost_estimator/tracked_allocator.h +++ b/lib/local-execution/include/local-execution/cost_estimator/tracked_allocator.h @@ -21,7 +21,7 @@ struct TrackedAllocator : public IAllocator { private: size_t current_mem_usage = 0; - std::unordered_map ptr_mem_usage; + std::map ptr_mem_usage; Allocator allocator; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(TrackedAllocator); diff --git a/lib/local-execution/include/local-execution/local_task_argument_accessor.h b/lib/local-execution/include/local-execution/local_task_argument_accessor.h index 12eab4a76d..4d1b5fc53a 100644 --- a/lib/local-execution/include/local-execution/local_task_argument_accessor.h +++ b/lib/local-execution/include/local-execution/local_task_argument_accessor.h @@ -6,14 +6,14 @@ #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" #include "task-spec/task_argument_accessor/itask_argument_accessor.h" #include "task-spec/task_argument_accessor/task_tensor_parameter.dtg.h" -#include +#include namespace FlexFlow { struct LocalTaskArgumentAccessor : public ITaskArgumentAccessor { explicit LocalTaskArgumentAccessor( Allocator const &allocator, - std::unordered_map const + std::map const &tensor_slots_backing, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, @@ -45,7 +45,7 @@ struct LocalTaskArgumentAccessor : public ITaskArgumentAccessor { private: Allocator allocator; - std::unordered_map + std::map tensor_slots_backing; ProfilingSettings profiling_settings; diff --git a/lib/local-execution/include/local-execution/tensor_allocation.h b/lib/local-execution/include/local-execution/tensor_allocation.h index 76fb3bbee6..ad2b4b0de5 100644 --- a/lib/local-execution/include/local-execution/tensor_allocation.h +++ b/lib/local-execution/include/local-execution/tensor_allocation.h @@ -16,7 +16,7 @@ DynamicValueAttrs perform_tensor_allocation_for_value(DynamicValueAttrs const &, DynamicOpenDataflowGraph perform_tensor_allocation( DynamicOpenDataflowGraph const &, - std::unordered_map const + std::map const &preallocated, Allocator &); diff --git a/lib/local-execution/src/local-execution/computation_graph_instance.cc b/lib/local-execution/src/local-execution/computation_graph_instance.cc index d2781472b0..963007b68c 100644 --- a/lib/local-execution/src/local-execution/computation_graph_instance.cc +++ b/lib/local-execution/src/local-execution/computation_graph_instance.cc @@ -15,7 +15,7 @@ #include "task-spec/per_device_op_state.h" #include "task-spec/task_argument_accessor/task_argument_accessor.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_map_from_pairs.h" +#include "utils/containers/map_from_pairs.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" #include "utils/optional.h" #include @@ -62,7 +62,7 @@ ComputationGraphInstance create_computation_graph_instance( ComputationGraph const &cg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::unordered_map const + std::map const &input_tensors, Allocator &allocator, ProfilingSettings const &profiling_settings, @@ -71,7 +71,7 @@ ComputationGraphInstance create_computation_graph_instance( DynamicOpenDataflowGraph dg = make_dynamic_open_dataflow_graph_from_cg(cg); dg = perform_pass_expansion(dg); - std::unordered_map inputs = + std::map inputs = input_tensors; std::optional logit_grad_value; if (loss.has_value()) { @@ -109,7 +109,7 @@ ComputationGraphInstance create_computation_graph_instance( invocation_topo_order, allocator, optimizer_attrs, logit_grad_tensor}; } -static std::unordered_map> +static std::map> execute_dynamic_node_invocation_set( std::vector const &invocations, Allocator &allocator, @@ -117,7 +117,7 @@ static std::unordered_map> ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, device_id_t device_idx) { - return unordered_map_from_pairs( + return map_from_pairs( transform(invocations, [&](DynamicNodeInvocation const &invocation) { std::optional timing = execute_dynamic_node_invocation( /*invocation=*/invocation, @@ -136,7 +136,7 @@ static std::unordered_map> })); } -std::unordered_map> +std::map> perform_all_passes_for_computation_graph_instance( ComputationGraphInstance &instance, ProfilingSettings const &profiling_settings, @@ -144,7 +144,7 @@ std::unordered_map> device_id_t device_idx) { std::vector execution_order = instance.get_execution_order(); - std::unordered_map> + std::map> result = execute_dynamic_node_invocation_set( /*invocations=*/execution_order, /*allocator=*/instance.get_allocator(), @@ -156,7 +156,7 @@ std::unordered_map> return result; } -std::unordered_map> +std::map> perform_forward_pass_for_computation_graph_instance( ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, @@ -179,7 +179,7 @@ std::unordered_map> /*device_idx=*/device_idx); } -std::unordered_map> +std::map> perform_backward_pass_for_computation_graph_instance( ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, diff --git a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc index 0c6107c4af..361ac149c2 100644 --- a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -19,7 +19,7 @@ #include "utils/containers/require_only_key.h" #include "utils/containers/sum.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/exception.h" #include "utils/optional.h" @@ -39,12 +39,12 @@ LocalCostEstimator::LocalCostEstimator( static ComputationGraph computation_graph_for_local_cost_estimation( ComputationGraphOpAttrs const &op, - std::unordered_map const &inputs, - std::unordered_map const &weights, - std::unordered_map const &outputs) { + std::map const &inputs, + std::map const &weights, + std::map const &outputs) { ComputationGraph computation_graph = make_empty_computation_graph(); - std::unordered_map input_tensors = + std::map input_tensors = map_values(inputs, [&](ParallelTensorShape const &shape) { LayerAddedResult inputs_layer = add_layer(computation_graph, @@ -56,7 +56,7 @@ static ComputationGraph computation_graph_for_local_cost_estimation( return require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); }); - std::unordered_map weight_tensors = + std::map weight_tensors = map_values(weights, [&](ParallelTensorShape const &shape) { LayerAddedResult weights_layer = add_layer(computation_graph, @@ -85,11 +85,11 @@ OpCostMetrics LocalCostEstimator::estimate_cost( OpCostEstimateKey const &op_cost_estimate_key) const { PCGOperatorAttrs op = op_cost_estimate_key.op_attrs; - std::unordered_map inputs = + std::map inputs = op_cost_estimate_key.input_shapes; - std::unordered_map weights = + std::map weights = op_cost_estimate_key.weight_shapes; - std::unordered_map outputs = + std::map outputs = op_cost_estimate_key.output_shapes; OptimizerAttrs optimizer_attrs = op_cost_estimate_key.optimizer_attrs; @@ -130,14 +130,14 @@ OpCostMetrics LocalCostEstimator::estimate_cost( // execute layer dynamic_layer_guid_t operator_layer_guid{get_layer_by_name(cg, "operator")}; - std::unordered_map> + std::map> fwd_timing = perform_forward_pass_for_computation_graph_instance( instance, this->profiling_settings, this->device_handle, this->device_idx); milliseconds_t fwd = fwd_timing.at(operator_layer_guid).value(); - std::unordered_map> + std::map> bwd_timing = perform_backward_pass_for_computation_graph_instance( instance, this->profiling_settings, @@ -170,7 +170,7 @@ milliseconds_t LocalCostEstimator::estimate_cost( }; return maximum( - transform(unordered_set_of(tensor_set_movement.edge_to_size), + transform(set_of(tensor_set_movement.edge_to_size), [&](std::pair const &p) { return estimate_single_comm_cost( p.first.get_src(), p.first.get_dst(), p.second); diff --git a/lib/local-execution/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/src/local-execution/local_task_argument_accessor.cc index b8feca720e..9e69c48cab 100644 --- a/lib/local-execution/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/src/local-execution/local_task_argument_accessor.cc @@ -10,7 +10,7 @@ namespace FlexFlow { LocalTaskArgumentAccessor::LocalTaskArgumentAccessor( Allocator const &allocator, - std::unordered_map const + std::map const &tensor_slots_backing, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, diff --git a/lib/local-execution/src/local-execution/task_execution.cc b/lib/local-execution/src/local-execution/task_execution.cc index b9ceef0e1b..ad2144b07f 100644 --- a/lib/local-execution/src/local-execution/task_execution.cc +++ b/lib/local-execution/src/local-execution/task_execution.cc @@ -14,7 +14,6 @@ #include "utils/optional.h" #include "utils/overload.h" #include -#include "utils/containers/unordered_map_from_map.h" namespace FlexFlow { @@ -57,10 +56,10 @@ TaskArgumentAccessor make_task_argument_accessor_for_invocation( auto get_accessor = [](DynamicValueAttrs const &value) { return assert_unwrap(value.accessor); }; - std::unordered_map - tensor_slots_backing = unordered_map_from_map(binary_merge_disjoint_maps( + std::map + tensor_slots_backing = binary_merge_disjoint_maps( map_keys_and_values(invocation.inputs, make_param, get_accessor), - map_keys_and_values(invocation.outputs, make_param, get_accessor))); + map_keys_and_values(invocation.outputs, make_param, get_accessor)); return TaskArgumentAccessor::create( /*allocator=*/allocator, diff --git a/lib/local-execution/src/local-execution/tensor_allocation.cc b/lib/local-execution/src/local-execution/tensor_allocation.cc index 203345a9af..bf3e4de4f4 100644 --- a/lib/local-execution/src/local-execution/tensor_allocation.cc +++ b/lib/local-execution/src/local-execution/tensor_allocation.cc @@ -6,7 +6,7 @@ #include "utils/containers/all_are_true.h" #include "utils/containers/contains_key.h" #include "utils/containers/map_values.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/optional.h" namespace FlexFlow { @@ -50,17 +50,17 @@ DynamicValueAttrs DynamicOpenDataflowGraph perform_tensor_allocation( DynamicOpenDataflowGraph const &g, - std::unordered_map const + std::map const &preallocated, Allocator &allocator) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); - for (DynamicValueAttrs const &v : unordered_keys(preallocated)) { + for (DynamicValueAttrs const &v : keys(preallocated)) { ASSERT(v.accessor == std::nullopt); } - std::unordered_set all_values = - unordered_set_of(get_dynamic_values(g)); + std::set all_values = + set_of(get_dynamic_values(g)); bidict unallocated_to_allocated = generate_bidict( diff --git a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc index 2a4e204d59..4aea9fbc28 100644 --- a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc +++ b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc @@ -143,7 +143,7 @@ TEST_SUITE(FF_TEST_SUITE) { device_id_t device_idx = make_device_id_t_from_idx(nonnegative_int{0}, DeviceType::CPU); - std::unordered_map input_tensors; + std::map input_tensors; ComputationGraphInstance computation_graph_instance = create_computation_graph_instance( @@ -312,7 +312,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { device_handle_t ff_handle = gpu_make_device_handle_t(managed_handle.raw_handle()); - std::unordered_map input_tensors; + std::map input_tensors; ComputationGraphInstance computation_graph_instance = create_computation_graph_instance( @@ -428,7 +428,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { device_handle_t ff_handle = gpu_make_device_handle_t(managed_handle.raw_handle()); - std::unordered_map input_tensors; + std::map input_tensors; auto compute_loss = [&](LossAttrs const &loss_attrs, GenericTensorAccessorR label_tensor) { diff --git a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc index 07bb869d5f..24362494aa 100644 --- a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc @@ -40,7 +40,7 @@ TEST_SUITE(FF_TEST_SUITE) { VARIADIC_TENSORS, }; - std::unordered_map + std::map tensor_slots_backing = { { make_task_tensor_parameter_fwd(TensorSlotName::LHS_INPUT), diff --git a/lib/op-attrs/include/op-attrs/ff_ordered/ff_ordered_from_map.h b/lib/op-attrs/include/op-attrs/ff_ordered/ff_ordered_from_map.h index 9232afddfb..e5be0f984d 100644 --- a/lib/op-attrs/include/op-attrs/ff_ordered/ff_ordered_from_map.h +++ b/lib/op-attrs/include/op-attrs/ff_ordered/ff_ordered_from_map.h @@ -8,7 +8,7 @@ namespace FlexFlow { template -FFOrdered ff_ordered_from_map(std::map const &m) { +FFOrdered ff_ordered_from_map(std::unordered_map const &m) { std::vector raw; for (int i = 0; i < m.size(); i++) { raw.push_back(m.at(ff_dim_t{nonnegative_int{i}})); @@ -17,7 +17,7 @@ FFOrdered ff_ordered_from_map(std::map const &m) { } template -FFOrdered ff_ordered_from_map(std::unordered_map const &m) { +FFOrdered ff_ordered_from_map(std::map const &m) { std::vector raw; for (int i = 0; i < m.size(); i++) { raw.push_back(m.at(ff_dim_t{nonnegative_int{i}})); diff --git a/lib/op-attrs/include/op-attrs/ff_ordered/map_from_ff_ordered.h b/lib/op-attrs/include/op-attrs/ff_ordered/map_from_ff_ordered.h index 4a7e564c20..1f88791a22 100644 --- a/lib/op-attrs/include/op-attrs/ff_ordered/map_from_ff_ordered.h +++ b/lib/op-attrs/include/op-attrs/ff_ordered/map_from_ff_ordered.h @@ -8,8 +8,8 @@ namespace FlexFlow { template -std::unordered_map map_from_ff_ordered(FFOrdered const &m) { - std::unordered_map result; +std::map map_from_ff_ordered(FFOrdered const &m) { + std::map result; for (ff_dim_t d : ff_dim_range(num_elements(m))) { result.insert({d, m.at(d)}); diff --git a/lib/op-attrs/include/op-attrs/get_incoming_tensor_roles.h b/lib/op-attrs/include/op-attrs/get_incoming_tensor_roles.h index 0ad9a9c062..409491950c 100644 --- a/lib/op-attrs/include/op-attrs/get_incoming_tensor_roles.h +++ b/lib/op-attrs/include/op-attrs/get_incoming_tensor_roles.h @@ -8,9 +8,9 @@ namespace FlexFlow { -std::unordered_map +std::map get_incoming_tensor_roles(ComputationGraphOpAttrs const &); -std::unordered_map +std::map get_incoming_tensor_roles(PCGOperatorAttrs const &); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.h b/lib/op-attrs/include/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.h index 3a6b4732e6..2dcd402f35 100644 --- a/lib/op-attrs/include/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.h +++ b/lib/op-attrs/include/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.h @@ -12,48 +12,48 @@ namespace FlexFlow { -std::unordered_map +std::map get_operator_to_incoming_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees); -std::unordered_map +std::map get_operator_to_incoming_mappings_for_role( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees, IncomingTensorRole role); -std::unordered_map +std::map get_operator_to_input_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees); -std::unordered_map +std::map get_operator_to_weight_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees); -std::unordered_map +std::map get_operator_to_output_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees); -std::unordered_map +std::map get_operator_to_ptensor_mappings_for_role( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees, TensorRole role); -std::unordered_map +std::map get_operator_to_ptensor_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/get_operator_task_space.h b/lib/op-attrs/include/op-attrs/get_operator_task_space.h index 9ee4a8779a..462cdddc4d 100644 --- a/lib/op-attrs/include/op-attrs/get_operator_task_space.h +++ b/lib/op-attrs/include/op-attrs/get_operator_task_space.h @@ -10,7 +10,7 @@ namespace FlexFlow { OperatorTaskSpace get_operator_task_space( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/operator_task_space.h b/lib/op-attrs/include/op-attrs/operator_task_space.h index 426fbc1850..81345c22a5 100644 --- a/lib/op-attrs/include/op-attrs/operator_task_space.h +++ b/lib/op-attrs/include/op-attrs/operator_task_space.h @@ -8,16 +8,16 @@ #include "utils/orthotope/dim_domain.dtg.h" #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/minimal_dim_domain.dtg.h" -#include +#include namespace FlexFlow { OperatorTaskSpace trivial_op_task_space(); -std::unordered_set +std::set operator_task_space_get_dim_idxs(OperatorTaskSpace const &); -std::unordered_set +std::set get_task_space_coordinates(OperatorTaskSpace const &operator_task_space); bool operator_task_space_contains_coord(OperatorTaskSpace const &, diff --git a/lib/op-attrs/include/op-attrs/ops/attention.h b/lib/op-attrs/include/op-attrs/ops/attention.h index fdd3f3775f..76f63f6780 100644 --- a/lib/op-attrs/include/op-attrs/ops/attention.h +++ b/lib/op-attrs/include/op-attrs/ops/attention.h @@ -39,7 +39,7 @@ positive_int get_kvSeqLength(MultiHeadAttentionInputs const &); positive_int get_num_samples(MultiHeadAttentionParallelInputs const &); positive_int get_num_samples(MultiHeadAttentionInputs const &); -std::unordered_map +std::map get_attention_incoming_tensor_roles(MultiHeadAttentionAttrs const &); tl::expected @@ -63,7 +63,7 @@ tl::expected TensorShape const &input_k, TensorShape const &input_v); -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(MultiHeadAttentionAttrs const &, TensorShape const &input_q, TensorShape const &input_k, @@ -106,14 +106,14 @@ tl::expected ParallelTensorShape const &input_k, ParallelTensorShape const &input_v); -tl::expected, +tl::expected, std::string> get_weight_shapes(MultiHeadAttentionAttrs const &, ParallelTensorShape const &input_q, ParallelTensorShape const &input_k, ParallelTensorShape const &input_v); -tl::expected, std::string> +tl::expected, std::string> get_initializers( MultiHeadAttentionAttrs const &, TensorShape const &input_q, diff --git a/lib/op-attrs/include/op-attrs/ops/batch_norm.h b/lib/op-attrs/include/op-attrs/ops/batch_norm.h index bbdb52cecc..9422c14c6c 100644 --- a/lib/op-attrs/include/op-attrs/ops/batch_norm.h +++ b/lib/op-attrs/include/op-attrs/ops/batch_norm.h @@ -12,7 +12,7 @@ namespace FlexFlow { -std::unordered_map +std::map get_batch_norm_incoming_tensor_roles(BatchNormAttrs const &); tl::expected get_output_shape(BatchNormAttrs const &, @@ -22,7 +22,7 @@ tl::expected tl::expected get_beta_weights_shape(BatchNormAttrs const &, TensorShape const &); -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(BatchNormAttrs const &attrs, TensorShape const &input_shape); @@ -36,7 +36,7 @@ tl::expected get_beta_weights_parallel_dim_degrees(BatchNormAttrs const &, ParallelTensorDimDegrees const &); -tl::expected, +tl::expected, std::string> get_weight_parallel_dim_degrees( BatchNormAttrs const &attrs, @@ -50,7 +50,7 @@ tl::expected tl::expected get_beta_weights_shape(BatchNormAttrs const &, ParallelTensorShape const &); -tl::expected, +tl::expected, std::string> get_weight_shapes(BatchNormAttrs const &attrs, ParallelTensorShape const &input_shape); @@ -61,7 +61,7 @@ tl::expected, * see * https://github.com/pytorch/pytorch/blob/1eba9b3aa3c43f86f4a2c807ac8e12c4a7767340/torch/nn/modules/batchnorm.py#L93-L97 */ -tl::expected, std::string> +tl::expected, std::string> get_initializers(BatchNormAttrs const &attrs); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/ops/conv_2d.h b/lib/op-attrs/include/op-attrs/ops/conv_2d.h index 0f27b00406..8ae5450e47 100644 --- a/lib/op-attrs/include/op-attrs/ops/conv_2d.h +++ b/lib/op-attrs/include/op-attrs/ops/conv_2d.h @@ -10,7 +10,7 @@ namespace FlexFlow { -std::unordered_map +std::map get_conv2d_incoming_tensor_roles(Conv2DAttrs const &); TensorShape get_kernel_shape(Conv2DAttrs const &attrs, @@ -19,7 +19,7 @@ TensorShape get_bias_shape(Conv2DAttrs const &attrs, TensorShape const &input); TensorShape get_output_shape(Conv2DAttrs const &attrs, TensorShape const &input); -std::unordered_map +std::map get_weight_shapes(Conv2DAttrs const &attrs, TensorShape const &input_shape); ParallelTensorShape get_kernel_shape(Conv2DAttrs const &attrs, @@ -29,11 +29,11 @@ ParallelTensorShape get_bias_shape(Conv2DAttrs const &attrs, ParallelTensorShape get_output_shape(Conv2DAttrs const &attrs, ParallelTensorShape const &input_shape); -std::unordered_map +std::map get_weight_shapes(Conv2DAttrs const &attrs, ParallelTensorShape const &input_shape); -std::unordered_map get_initializers( +std::map get_initializers( Conv2DAttrs const &attrs, TensorShape const &input_shape, std::optional kernel_initializer = std::nullopt, diff --git a/lib/op-attrs/include/op-attrs/ops/embedding.h b/lib/op-attrs/include/op-attrs/ops/embedding.h index 7ae8350dee..97f9bca7b5 100644 --- a/lib/op-attrs/include/op-attrs/ops/embedding.h +++ b/lib/op-attrs/include/op-attrs/ops/embedding.h @@ -27,7 +27,7 @@ tl::expected * see * https://github.com/pytorch/pytorch/blob/1eba9b3aa3c43f86f4a2c807ac8e12c4a7767340/torch/nn/modules/sparse.py#L180-L182 */ -std::unordered_map get_initializers( +std::map get_initializers( EmbeddingAttrs const &, std::optional const &initializer_attrs = std::nullopt); diff --git a/lib/op-attrs/include/op-attrs/ops/index.dox b/lib/op-attrs/include/op-attrs/ops/index.dox index 6e5465ca68..b43aba0e70 100644 --- a/lib/op-attrs/include/op-attrs/ops/index.dox +++ b/lib/op-attrs/include/op-attrs/ops/index.dox @@ -39,7 +39,7 @@ More specifically, this consists of the following pieces: Note that as different operators have different numbers of inputs, etc. the number and signatures of these functions may be different for different operators. While keeping the structure of the various operators similar is makes it easier to understand, it's not strictly necessary: the code that calls these functions for a generic operator allows custom behavior for each operator, which allows us to have a bit more freedom to evolve operator definitions over time: - \ref get_operator_to_ptensor_mappings (and associated functions in \ref get_operator_space_to_parallel_tensor_space_mappings.h) - \ref "get_incoming_tensor_roles(ComputationGraphOpAttrs const &)" (and associated functions in \ref get_incoming_tensor_roles.h) -- \ref "get_output_shapes(ComputationGraphOpAttrs const &, std::unordered_map const &input_shapes)" (and associated functions in \ref op-attrs/shape_inference.h) +- \ref "get_output_shapes(ComputationGraphOpAttrs const &, std::map const &input_shapes)" (and associated functions in \ref op-attrs/shape_inference.h) */ } diff --git a/lib/op-attrs/include/op-attrs/ops/layer_norm.h b/lib/op-attrs/include/op-attrs/ops/layer_norm.h index 00c1ad9b12..7e1b06483d 100644 --- a/lib/op-attrs/include/op-attrs/ops/layer_norm.h +++ b/lib/op-attrs/include/op-attrs/ops/layer_norm.h @@ -11,7 +11,7 @@ namespace FlexFlow { -std::unordered_map +std::map get_layer_norm_incoming_tensor_roles(LayerNormAttrs const &); tl::expected get_output_shape(LayerNormAttrs const &, @@ -21,7 +21,7 @@ tl::expected tl::expected get_beta_weights_shape(LayerNormAttrs const &, TensorShape const &); -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(LayerNormAttrs const &attrs, TensorShape const &input_shape); @@ -33,7 +33,7 @@ tl::expected tl::expected get_beta_weights_shape(LayerNormAttrs const &, ParallelTensorShape const &); -tl::expected, +tl::expected, std::string> get_weight_shapes(LayerNormAttrs const &attrs, ParallelTensorShape const &input_shape); @@ -44,7 +44,7 @@ tl::expected, * see * https://github.com/pytorch/pytorch/blob/1eba9b3aa3c43f86f4a2c807ac8e12c4a7767340/torch/nn/modules/normalization.py#L210-L214 */ -std::unordered_map +std::map get_initializers(LayerNormAttrs const &attrs); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/ops/linear.h b/lib/op-attrs/include/op-attrs/ops/linear.h index b5010c7186..817652abc6 100644 --- a/lib/op-attrs/include/op-attrs/ops/linear.h +++ b/lib/op-attrs/include/op-attrs/ops/linear.h @@ -17,7 +17,7 @@ namespace FlexFlow { -std::unordered_map +std::map get_linear_incoming_tensor_roles(LinearAttrs const &); tl::expected @@ -27,7 +27,7 @@ tl::expected get_bias_shape(LinearAttrs const &attrs, tl::expected get_output_shape(LinearAttrs const &attrs, TensorShape const &input); -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(LinearAttrs const &attrs, TensorShape const &input_shape); ParallelTensorDimDegrees @@ -49,12 +49,12 @@ tl::expected get_output_shape(LinearAttrs const &attrs, ParallelTensorShape const &input); -tl::expected, +tl::expected, std::string> get_weight_shapes(LinearAttrs const &attrs, ParallelTensorShape const &input_shape); -tl::expected, std::string> +tl::expected, std::string> get_initializers(LinearAttrs const &, TensorShape const &input_shape, std::optional const diff --git a/lib/op-attrs/include/op-attrs/parallel_tensor_dim_degrees.h b/lib/op-attrs/include/op-attrs/parallel_tensor_dim_degrees.h index 5582bf6e07..f15c8a7552 100644 --- a/lib/op-attrs/include/op-attrs/parallel_tensor_dim_degrees.h +++ b/lib/op-attrs/include/op-attrs/parallel_tensor_dim_degrees.h @@ -16,7 +16,7 @@ num_ptensor_shard_dims_t num_tensor_dims_t get_ptensor_dim_degrees_num_tensor_dims(ParallelTensorDimDegrees const &); -std::unordered_set +std::set get_parallel_tensor_dim_indices(ParallelTensorDimDegrees const &); std::set get_nontrivial_parallel_tensor_dim_indices( @@ -26,10 +26,10 @@ positive_int get_degree_for_parallel_tensor_dim_idx(ParallelTensorDimDegrees const &, parallel_tensor_dim_idx_t const &); -std::unordered_map +std::map get_parallel_tensor_degree_map(ParallelTensorDimDegrees const &); -std::unordered_set +std::set get_parallel_tensor_space_coordinates(ParallelTensorDimDegrees const &); DimDomain diff --git a/lib/op-attrs/include/op-attrs/parallel_tensor_dims.dtg.toml b/lib/op-attrs/include/op-attrs/parallel_tensor_dims.dtg.toml index 33e2e29db1..1050023249 100644 --- a/lib/op-attrs/include/op-attrs/parallel_tensor_dims.dtg.toml +++ b/lib/op-attrs/include/op-attrs/parallel_tensor_dims.dtg.toml @@ -14,8 +14,8 @@ includes = [ "op-attrs/ff_ordered/ff_ordered.h", "op-attrs/shard_parallel_dim.dtg.h", "op-attrs/replica_parallel_dim_set.dtg.h", - "", - "utils/fmt/unordered_map.h", + "", + "utils/fmt/map.h", "utils/fmt/pair.h", ] diff --git a/lib/op-attrs/include/op-attrs/parallel_tensor_dims.h b/lib/op-attrs/include/op-attrs/parallel_tensor_dims.h index 9e71785013..0283d5bc7d 100644 --- a/lib/op-attrs/include/op-attrs/parallel_tensor_dims.h +++ b/lib/op-attrs/include/op-attrs/parallel_tensor_dims.h @@ -11,7 +11,7 @@ namespace FlexFlow { FFOrdered ff_ordered_shard_dims(ParallelTensorDims const &); FFOrdered ff_ordered_shard_degrees(ParallelTensorDims const &); -std::unordered_set replica_dims(ParallelTensorDims const &); +std::set replica_dims(ParallelTensorDims const &); /* size_t get_volume(ParallelTensorDims const &); */ num_ptensor_shard_dims_t num_shard_dims(ParallelTensorDims const &); diff --git a/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h b/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h index e23ae33cbf..e48798a860 100644 --- a/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h +++ b/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h @@ -38,7 +38,7 @@ ParallelTensorShape TensorShape get_piece_shape(ParallelTensorShape const &); num_bytes_t get_piece_size_in_bytes(ParallelTensorShape const &); -std::unordered_set +std::set replica_dims(ParallelTensorShape const &); positive_int get_num_replica_dims(ParallelTensorShape const &); @@ -60,7 +60,7 @@ TensorShape get_reduced_shape(ParallelTensorShape const &); ParallelDim get_parallel_dim_at_idx(ParallelTensorShape const &shape, parallel_tensor_dim_idx_t idx); -std::unordered_set +std::set get_parallel_tensor_dim_indices(ParallelTensorShape const &shape); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/parallel_tensor_space_coordinate.h b/lib/op-attrs/include/op-attrs/parallel_tensor_space_coordinate.h index 3fd684c5ef..8f2289b18f 100644 --- a/lib/op-attrs/include/op-attrs/parallel_tensor_space_coordinate.h +++ b/lib/op-attrs/include/op-attrs/parallel_tensor_space_coordinate.h @@ -14,14 +14,14 @@ num_ptensor_parallel_dims_t num_ptensor_shard_dims_t ptensor_coord_num_shard_dims(ParallelTensorSpaceCoordinate const &); -std::unordered_set +std::set get_dim_idxs_in_ptensor_space_coord(ParallelTensorSpaceCoordinate const &); nonnegative_int ptensor_coord_component_for_ptensor_dim_idx( ParallelTensorSpaceCoordinate const &, parallel_tensor_dim_idx_t); ParallelTensorSpaceCoordinate parallel_tensor_space_coord_from_map( - std::unordered_map const &); + std::map const &); ParallelTensorSpaceCoordinate parallel_tensor_space_coord_from_dim_coord( DimCoord const &); diff --git a/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h b/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h index 28c48620a9..02ec2d4611 100644 --- a/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h +++ b/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h @@ -10,7 +10,7 @@ namespace FlexFlow { ReplicaParallelDimSet empty_replica_parallel_dim_set(); positive_int get_degree_of_replica_type(ReplicaParallelDimSet const &, ReplicaType); -std::unordered_set +std::set get_replica_dims(ReplicaParallelDimSet const &); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/shape_inference.h b/lib/op-attrs/include/op-attrs/shape_inference.h index 14184fac92..37c0c8536a 100644 --- a/lib/op-attrs/include/op-attrs/shape_inference.h +++ b/lib/op-attrs/include/op-attrs/shape_inference.h @@ -9,22 +9,22 @@ namespace FlexFlow { -std::unordered_map get_output_shapes( +std::map get_output_shapes( ComputationGraphOpAttrs const &, - std::unordered_map const &input_shapes); + std::map const &input_shapes); -std::unordered_map get_weight_shapes( +std::map get_weight_shapes( ComputationGraphOpAttrs const &, - std::unordered_map const &input_shapes); + std::map const &input_shapes); -std::unordered_map get_output_shapes( +std::map get_output_shapes( PCGOperatorAttrs const &, - std::unordered_map const + std::map const &input_shapes); -std::unordered_map get_weight_shapes( +std::map get_weight_shapes( PCGOperatorAttrs const &, - std::unordered_map const + std::map const &input_shapes); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/tensor_dims.h b/lib/op-attrs/include/op-attrs/tensor_dims.h index e0c8aa2dc6..9e9d6faf5d 100644 --- a/lib/op-attrs/include/op-attrs/tensor_dims.h +++ b/lib/op-attrs/include/op-attrs/tensor_dims.h @@ -37,13 +37,13 @@ TensorDimsCoord get_broadcast_src_coord(TensorDims const &input_dims, TensorDims const &output_dims, TensorDimsCoord const &dst_coord); -std::unordered_set +std::set get_tensor_dims_coord_set(TensorDims const &tensor_dims); -std::unordered_set get_ff_dim_t_set(TensorDims const &); +std::set get_ff_dim_t_set(TensorDims const &); std::optional - get_broadcast_target_dims(std::unordered_set const &); + get_broadcast_target_dims(std::set const &); TensorDims tensor_dims_drop_dims(TensorDims const &dims, diff --git a/lib/op-attrs/src/op-attrs/datatype.cc b/lib/op-attrs/src/op-attrs/datatype.cc index d9e4a65f13..c5e383c0fa 100644 --- a/lib/op-attrs/src/op-attrs/datatype.cc +++ b/lib/op-attrs/src/op-attrs/datatype.cc @@ -25,7 +25,7 @@ positive_int size_of_datatype(DataType data_type) { } bool can_strictly_promote_datatype_from_to(DataType src, DataType dst) { - std::unordered_set allowed; + std::set allowed; switch (src) { case DataType::BOOL: allowed = {DataType::INT32, @@ -55,7 +55,7 @@ bool can_strictly_promote_datatype_from_to(DataType src, DataType dst) { } bool can_torch_strictly_promote_datatype_from_to(DataType src, DataType dst) { - std::unordered_set allowed; + std::set allowed; switch (src) { case DataType::BOOL: allowed = {DataType::INT32, diff --git a/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc b/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc index e39fedb858..c9f851369e 100644 --- a/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc +++ b/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc @@ -5,9 +5,8 @@ namespace FlexFlow { using T = value_type<0>; -template FFOrdered ff_ordered_from_map(std::map const &); +template FFOrdered ff_ordered_from_map(std::unordered_map const &); -template FFOrdered - ff_ordered_from_map(std::unordered_map const &); +template FFOrdered ff_ordered_from_map(std::map const &); } // namespace FlexFlow diff --git a/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_of.cc b/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_of.cc index 0e0e8711d6..8a7a79e7aa 100644 --- a/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_of.cc +++ b/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_of.cc @@ -7,6 +7,6 @@ using T = value_type<0>; template FFOrdered ff_ordered_of(std::vector const &); -template FFOrdered ff_ordered_of(std::unordered_set const &); +template FFOrdered ff_ordered_of(std::set const &); } // namespace FlexFlow diff --git a/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc b/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc index 8c4e6c2e37..f698dce0c2 100644 --- a/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc +++ b/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using T = value_type<0>; -template std::unordered_map +template std::map map_from_ff_ordered(FFOrdered const &); } // namespace FlexFlow diff --git a/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc b/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc index 3f800fcdc6..1df85a7134 100644 --- a/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc +++ b/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc @@ -10,37 +10,37 @@ namespace FlexFlow { -std::unordered_map +std::map get_incoming_tensor_roles( ComputationGraphOpAttrs const &comp_graph_op_attrs) { return get_incoming_tensor_roles( pcg_op_attrs_from_compgraph_op_attrs(comp_graph_op_attrs)); } -std::unordered_map +std::map get_incoming_tensor_roles(PCGOperatorAttrs const &pcg_op_attrs) { return pcg_op_attrs - .visit>(overload{ + .visit>(overload{ [](BatchNormAttrs const &attrs) { return get_batch_norm_incoming_tensor_roles(attrs); }, [](BroadcastAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](CastAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](CombineAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [&](ConcatAttrs const &) { - return generate_unordered_map(get_variadic_inputs_slot_name_sequence(), + return generate_map(get_variadic_inputs_slot_name_sequence(), [](TensorSlotName) -> IncomingTensorRole { return IncomingTensorRole::INPUT; }); @@ -49,39 +49,39 @@ std::unordered_map return get_conv2d_incoming_tensor_roles(attrs); }, [](DropoutAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](ElementBinaryAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::LHS_INPUT, IncomingTensorRole::INPUT}, {TensorSlotName::RHS_INPUT, IncomingTensorRole::INPUT}, }; }, [](ElementUnaryAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](EmbeddingAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, {TensorSlotName::WEIGHT, IncomingTensorRole::WEIGHT}, }; }, [](FlatAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](GatherAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](InputAttrs const &) { - return std::unordered_map{}; + return std::map{}; }, [](LayerNormAttrs const &attrs) { return get_layer_norm_incoming_tensor_roles(attrs); @@ -93,67 +93,67 @@ std::unordered_map return get_attention_incoming_tensor_roles(attrs); }, [](NoopAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](Pool2DAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](ReduceAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](ReductionAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](RepartitionAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](ReplicateAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](ReverseAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](ReshapeAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](SplitAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](SoftmaxAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](TopKAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](TransposeAttrs const &) { - return std::unordered_map{ + return std::map{ {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; }, [](WeightAttrs const &) { - return std::unordered_map{}; + return std::map{}; }, }); } diff --git a/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc b/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc index 1a97f8b38b..0e180dc820 100644 --- a/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc +++ b/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc @@ -12,20 +12,20 @@ #include "utils/containers/require_two_keys.h" #include "utils/containers/zip_values_strict.h" #include "utils/overload.h" -#include "utils/containers/merge_disjoint_unordered_maps.h" +#include "utils/containers/merge_disjoint_maps.h" namespace FlexFlow { -std::unordered_map +std::map get_operator_to_incoming_mappings( ComputationGraphOpAttrs const &comp_graph_op_attrs, - std::unordered_map const + std::map const &inputs_degrees) { return comp_graph_op_attrs.visit< - std::unordered_map>(overload{ [&](ElementBinaryAttrs const &attrs) - -> std::unordered_map std::map { ASSERT(inputs_degrees.size() == 2); @@ -48,7 +48,7 @@ std::unordered_map }; }, [&](ElementUnaryAttrs const &attrs) - -> std::unordered_map std::map { ParallelTensorDimDegrees input_degrees = require_only_key(inputs_degrees, TensorSlotName::INPUT); @@ -63,16 +63,16 @@ std::unordered_map [&](InputAttrs const &) { ASSERT(inputs_degrees.size() == 0); - return std::unordered_map{}; }, [&](LinearAttrs const &attrs) - -> std::unordered_map std::map { ParallelTensorDimDegrees input_degrees = require_only_key(inputs_degrees, TensorSlotName::INPUT); - std::unordered_map result = { {TensorSlotName::INPUT, @@ -89,7 +89,7 @@ std::unordered_map return result; }, [&](TransposeAttrs const &attrs) - -> std::unordered_map std::map { ParallelTensorDimDegrees input_degrees = require_only_key(inputs_degrees, TensorSlotName::INPUT); @@ -104,29 +104,29 @@ std::unordered_map [&](WeightAttrs const &) { ASSERT(inputs_degrees.size() == 0); - return std::unordered_map{}; }, [](auto const &attrs) - -> std::unordered_map std::map { PANIC("Missing implmentation of get_operator_to_input_mappings", attrs); }, }); } -std::unordered_map +std::map get_operator_to_incoming_mappings_for_role( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees, IncomingTensorRole incoming_tensor_role) { - std::unordered_map + std::map incoming_mappings = get_operator_to_incoming_mappings(attrs, inputs_degrees); - std::unordered_map incoming_tensor_roles = + std::map incoming_tensor_roles = get_incoming_tensor_roles(attrs); return filtermap_values( @@ -144,36 +144,36 @@ std::unordered_map }); } -std::unordered_map +std::map get_operator_to_input_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees) { return get_operator_to_incoming_mappings_for_role( attrs, inputs_degrees, IncomingTensorRole::INPUT); } -std::unordered_map +std::map get_operator_to_weight_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees) { return get_operator_to_incoming_mappings_for_role( attrs, inputs_degrees, IncomingTensorRole::WEIGHT); } -std::unordered_map +std::map get_operator_to_output_mappings( ComputationGraphOpAttrs const &comp_graph_op_attrs, - std::unordered_map const + std::map const &inputs_degrees) { return comp_graph_op_attrs.visit< - std::unordered_map>(overload{ [&](ElementBinaryAttrs const &attrs) - -> std::unordered_map std::map { auto [lhs_degrees, rhs_degrees] = require_two_keys(inputs_degrees, @@ -188,7 +188,7 @@ std::unordered_map }; }, [&](ElementUnaryAttrs const &attrs) - -> std::unordered_map std::map { ParallelTensorDimDegrees input_degrees = require_only_key(inputs_degrees, TensorSlotName::INPUT); @@ -201,7 +201,7 @@ std::unordered_map }; }, [&](LinearAttrs const &attrs) - -> std::unordered_map std::map { ParallelTensorDimDegrees input_degrees = require_only_key(inputs_degrees, TensorSlotName::INPUT); @@ -214,7 +214,7 @@ std::unordered_map }; }, [&](InputAttrs const &attrs) - -> std::unordered_map std::map { ASSERT(inputs_degrees.size() == 0); @@ -226,7 +226,7 @@ std::unordered_map }; }, [&](TransposeAttrs const &attrs) - -> std::unordered_map std::map { ParallelTensorDimDegrees input_degrees = require_only_key(inputs_degrees, TensorSlotName::INPUT); @@ -239,7 +239,7 @@ std::unordered_map }; }, [&](WeightAttrs const &attrs) - -> std::unordered_map std::map { ASSERT(inputs_degrees.size() == 0); @@ -251,17 +251,17 @@ std::unordered_map }; }, [](auto const &attrs) - -> std::unordered_map std::map { PANIC("Missing implmentation of get_operator_to_input_mappings", attrs); }, }); } -std::unordered_map +std::map get_operator_to_ptensor_mappings_for_role( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees, TensorRole role) { switch (role) { @@ -276,12 +276,12 @@ std::unordered_map } } -std::unordered_map +std::map get_operator_to_ptensor_mappings( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees) { - return merge_disjoint_unordered_maps(std::vector{ + return merge_disjoint_maps(std::vector{ get_operator_to_input_mappings(attrs, inputs_degrees), get_operator_to_weight_mappings(attrs, inputs_degrees), get_operator_to_output_mappings(attrs, inputs_degrees), diff --git a/lib/op-attrs/src/op-attrs/get_operator_task_space.cc b/lib/op-attrs/src/op-attrs/get_operator_task_space.cc index 40ec0b964c..f6b0733328 100644 --- a/lib/op-attrs/src/op-attrs/get_operator_task_space.cc +++ b/lib/op-attrs/src/op-attrs/get_operator_task_space.cc @@ -15,7 +15,7 @@ namespace FlexFlow { OperatorTaskSpace get_operator_task_space( ComputationGraphOpAttrs const &attrs, - std::unordered_map const + std::map const &inputs_degrees) { return attrs.visit(overload{ [&](ElementUnaryAttrs const &attrs) { diff --git a/lib/op-attrs/src/op-attrs/operator_space_to_parallel_tensor_space_mapping.cc b/lib/op-attrs/src/op-attrs/operator_space_to_parallel_tensor_space_mapping.cc index c88043f6ce..88c0c2e07b 100644 --- a/lib/op-attrs/src/op-attrs/operator_space_to_parallel_tensor_space_mapping.cc +++ b/lib/op-attrs/src/op-attrs/operator_space_to_parallel_tensor_space_mapping.cc @@ -108,8 +108,8 @@ ParallelTensorSpaceCoordinate ptensor_coord_for_task_space_coord( TaskSpaceCoordinate const &task_space_coordinate, num_ptensor_shard_dims_t num_dims) { - std::unordered_set ptensor_dim_idxs = - unordered_set_of(dim_idxs_for_num_shard_dims(num_dims)); + std::set ptensor_dim_idxs = + dim_idxs_for_num_shard_dims(num_dims); DimCoord mapped_dim_coord = mapping.raw_mapping.at_l( diff --git a/lib/op-attrs/src/op-attrs/operator_task_space.cc b/lib/op-attrs/src/op-attrs/operator_task_space.cc index 98f7525564..ef7c8dde09 100644 --- a/lib/op-attrs/src/op-attrs/operator_task_space.cc +++ b/lib/op-attrs/src/op-attrs/operator_task_space.cc @@ -11,9 +11,9 @@ #include "utils/containers/product.h" #include "utils/containers/range.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/vector_of.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/nonnegative_int/nonnegative_range.h" #include "utils/nonnegative_int/num_elements.h" #include "utils/orthotope/dim_domain.h" @@ -29,13 +29,13 @@ OperatorTaskSpace trivial_op_task_space() { return OperatorTaskSpace{MinimalOrthotope{{}}}; } -std::unordered_set +std::set operator_task_space_get_dim_idxs(OperatorTaskSpace const &op_task_space) { return get_minimal_domain_dims( minimal_dim_domain_from_operator_task_space(op_task_space)); } -std::unordered_set +std::set get_task_space_coordinates(OperatorTaskSpace const &task) { std::vector> coordinate_ranges = @@ -43,9 +43,9 @@ std::unordered_set return nonnegative_range(num_points.nonnegative_int_from_int_ge_two()); }); - std::unordered_set> raw_coordinates = - unordered_set_of(cartesian_product(coordinate_ranges)); - std::unordered_set task_space_coordinates = + std::set> raw_coordinates = + set_of(cartesian_product(coordinate_ranges)); + std::set task_space_coordinates = transform(raw_coordinates, [](std::vector const &point) { return TaskSpaceCoordinate{OrthotopeCoord{point}}; }); @@ -78,7 +78,7 @@ MinimalDimDomain return minimal_dim_domain_from_minimal_orthotope( minimal_orthotope, - unordered_set_of(operator_task_space_dim_idx_range( + set_of(operator_task_space_dim_idx_range( minimal_orthotope_get_num_dims(minimal_orthotope))), get_operator_task_space_dim_ordering()); } diff --git a/lib/op-attrs/src/op-attrs/ops/attention.cc b/lib/op-attrs/src/op-attrs/ops/attention.cc index 816c2787cd..aa86b9e673 100644 --- a/lib/op-attrs/src/op-attrs/ops/attention.cc +++ b/lib/op-attrs/src/op-attrs/ops/attention.cc @@ -102,12 +102,12 @@ static void check_attrs(MultiHeadAttentionAttrs const &attrs) { "functionality, please create an issue."); } -std::unordered_map +std::map get_attention_incoming_tensor_roles(MultiHeadAttentionAttrs const &attrs) { check_attrs(attrs); - std::unordered_map roles = { + std::map roles = { {TensorSlotName::QUERY, IncomingTensorRole::INPUT}, {TensorSlotName::KEY, IncomingTensorRole::INPUT}, {TensorSlotName::VALUE, IncomingTensorRole::INPUT}, @@ -233,13 +233,13 @@ tl::expected }; } -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(MultiHeadAttentionAttrs const &attrs, TensorShape const &input_q, TensorShape const &input_k, TensorShape const &input_v) { - std::unordered_map weight_shapes = { + std::map weight_shapes = { { TensorSlotName::WEIGHT, PROPAGATE_ERR(get_weights_shape(attrs, input_q, input_k, input_v)), @@ -416,14 +416,14 @@ positive_int get_oSize(TensorShape const &) { NOT_IMPLEMENTED(); } -tl::expected, +tl::expected, std::string> get_weight_shapes(MultiHeadAttentionAttrs const &attrs, ParallelTensorShape const &input_q, ParallelTensorShape const &input_k, ParallelTensorShape const &input_v) { - std::unordered_map weight_shapes = { + std::map weight_shapes = { { TensorSlotName::WEIGHT, PROPAGATE_ERR(get_weights_shape(attrs, input_q, input_k, input_v)), @@ -445,7 +445,7 @@ tl::expected, return weight_shapes; } -tl::expected, std::string> +tl::expected, std::string> get_initializers( MultiHeadAttentionAttrs const &attrs, TensorShape const &input_q, @@ -492,13 +492,13 @@ tl::expected, std::string> maybe_output_bias_initializer.value_or(default_output_bias_initializer); if (attrs.bias) { - return std::unordered_map{ + return std::map{ {TensorSlotName::WEIGHT, weights_initializer}, {TensorSlotName::INPUT_BIAS, input_bias_initializer}, {TensorSlotName::OUTPUT_BIAS, output_bias_initializer}, }; } else { - return std::unordered_map{ + return std::map{ {TensorSlotName::WEIGHT, weights_initializer}, }; } diff --git a/lib/op-attrs/src/op-attrs/ops/batch_norm.cc b/lib/op-attrs/src/op-attrs/ops/batch_norm.cc index 5d451c617d..d045b829d4 100644 --- a/lib/op-attrs/src/op-attrs/ops/batch_norm.cc +++ b/lib/op-attrs/src/op-attrs/ops/batch_norm.cc @@ -10,9 +10,9 @@ namespace FlexFlow { -std::unordered_map +std::map get_batch_norm_incoming_tensor_roles(BatchNormAttrs const &attrs) { - std::unordered_map result = { + std::map result = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, @@ -96,7 +96,7 @@ tl::expected return get_gamma_weights_shape(attrs, input_shape); } -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(BatchNormAttrs const &attrs, TensorShape const &input_shape) { @@ -105,7 +105,7 @@ tl::expected, std::string> TensorShape beta_shape = PROPAGATE_ERR(get_beta_weights_shape(attrs, input_shape)); - return std::unordered_map{ + return std::map{ { TensorSlotName::GAMMA, gamma_shape, @@ -210,7 +210,7 @@ tl::expected return get_gamma_weights_parallel_dim_degrees(attrs, input_degrees); } -tl::expected, +tl::expected, std::string> get_weight_parallel_dim_degrees( BatchNormAttrs const &attrs, @@ -221,7 +221,7 @@ tl::expected, ParallelTensorDimDegrees beta_degrees = PROPAGATE_ERR( get_beta_weights_parallel_dim_degrees(attrs, input_degrees)); - return std::unordered_map{ + return std::map{ { TensorSlotName::GAMMA, gamma_degrees, @@ -310,7 +310,7 @@ tl::expected return lift_to_parallel_with_degrees(unpar, degrees); } -tl::expected, +tl::expected, std::string> get_weight_shapes(BatchNormAttrs const &attrs, ParallelTensorShape const &input_shape) { @@ -320,7 +320,7 @@ tl::expected, ParallelTensorShape beta_shape = PROPAGATE_ERR(get_beta_weights_shape(attrs, input_shape)); - return std::unordered_map{ + return std::map{ { TensorSlotName::GAMMA, gamma_shape, @@ -332,7 +332,7 @@ tl::expected, }; } -tl::expected, std::string> +tl::expected, std::string> get_initializers(BatchNormAttrs const &attrs) { if (attrs.affine) { InitializerAttrs gamma_initializer = @@ -341,7 +341,7 @@ tl::expected, std::string> InitializerAttrs beta_initializer = InitializerAttrs{ConstantInitializerAttrs{DataTypeValue{float{0}}}}; - return std::unordered_map{ + return std::map{ { TensorSlotName::GAMMA, gamma_initializer, @@ -352,7 +352,7 @@ tl::expected, std::string> }, }; } else { - return std::unordered_map{}; + return std::map{}; } } diff --git a/lib/op-attrs/src/op-attrs/ops/conv_2d.cc b/lib/op-attrs/src/op-attrs/ops/conv_2d.cc index b50446a693..4de45d72a1 100644 --- a/lib/op-attrs/src/op-attrs/ops/conv_2d.cc +++ b/lib/op-attrs/src/op-attrs/ops/conv_2d.cc @@ -8,9 +8,9 @@ namespace FlexFlow { -std::unordered_map +std::map get_conv2d_incoming_tensor_roles(Conv2DAttrs const &attrs) { - std::unordered_map result = { + std::map result = { {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, {TensorSlotName::FILTER, IncomingTensorRole::WEIGHT}, }; @@ -89,10 +89,10 @@ TensorShape get_output_shape(Conv2DAttrs const &attrs, input.datatype}; } -std::unordered_map +std::map get_weight_shapes(Conv2DAttrs const &attrs, TensorShape const &input_shape) { - std::unordered_map weight_shapes = { + std::map weight_shapes = { { TensorSlotName::FILTER, get_kernel_shape(attrs, input_shape), @@ -180,10 +180,10 @@ ParallelTensorShape get_output_shape(Conv2DAttrs const &attrs, unpar, sum_degree, discard_copy_degree, shard_degrees); } -std::unordered_map +std::map get_weight_shapes(Conv2DAttrs const &attrs, ParallelTensorShape const &input_shape) { - std::unordered_map weight_shapes = { + std::map weight_shapes = { { TensorSlotName::FILTER, get_kernel_shape(attrs, input_shape), @@ -206,7 +206,7 @@ std::unordered_map * see * https://github.com/pytorch/pytorch/blob/1eba9b3aa3c43f86f4a2c807ac8e12c4a7767340/torch/nn/modules/conv.py#L178-L187 */ -std::unordered_map +std::map get_initializers(Conv2DAttrs const &attrs, TensorShape const &input_shape, std::optional maybe_kernel_initializer, diff --git a/lib/op-attrs/src/op-attrs/ops/embedding.cc b/lib/op-attrs/src/op-attrs/ops/embedding.cc index b400c6263a..d192a2d546 100644 --- a/lib/op-attrs/src/op-attrs/ops/embedding.cc +++ b/lib/op-attrs/src/op-attrs/ops/embedding.cc @@ -114,7 +114,7 @@ tl::expected unpar, sum_degree, discard_copy_degree, shard_degrees); } -std::unordered_map get_initializers( +std::map get_initializers( EmbeddingAttrs const &, std::optional const &maybe_initializer_attrs) { InitializerAttrs default_initializer_attrs = InitializerAttrs{ diff --git a/lib/op-attrs/src/op-attrs/ops/layer_norm.cc b/lib/op-attrs/src/op-attrs/ops/layer_norm.cc index 81aa2d8a52..c732ace44e 100644 --- a/lib/op-attrs/src/op-attrs/ops/layer_norm.cc +++ b/lib/op-attrs/src/op-attrs/ops/layer_norm.cc @@ -15,9 +15,9 @@ namespace FlexFlow { -std::unordered_map +std::map get_layer_norm_incoming_tensor_roles(LayerNormAttrs const &attrs) { - std::unordered_map result = { + std::map result = { {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, }; @@ -100,7 +100,7 @@ tl::expected return get_gamma_weights_shape(attrs, input_shape); } -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(LayerNormAttrs const &attrs, TensorShape const &input_shape) { @@ -109,7 +109,7 @@ tl::expected, std::string> TensorShape beta_shape = PROPAGATE_ERR(get_beta_weights_shape(attrs, input_shape)); - return std::unordered_map{ + return std::map{ { TensorSlotName::GAMMA, gamma_shape, @@ -220,7 +220,7 @@ tl::expected return get_gamma_weights_shape(attrs, input_shape); } -tl::expected, +tl::expected, std::string> get_weight_shapes(LayerNormAttrs const &attrs, ParallelTensorShape const &input_shape) { @@ -230,7 +230,7 @@ tl::expected, ParallelTensorShape beta_shape = PROPAGATE_ERR(get_beta_weights_shape(attrs, input_shape)); - return std::unordered_map{ + return std::map{ { TensorSlotName::GAMMA, gamma_shape, @@ -242,7 +242,7 @@ tl::expected, }; } -std::unordered_map +std::map get_initializers(LayerNormAttrs const &attrs) { if (attrs.elementwise_affine) { InitializerAttrs gamma_initializer = diff --git a/lib/op-attrs/src/op-attrs/ops/linear.cc b/lib/op-attrs/src/op-attrs/ops/linear.cc index 9099c7dac6..358929f1af 100644 --- a/lib/op-attrs/src/op-attrs/ops/linear.cc +++ b/lib/op-attrs/src/op-attrs/ops/linear.cc @@ -14,7 +14,7 @@ #include "op-attrs/tensor_dims.h" #include "op-attrs/tensor_shape.h" #include "utils/containers/product.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/expected.h" #include "utils/fmt/optional.h" #include "utils/integer_conversions.h" @@ -26,9 +26,9 @@ namespace FlexFlow { -std::unordered_map +std::map get_linear_incoming_tensor_roles(LinearAttrs const &attrs) { - std::unordered_map result = { + std::map result = { {TensorSlotName::INPUT, IncomingTensorRole::INPUT}, {TensorSlotName::WEIGHT, IncomingTensorRole::WEIGHT}, }; @@ -72,11 +72,11 @@ tl::expected return output_shape; } -tl::expected, std::string> +tl::expected, std::string> get_weight_shapes(LinearAttrs const &attrs, TensorShape const &input_shape) { - std::unordered_map weight_shapes = { + std::map weight_shapes = { { TensorSlotName::WEIGHT, PROPAGATE_ERR(get_projection_shape(attrs, input_shape)), @@ -205,12 +205,12 @@ ParallelTensorDimDegrees }; } -tl::expected, +tl::expected, std::string> get_weight_shapes(LinearAttrs const &attrs, ParallelTensorShape const &input_shape) { - std::unordered_map weight_shapes = { + std::map weight_shapes = { { TensorSlotName::WEIGHT, PROPAGATE_ERR(get_projection_shape(attrs, input_shape)), @@ -231,7 +231,7 @@ tl::expected, * see * https://github.com/pytorch/pytorch/blob/1eba9b3aa3c43f86f4a2c807ac8e12c4a7767340/torch/nn/modules/linear.py#L114-L122 */ -tl::expected, std::string> +tl::expected, std::string> get_initializers( LinearAttrs const &attrs, TensorShape const &input_shape, @@ -275,12 +275,12 @@ tl::expected, std::string> maybe_bias_initializer.value_or(bias_default_initializer); if (attrs.use_bias) { - return std::unordered_map{ + return std::map{ {TensorSlotName::WEIGHT, projection_initializer}, {TensorSlotName::BIAS, bias_initializer}, }; } else { - return std::unordered_map{ + return std::map{ {TensorSlotName::WEIGHT, projection_initializer}, }; } @@ -356,8 +356,8 @@ static ParallelTensorSpaceToParallelTensorSpaceMapping }; { - std::unordered_set dims_from = - unordered_set_of(dim_idxs_for_num_shard_dims(input_num_shard_dims)); + std::set dims_from = + set_of(dim_idxs_for_num_shard_dims(input_num_shard_dims)); dims_from.insert(sum_dim_idx()); dims_from.erase(input_channel_dim); dims_from.erase(discard_copy_dim_idx()); @@ -412,8 +412,8 @@ static ParallelTensorSpaceToParallelTensorSpaceMapping }; { - std::unordered_set dims_from = - unordered_set_of(dim_idxs_for_num_shard_dims(input_num_shard_dims)); + std::set dims_from = + set_of(dim_idxs_for_num_shard_dims(input_num_shard_dims)); dims_from.erase(input_channel_dim); dims_from.erase(sum_dim_idx()); diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc index 5b5d0b514f..a334fac056 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc @@ -7,18 +7,19 @@ #include "op-attrs/parallel_tensor_space_coordinate.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/filtrans.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/map_keys.h" #include "utils/containers/map_values.h" #include "utils/containers/range.h" #include "utils/containers/set_union.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/nonnegative_int/nonnegative_range.h" #include "utils/nonnegative_int/num_elements.h" #include "utils/orthotope/minimal_dim_domain.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { @@ -35,11 +36,11 @@ num_tensor_dims_t get_ptensor_dim_degrees_num_tensor_dims( get_ptensor_dim_degrees_num_shard_dims(degrees)); } -std::unordered_set +std::set get_parallel_tensor_dim_indices(ParallelTensorDimDegrees const °rees) { - std::unordered_set result = - unordered_set_of(dim_idxs_for_num_shard_dims( + std::set result = + set_of(dim_idxs_for_num_shard_dims( get_ptensor_dim_degrees_num_shard_dims(degrees))); result.insert(sum_dim_idx()); result.insert(discard_copy_dim_idx()); @@ -84,10 +85,10 @@ positive_int get_degree_for_parallel_tensor_dim_idx( } } -std::unordered_map +std::map get_parallel_tensor_degree_map(ParallelTensorDimDegrees const °rees) { - std::unordered_map + std::map replica_dim_degrees = { {parallel_tensor_dim_idx_t{ReplicaType::SUM}, degrees.sum_degree.value}, @@ -95,34 +96,34 @@ std::unordered_map degrees.discard_copy_degree.value}, }; - std::unordered_map shard_dim_degrees = - generate_unordered_map(get_idxs(degrees.shard_degrees), [&](ff_dim_t const &dim) { + std::map shard_dim_degrees = + generate_map(get_idxs(degrees.shard_degrees), [&](ff_dim_t const &dim) { return degrees.shard_degrees.at(dim); }); - return binary_merge_disjoint_unordered_maps( + return binary_merge_disjoint_maps( /*lhs=*/replica_dim_degrees, /*rhs=*/map_keys(shard_dim_degrees, [](ff_dim_t const &dim) { return parallel_tensor_dim_idx_t{dim}; })); } -std::unordered_set +std::set get_parallel_tensor_space_coordinates( ParallelTensorDimDegrees const °rees) { - std::unordered_map degree_map = + std::map degree_map = get_parallel_tensor_degree_map(degrees); - std::unordered_map> + std::map> possible_per_dim_coords = map_values(degree_map, [](positive_int degree) { - return unordered_set_of(nonnegative_range(degree)); + return set_of(nonnegative_range(degree)); }); return transform( get_all_assignments(possible_per_dim_coords), - [](std::unordered_map const + [](std::map const &m) { return parallel_tensor_space_coord_from_map(m); }); } @@ -131,7 +132,7 @@ DimDomain ParallelTensorDimDegrees const &dim_degrees) { return DimDomain{ - generate_unordered_map(get_parallel_tensor_dim_indices(dim_degrees), + generate_map(get_parallel_tensor_dim_indices(dim_degrees), [&](parallel_tensor_dim_idx_t idx) { return get_degree_for_parallel_tensor_dim_idx(dim_degrees, idx); @@ -142,7 +143,7 @@ DimDomain ParallelTensorDimDegrees parallel_tensor_dim_degrees_from_dim_domain( DimDomain const &dim_domain) { - std::unordered_map shard_dims = + std::map shard_dims = filtermap_keys(dim_domain.dims, [](parallel_tensor_dim_idx_t dim_idx) { return dim_idx.try_require_shard_dim(); }); diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc index 71419e4a57..73fd7fcbec 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc @@ -24,7 +24,7 @@ FFOrdered ff_ordered_shard_degrees(ParallelTensorDims const &d) { [](ShardParallelDim const &d) { return d.degree; }); } -std::unordered_set +std::set replica_dims(ParallelTensorDims const &d) { return get_replica_dims(d.replica_dims); } diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc index cc88692124..86788671a9 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc @@ -19,7 +19,7 @@ num_ptensor_shard_dims_t num_shard_dims(ParallelTensorShape const &s) { return num_shard_dims(s.dims); } -std::unordered_set +std::set replica_dims(ParallelTensorShape const &s) { return replica_dims(s.dims); } @@ -139,9 +139,9 @@ ParallelDim get_parallel_dim_at_idx(ParallelTensorShape const &shape, }}); } -std::unordered_set +std::set get_parallel_tensor_dim_indices(ParallelTensorShape const &shape) { - std::unordered_set indices; + std::set indices; extend(indices, transform(nonnegative_range(num_shard_dims(shape.dims).value), [](nonnegative_int idx) { diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc index 79d765b02c..101ac648c6 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc @@ -3,9 +3,8 @@ #include "op-attrs/parallel_tensor_dim_idx_t.h" #include "utils/containers/contains_key.h" #include "utils/containers/filtermap_keys.h" -#include "utils/containers/generate_unordered_map.h" -#include "utils/containers/unordered_set_of.h" #include "utils/nonnegative_int/num_elements.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { @@ -23,12 +22,12 @@ num_ptensor_shard_dims_t }; } -std::unordered_set +std::set get_dim_idxs_in_ptensor_space_coord( ParallelTensorSpaceCoordinate const &coord) { - std::unordered_set result = unordered_set_of( - dim_idxs_for_num_shard_dims(ptensor_coord_num_shard_dims(coord))); + std::set result = + dim_idxs_for_num_shard_dims(ptensor_coord_num_shard_dims(coord)); result.insert(sum_dim_idx()); result.insert(discard_copy_dim_idx()); return result; @@ -47,11 +46,11 @@ nonnegative_int ptensor_coord_component_for_ptensor_dim_idx( } ParallelTensorSpaceCoordinate parallel_tensor_space_coord_from_map( - std::unordered_map const &m) { + std::map const &m) { ASSERT(contains_key(m, sum_dim_idx())); ASSERT(contains_key(m, discard_copy_dim_idx())); - std::unordered_map shard_map = + std::map shard_map = filtermap_keys(m, [](parallel_tensor_dim_idx_t const &d) { return d.try_require_shard_dim(); }); @@ -73,7 +72,7 @@ DimCoord dim_coord_from_parallel_tensor_space_coord( ParallelTensorSpaceCoordinate const &coord) { return DimCoord{ - generate_unordered_map(get_dim_idxs_in_ptensor_space_coord(coord), + generate_map(get_dim_idxs_in_ptensor_space_coord(coord), [&](parallel_tensor_dim_idx_t idx) { return ptensor_coord_component_for_ptensor_dim_idx(coord, idx); diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc index 2a161838cd..e4b33dde7d 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc @@ -13,20 +13,20 @@ ParallelTensorSpaceToParallelTensorSpaceMapping // TODO(@lockshaw)(#pr): // { - // std::unordered_set + // std::set // l_dims = - // unordered_set_of(get_nontrivial_parallel_tensor_dim_indices(l_degrees)); - // std::unordered_set + // set_of(get_nontrivial_parallel_tensor_dim_indices(l_degrees)); + // std::set // projection_input_dims = input_dims_of_projection(projection); // // ASSERT(l_dims == projection_input_dims); // } // // { - // std::unordered_set + // std::set // r_dims = - // unordered_set_of(get_nontrivial_parallel_tensor_dim_indices(r_degrees)); - // std::unordered_set + // set_of(get_nontrivial_parallel_tensor_dim_indices(r_degrees)); + // std::set // projection_output_dims = output_dims_of_projection(projection); // // ASSERT(r_dims == projection_output_dims); diff --git a/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc b/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc index 871a39f91f..36b24d8eaf 100644 --- a/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc +++ b/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc @@ -20,9 +20,9 @@ positive_int get_degree_of_replica_type(ReplicaParallelDimSet const &s, } } -std::unordered_set +std::set get_replica_dims(ReplicaParallelDimSet const &s) { - return std::unordered_set{ + return std::set{ ReplicaParallelDim{s.sum_degree.value, ReplicaType::SUM}, ReplicaParallelDim{s.discard_copy_degree.value, ReplicaType::DISCARD_COPY}, diff --git a/lib/op-attrs/src/op-attrs/shape_inference.cc b/lib/op-attrs/src/op-attrs/shape_inference.cc index 7e70cb5500..0d1e0ee82f 100644 --- a/lib/op-attrs/src/op-attrs/shape_inference.cc +++ b/lib/op-attrs/src/op-attrs/shape_inference.cc @@ -32,7 +32,7 @@ namespace FlexFlow { template static std::tuple - require_3(std::unordered_map const &v, + require_3(std::map const &v, TensorSlotName k1, TensorSlotName k2, TensorSlotName k3) { @@ -43,7 +43,7 @@ static std::tuple template static std::vector - require_only_slots_sequence(std::unordered_map const &v, + require_only_slots_sequence(std::map const &v, std::vector const &slots) { nonnegative_int v_num_slots = num_elements(v); ASSERT(v_num_slots <= slots.size()); @@ -51,20 +51,20 @@ static std::vector std::vector expected_slots = slice(slots, 0, v_num_slots.unwrap_nonnegative()); - ASSERT(unordered_set_of(expected_slots) == unordered_keys(v)); + ASSERT(set_of(expected_slots) == keys(v)); return transform(expected_slots, [&](TensorSlotName const &slot_name) { return v.at(slot_name); }); }; -std::unordered_map get_output_shapes( +std::map get_output_shapes( ComputationGraphOpAttrs const &op_attrs, - std::unordered_map const &input_shapes) { - return op_attrs.visit>( + std::map const &input_shapes) { + return op_attrs.visit>( overload{ [&](BatchNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -76,7 +76,7 @@ std::unordered_map get_output_shapes( }; }, [&](CastAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -88,7 +88,7 @@ std::unordered_map get_output_shapes( }; }, [&](ConcatAttrs const &attrs) - -> std::unordered_map { + -> std::map { std::vector inputs = require_only_slots_sequence( input_shapes, get_variadic_inputs_slot_name_sequence()); @@ -100,7 +100,7 @@ std::unordered_map get_output_shapes( }; }, [&](Conv2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -112,7 +112,7 @@ std::unordered_map get_output_shapes( }; }, [&](DropoutAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -124,7 +124,7 @@ std::unordered_map get_output_shapes( }; }, [&](ElementBinaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [lhs, rhs] = require_two_keys(input_shapes, TensorSlotName::LHS_INPUT, TensorSlotName::RHS_INPUT); @@ -137,7 +137,7 @@ std::unordered_map get_output_shapes( }; }, [&](ElementUnaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -149,7 +149,7 @@ std::unordered_map get_output_shapes( }; }, [&](EmbeddingAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -161,7 +161,7 @@ std::unordered_map get_output_shapes( }; }, [&](FlatAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -173,7 +173,7 @@ std::unordered_map get_output_shapes( }; }, [&](GatherAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [input, index] = require_two_keys( input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); @@ -185,7 +185,7 @@ std::unordered_map get_output_shapes( }; }, [&](InputAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return { @@ -196,7 +196,7 @@ std::unordered_map get_output_shapes( }; }, [&](LayerNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -208,7 +208,7 @@ std::unordered_map get_output_shapes( }; }, [&](LinearAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -220,7 +220,7 @@ std::unordered_map get_output_shapes( }; }, [&](MultiHeadAttentionAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [query, key, value] = require_3(input_shapes, TensorSlotName::QUERY, TensorSlotName::KEY, @@ -233,7 +233,7 @@ std::unordered_map get_output_shapes( }; }, [&](Pool2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -243,7 +243,7 @@ std::unordered_map get_output_shapes( }; }, [&](SoftmaxAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -253,7 +253,7 @@ std::unordered_map get_output_shapes( }; }, [&](TransposeAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -265,7 +265,7 @@ std::unordered_map get_output_shapes( }; }, [&](WeightAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return { @@ -276,62 +276,62 @@ std::unordered_map get_output_shapes( }; }, [&](auto const &attrs) - -> std::unordered_map { + -> std::map { NOT_IMPLEMENTED(); }, }); } -std::unordered_map get_weight_shapes( +std::map get_weight_shapes( ComputationGraphOpAttrs const &op_attrs, - std::unordered_map const &input_shapes) { - return op_attrs.visit>( + std::map const &input_shapes) { + return op_attrs.visit>( overload{ [&](BatchNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return throw_if_unexpected(get_weight_shapes(attrs, input)); }, [&](CastAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](ConcatAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_slots_sequence( input_shapes, get_variadic_inputs_slot_name_sequence()); return {}; }, [&](Conv2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return get_weight_shapes(attrs, input); }, [&](DropoutAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](ElementBinaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_two_keys(input_shapes, TensorSlotName::LHS_INPUT, TensorSlotName::RHS_INPUT); return {}; }, [&](ElementUnaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](EmbeddingAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -345,37 +345,37 @@ std::unordered_map get_weight_shapes( }; }, [&](FlatAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](GatherAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_two_keys( input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); return {}; }, [&](InputAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return {}; }, [&](LayerNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return throw_if_unexpected(get_weight_shapes(attrs, input)); }, [&](LinearAttrs const &attrs) - -> std::unordered_map { + -> std::map { TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return throw_if_unexpected(get_weight_shapes(attrs, input)); }, [&](MultiHeadAttentionAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [query, key, value] = require_3(input_shapes, TensorSlotName::QUERY, TensorSlotName::KEY, @@ -385,37 +385,37 @@ std::unordered_map get_weight_shapes( get_weight_shapes(attrs, query, key, value)); }, [&](Pool2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](SoftmaxAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](WeightAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return {}; }, [&](auto const &attrs) - -> std::unordered_map { + -> std::map { NOT_IMPLEMENTED(); }, }); } -std::unordered_map get_output_shapes( +std::map get_output_shapes( PCGOperatorAttrs const &pcg_op_attrs, - std::unordered_map const + std::map const &input_shapes) { return pcg_op_attrs - .visit>(overload{ + .visit>(overload{ [&](BatchNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -427,7 +427,7 @@ std::unordered_map get_output_shapes( }; }, [&](CastAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -437,7 +437,7 @@ std::unordered_map get_output_shapes( }; }, [&](CombineAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -447,7 +447,7 @@ std::unordered_map get_output_shapes( }; }, [&](ConcatAttrs const &attrs) - -> std::unordered_map { + -> std::map { std::vector inputs = require_only_slots_sequence( input_shapes, get_variadic_inputs_slot_name_sequence()); @@ -458,7 +458,7 @@ std::unordered_map get_output_shapes( }; }, [&](Conv2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -467,7 +467,7 @@ std::unordered_map get_output_shapes( }; }, [&](DropoutAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -479,7 +479,7 @@ std::unordered_map get_output_shapes( }; }, [&](ElementBinaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [lhs, rhs] = require_two_keys(input_shapes, TensorSlotName::LHS_INPUT, TensorSlotName::RHS_INPUT); @@ -492,7 +492,7 @@ std::unordered_map get_output_shapes( }; }, [&](ElementUnaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -504,7 +504,7 @@ std::unordered_map get_output_shapes( }; }, [&](EmbeddingAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -516,7 +516,7 @@ std::unordered_map get_output_shapes( }; }, [&](FlatAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -528,7 +528,7 @@ std::unordered_map get_output_shapes( }; }, [&](GatherAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [input, index] = require_two_keys( input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); @@ -540,7 +540,7 @@ std::unordered_map get_output_shapes( }; }, [&](InputAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return { @@ -551,7 +551,7 @@ std::unordered_map get_output_shapes( }; }, [&](LayerNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -563,7 +563,7 @@ std::unordered_map get_output_shapes( }; }, [&](LinearAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -575,7 +575,7 @@ std::unordered_map get_output_shapes( }; }, [&](MultiHeadAttentionAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [i1, i2, i3] = require_3(input_shapes, TensorSlotName::QUERY, TensorSlotName::KEY, @@ -587,7 +587,7 @@ std::unordered_map get_output_shapes( }; }, [&](Pool2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -599,7 +599,7 @@ std::unordered_map get_output_shapes( }; }, [&](ReductionAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -611,7 +611,7 @@ std::unordered_map get_output_shapes( }; }, [&](RepartitionAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -623,7 +623,7 @@ std::unordered_map get_output_shapes( }; }, [&](ReplicateAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -635,7 +635,7 @@ std::unordered_map get_output_shapes( }; }, [&](SoftmaxAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -647,7 +647,7 @@ std::unordered_map get_output_shapes( }; }, [&](TransposeAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -659,7 +659,7 @@ std::unordered_map get_output_shapes( }; }, [&](WeightAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return { @@ -670,58 +670,58 @@ std::unordered_map get_output_shapes( }; }, [&](auto const &attrs) - -> std::unordered_map { + -> std::map { NOT_IMPLEMENTED(); }, }); } -std::unordered_map get_weight_shapes( +std::map get_weight_shapes( PCGOperatorAttrs const &pcg_op_attrs, - std::unordered_map const + std::map const &input_shapes) { return pcg_op_attrs - .visit>(overload{ + .visit>(overload{ [&](BatchNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return throw_if_unexpected(get_weight_shapes(attrs, input)); }, [&](CastAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](CombineAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](ConcatAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](Conv2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return get_weight_shapes(attrs, input); }, [&](DropoutAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](ElementBinaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_two_keys(input_shapes, TensorSlotName::LHS_INPUT, TensorSlotName::RHS_INPUT); @@ -729,13 +729,13 @@ std::unordered_map get_weight_shapes( return {}; }, [&](ElementUnaryAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](EmbeddingAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); @@ -747,40 +747,40 @@ std::unordered_map get_weight_shapes( }; }, [&](FlatAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](GatherAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_two_keys( input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); return {}; }, [&](InputAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return {}; }, [&](LayerNormAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return throw_if_unexpected(get_weight_shapes(attrs, input)); }, [&](LinearAttrs const &attrs) - -> std::unordered_map { + -> std::map { ParallelTensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); return throw_if_unexpected(get_weight_shapes(attrs, input)); }, [&](MultiHeadAttentionAttrs const &attrs) - -> std::unordered_map { + -> std::map { auto [query, key, value] = require_3(input_shapes, TensorSlotName::QUERY, TensorSlotName::KEY, @@ -790,49 +790,49 @@ std::unordered_map get_weight_shapes( get_weight_shapes(attrs, query, key, value)); }, [&](Pool2DAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](RepartitionAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](ReplicateAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](ReductionAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](SoftmaxAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](TransposeAttrs const &attrs) - -> std::unordered_map { + -> std::map { require_only_key(input_shapes, TensorSlotName::INPUT); return {}; }, [&](WeightAttrs const &attrs) - -> std::unordered_map { + -> std::map { ASSERT(input_shapes.size() == 0); return {}; }, [&](auto const &attrs) - -> std::unordered_map { + -> std::map { NOT_IMPLEMENTED(); }, }); diff --git a/lib/op-attrs/src/op-attrs/task_space_coordinate.cc b/lib/op-attrs/src/op-attrs/task_space_coordinate.cc index 302825f27e..b94391b793 100644 --- a/lib/op-attrs/src/op-attrs/task_space_coordinate.cc +++ b/lib/op-attrs/src/op-attrs/task_space_coordinate.cc @@ -3,7 +3,6 @@ #include "op-attrs/operator_task_space_dim_idx_t.h" #include "utils/containers/map_keys.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" #include "utils/containers/vector_from_idx_map.h" #include "utils/nonnegative_int/nonnegative_range.h" #include "utils/nonnegative_int/num_elements.h" @@ -23,15 +22,15 @@ TaskSpaceCoordinate TaskSpaceCoordinate task_space_coordinate_from_dim_coord( DimCoord const &dim_coord) { - std::unordered_set coord_dims = + std::set coord_dims = get_coord_dims(dim_coord); std::set dims = operator_task_space_dim_idx_range(num_elements(coord_dims)); - ASSERT(coord_dims == unordered_set_of(dims)); + ASSERT(coord_dims == dims); - std::unordered_map idx_map = + std::map idx_map = map_keys(dim_coord.raw, [](operator_task_space_dim_idx_t idx) { return idx.raw_idx; }); @@ -47,8 +46,8 @@ DimCoord return dim_coord_from_orthotope_coord( coord.orthotope_coord, - unordered_set_of(operator_task_space_dim_idx_range( - orthotope_coord_num_dims(coord.orthotope_coord))), + operator_task_space_dim_idx_range( + orthotope_coord_num_dims(coord.orthotope_coord)), get_operator_task_space_dim_ordering()); } diff --git a/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc b/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc index 1f6fa4b5d4..bd75207834 100644 --- a/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc +++ b/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc @@ -11,13 +11,13 @@ #include "utils/containers/minimum.h" #include "utils/containers/permute_with_key.h" #include "utils/containers/require_same.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/hash/tuple.h" namespace FlexFlow { static void - check_are_contiguous_from_one(std::unordered_set const &idxs) { + check_are_contiguous_from_one(std::set const &idxs) { if (idxs.empty()) { return; } diff --git a/lib/op-attrs/src/op-attrs/tensor_dims.cc b/lib/op-attrs/src/op-attrs/tensor_dims.cc index c69418c90c..1693cc92d2 100644 --- a/lib/op-attrs/src/op-attrs/tensor_dims.cc +++ b/lib/op-attrs/src/op-attrs/tensor_dims.cc @@ -14,7 +14,7 @@ #include "utils/containers/product.h" #include "utils/containers/reversed.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/vector_of.h" #include "utils/containers/zip.h" #include "utils/integer_conversions.h" @@ -147,7 +147,7 @@ TensorDimsCoord get_broadcast_src_coord(TensorDims const &input_dims, return result; } -std::unordered_set +std::set get_tensor_dims_coord_set(TensorDims const &tensor_dims) { std::vector> per_dim_ranges = transform( vector_of(tensor_dims.ff_ordered), @@ -155,8 +155,8 @@ std::unordered_set return nonnegative_range(dim_size.nonnegative_int_from_positive_int()); }); - std::unordered_set> raw_points = - unordered_set_of(cartesian_product(per_dim_ranges)); + std::set> raw_points = + set_of(cartesian_product(per_dim_ranges)); return transform(raw_points, [](std::vector const &raw_point) { @@ -164,12 +164,12 @@ std::unordered_set }); } -std::unordered_set get_ff_dim_t_set(TensorDims const &tensor_dims) { - return unordered_set_of(get_idxs(tensor_dims.ff_ordered)); +std::set get_ff_dim_t_set(TensorDims const &tensor_dims) { + return set_of(get_idxs(tensor_dims.ff_ordered)); } std::optional - get_broadcast_target_dims(std::unordered_set const &dims) { + get_broadcast_target_dims(std::set const &dims) { for (TensorDims target_candidate : dims) { if (all_of(dims, [&](TensorDims const &d) { return tensor_dims_is_broadcastable_to(d, target_candidate); diff --git a/lib/op-attrs/test/src/op-attrs/ff_ordered/ff_ordered_from_map.cc b/lib/op-attrs/test/src/op-attrs/ff_ordered/ff_ordered_from_map.cc index 49bc13cf8e..366a85731b 100644 --- a/lib/op-attrs/test/src/op-attrs/ff_ordered/ff_ordered_from_map.cc +++ b/lib/op-attrs/test/src/op-attrs/ff_ordered/ff_ordered_from_map.cc @@ -7,7 +7,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE_TEMPLATE("ff_ordered_from_map", T, std::map, - std::unordered_map) { + std::map) { SUBCASE("input is empty") { T m = {}; diff --git a/lib/op-attrs/test/src/op-attrs/get_incoming_tensor_roles.cc b/lib/op-attrs/test/src/op-attrs/get_incoming_tensor_roles.cc index e03970b039..34a22fa833 100644 --- a/lib/op-attrs/test/src/op-attrs/get_incoming_tensor_roles.cc +++ b/lib/op-attrs/test/src/op-attrs/get_incoming_tensor_roles.cc @@ -14,9 +14,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_map result = + std::map result = get_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT_0, IncomingTensorRole::INPUT, diff --git a/lib/op-attrs/test/src/op-attrs/operator_task_space.cc b/lib/op-attrs/test/src/op-attrs/operator_task_space.cc index faa04a7bba..d4bd194ce9 100644 --- a/lib/op-attrs/test/src/op-attrs/operator_task_space.cc +++ b/lib/op-attrs/test/src/op-attrs/operator_task_space.cc @@ -1,5 +1,5 @@ #include "op-attrs/operator_task_space.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include using namespace FlexFlow; @@ -10,9 +10,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("OperatorTaskSpace has 0 dimensions") { OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{}}}; - std::unordered_set correct = { + std::set correct = { TaskSpaceCoordinate{OrthotopeCoord{{}}}}; - std::unordered_set result = + std::set result = get_task_space_coordinates(task); CHECK(correct == result); } @@ -22,13 +22,13 @@ TEST_SUITE(FF_TEST_SUITE) { OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{2_ge2, 2_ge2}}}; - std::unordered_set correct = {{ + std::set correct = {{ TaskSpaceCoordinate{OrthotopeCoord{{0_n, 0_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{0_n, 1_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{1_n, 0_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{1_n, 1_n}}}, }}; - std::unordered_set result = + std::set result = get_task_space_coordinates(task); CHECK(correct == result); } @@ -38,7 +38,7 @@ TEST_SUITE(FF_TEST_SUITE) { OperatorTaskSpace task = OperatorTaskSpace{MinimalOrthotope{{3_ge2, 2_ge2, 2_ge2}}}; - std::unordered_set correct = {{ + std::set correct = {{ TaskSpaceCoordinate{OrthotopeCoord{{0_n, 0_n, 0_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{0_n, 0_n, 1_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{0_n, 1_n, 0_n}}}, @@ -52,7 +52,7 @@ TEST_SUITE(FF_TEST_SUITE) { TaskSpaceCoordinate{OrthotopeCoord{{2_n, 1_n, 0_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{2_n, 1_n, 1_n}}}, }}; - std::unordered_set result = + std::set result = get_task_space_coordinates(task); CHECK(correct == result); } diff --git a/lib/op-attrs/test/src/op-attrs/ops/attention.cc b/lib/op-attrs/test/src/op-attrs/ops/attention.cc index 5de69360f8..9a5c6cdbe3 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/attention.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/attention.cc @@ -24,10 +24,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("without bias") { MultiHeadAttentionAttrs attrs = make_attrs(/*bias=*/false); - std::unordered_map result = + std::map result = get_attention_incoming_tensor_roles(attrs); - std::unordered_map correct = - std::unordered_map{ + std::map correct = + std::map{ { TensorSlotName::KEY, IncomingTensorRole::INPUT, @@ -52,10 +52,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("with bias") { MultiHeadAttentionAttrs attrs = make_attrs(/*bias=*/true); - std::unordered_map result = + std::map result = get_attention_incoming_tensor_roles(attrs); - std::unordered_map correct = - std::unordered_map{ + std::map correct = + std::map{ { TensorSlotName::KEY, IncomingTensorRole::INPUT, diff --git a/lib/op-attrs/test/src/op-attrs/ops/batch_norm.cc b/lib/op-attrs/test/src/op-attrs/ops/batch_norm.cc index e39649f9bd..fe224c6674 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/batch_norm.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/batch_norm.cc @@ -21,9 +21,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("affine = true") { BatchNormAttrs attrs = make_attrs(/*affine=*/true); - std::unordered_map result = + std::map result = get_batch_norm_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, @@ -44,9 +44,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("affine = false") { BatchNormAttrs attrs = make_attrs(/*affine=*/false); - std::unordered_map result = + std::map result = get_batch_norm_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, diff --git a/lib/op-attrs/test/src/op-attrs/ops/conv_2d.cc b/lib/op-attrs/test/src/op-attrs/ops/conv_2d.cc index 9c5cd9009b..ad3799b2c5 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/conv_2d.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/conv_2d.cc @@ -22,9 +22,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("with bias") { Conv2DAttrs attrs = make_attrs(/*use_bias=*/true); - std::unordered_map result = + std::map result = get_conv2d_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, @@ -45,9 +45,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("without bias") { Conv2DAttrs attrs = make_attrs(/*use_bias=*/false); - std::unordered_map result = + std::map result = get_conv2d_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, diff --git a/lib/op-attrs/test/src/op-attrs/ops/layer_norm.cc b/lib/op-attrs/test/src/op-attrs/ops/layer_norm.cc index 14591cb3d6..01a5f52c7b 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/layer_norm.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/layer_norm.cc @@ -20,9 +20,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("elementwise_affine = true") { LayerNormAttrs attrs = make_attrs(/*elementwise_affine=*/true); - std::unordered_map result = + std::map result = get_layer_norm_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, @@ -43,9 +43,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("elementwise_affine = false") { LayerNormAttrs attrs = make_attrs(/*elementwise_affine=*/false); - std::unordered_map result = + std::map result = get_layer_norm_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, diff --git a/lib/op-attrs/test/src/op-attrs/ops/linear.cc b/lib/op-attrs/test/src/op-attrs/ops/linear.cc index c46e36bf7b..99f0926d2b 100644 --- a/lib/op-attrs/test/src/op-attrs/ops/linear.cc +++ b/lib/op-attrs/test/src/op-attrs/ops/linear.cc @@ -21,9 +21,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("use_bias = true") { LinearAttrs attrs = make_attrs(/*use_bias=*/true); - std::unordered_map result = + std::map result = get_linear_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, @@ -44,9 +44,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("use_bias = false") { LinearAttrs attrs = make_attrs(/*use_bias=*/false); - std::unordered_map result = + std::map result = get_linear_incoming_tensor_roles(attrs); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, IncomingTensorRole::INPUT, diff --git a/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc b/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc index 6d0e072db5..a7c01f6431 100644 --- a/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc +++ b/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc @@ -1,8 +1,8 @@ #include "op-attrs/parallel_tensor_dim_degrees.h" #include "op-attrs/parallel_tensor_dim_idx_t.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_map.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; @@ -23,9 +23,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_map result = + std::map result = get_parallel_tensor_degree_map(degrees); - std::unordered_map correct = { + std::map correct = { {parallel_tensor_dim_idx_t{ReplicaType::SUM}, 3_p}, {parallel_tensor_dim_idx_t{ReplicaType::DISCARD_COPY}, 1_p}, {shard_dim_idx_from_raw(0), 1_p}, @@ -47,9 +47,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set result = + std::set result = get_parallel_tensor_space_coordinates(degrees); - std::unordered_set correct = { + std::set correct = { ParallelTensorSpaceCoordinate{ /*sum_idx=*/0_n, /*discard_copy_idx=*/0_n, diff --git a/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_idx_t.cc b/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_idx_t.cc index 8edb5d19a9..7a2bad629c 100644 --- a/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_idx_t.cc +++ b/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_idx_t.cc @@ -57,7 +57,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("properly sorts a set of dimensions") { - std::unordered_set input = { + std::set input = { sum_dim_idx(), shard_dim_idx(ff_dim_t{1_n}), shard_dim_idx(ff_dim_t{0_n}), diff --git a/lib/op-attrs/test/src/op-attrs/tensor_dims.cc b/lib/op-attrs/test/src/op-attrs/tensor_dims.cc index fc501873d9..d191e8e482 100644 --- a/lib/op-attrs/test/src/op-attrs/tensor_dims.cc +++ b/lib/op-attrs/test/src/op-attrs/tensor_dims.cc @@ -1,6 +1,6 @@ #include "op-attrs/tensor_dims.h" #include "test/utils/doctest/fmt/optional.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; @@ -121,9 +121,9 @@ TEST_SUITE(FF_TEST_SUITE) { FFOrdered{3_p, 1_p, 2_p}, }; - std::unordered_set result = + std::set result = get_tensor_dims_coord_set(input); - std::unordered_set correct = { + std::set correct = { TensorDimsCoord{FFOrdered{0_n, 0_n, 0_n}}, TensorDimsCoord{FFOrdered{0_n, 0_n, 1_n}}, TensorDimsCoord{FFOrdered{1_n, 0_n, 0_n}}, @@ -138,9 +138,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("TensorDims is zero-dimensional") { TensorDims input = TensorDims{FFOrdered{}}; - std::unordered_set result = + std::set result = get_tensor_dims_coord_set(input); - std::unordered_set correct = { + std::set correct = { TensorDimsCoord{FFOrdered{}}, }; @@ -148,7 +148,7 @@ TEST_SUITE(FF_TEST_SUITE) { } } - TEST_CASE("get_broadcast_target_dims(std::unordered_set)") { + TEST_CASE("get_broadcast_target_dims(std::set)") { TensorDims d1 = TensorDims{FFOrdered{1_p, 10_p, 4_p, 3_p}}; TensorDims d2 = TensorDims{FFOrdered{10_p, 4_p, 1_p}}; diff --git a/lib/pcg/include/pcg/computation_graph.h b/lib/pcg/include/pcg/computation_graph.h index 8dfb2eedb4..286ac6d027 100644 --- a/lib/pcg/include/pcg/computation_graph.h +++ b/lib/pcg/include/pcg/computation_graph.h @@ -14,14 +14,14 @@ namespace FlexFlow { ComputationGraph make_empty_computation_graph(); -std::unordered_set get_layers(ComputationGraph const &); +std::set get_layers(ComputationGraph const &); LayerAddedResult add_layer( ComputationGraph &computation_graph, LayerAttrs const &attrs, - std::unordered_map const &inputs, - std::unordered_map const &weights, - std::optional> const + std::map const &inputs, + std::map const &weights, + std::optional> const &outputs = std::nullopt); LayerAddedResult add_input_layer(ComputationGraph &computation_graph, @@ -35,38 +35,38 @@ bool are_tensor_guid_shapes_equivalent(ComputationGraph const &cg, std::vector topological_ordering(ComputationGraph const &cg); -std::unordered_map +std::map get_outgoing_tensors(ComputationGraph const &cg, layer_guid_t n); -std::unordered_map +std::map get_incoming_tensors(ComputationGraph const &cg, layer_guid_t n); -std::unordered_map +std::map get_incoming_inputs(ComputationGraph const &, layer_guid_t const &); -std::unordered_map +std::map get_incoming_input_shapes(ComputationGraph const &, layer_guid_t const &); -std::unordered_map +std::map get_incoming_weights(ComputationGraph const &, layer_guid_t const &); -std::unordered_set get_all_tensors(ComputationGraph const &); -std::unordered_map +std::set get_all_tensors(ComputationGraph const &); +std::map get_all_tensor_attrs(ComputationGraph const &); -std::unordered_set +std::set get_subgraph_incoming_edges(ComputationGraph const &, - std::unordered_set const &); -std::unordered_set + std::set const &); +std::set get_subgraph_outgoing_edges(ComputationGraph const &, - std::unordered_set const &); -std::unordered_set + std::set const &); +std::set get_subgraph_successors(ComputationGraph const &, - std::unordered_set const &); + std::set const &); LayerAttrs get_layer_attrs(ComputationGraph const &cg, layer_guid_t const &n); -std::unordered_map +std::map get_layer_attrs_mapping(ComputationGraph const &cg); layer_guid_t get_layer_by_name(ComputationGraph const &cg, diff --git a/lib/pcg/include/pcg/computation_graph/layer_added_result.dtg.toml b/lib/pcg/include/pcg/computation_graph/layer_added_result.dtg.toml index 17256abe5a..7bc0df2088 100644 --- a/lib/pcg/include/pcg/computation_graph/layer_added_result.dtg.toml +++ b/lib/pcg/include/pcg/computation_graph/layer_added_result.dtg.toml @@ -13,7 +13,7 @@ includes = [ ] src_includes = [ - "utils/fmt/unordered_map.h", + "utils/fmt/map.h", ] [[fields]] @@ -22,4 +22,4 @@ type = "::FlexFlow::layer_guid_t" [[fields]] name = "outputs" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::tensor_guid_t>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::tensor_guid_t>" diff --git a/lib/pcg/include/pcg/computation_graph_builder.h b/lib/pcg/include/pcg/computation_graph_builder.h index 4e4cacc731..f183aabbe1 100644 --- a/lib/pcg/include/pcg/computation_graph_builder.h +++ b/lib/pcg/include/pcg/computation_graph_builder.h @@ -262,11 +262,11 @@ struct ComputationGraphBuilder { TensorShape get_shape(tensor_guid_t const &) const; private: - std::unordered_map add_layer( + std::map add_layer( LayerAttrs const &layer, - std::unordered_map const &inputs, - std::unordered_map const &weights, - std::optional> const + std::map const &inputs, + std::map const &weights, + std::optional> const &outputs = std::nullopt); tensor_guid_t diff --git a/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.dtg.toml b/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.dtg.toml index fa32a3cbce..73bcd255cd 100644 --- a/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.dtg.toml +++ b/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.dtg.toml @@ -16,7 +16,7 @@ template_params = [ includes = [ "", - "", + "", "pcg/file_format/v1/graphs/v1_kwarg_graph_edge.dtg.h", "pcg/file_format/v1/graphs/v1_kwarg_graph_output.dtg.h", "utils/nonnegative_int/nonnegative_int.h", @@ -25,8 +25,8 @@ includes = [ src_includes = [ "utils/fmt/vector.h", "utils/hash/vector.h", - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h", + "utils/fmt/set.h", + "utils/hash/set.h", ] [[fields]] @@ -35,8 +35,8 @@ type = "std::vector<::FlexFlow::nonnegative_int>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::V1KwargGraphEdge>" +type = "std::set<::FlexFlow::V1KwargGraphEdge>" [[fields]] name = "outputs" -type = "std::unordered_set<::FlexFlow::V1KwargGraphOutput>" +type = "std::set<::FlexFlow::V1KwargGraphOutput>" diff --git a/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h b/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h index 17d5c723bd..c1fed1ffb2 100644 --- a/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h +++ b/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h @@ -9,7 +9,7 @@ #include "utils/containers/generate_map.h" #include "utils/containers/sorted.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" @@ -26,17 +26,17 @@ V1KwargDataflowGraph to_v1(KwargDataflowGraphView const &g) { bidict node_enumeration_bidict = bidict_from_enumerating(get_nodes(g)); - std::unordered_map node_enumeration = - node_enumeration_bidict.reversed().as_unordered_map(); + std::map node_enumeration = + node_enumeration_bidict.reversed().as_map(); return to_v1(g, node_enumeration); } template V1KwargDataflowGraph to_v1(KwargDataflowGraphView const &g, - std::unordered_map const &nodes) { + std::map const &nodes) { - std::unordered_set> edges = + std::set> edges = transform(get_all_kwarg_dataflow_edges(g), [&](KwargDataflowEdge const &e) { return V1KwargGraphEdge{nodes.at(e.src.node), @@ -45,7 +45,7 @@ V1KwargDataflowGraph e.dst.slot_name}; }); - std::unordered_set> outputs = + std::set> outputs = transform(get_all_kwarg_dataflow_outputs(g), [&](KwargDataflowOutput const &o) { return V1KwargGraphOutput{nodes.at(o.node), o.slot_name}; @@ -60,15 +60,15 @@ V1KwargDataflowGraph template std::pair, - std::unordered_map> + std::map> from_v1_including_node_numbering(V1KwargDataflowGraph const &v1) { - std::unordered_map node_map = + std::map node_map = generate_map(v1.nodes, [](nonnegative_int n) { return Node{n.size_t_from_nonnegative_int()}; }); - std::unordered_set node_set = unordered_set_of(values(node_map)); + std::set node_set = set_of(values(node_map)); - std::unordered_set> edges = + std::set> edges = transform(v1.edges, [](V1KwargGraphEdge const &e) { Node srcNode = Node{e.srcNode.size_t_from_nonnegative_int()}; Node dstNode = Node{e.dstNode.size_t_from_nonnegative_int()}; @@ -78,7 +78,7 @@ std::pair, }}; }); - std::unordered_set> outputs = + std::set> outputs = transform(v1.outputs, [](V1KwargGraphOutput const &o) { Node n = Node{o.node.size_t_from_nonnegative_int()}; return KwargDataflowOutput{n, o.slot_name}; @@ -86,10 +86,10 @@ std::pair, OpenKwargDataflowGraphData graph_data = OpenKwargDataflowGraphData{ - /*nodes=*/node_set, - /*edges=*/edges, + /*nodes=*/set_of(node_set), + /*edges=*/set_of(edges), /*inputs=*/{}, - /*outputs=*/outputs, + /*outputs=*/set_of(outputs), }; return std::pair{view_from_open_kwarg_dataflow_graph_data(graph_data), node_map}; diff --git a/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.dtg.toml b/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.dtg.toml index 76113dfea6..f71d8cfdb5 100644 --- a/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.dtg.toml +++ b/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.dtg.toml @@ -17,26 +17,26 @@ template_params = [ ] includes = [ - "", + "", "pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.dtg.h", "pcg/file_format/v1/graphs/v1_kwarg_graph_output.dtg.h", "utils/nonnegative_int/nonnegative_int.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", "utils/fmt/vector.h", "utils/hash/vector.h", ] [[fields]] name = "node_labels" -type = "std::unordered_map<::FlexFlow::nonnegative_int, NodeLabel>" +type = "std::map<::FlexFlow::nonnegative_int, NodeLabel>" [[fields]] name = "output_labels" -type = "std::unordered_map<::FlexFlow::V1KwargGraphOutput, OutputLabel>" +type = "std::map<::FlexFlow::V1KwargGraphOutput, OutputLabel>" [[fields]] name = "graph" diff --git a/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h b/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h index 0d876c08dc..eee45c900d 100644 --- a/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h +++ b/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h @@ -7,7 +7,7 @@ #include "utils/containers/map_keys.h" #include "utils/containers/map_values.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_map_from_pairs.h" +#include "utils/containers/map_from_pairs.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h" #include "utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph_view.h" @@ -26,11 +26,11 @@ std::pair, V1KwargDataflowGraph unlabelled = to_v1(g, nodes.reversed()); - std::unordered_map node_labels = map_values( - nodes.as_unordered_map(), [&](Node const &n) { return g.at(n); }); + std::map node_labels = map_values( + nodes.as_map(), [&](Node const &n) { return g.at(n); }); - std::unordered_map, OutputLabel> output_labels = - unordered_map_from_pairs( + std::map, OutputLabel> output_labels = + map_from_pairs( transform(get_all_kwarg_dataflow_outputs(g), [&](KwargDataflowOutput const &o) { return std::pair{ @@ -54,16 +54,16 @@ V1LabelledKwargDataflowGraph to_v1( template std::pair, - std::unordered_map> + std::map> from_v1_including_node_numbering( V1LabelledKwargDataflowGraph const &v1) { auto [graph_view, node_map] = from_v1_including_node_numbering(v1.graph); - std::unordered_map node_labels = map_keys( + std::map node_labels = map_keys( v1.node_labels, [&](nonnegative_int n) { return node_map.at(n); }); - std::unordered_map, OutputLabel> value_labels = + std::map, OutputLabel> value_labels = map_keys(v1.output_labels, [&](V1KwargGraphOutput const &o) { return KwargDataflowOutput{node_map.at(o.node), o.slot_name}; }); diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.toml b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.toml index a4f53166d2..79f5ae7e13 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.toml +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.toml @@ -7,7 +7,7 @@ includes = [ "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h", "pcg/parallel_computation_graph/parallel_computation_graph.h", "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_attrs.dtg.h", - "", + "", ] [[fields]] diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h index 2e789e14c9..c33bc7e7a4 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h @@ -7,7 +7,7 @@ namespace FlexFlow { -std::unordered_set +std::set mpcg_get_parallel_layers(MappedParallelComputationGraph const &); std::set @@ -30,11 +30,11 @@ ParallelTensorAttrs mpcg_get_parallel_tensor_attrs(MappedParallelComputationGraph const &, parallel_tensor_guid_t const &); -std::unordered_map +std::map mpcg_get_incoming_edges(MappedParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_set +std::set mpcg_get_outgoing_edges(MappedParallelComputationGraph const &, parallel_layer_guid_t const &); @@ -46,16 +46,16 @@ bidict mpcg_get_outgoing_tensors(MappedParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_set +std::set mpcg_get_edges(MappedParallelComputationGraph const &); -std::unordered_set +std::set mpcg_get_parallel_tensor_uses(MappedParallelComputationGraph const &, parallel_tensor_guid_t const &); MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( ParallelComputationGraph const &pcg, - std::unordered_map const + std::map const &mapped_op_task_groups); MappedParallelComputationGraph diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.toml b/lib/pcg/include/pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.toml index c06eff0375..cee6696bee 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.toml +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.toml @@ -15,11 +15,10 @@ includes = [ ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", - "utils/ord/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", ] [[fields]] name = "tensor_coords" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorSpaceCoordinate>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::ParallelTensorSpaceCoordinate>" diff --git a/lib/pcg/include/pcg/metric_attrs.h b/lib/pcg/include/pcg/metric_attrs.h index 21f9115a67..c1cbdd4e96 100644 --- a/lib/pcg/include/pcg/metric_attrs.h +++ b/lib/pcg/include/pcg/metric_attrs.h @@ -4,14 +4,14 @@ #include "op-attrs/ops/loss_functions/loss_function.dtg.h" #include "pcg/metric.dtg.h" #include "utils/fmt.h" -#include +#include namespace FlexFlow { class MetricsAttrs { public: MetricsAttrs() = delete; - MetricsAttrs(LossFunction, std::unordered_set const &); + MetricsAttrs(LossFunction, std::set const &); public: LossFunction loss_type; diff --git a/lib/pcg/include/pcg/operator_space_to_machine_space_mapping.dtg.toml b/lib/pcg/include/pcg/operator_space_to_machine_space_mapping.dtg.toml index a97d84da12..e93c95116e 100644 --- a/lib/pcg/include/pcg/operator_space_to_machine_space_mapping.dtg.toml +++ b/lib/pcg/include/pcg/operator_space_to_machine_space_mapping.dtg.toml @@ -3,6 +3,7 @@ name = "OperatorSpaceToMachineSpaceMapping" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/pcg/include/pcg/optimizer_attrs.h b/lib/pcg/include/pcg/optimizer_attrs.h index b554b68284..2a428223c6 100644 --- a/lib/pcg/include/pcg/optimizer_attrs.h +++ b/lib/pcg/include/pcg/optimizer_attrs.h @@ -9,7 +9,7 @@ namespace FlexFlow { OptimizerAttrs get_optimizer_attrs_for_next_iter(OptimizerAttrs const &old); -std::unordered_set +std::set get_slot_names_for_optimizer(OptimizerAttrs const &); } // namespace FlexFlow diff --git a/lib/pcg/include/pcg/parallel_computation_graph/generate_weight_transform.h b/lib/pcg/include/pcg/parallel_computation_graph/generate_weight_transform.h index eb4928deaa..442da1376e 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/generate_weight_transform.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/generate_weight_transform.h @@ -7,7 +7,7 @@ namespace FlexFlow { -std::unordered_set +std::set generate_weight_transform(TensorShape const ¤t, ParallelTensorShape const &goal); diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h index 7c5a825420..d3cc9f0149 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h @@ -11,24 +11,24 @@ #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h" #include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" -#include +#include #include "pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.h" namespace FlexFlow { ParallelComputationGraph empty_parallel_computation_graph(); -std::unordered_set - get_parallel_layers(ParallelComputationGraph const &); -std::unordered_set +std::set + pcg_get_parallel_layers(ParallelComputationGraph const &); +std::set get_parallel_tensors(ParallelComputationGraph const &); ParallelLayerAddedResult add_parallel_layer( ParallelComputationGraph &pcg, ParallelLayerAttrs const &layer_attrs, - std::unordered_map const &inputs, - std::unordered_map const &weights, - std::optional> const + std::map const &inputs, + std::map const &weights, + std::optional> const &outputs = std::nullopt); ParallelLayerAddedResult @@ -46,41 +46,41 @@ ParallelLayerInvocationInfo pcg_get_invocation_info_for_layer(ParallelComputationGraph const &, parallel_layer_guid_t); -std::unordered_set +std::set get_pcg_edges_from_layer_to_layer(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &src, parallel_layer_guid_t const &dst); -std::unordered_set +std::set get_edges(ParallelComputationGraph const &); -std::unordered_set +std::set get_outgoing_edges(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_map +std::map get_incoming_edges(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_set +std::set pcg_get_parallel_tensor_uses(ParallelComputationGraph const &, parallel_tensor_guid_t const &); -std::unordered_set +std::set get_initial_layers(ParallelComputationGraph const &); -std::unordered_map +std::map get_outgoing_tensors(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_map +std::map get_incoming_tensors(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_map +std::map pcg_get_operator_to_incoming_mappings(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_map +std::map pcg_get_operator_to_output_mappings(ParallelComputationGraph const &, parallel_layer_guid_t const &); @@ -88,24 +88,24 @@ OperatorTaskSpaceToOperatorTaskSpaceMapping pcg_get_mapping_along_edge(ParallelComputationGraph const &, ParallelComputationGraphEdge const &); -std::unordered_map +std::map get_incoming_inputs(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_map +std::map get_incoming_weights(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_map +std::map get_incoming_input_degrees(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_set +std::set get_successors(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_set +std::set get_subgraph_successors(ParallelComputationGraph const &, - std::unordered_set const &); + std::set const &); parallel_layer_guid_t get_source_layer(ParallelComputationGraph const &g, parallel_tensor_guid_t const &t); @@ -122,7 +122,7 @@ ParallelTensorShape get_parallel_tensor_shape(ParallelComputationGraph const &, std::vector topological_ordering(ParallelComputationGraph const &); -std::unordered_map +std::map get_parallel_layer_attrs_mapping(ParallelComputationGraph const &pcg); parallel_layer_guid_t diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h index 88df8128e7..23d1e0ff9c 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h @@ -138,10 +138,10 @@ struct ParallelComputationGraphBuilder { std::string const &name); private: - std::unordered_map add_layer( + std::map add_layer( ParallelLayerAttrs const &layer, - std::unordered_map const &inputs, - std::unordered_map const + std::map const &inputs, + std::map const &weight_initializers); parallel_tensor_guid_t diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_added_result.dtg.toml b/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_added_result.dtg.toml index 455e61c783..4db23ca20d 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_added_result.dtg.toml +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_layer_added_result.dtg.toml @@ -15,7 +15,7 @@ includes = [ ] src_includes = [ - "utils/fmt/unordered_map.h", + "utils/fmt/map.h", ] [[fields]] @@ -24,4 +24,4 @@ type = "::FlexFlow::parallel_layer_guid_t" [[fields]] name = "outputs" -type = "std::unordered_map<::FlexFlow::TensorSlotName, ::FlexFlow::parallel_tensor_guid_t>" +type = "std::map<::FlexFlow::TensorSlotName, ::FlexFlow::parallel_tensor_guid_t>" diff --git a/lib/pcg/src/pcg/computation_graph.cc b/lib/pcg/src/pcg/computation_graph.cc index 1b0b3b3204..1ea5420951 100644 --- a/lib/pcg/src/pcg/computation_graph.cc +++ b/lib/pcg/src/pcg/computation_graph.cc @@ -34,7 +34,7 @@ #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/node/algorithms.h" #include "utils/record_formatter.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -47,7 +47,7 @@ ComputationGraph make_empty_computation_graph() { TensorSlotName>>()}; } -std::unordered_set get_layers(ComputationGraph const &cg) { +std::set get_layers(ComputationGraph const &cg) { return transform(get_nodes(cg.raw_graph), [&](Node const &n) { return layer_guid_t{n}; }); } @@ -55,39 +55,39 @@ std::unordered_set get_layers(ComputationGraph const &cg) { LayerAddedResult add_layer( ComputationGraph &computation_graph, LayerAttrs const &layer_attrs, - std::unordered_map const &inputs, - std::unordered_map const &weights, - std::optional> const + std::map const &inputs, + std::map const &weights, + std::optional> const &maybe_output_flags) { - std::unordered_map input_shapes = + std::map input_shapes = map_values(inputs, [&](tensor_guid_t const &i) { return get_tensor_attrs(computation_graph, i).shape; }); - std::unordered_map provided_weight_shapes = + std::map provided_weight_shapes = map_values(weights, [&](tensor_guid_t const &w) { return get_tensor_attrs(computation_graph, w).shape; }); - std::unordered_map expected_weight_shapes = + std::map expected_weight_shapes = get_weight_shapes(layer_attrs.op_attrs, input_shapes); - std::unordered_map> + std::map> raw_inputs = map_values( inputs, [&](tensor_guid_t const &t) { return t.raw_graph_output; }); - std::unordered_map> + std::map> raw_weights = map_values( weights, [&](tensor_guid_t const &t) { return t.raw_graph_output; }); - std::unordered_map output_shapes = + std::map output_shapes = get_output_shapes(layer_attrs.op_attrs, input_shapes); - std::unordered_map output_flags = + std::map output_flags = maybe_output_flags.value_or(map_values( output_shapes, [&](TensorShape const &) { return CreateGrad::YES; })); - std::unordered_map output_attrs = + std::map output_attrs = zip_values_strict_with(output_shapes, output_flags, [](TensorShape const &shape, @@ -101,7 +101,7 @@ LayerAddedResult add_layer( KwargNodeAddedResult added = computation_graph.raw_graph.add_node( layer_attrs, - binary_merge_disjoint_unordered_maps(raw_inputs, raw_weights), + binary_merge_disjoint_maps(raw_inputs, raw_weights), output_attrs); return LayerAddedResult{ @@ -126,7 +126,7 @@ LayerAddedResult add_input_layer(ComputationGraph &cg, /*inputs=*/{}, /*weights=*/{}, /*outputs=*/ - std::unordered_map{ + std::map{ {TensorSlotName::OUTPUT, create_grad}, }); } @@ -155,7 +155,7 @@ std::vector layers, [&](Node const &e) -> layer_guid_t { return layer_guid_t{e}; }); } -std::unordered_map +std::map get_outgoing_tensors(ComputationGraph const &cg, layer_guid_t n) { return map_values( get_outgoing_kwarg_dataflow_outputs_for_node(cg.raw_graph, n.raw_node), @@ -164,7 +164,7 @@ std::unordered_map }); } -std::unordered_map +std::map get_incoming_tensors(ComputationGraph const &cg, layer_guid_t n) { return map_values( get_incoming_kwarg_dataflow_outputs_for_node(cg.raw_graph, n.raw_node), @@ -173,7 +173,7 @@ std::unordered_map }); } -std::unordered_map +std::map get_incoming_input_shapes(ComputationGraph const &cg, layer_guid_t const &n) { return map_values(get_incoming_inputs(cg, n), [&](tensor_guid_t const &t) { @@ -181,62 +181,62 @@ std::unordered_map }); } -static std::unordered_map +static std::map get_incoming_tensors_with_role(ComputationGraph const &cg, layer_guid_t const &l, IncomingTensorRole desired_role) { ComputationGraphOpAttrs attrs = get_layer_attrs(cg, l).op_attrs; - std::unordered_map incoming_tensors = + std::map incoming_tensors = get_incoming_tensors(cg, l); - std::unordered_map incoming_slot_roles = + std::map incoming_slot_roles = get_incoming_tensor_roles(attrs); ASSERT(incoming_tensors.size() == incoming_slot_roles.size()); - std::unordered_set slots_with_desired_role = - unordered_keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { + std::set slots_with_desired_role = + keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { return role == desired_role; })); return restrict_keys(incoming_tensors, slots_with_desired_role); } -std::unordered_map +std::map get_incoming_inputs(ComputationGraph const &cg, layer_guid_t const &l) { return get_incoming_tensors_with_role(cg, l, IncomingTensorRole::INPUT); } -std::unordered_map +std::map get_incoming_weights(ComputationGraph const &cg, layer_guid_t const &l) { return get_incoming_tensors_with_role(cg, l, IncomingTensorRole::WEIGHT); } -std::unordered_set get_all_tensors(ComputationGraph const &cg) { +std::set get_all_tensors(ComputationGraph const &cg) { return transform(get_all_kwarg_dataflow_outputs(cg.raw_graph), [](KwargDataflowOutput const &t) { return tensor_guid_t(t); }); } -std::unordered_map +std::map get_all_tensor_attrs(ComputationGraph const &cg) { - std::unordered_set all_tensors = get_all_tensors(cg); - std::unordered_map all_tensor_attrs; + std::set all_tensors = get_all_tensors(cg); + std::map all_tensor_attrs; for (tensor_guid_t const &tensor_guid : all_tensors) { all_tensor_attrs.insert({tensor_guid, get_tensor_attrs(cg, tensor_guid)}); } return all_tensor_attrs; } -std::unordered_set get_subgraph_incoming_edges( +std::set get_subgraph_incoming_edges( ComputationGraph const &cg, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { - std::unordered_set raw_subgraph_nodes = transform( + std::set raw_subgraph_nodes = transform( subgraph_nodes, [](layer_guid_t const &l) { return l.raw_node; }); - std::unordered_set> raw_incoming_edges = + std::set> raw_incoming_edges = get_kwarg_dataflow_subgraph_incoming_edges(cg.raw_graph, raw_subgraph_nodes); @@ -246,13 +246,13 @@ std::unordered_set get_subgraph_incoming_edges( }); } -std::unordered_set get_subgraph_outgoing_edges( +std::set get_subgraph_outgoing_edges( ComputationGraph const &cg, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { - std::unordered_set raw_subgraph_nodes = transform( + std::set raw_subgraph_nodes = transform( subgraph_nodes, [](layer_guid_t const &l) { return l.raw_node; }); - std::unordered_set> raw_outgoing_edges = + std::set> raw_outgoing_edges = get_kwarg_dataflow_subgraph_outgoing_edges(cg.raw_graph, raw_subgraph_nodes); @@ -262,13 +262,13 @@ std::unordered_set get_subgraph_outgoing_edges( }); } -std::unordered_set get_subgraph_successors( +std::set get_subgraph_successors( ComputationGraph const &cg, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { - std::unordered_set raw_subgraph_nodes = transform( + std::set raw_subgraph_nodes = transform( subgraph_nodes, [](layer_guid_t const &l) { return l.raw_node; }); - std::unordered_set raw_successors = + std::set raw_successors = get_subgraph_successors(cg.raw_graph, raw_subgraph_nodes); return transform(raw_successors, @@ -279,9 +279,9 @@ LayerAttrs get_layer_attrs(ComputationGraph const &cg, layer_guid_t const &n) { return cg.raw_graph.at(n.raw_node); } -std::unordered_map +std::map get_layer_attrs_mapping(ComputationGraph const &cg) { - std::unordered_map layer_attrs_mapping; + std::map layer_attrs_mapping; for (layer_guid_t const &layer_guid : get_layers(cg)) { layer_attrs_mapping.insert({layer_guid, get_layer_attrs(cg, layer_guid)}); } @@ -290,7 +290,7 @@ std::unordered_map layer_guid_t get_layer_by_name(ComputationGraph const &cg, std::string const &name) { - std::unordered_set found = + std::set found = filter(get_layers(cg), [&](layer_guid_t const &l) { return get_layer_attrs(cg, l).name == name; }); @@ -348,8 +348,8 @@ std::string as_dot(ComputationGraph const &cg) { }; std::function( - std::unordered_set const &)> - order_slots = [](std::unordered_set const &unordered) + std::set const &)> + order_slots = [](std::set const &unordered) -> nlohmann::json { return sorted(unordered); }; return labelled_open_kwarg_dataflow_graph_view_as_dot( diff --git a/lib/pcg/src/pcg/computation_graph_builder.cc b/lib/pcg/src/pcg/computation_graph_builder.cc index 2eb140fa58..7ca786661a 100644 --- a/lib/pcg/src/pcg/computation_graph_builder.cc +++ b/lib/pcg/src/pcg/computation_graph_builder.cc @@ -48,7 +48,7 @@ #include "utils/fmt/set.h" #include "utils/stack_vector/stack_vector_of.h" #include -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -82,7 +82,7 @@ tensor_guid_t ComputationGraphBuilder::create_input( /*inputs=*/{}, /*weights=*/{}, /*outputs=*/ - std::unordered_map{ + std::map{ { TensorSlotName::OUTPUT, create_grad, @@ -109,17 +109,17 @@ tensor_guid_t ComputationGraphBuilder::create_weight( static void check_incoming_tensor_roles( LayerAttrs const &layer, - std::unordered_set const &input_slots, - std::unordered_set const &weight_slots) { - std::unordered_map correct = + std::set const &input_slots, + std::set const &weight_slots) { + std::map correct = restrict_keys(get_incoming_tensor_roles(layer.op_attrs), set_union(input_slots, weight_slots)); - std::unordered_map current = - binary_merge_disjoint_unordered_maps( - generate_unordered_map( + std::map current = + binary_merge_disjoint_maps( + generate_map( input_slots, [](TensorSlotName) { return IncomingTensorRole::INPUT; }), - generate_unordered_map(weight_slots, [](TensorSlotName) { + generate_map(weight_slots, [](TensorSlotName) { return IncomingTensorRole::WEIGHT; })); @@ -127,24 +127,24 @@ static void check_incoming_tensor_roles( "check_incoming_tensor_roles found deviation in incoming tensors"); } -std::unordered_map +std::map ComputationGraphBuilder::add_layer( LayerAttrs const &layer, - std::unordered_map const &inputs, - std::unordered_map const + std::map const &inputs, + std::map const &weight_initializers, - std::optional> const + std::optional> const &outputs) { - ASSERT(are_disjoint(unordered_keys(inputs), unordered_keys(weight_initializers))); - check_incoming_tensor_roles(layer, unordered_keys(inputs), unordered_keys(weight_initializers)); + ASSERT(are_disjoint(keys(inputs), keys(weight_initializers))); + check_incoming_tensor_roles(layer, keys(inputs), keys(weight_initializers)); - std::unordered_map input_shapes = map_values( + std::map input_shapes = map_values( inputs, [&](tensor_guid_t const &t) { return this->get_shape(t); }); - std::unordered_map weight_shapes = + std::map weight_shapes = get_weight_shapes(layer.op_attrs, input_shapes); - std::unordered_map weights = + std::map weights = zip_values_strict_with( weight_shapes, weight_initializers, @@ -471,7 +471,7 @@ tensor_guid_t ComputationGraphBuilder::conv2d( LayerAttrs layer = LayerAttrs{ComputationGraphOpAttrs{attrs}, name}; - std::unordered_map initializers = + std::map initializers = get_initializers(attrs, this->get_shape(input), maybe_kernel_initializer, @@ -533,7 +533,7 @@ tensor_guid_t ComputationGraphBuilder::embedding( TensorShape input_shape = this->get_shape(input); - std::unordered_map initializers = + std::map initializers = get_initializers(attrs, initializer); return require_only_key(this->add_layer(layer, @@ -687,7 +687,7 @@ tensor_guid_t ComputationGraphBuilder::batch_norm( TensorShape input_shape = this->get_shape(input); - std::unordered_map initializers = + std::map initializers = throw_if_unexpected(get_initializers(attrs)); return require_only_key(this->add_layer(layer, @@ -742,7 +742,7 @@ tensor_guid_t ComputationGraphBuilder::multihead_attention( LayerAttrs layer = LayerAttrs{ComputationGraphOpAttrs{attrs}, name}; - std::unordered_map initializers = + std::map initializers = throw_if_unexpected(get_initializers(attrs, this->get_shape(query), this->get_shape(key), @@ -779,7 +779,7 @@ TensorDims ComputationGraphBuilder::get_broadcast_target_dims( TensorDims ComputationGraphBuilder::get_broadcast_target_dims( std::vector const &inputs_dims) { std::optional maybe_result = - ::FlexFlow::get_broadcast_target_dims(unordered_set_of(inputs_dims)); + ::FlexFlow::get_broadcast_target_dims(set_of(inputs_dims)); if (maybe_result.has_value()) { return maybe_result.value(); @@ -813,7 +813,7 @@ tensor_guid_t ComputationGraphBuilder::dense( LayerAttrs layer = LayerAttrs{ComputationGraphOpAttrs{attrs}, name}; - std::unordered_map initializers = + std::map initializers = throw_if_unexpected(get_initializers(attrs, this->get_shape(input), maybe_projection_initializer, @@ -854,7 +854,7 @@ tensor_guid_t ComputationGraphBuilder::concat( return require_only_key( this->add_layer( - layer, unordered_map_from_pairs(zip(input_slot_names, inputs)), {}), + layer, map_from_pairs(zip(input_slot_names, inputs)), {}), TensorSlotName::OUTPUT); } @@ -931,7 +931,7 @@ tensor_guid_t ComputationGraphBuilder::layer_norm( LayerAttrs layer = LayerAttrs{ComputationGraphOpAttrs{attrs}, name}; - std::unordered_map initializers = + std::map initializers = get_initializers(attrs); return require_only_key(this->add_layer(layer, diff --git a/lib/pcg/src/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.cc b/lib/pcg/src/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.cc index cc10bbf4cb..326ecb116f 100644 --- a/lib/pcg/src/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.cc +++ b/lib/pcg/src/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.cc @@ -10,10 +10,10 @@ template V1KwargDataflowGraph template V1KwargDataflowGraph to_v1(KwargDataflowGraphView const &, - std::unordered_map const &); + std::map const &); template std::pair, - std::unordered_map> + std::map> from_v1_including_node_numbering(V1KwargDataflowGraph const &); template KwargDataflowGraphView diff --git a/lib/pcg/src/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.cc b/lib/pcg/src/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.cc index 4e50949e3f..82b8ed58c4 100644 --- a/lib/pcg/src/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.cc +++ b/lib/pcg/src/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using NodeLabel = value_type<0>; -using OutputLabel = value_type<1>; +using OutputLabel = ordered_value_type<1>; using SlotName = ordered_value_type<2>; template std::pair< @@ -20,7 +20,7 @@ template V1LabelledKwargDataflowGraph to_v1( template std::pair< LabelledKwargDataflowGraphView, - std::unordered_map> + std::map> from_v1_including_node_numbering( V1LabelledKwargDataflowGraph const &); diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc index 72b3aac5dd..88ccb50647 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc @@ -18,7 +18,6 @@ #include "utils/containers/contains.h" #include "utils/bidict/algorithms/right_entries.h" #include "utils/containers/map_values.h" -#include "utils/containers/unordered_set_of.h" #include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" namespace FlexFlow { @@ -27,14 +26,14 @@ MappedOperatorTaskGroup::MappedOperatorTaskGroup( bidict const &shard_bindings) : shard_bindings(shard_bindings) { - std::vector> binding_slot_sets = + std::vector> binding_slot_sets = transform(vector_of(shard_bindings.right_values()), [&](OperatorAtomicTaskShardBinding const &s) - -> std::unordered_set { - return unordered_keys(s.tensor_coords); + -> std::set { + return keys(s.tensor_coords); }); - std::unordered_set slot_names = + std::set slot_names = require_all_same(binding_slot_sets).value(); for (TensorSlotName const &slot_name : slot_names) { @@ -104,13 +103,13 @@ bidict std::set slot_names = get_slot_names_for_task_group(task_group); ASSERT(contains(slot_names, slot_name)); - std::unordered_map m = - map_values(task_group.get_shard_bindings().as_unordered_map(), + std::map m = + map_values(task_group.get_shard_bindings().as_map(), [&](OperatorAtomicTaskShardBinding const &b) -> ParallelTensorSpaceCoordinate { return ptensor_space_coord_for_slot_name(b, slot_name); }); - return bidict_from_unstructured_relation(unordered_set_of(m)).reversed(); + return bidict_from_unstructured_relation(set_of(m)).reversed(); } std::set get_slot_names_for_task_group(MappedOperatorTaskGroup const &g) { diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 1bece17c9a..7571daecfa 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -19,9 +19,9 @@ namespace FlexFlow { -std::unordered_set +std::set mpcg_get_parallel_layers(MappedParallelComputationGraph const &mpcg) { - return get_parallel_layers(pcg_from_mpcg(mpcg)); + return pcg_get_parallel_layers(pcg_from_mpcg(mpcg)); } std::set @@ -88,13 +88,13 @@ ParallelTensorAttrs return get_parallel_tensor_attrs(pcg_from_mpcg(mpcg), t); } -std::unordered_map +std::map mpcg_get_incoming_edges(MappedParallelComputationGraph const &mpcg, parallel_layer_guid_t const &l) { return get_incoming_edges(pcg_from_mpcg(mpcg), l); } -std::unordered_set +std::set mpcg_get_outgoing_edges(MappedParallelComputationGraph const &mpcg, parallel_layer_guid_t const &l) { return get_outgoing_edges(pcg_from_mpcg(mpcg), l); @@ -112,12 +112,12 @@ bidict return bidict_from_map(get_outgoing_tensors(pcg_from_mpcg(mpcg), l)); } -std::unordered_set +std::set mpcg_get_edges(MappedParallelComputationGraph const &mpcg) { return get_edges(pcg_from_mpcg(mpcg)); } -std::unordered_set +std::set mpcg_get_parallel_tensor_uses(MappedParallelComputationGraph const &mpcg, parallel_tensor_guid_t const &t) { return pcg_get_parallel_tensor_uses(pcg_from_mpcg(mpcg), t); @@ -125,7 +125,7 @@ std::unordered_set MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( ParallelComputationGraph const &pcg, - std::unordered_map const + std::map const &mapped_op_task_groups) { auto mapping_for_layer = [&](parallel_layer_guid_t l) -> MappedOperatorTaskGroup { @@ -143,7 +143,7 @@ MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( }; require_all_of( - get_parallel_layers(pcg), + pcg_get_parallel_layers(pcg), [&](parallel_layer_guid_t l) -> void { std::set for_layer = slot_names_for_layer(l); std::set for_layer_mapping = slot_names_for_layer_mapping(l); @@ -237,8 +237,8 @@ std::string mapped_pcg_as_dot(MappedParallelComputationGraph const &mpcg) { }; std::function( - std::unordered_set const &)> - order_slots = [](std::unordered_set const &slot_names) + std::set const &)> + order_slots = [](std::set const &slot_names) -> std::vector { return sorted(slot_names); }; return labelled_kwarg_dataflow_graph_view_as_dot(mpcg.raw_graph, diff --git a/lib/pcg/src/pcg/metric_attrs.cc b/lib/pcg/src/pcg/metric_attrs.cc index 5357775149..6e6ad6e0ee 100644 --- a/lib/pcg/src/pcg/metric_attrs.cc +++ b/lib/pcg/src/pcg/metric_attrs.cc @@ -3,7 +3,7 @@ namespace FlexFlow { MetricsAttrs::MetricsAttrs(LossFunction _loss_type, - std::unordered_set const &metrics) + std::set const &metrics) : loss_type(_loss_type), measure_accuracy(false), measure_categorical_crossentropy(false), measure_sparse_categorical_crossentropy(false), diff --git a/lib/pcg/src/pcg/optimizer_attrs.cc b/lib/pcg/src/pcg/optimizer_attrs.cc index 46192c2103..05c5dd95ac 100644 --- a/lib/pcg/src/pcg/optimizer_attrs.cc +++ b/lib/pcg/src/pcg/optimizer_attrs.cc @@ -23,18 +23,18 @@ OptimizerAttrs } } -std::unordered_set +std::set get_slot_names_for_optimizer(OptimizerAttrs const &attrs) { - return attrs.visit>(overload{ + return attrs.visit>(overload{ [](SGDOptimizerAttrs const &sgd_attrs) - -> std::unordered_set { + -> std::set { if (sgd_attrs.momentum > 0.0f) { return {OptimizerSlotName::SGD_V}; } else { return {}; } }, - [](AdamOptimizerAttrs const &) -> std::unordered_set { + [](AdamOptimizerAttrs const &) -> std::set { return { OptimizerSlotName::ADAM_M, OptimizerSlotName::ADAM_V, diff --git a/lib/pcg/src/pcg/parallel_computation_graph/generate_weight_transform.cc b/lib/pcg/src/pcg/parallel_computation_graph/generate_weight_transform.cc index 50cbea9ca0..e5519ca7ed 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/generate_weight_transform.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/generate_weight_transform.cc @@ -5,10 +5,10 @@ namespace FlexFlow { -std::unordered_set +std::set generate_weight_transform(TensorShape const ¤t, ParallelTensorShape const &goal) { - std::unordered_set result; + std::set result; positive_int sum_degree = get_sum_degree(goal); ASSERT(sum_degree == 1, diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index 40e8cb9e5e..1a27c88303 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -17,7 +17,7 @@ #include "utils/containers/get_only.h" #include "utils/containers/repeat_element.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/zip_values_strict_with.h" #include "utils/containers/zip_with_strict.h" #include "utils/graph/digraph/algorithms/get_initial_nodes.h" @@ -36,9 +36,8 @@ #include "utils/graph/node/algorithms.h" #include "utils/graph/node/node.dtg.h" #include "utils/record_formatter.h" -#include -#include "utils/containers/map_from_unordered.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -53,8 +52,8 @@ ParallelComputationGraph empty_parallel_computation_graph() { TensorSlotName>>()}; } -std::unordered_set - get_parallel_layers(ParallelComputationGraph const &pcg) { +std::set + pcg_get_parallel_layers(ParallelComputationGraph const &pcg) { return transform(get_nodes(pcg.raw_graph), [&](Node const &n) { return parallel_layer_guid_t{n}; }); } @@ -62,48 +61,48 @@ std::unordered_set ParallelLayerAddedResult add_parallel_layer( ParallelComputationGraph &pcg, ParallelLayerAttrs const &layer_attrs, - std::unordered_map const &inputs, - std::unordered_map const &weights, - std::optional> const + std::map const &inputs, + std::map const &weights, + std::optional> const &maybe_output_flags) { - std::unordered_map input_shapes = + std::map input_shapes = map_values(inputs, [&](parallel_tensor_guid_t const &i) { return get_parallel_tensor_shape(pcg, i); }); - std::unordered_map weight_shapes = + std::map weight_shapes = map_values(weights, [&](parallel_tensor_guid_t const &i) { return get_parallel_tensor_shape(pcg, i); }); - std::unordered_map + std::map correct_weight_shapes = get_weight_shapes(layer_attrs.op_attrs, input_shapes); ASSERT(weight_shapes == correct_weight_shapes, "add_parallel_layer received incorrect weight shapes"); - std::unordered_map output_shapes = + std::map output_shapes = get_output_shapes(layer_attrs.op_attrs, input_shapes); - std::unordered_map> + std::map> unwrapped_inputs = map_values(inputs, [](parallel_tensor_guid_t const &t) { return t.raw_graph_output; }); - std::unordered_map> + std::map> unwrapped_weights = map_values(weights, [](parallel_tensor_guid_t const &t) { return t.raw_graph_output; }); - std::unordered_map output_flags = + std::map output_flags = maybe_output_flags.value_or( - generate_unordered_map(unordered_keys(output_shapes), + generate_map(keys(output_shapes), [](TensorSlotName const &) { return CreateGrad::YES; })); - std::unordered_map output_attrs = + std::map output_attrs = zip_values_strict_with( output_shapes, output_flags, @@ -113,7 +112,7 @@ ParallelLayerAddedResult add_parallel_layer( KwargNodeAddedResult op_added = pcg.raw_graph.add_node( layer_attrs, - binary_merge_disjoint_unordered_maps(unwrapped_inputs, unwrapped_weights), + binary_merge_disjoint_maps(unwrapped_inputs, unwrapped_weights), output_attrs); return ParallelLayerAddedResult{ @@ -138,7 +137,7 @@ ParallelLayerAddedResult pcg_add_input_layer(ParallelComputationGraph &pcg, /*inputs=*/{}, /*weights=*/{}, /*output_flags=*/ - std::unordered_map{ + std::map{ { TensorSlotName::OUTPUT, create_grad, @@ -152,10 +151,10 @@ OperatorTaskSpace get_operator_task_space(ParallelComputationGraph const &pcg, ASSERT(!is_parallel_op(op_attrs)); - std::unordered_map inputs = + std::map inputs = get_incoming_inputs(pcg, layer); - std::unordered_map input_degrees = + std::map input_degrees = map_values(get_incoming_inputs(pcg, layer), [&](parallel_tensor_guid_t input_guid) { return get_parallel_degrees( @@ -169,7 +168,7 @@ OperatorTaskSpace get_operator_task_space(ParallelComputationGraph const &pcg, std::set pcg_get_invocation_info_set(ParallelComputationGraph const &pcg) { - return transform(set_of(get_parallel_layers(pcg)), + return transform(set_of(pcg_get_parallel_layers(pcg)), [&](parallel_layer_guid_t l) -> ParallelLayerInvocationInfo { return pcg_get_invocation_info_for_layer(pcg, l); }); @@ -182,10 +181,10 @@ ParallelLayerInvocationInfo ParallelLayerAttrs l_attrs = get_parallel_layer_attrs(pcg, l); std::map incoming = - map_from_unordered(get_incoming_tensors(pcg, l)); + get_incoming_tensors(pcg, l); std::map outgoing = - map_from_unordered(get_outgoing_tensors(pcg, l)); + get_outgoing_tensors(pcg, l); auto get_parallel_tensor_info = [&](parallel_tensor_guid_t t) -> ParallelTensorInfo { ParallelTensorAttrs t_attrs = get_parallel_tensor_attrs(pcg, t); @@ -206,7 +205,7 @@ ParallelLayerInvocationInfo }; } -std::unordered_set +std::set get_edges(ParallelComputationGraph const &pcg) { return transform(get_all_kwarg_dataflow_edges(pcg.raw_graph), [](KwargDataflowEdge const &e) { @@ -214,11 +213,11 @@ std::unordered_set }); } -std::unordered_set +std::set get_pcg_edges_from_layer_to_layer(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &src, parallel_layer_guid_t const &dst) { - std::unordered_set> raw_edges = + std::set> raw_edges = get_kwarg_dataflow_edges_from_node_to_node( pcg.raw_graph, src.raw_graph_node, dst.raw_graph_node); return transform(raw_edges, [](KwargDataflowEdge const &e) { @@ -226,11 +225,11 @@ std::unordered_set }); } -std::unordered_set +std::set get_outgoing_edges(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { - std::unordered_set> raw_edges = - unordered_set_of( + std::set> raw_edges = + set_of( get_outgoing_kwarg_dataflow_edges_for_node(pcg.raw_graph, l.raw_graph_node) .right_values()); @@ -239,10 +238,10 @@ std::unordered_set }); } -std::unordered_map +std::map get_incoming_edges(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { - std::unordered_map> + std::map> raw_edges = get_incoming_kwarg_dataflow_edges_for_node(pcg.raw_graph, l.raw_graph_node); return map_values(raw_edges, [](KwargDataflowEdge const &e) { @@ -250,10 +249,10 @@ std::unordered_map }); } -std::unordered_set +std::set pcg_get_parallel_tensor_uses(ParallelComputationGraph const &pcg, parallel_tensor_guid_t const &t) { - std::unordered_set> raw_uses = + std::set> raw_uses = get_kwarg_dataflow_value_uses(pcg.raw_graph, t.raw_graph_output); return transform(raw_uses, [](KwargDataflowInput const &i) { @@ -261,14 +260,14 @@ std::unordered_set }); } -std::unordered_set +std::set get_initial_layers(ParallelComputationGraph const &pcg) { - std::unordered_set raw_sources = get_initial_nodes(pcg.raw_graph); + std::set raw_sources = get_initial_nodes(pcg.raw_graph); return transform(raw_sources, [](Node const &n) { return parallel_layer_guid_t{n}; }); } -std::unordered_map +std::map get_outgoing_tensors(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { return map_values(get_outgoing_kwarg_dataflow_outputs_for_node( @@ -278,7 +277,7 @@ std::unordered_map }); } -std::unordered_map +std::map get_incoming_tensors(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { return map_values(get_incoming_kwarg_dataflow_outputs_for_node( @@ -288,7 +287,7 @@ std::unordered_map }); } -std::unordered_map +std::map pcg_get_operator_to_incoming_mappings(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { ComputationGraphOpAttrs op_attrs = @@ -299,7 +298,7 @@ std::unordered_map /*input_degrees=*/get_incoming_input_degrees(pcg, l)); } -std::unordered_map +std::map pcg_get_operator_to_output_mappings(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { ComputationGraphOpAttrs op_attrs = @@ -336,41 +335,41 @@ OperatorTaskSpaceToOperatorTaskSpaceMapping src_to_tensor_mapping, dst_to_tensor_mapping); } -static std::unordered_map +static std::map get_incoming_tensors_with_role(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l, IncomingTensorRole desired_role) { PCGOperatorAttrs attrs = get_parallel_layer_attrs(pcg, l).op_attrs; - std::unordered_map incoming_tensors = + std::map incoming_tensors = get_incoming_tensors(pcg, l); - std::unordered_map incoming_slot_roles = + std::map incoming_slot_roles = get_incoming_tensor_roles(attrs); ASSERT(incoming_tensors.size() == incoming_slot_roles.size()); - std::unordered_set slots_with_desired_role = - unordered_keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { + std::set slots_with_desired_role = + keys(filter_values(incoming_slot_roles, [&](IncomingTensorRole role) { return role == desired_role; })); return restrict_keys(incoming_tensors, slots_with_desired_role); } -std::unordered_map +std::map get_incoming_inputs(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { return get_incoming_tensors_with_role(pcg, l, IncomingTensorRole::INPUT); } -std::unordered_map +std::map get_incoming_weights(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { return get_incoming_tensors_with_role(pcg, l, IncomingTensorRole::WEIGHT); } -std::unordered_map +std::map get_incoming_input_degrees(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { @@ -379,22 +378,22 @@ std::unordered_map }); } -std::unordered_set +std::set get_successors(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { return transform(get_successors(pcg.raw_graph, l.raw_graph_node), [](Node const &n) { return parallel_layer_guid_t{n}; }); } -std::unordered_set get_subgraph_successors( +std::set get_subgraph_successors( ParallelComputationGraph const &pcg, - std::unordered_set const &subgraph_layers) { + std::set const &subgraph_layers) { - std::unordered_set raw_subgraph_nodes = + std::set raw_subgraph_nodes = transform(subgraph_layers, [](parallel_layer_guid_t const &l) { return l.raw_graph_node; }); - std::unordered_set raw_successors = + std::set raw_successors = get_subgraph_successors(pcg.raw_graph, raw_subgraph_nodes); return transform(raw_successors, @@ -434,11 +433,11 @@ std::vector [](Node const &n) { return parallel_layer_guid_t{n}; }); } -std::unordered_map +std::map get_parallel_layer_attrs_mapping(ParallelComputationGraph const &pcg) { - std::unordered_map + std::map layer_attrs_mapping; - for (parallel_layer_guid_t const &layer_guid : get_parallel_layers(pcg)) { + for (parallel_layer_guid_t const &layer_guid : pcg_get_parallel_layers(pcg)) { layer_attrs_mapping.insert( {layer_guid, get_parallel_layer_attrs(pcg, layer_guid)}); } @@ -448,8 +447,8 @@ std::unordered_map parallel_layer_guid_t get_parallel_layer_by_name(ParallelComputationGraph const &pcg, std::string const &name) { - std::unordered_set found = - filter(get_parallel_layers(pcg), [&](parallel_layer_guid_t const &l) { + std::set found = + filter(pcg_get_parallel_layers(pcg), [&](parallel_layer_guid_t const &l) { return get_parallel_layer_attrs(pcg, l).name == name; }); return get_only(found); @@ -514,8 +513,8 @@ std::string pcg_as_dot(ParallelComputationGraph const &cg) { }; std::function( - std::unordered_set const &)> - order_slots = [](std::unordered_set const &slot_names) + std::set const &)> + order_slots = [](std::set const &slot_names) -> std::vector { return sorted(slot_names); }; return labelled_kwarg_dataflow_graph_view_as_dot(cg.raw_graph, diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc index f19cb100bb..0943800fd4 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc @@ -33,7 +33,7 @@ #include "utils/containers/transform.h" #include "utils/containers/zip_values_strict_with.h" #include "utils/containers/zip_with.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -61,7 +61,7 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::create_input_tensor( layer_attrs, {}, {}, - std::unordered_map{ + std::map{ { TensorSlotName::OUTPUT, CreateGrad::NO, @@ -179,7 +179,7 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::conv2d( ParallelTensorShape input_shape = this->get_shape(input); - std::unordered_map initializers = + std::map initializers = get_initializers(attrs, get_reduced_shape(input_shape), maybe_kernel_initializer, @@ -220,7 +220,7 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::dense( ParallelTensorShape input_shape = this->get_shape(input); - std::unordered_map initializers = + std::map initializers = throw_if_unexpected(get_initializers(attrs, get_reduced_shape(input_shape), maybe_projection_initializer, @@ -258,7 +258,7 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::embedding( ParallelLayerAttrs layer = ParallelLayerAttrs{PCGOperatorAttrs{attrs}, name}; - std::unordered_map initializers = + std::map initializers = get_initializers(attrs, maybe_kernel_initializer); return require_only_key(this->add_layer(layer, @@ -308,7 +308,7 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::multihead_attention( ParallelLayerAttrs layer = ParallelLayerAttrs{PCGOperatorAttrs{attrs}, name}; - std::unordered_map initializers = + std::map initializers = throw_if_unexpected( get_initializers(attrs, get_reduced_shape(this->get_shape(query)), @@ -370,7 +370,7 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::batch_norm( std::vector weights; - std::unordered_map initializers = + std::map initializers = throw_if_unexpected(get_initializers(attrs)); return require_only_key(this->add_layer(layer, @@ -645,16 +645,16 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::add_weight( static void check_incoming_tensor_roles( ParallelLayerAttrs const &layer, - std::unordered_set const &input_slots, - std::unordered_set const &weight_slots) { - std::unordered_map correct = + std::set const &input_slots, + std::set const &weight_slots) { + std::map correct = get_incoming_tensor_roles(layer.op_attrs); - std::unordered_map current = - binary_merge_disjoint_unordered_maps( - generate_unordered_map( + std::map current = + binary_merge_disjoint_maps( + generate_map( input_slots, [](TensorSlotName) { return IncomingTensorRole::INPUT; }), - generate_unordered_map(weight_slots, [](TensorSlotName) { + generate_map(weight_slots, [](TensorSlotName) { return IncomingTensorRole::WEIGHT; })); @@ -662,25 +662,25 @@ static void check_incoming_tensor_roles( "check_incoming_tensor_roles found deviation in incoming tensors"); } -std::unordered_map +std::map ParallelComputationGraphBuilder::add_layer( ParallelLayerAttrs const &layer, - std::unordered_map const + std::map const &inputs, - std::unordered_map const + std::map const &weight_initializers) { - ASSERT(are_disjoint(unordered_keys(inputs), unordered_keys(weight_initializers))); - check_incoming_tensor_roles(layer, unordered_keys(inputs), unordered_keys(weight_initializers)); + ASSERT(are_disjoint(keys(inputs), keys(weight_initializers))); + check_incoming_tensor_roles(layer, keys(inputs), keys(weight_initializers)); - std::unordered_map input_shapes = + std::map input_shapes = map_values(inputs, [&](parallel_tensor_guid_t const &i) { return this->get_shape(i); }); - std::unordered_map weight_shapes = + std::map weight_shapes = get_weight_shapes(layer.op_attrs, input_shapes); - std::unordered_map weight_tensors = + std::map weight_tensors = zip_values_strict_with(weight_shapes, weight_initializers, [&](ParallelTensorShape const &weight_shape, diff --git a/lib/pcg/test/src/pcg/computation_graph.cc b/lib/pcg/test/src/pcg/computation_graph.cc index 721179b647..50b401a810 100644 --- a/lib/pcg/test/src/pcg/computation_graph.cc +++ b/lib/pcg/test/src/pcg/computation_graph.cc @@ -29,9 +29,9 @@ TEST_SUITE(FF_TEST_SUITE) { layer_guid_t input_layer = get_layer_by_name(cg, input_name); - std::unordered_map result = + std::map result = get_incoming_inputs(cg, input_layer); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } @@ -56,9 +56,9 @@ TEST_SUITE(FF_TEST_SUITE) { layer_guid_t layer = get_layer_by_name(cg, layer_name); - std::unordered_map result = + std::map result = get_incoming_inputs(cg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -95,9 +95,9 @@ TEST_SUITE(FF_TEST_SUITE) { layer_guid_t dense_layer = get_layer_by_name(cg, layer_name); - std::unordered_map result = + std::map result = get_incoming_inputs(cg, dense_layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -130,9 +130,9 @@ TEST_SUITE(FF_TEST_SUITE) { layer_guid_t input_layer = get_layer_by_name(cg, input_name); - std::unordered_map result = + std::map result = get_incoming_weights(cg, input_layer); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } @@ -159,9 +159,9 @@ TEST_SUITE(FF_TEST_SUITE) { layer_guid_t layer = get_layer_by_name(cg, layer_name); - std::unordered_map result = + std::map result = get_incoming_weights(cg, layer); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } @@ -239,9 +239,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, }); - std::unordered_map result = + std::map result = get_incoming_weights(cg, linear_added.layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::WEIGHT, t_projection_weight, diff --git a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index a53c67c336..d6618ca5a0 100644 --- a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -100,7 +100,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_map + std::map mapped_tasks = { { l_input1, diff --git a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index 97ef5fa46e..22c294dbce 100644 --- a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -94,9 +94,9 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelLayerAddedResult input_added = pcg_add_input_layer(pcg, input_shape); - std::unordered_map result = + std::map result = get_incoming_inputs(pcg, input_added.parallel_layer); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } @@ -161,9 +161,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, }); - std::unordered_map result = + std::map result = get_incoming_inputs(pcg, linear_added.parallel_layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, t_input, @@ -294,9 +294,9 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelLayerAddedResult input_added = pcg_add_input_layer(pcg, input_shape); - std::unordered_map result = + std::map result = get_incoming_weights(pcg, input_added.parallel_layer); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } @@ -319,9 +319,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*weights=*/{}); - std::unordered_map result = + std::map result = get_incoming_weights(pcg, relu_added.parallel_layer); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } @@ -413,9 +413,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, }); - std::unordered_map result = + std::map result = get_incoming_weights(pcg, linear_added.parallel_layer); - std::unordered_map correct = { + std::map correct = { {TensorSlotName::WEIGHT, t_replicated_projection_weight}, }; @@ -453,7 +453,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*inputs=*/{}, /*weights=*/{}, /*output_labels=*/ - std::unordered_map{ + std::map{ { TensorSlotName::OUTPUT, CreateGrad::NO, @@ -562,9 +562,9 @@ TEST_SUITE(FF_TEST_SUITE) { DimDomain layer_2_task_space = layer_1_task_space; - auto make_coord = [](nonnegative_int x) { + auto make_coord = [](nonnegative_int x) -> DimCoord { return DimCoord{ - std::unordered_map{ + std::map{ {operator_task_space_dim_idx_t{0_n}, x}, }, }; @@ -663,7 +663,7 @@ TEST_SUITE(FF_TEST_SUITE) { auto make_coord = [](nonnegative_int x) { return DimCoord{ - std::unordered_map{ + std::map{ {operator_task_space_dim_idx_t{0_n}, x}, }, }; diff --git a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc index 8d07de9ea1..662bac0a7c 100644 --- a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc +++ b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc @@ -5,7 +5,7 @@ #include "pcg/parallel_computation_graph/parallel_layer_attrs.h" #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.h" #include "utils/containers/count.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_only.h" #include "utils/containers/items.h" #include "utils/containers/require_only_key.h" @@ -51,9 +51,9 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(out); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::LHS_INPUT, lhs, @@ -68,9 +68,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, out, @@ -114,9 +114,9 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -127,9 +127,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -175,8 +175,8 @@ TEST_SUITE(FF_TEST_SUITE) { /*paddingH=*/paddingH, /*paddingW=*/paddingW); - std::unordered_map layers = - generate_unordered_map(get_parallel_layers(b.pcg), + std::map layers = + generate_map(pcg_get_parallel_layers(b.pcg), [&](parallel_layer_guid_t const &l) { return get_parallel_layer_attrs(b.pcg, l); }); @@ -237,7 +237,7 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelTensorShape correct_bias_shape = get_bias_shape(correct_attrs, par_input_shape); - std::unordered_map conv_incoming = + std::map conv_incoming = get_incoming_tensors(b.pcg, conv_guid); parallel_tensor_guid_t conv_input = conv_incoming.at(TensorSlotName::INPUT); @@ -256,7 +256,7 @@ TEST_SUITE(FF_TEST_SUITE) { get_parallel_tensor_attrs(b.pcg, conv_bias).shape; CHECK(conv_bias_shape == correct_bias_shape); - std::unordered_map conv_outputs = + std::map conv_outputs = get_outgoing_tensors(b.pcg, conv_guid); CHECK(conv_outputs.size() == 1); @@ -290,7 +290,7 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); CHECK(result.at(TensorSlotName::INPUT) == input); @@ -298,9 +298,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -332,7 +332,7 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); CHECK(result.size() == 2); @@ -340,9 +340,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -381,7 +381,7 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); CHECK(result.size() == 6); @@ -391,9 +391,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -422,9 +422,9 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -435,9 +435,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -470,9 +470,9 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -483,9 +483,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -516,9 +516,9 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -529,9 +529,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -560,9 +560,9 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -573,9 +573,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, @@ -609,9 +609,9 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_layer_guid_t layer = get_source_layer(output); SUBCASE("incoming") { - std::unordered_map result = + std::map result = get_incoming_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::INPUT, input, @@ -622,9 +622,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("outputs") { - std::unordered_map result = + std::map result = get_outgoing_tensors(b.pcg, layer); - std::unordered_map correct = { + std::map correct = { { TensorSlotName::OUTPUT, output, diff --git a/lib/realm-execution/include/realm-execution/dependency_set.h b/lib/realm-execution/include/realm-execution/dependency_set.h index ba8e9dc9b5..1afee5b915 100644 --- a/lib/realm-execution/include/realm-execution/dependency_set.h +++ b/lib/realm-execution/include/realm-execution/dependency_set.h @@ -4,7 +4,7 @@ #include "realm-execution/atomic_dependency_set.h" #include "realm-execution/realm.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" -#include +#include namespace FlexFlow { @@ -29,7 +29,7 @@ struct DependencySet { private: Realm::Event precondition; - std::unordered_map + std::map atomic_dependencies; }; diff --git a/lib/realm-execution/include/realm-execution/distributed_ff_handle.h b/lib/realm-execution/include/realm-execution/distributed_ff_handle.h index 8409a234a7..e2bc62717a 100644 --- a/lib/realm-execution/include/realm-execution/distributed_ff_handle.h +++ b/lib/realm-execution/include/realm-execution/distributed_ff_handle.h @@ -4,7 +4,7 @@ #include "realm-execution/device_specific_managed_per_device_ff_handle.h" #include "realm-execution/realm.h" #include "realm-execution/realm_context.h" -#include +#include namespace FlexFlow { @@ -16,7 +16,7 @@ struct DistributedFfHandle { public: DistributedFfHandle() = delete; explicit DistributedFfHandle( - std::unordered_map> const &handles); @@ -24,7 +24,7 @@ struct DistributedFfHandle { at(Realm::Processor processor) const; private: - std::unordered_map> handles; }; diff --git a/lib/realm-execution/include/realm-execution/instance_allocation.h b/lib/realm-execution/include/realm-execution/instance_allocation.h index 66cc07af75..330c265ca0 100644 --- a/lib/realm-execution/include/realm-execution/instance_allocation.h +++ b/lib/realm-execution/include/realm-execution/instance_allocation.h @@ -26,7 +26,7 @@ std::pair */ TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, - std::unordered_map const + std::map const &preallocated, RealmContext &ctx); diff --git a/lib/realm-execution/include/realm-execution/pcg_instance.h b/lib/realm-execution/include/realm-execution/pcg_instance.h index 7b86d6d383..1894de5416 100644 --- a/lib/realm-execution/include/realm-execution/pcg_instance.h +++ b/lib/realm-execution/include/realm-execution/pcg_instance.h @@ -83,7 +83,7 @@ PCGInstance create_pcg_instance( MappedParallelComputationGraph const &mpcg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::unordered_map const + std::map const &input_tensors, ProfilingSettings const &profiling_settings, DistributedFfHandle const &ff_handle); @@ -99,25 +99,25 @@ PCGInstance create_pcg_instance( * * \relates PCGInstance */ -std::unordered_map +std::map perform_all_passes_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, DistributedFfHandle const &ff_handle); -std::unordered_map +std::map perform_forward_pass_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, DistributedFfHandle const &ff_handle); -std::unordered_map +std::map perform_backward_pass_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, DistributedFfHandle const &ff_handle); -std::unordered_map +std::map perform_update_pass_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, diff --git a/lib/realm-execution/include/realm-execution/per_device_op_state_backing.dtg.toml b/lib/realm-execution/include/realm-execution/per_device_op_state_backing.dtg.toml index 92e9de8145..5a7e843f69 100644 --- a/lib/realm-execution/include/realm-execution/per_device_op_state_backing.dtg.toml +++ b/lib/realm-execution/include/realm-execution/per_device_op_state_backing.dtg.toml @@ -10,7 +10,7 @@ docstring = ''' includes = [ - "", + "", "realm-execution/device_specific_ptr.h", "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", "task-spec/per_device_op_state.dtg.h", @@ -18,4 +18,4 @@ includes = [ [[fields]] name = "backing" -type = "std::unordered_map<::FlexFlow::DynamicNodeInvocation, ::FlexFlow::DeviceSpecificPtr<::FlexFlow::PerDeviceOpState>>" +type = "std::map<::FlexFlow::DynamicNodeInvocation, ::FlexFlow::DeviceSpecificPtr<::FlexFlow::PerDeviceOpState>>" diff --git a/lib/realm-execution/include/realm-execution/realm_allocator.h b/lib/realm-execution/include/realm-execution/realm_allocator.h index 77af4a742c..9f234957e8 100644 --- a/lib/realm-execution/include/realm-execution/realm_allocator.h +++ b/lib/realm-execution/include/realm-execution/realm_allocator.h @@ -30,7 +30,7 @@ struct RealmAllocator : public IAllocator { private: Realm::Processor processor; Realm::Memory memory; - std::unordered_map ptr_instances; + std::map ptr_instances; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(RealmAllocator); diff --git a/lib/realm-execution/include/realm-execution/realm_context.h b/lib/realm-execution/include/realm-execution/realm_context.h index 5b76d52e2c..e9b33c95e1 100644 --- a/lib/realm-execution/include/realm-execution/realm_context.h +++ b/lib/realm-execution/include/realm-execution/realm_context.h @@ -12,7 +12,7 @@ #include "realm-execution/redops/redop_id_t.dtg.h" #include "realm-execution/tasks/task_id_t.dtg.h" #include -#include +#include namespace FlexFlow { @@ -128,7 +128,7 @@ struct RealmContext { Realm::Processor processor; Allocator allocator; std::vector outstanding_events; - std::unordered_map, + std::map, std::vector> processors; }; diff --git a/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_tensor_instance_backing.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_tensor_instance_backing.dtg.toml index 75a796b2ee..a40c6df5e0 100644 --- a/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_tensor_instance_backing.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_tensor_instance_backing.dtg.toml @@ -9,18 +9,18 @@ features = [ ] includes = [ - "", + "", "realm-execution/tasks/serializer/serializable_realm_event.dtg.h", "realm-execution/tasks/serializer/serializable_realm_instance.dtg.h", "task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.h", ] src_includes = [ - "utils/hash/unordered_map.h", + "utils/hash/map.h", "utils/fmt/pair.h", - "utils/fmt/unordered_map.h", + "utils/fmt/map.h", ] [[fields]] name = "backing" -type = "std::unordered_map<::FlexFlow::SerializableDynamicValueAttrs, std::pair<::FlexFlow::SerializableRealmInstance, ::FlexFlow::SerializableRealmEvent>>" +type = "std::map<::FlexFlow::SerializableDynamicValueAttrs, std::pair<::FlexFlow::SerializableRealmInstance, ::FlexFlow::SerializableRealmEvent>>" diff --git a/lib/realm-execution/include/realm-execution/tensor_instance_backing.dtg.toml b/lib/realm-execution/include/realm-execution/tensor_instance_backing.dtg.toml index 051edb0b9f..30a79a38e9 100644 --- a/lib/realm-execution/include/realm-execution/tensor_instance_backing.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tensor_instance_backing.dtg.toml @@ -14,7 +14,7 @@ and \ref destroy_instances, respectively. ''' includes = [ - "", + "", "realm-execution/realm.h", "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h", ] @@ -22,13 +22,13 @@ includes = [ src_includes = [ "realm-execution/fmt/realm_event.h", "realm-execution/fmt/realm_instance.h", - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "backing" -type = "std::unordered_map<::FlexFlow::DynamicValueAttrs, std::pair<::FlexFlow::Realm::RegionInstance, ::FlexFlow::Realm::Event>>" +type = "std::map<::FlexFlow::DynamicValueAttrs, std::pair<::FlexFlow::Realm::RegionInstance, ::FlexFlow::Realm::Event>>" docstring = ''' The need to track a pair of RegionInstance and Event, rather than just a RegionInstance, is due to the fact that unlike in Legion, RegionInstance does not natively encode its own ready state. This gives you have a choice: diff --git a/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc b/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc index 2fdc31d2c5..01c1bf909f 100644 --- a/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc +++ b/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc @@ -6,7 +6,7 @@ namespace FlexFlow { DistributedFfHandle::DistributedFfHandle( - std::unordered_map> const &handles) : handles(handles) {} @@ -21,7 +21,7 @@ DistributedFfHandle size_t workSpaceSize, bool allowTensorOpMathConversion, Realm::Event precondition) { - std::unordered_map> handles; diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index 3240027555..10d8647b29 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -11,7 +11,7 @@ #include "utils/containers/values.h" #include "utils/optional.h" #include -#include +#include #include namespace FlexFlow { @@ -28,7 +28,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( // Initialize all operators and save the per-device op state ASSERT(no_nodes_are_initialized(dg)); - std::unordered_map *> device_state_map; for (DynamicNodeInvocation const &invocation : dg.invocations) { @@ -73,7 +73,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( ctx.get_outstanding_events().wait(); auto deref = [](DeviceSpecificPtr *const &p) { return *p; }; - std::unordered_map> + std::map> result = map_values(device_state_map, deref); for (DeviceSpecificPtr *device_state_ptr : diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index f1d88f672c..5c3656458a 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -14,7 +14,7 @@ #include "utils/containers/contains_key.h" #include "utils/containers/make.h" #include "utils/containers/map_values.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/exception.h" #include "utils/optional.h" @@ -37,12 +37,12 @@ std::pair TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, - std::unordered_map const + std::map const &preallocated, RealmContext &ctx) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); - for (DynamicValueAttrs const &v : unordered_keys(preallocated)) { + for (DynamicValueAttrs const &v : keys(preallocated)) { ASSERT(v.accessor == std::nullopt); } diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index a9cdf394a8..6d44c9bbf4 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -83,7 +83,7 @@ PCGInstance create_pcg_instance( MappedParallelComputationGraph const &mpcg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::unordered_map const + std::map const &input_tensors, ProfilingSettings const &profiling_settings, DistributedFfHandle const &device_handle) { @@ -92,7 +92,7 @@ PCGInstance create_pcg_instance( make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg); dg = perform_pass_expansion(dg); - std::unordered_map inputs = + std::map inputs = input_tensors; std::optional logit_grad_value; if (loss.has_value()) { @@ -275,7 +275,7 @@ static Realm::Event spawn_dynamic_node_invocation( }); } -static std::unordered_map +static std::map execute_distributed_dynamic_node_invocation_set( RealmContext &ctx, std::vector const &invocations, @@ -287,7 +287,7 @@ static std::unordered_map // For simplicity we'll track a dependency on all outstanding operations up to // this point. This will create an effective barrier between phases. DependencySet dependency_set{ctx.get_outstanding_events()}; - return unordered_map_from_pairs( + return map_from_pairs( transform(invocations, [&](DynamicNodeInvocation const &invocation) { std::vector input_dependencies = transform(vector_of(values(invocation.inputs)), @@ -321,14 +321,14 @@ static std::unordered_map })); } -std::unordered_map +std::map perform_all_passes_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, DistributedFfHandle const &device_handle) { std::vector execution_order = pcg_instance.get_execution_order(); - std::unordered_map result = + std::map result = execute_distributed_dynamic_node_invocation_set( /*ctx=*/pcg_instance.get_realm_context(), /*invocations=*/execution_order, @@ -342,7 +342,7 @@ std::unordered_map return result; } -std::unordered_map +std::map perform_forward_pass_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, @@ -365,7 +365,7 @@ std::unordered_map /*device_handle=*/device_handle); } -std::unordered_map +std::map perform_backward_pass_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, @@ -388,7 +388,7 @@ std::unordered_map /*device_handle=*/device_handle); } -std::unordered_map +std::map perform_update_pass_for_pcg_instance( PCGInstance &pcg_instance, ProfilingSettings const &profiling_settings, @@ -401,7 +401,7 @@ std::unordered_map return task_type == DynamicTaskType::UPD; }); - std::unordered_map result = + std::map result = execute_distributed_dynamic_node_invocation_set( /*ctx=*/pcg_instance.get_realm_context(), /*invocations=*/execution_order, diff --git a/lib/realm-execution/test/src/realm-execution/test_e2e.cc b/lib/realm-execution/test/src/realm-execution/test_e2e.cc index 39eae4e1cb..0b4854f78b 100644 --- a/lib/realm-execution/test/src/realm-execution/test_e2e.cc +++ b/lib/realm-execution/test/src/realm-execution/test_e2e.cc @@ -204,7 +204,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*nesterov=*/false, /*weight_decay=*/0.001}}; - std::unordered_map + std::map input_tensors; DistributedFfHandle device_handle = @@ -433,7 +433,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { GenericTensorAccessorW label_tensor = allocator.allocate_tensor(label_tensor_shape); - std::unordered_map + std::map input_tensors; DistributedFfHandle device_handle = create_distributed_ff_handle( diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index 6efbb17eb3..fac4b86871 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -252,7 +252,7 @@ TEST_SUITE(FF_TEST_SUITE) { MappedParallelComputationGraph mpcg = make_test_mpcg_for_device_type(DeviceType::CPU); - std::unordered_map + std::map input_tensors; OptimizerAttrs optimizer_attrs = OptimizerAttrs{ @@ -276,8 +276,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*loss=*/std::nullopt, /*input_tensors=*/input_tensors, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*iteration_config=*/FFIterationConfig{1_p}); + /*device_handle=*/device_handle); // begin training loop int num_epochs = 1; @@ -285,8 +284,7 @@ TEST_SUITE(FF_TEST_SUITE) { perform_all_passes_for_pcg_instance( /*instance=*/pcg_instance, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*iteration_config=*/FFIterationConfig{1_p}); + /*device_handle=*/device_handle); } }); result.wait(); @@ -318,7 +316,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { }, }; - std::unordered_map + std::map input_tensors; DistributedFfHandle device_handle = create_distributed_ff_handle( @@ -333,8 +331,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { /*loss=*/std::nullopt, /*input_tensors=*/input_tensors, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*iteration_config=*/FFIterationConfig{1_p}); + /*device_handle=*/device_handle); // begin training loop int num_epochs = 1; @@ -342,8 +339,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { perform_all_passes_for_pcg_instance( /*instance=*/pcg_instance, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*iteration_config=*/FFIterationConfig{1_p}); + /*device_handle=*/device_handle); } }); result.wait(); diff --git a/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h b/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h index c3ebc0f77f..3276e704c7 100644 --- a/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h +++ b/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h @@ -33,7 +33,7 @@ LabelledOpenKwargDataflowGraphView const &g, - std::unordered_map, + std::map, ParallelTensorShape> const &input_shapes); } // namespace FlexFlow diff --git a/lib/substitutions/include/substitutions/operator_pattern/get_attribute_map.h b/lib/substitutions/include/substitutions/operator_pattern/get_attribute_map.h index 2b31dada04..bddeb4a5ca 100644 --- a/lib/substitutions/include/substitutions/operator_pattern/get_attribute_map.h +++ b/lib/substitutions/include/substitutions/operator_pattern/get_attribute_map.h @@ -7,7 +7,7 @@ namespace FlexFlow { -std::unordered_map +std::map get_attribute_map(PCGOperatorAttrs const &); } // namespace FlexFlow diff --git a/lib/substitutions/include/substitutions/operator_pattern/operator_attribute_pattern.dtg.toml b/lib/substitutions/include/substitutions/operator_pattern/operator_attribute_pattern.dtg.toml index 44dfd9e61d..2c607e47ef 100644 --- a/lib/substitutions/include/substitutions/operator_pattern/operator_attribute_pattern.dtg.toml +++ b/lib/substitutions/include/substitutions/operator_pattern/operator_attribute_pattern.dtg.toml @@ -11,12 +11,12 @@ features = [ ] includes = [ - "", - "utils/fmt/unordered_set.h", + "", + "utils/fmt/set.h", "substitutions/operator_pattern/operator_attribute_constraint.dtg.h", - "utils/hash/unordered_set.h", + "utils/hash/set.h", ] [[fields]] name = "attribute_constraints" -type = "std::unordered_set<::FlexFlow::OperatorAttributeConstraint>" +type = "std::set<::FlexFlow::OperatorAttributeConstraint>" diff --git a/lib/substitutions/include/substitutions/output_graph/materialize_operator_from_attrs_map.h b/lib/substitutions/include/substitutions/output_graph/materialize_operator_from_attrs_map.h index cc2fac4805..f67fc2965e 100644 --- a/lib/substitutions/include/substitutions/output_graph/materialize_operator_from_attrs_map.h +++ b/lib/substitutions/include/substitutions/output_graph/materialize_operator_from_attrs_map.h @@ -8,7 +8,7 @@ namespace FlexFlow { PCGOperatorAttrs materialize_operator_from_attrs_map( - std::unordered_map const &); + std::map const &); } // namespace FlexFlow diff --git a/lib/substitutions/include/substitutions/output_graph/output_graph_expr.h b/lib/substitutions/include/substitutions/output_graph/output_graph_expr.h index e5a897330f..63693fc54f 100644 --- a/lib/substitutions/include/substitutions/output_graph/output_graph_expr.h +++ b/lib/substitutions/include/substitutions/output_graph/output_graph_expr.h @@ -8,12 +8,12 @@ namespace FlexFlow { -std::unordered_set get_nodes(OutputGraphExpr const &); +std::set get_nodes(OutputGraphExpr const &); -std::unordered_map +std::map get_node_outputs(OutputGraphExpr const &, OutputGraphExprNode const &); -std::unordered_set get_inputs(OutputGraphExpr const &); +std::set get_inputs(OutputGraphExpr const &); } // namespace FlexFlow diff --git a/lib/substitutions/include/substitutions/output_graph/output_operator_attribute_expr.h b/lib/substitutions/include/substitutions/output_graph/output_operator_attribute_expr.h index cba095b444..df5b09163d 100644 --- a/lib/substitutions/include/substitutions/output_graph/output_operator_attribute_expr.h +++ b/lib/substitutions/include/substitutions/output_graph/output_operator_attribute_expr.h @@ -8,7 +8,7 @@ namespace FlexFlow { OperatorAttributeValue evaluate_output_operator_attribute_expr( OutputOperatorAttributeExpr const &, - std::unordered_map const &node_match); + std::map const &node_match); } // namespace FlexFlow diff --git a/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.dtg.toml b/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.dtg.toml index ca613dad91..753348c316 100644 --- a/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.dtg.toml +++ b/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.dtg.toml @@ -13,12 +13,12 @@ includes = [ "substitutions/operator_pattern/operator_attribute_key.dtg.h", "substitutions/output_graph/output_operator_attribute_expr.dtg.h", "substitutions/unlabelled/pattern_node.dtg.h", - "", + "", ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/fmt/unordered_map.h", + "utils/hash/map.h", + "utils/fmt/map.h", "utils/fmt/optional.h", ] @@ -30,4 +30,4 @@ type = "std::optional<::FlexFlow::PatternNode>" # define the assignment for each operator type. [[fields]] name = "assignments" -type = "std::unordered_map<::FlexFlow::OperatorAttributeKey, ::FlexFlow::OutputOperatorAttributeExpr>" +type = "std::map<::FlexFlow::OperatorAttributeKey, ::FlexFlow::OutputOperatorAttributeExpr>" diff --git a/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.h b/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.h index 0921569d62..b1928bf428 100644 --- a/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.h +++ b/lib/substitutions/include/substitutions/output_graph/output_operator_attrs_assignment.h @@ -11,7 +11,7 @@ OutputOperatorAttrsAssignment output_operator_clone_node(PatternNode const &); PCGOperatorAttrs materialize_output_operator_from_attrs_assignment( OutputOperatorAttrsAssignment const &attrs_assignment, - std::unordered_map const &node_match); + std::map const &node_match); std::pair copy_attr_from_pattern_node(OperatorAttributeKey key, diff --git a/lib/substitutions/include/substitutions/pcg_pattern.h b/lib/substitutions/include/substitutions/pcg_pattern.h index 8a4266fc5e..eafd2a10a9 100644 --- a/lib/substitutions/include/substitutions/pcg_pattern.h +++ b/lib/substitutions/include/substitutions/pcg_pattern.h @@ -10,7 +10,7 @@ namespace FlexFlow { -std::unordered_set get_nodes(PCGPattern const &); +std::set get_nodes(PCGPattern const &); std::optional get_random_pattern_match(PCGPattern const &pattern, @@ -29,8 +29,8 @@ TensorAttributePattern get_tensor_pattern(PCGPattern const &, PatternValue const &); OperatorAttributePattern get_operator_pattern(PCGPattern const &, PatternNode const &); -std::unordered_set get_inputs(PCGPattern const &); -std::unordered_map +std::set get_inputs(PCGPattern const &); +std::map get_pattern_node_outputs(PCGPattern const &, PatternNode const &); bool assignment_satisfies(SubParallelComputationGraph const &, diff --git a/lib/substitutions/include/substitutions/pcg_pattern_match.dtg.toml b/lib/substitutions/include/substitutions/pcg_pattern_match.dtg.toml index 5e10f5963c..a2fc5be589 100644 --- a/lib/substitutions/include/substitutions/pcg_pattern_match.dtg.toml +++ b/lib/substitutions/include/substitutions/pcg_pattern_match.dtg.toml @@ -3,6 +3,7 @@ name = "PCGPatternMatch" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -13,12 +14,12 @@ includes = [ "substitutions/unlabelled/pattern_input.dtg.h", "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h", "substitutions/open_parallel_tensor_guid_t.dtg.h", - "", + "", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] @@ -27,4 +28,4 @@ type = "::FlexFlow::bidict<::FlexFlow::PatternNode, ::FlexFlow::parallel_layer_g [[fields]] name = "input_assignment" -type = "std::unordered_map<::FlexFlow::PatternInput, ::FlexFlow::open_parallel_tensor_guid_t>" +type = "std::map<::FlexFlow::PatternInput, ::FlexFlow::open_parallel_tensor_guid_t>" diff --git a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h index 2a3dc8bbb8..178ba80bbf 100644 --- a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h +++ b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h @@ -13,9 +13,9 @@ namespace FlexFlow { -std::unordered_set - get_parallel_layers(SubParallelComputationGraph const &); -std::unordered_set +std::set + spcg_get_parallel_layers(SubParallelComputationGraph const &); +std::set get_parallel_tensors(SubParallelComputationGraph const &); ParallelLayerAttrs get_parallel_layer_attrs(SubParallelComputationGraph const &, parallel_layer_guid_t const &); @@ -33,21 +33,21 @@ parallel_layer_guid_t get_parallel_layer_by_name(SubParallelComputationGraph const &pcg, std::string const &name); -std::unordered_map +std::map get_layer_inputs(SubParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_map +std::map get_outgoing_tensors(SubParallelComputationGraph const &, parallel_layer_guid_t const &); -std::unordered_set get_subgraph_incoming_edges( +std::set get_subgraph_incoming_edges( SubParallelComputationGraph const &, - std::unordered_set const &); -std::unordered_set get_subgraph_outgoing_edges( + std::set const &); +std::set get_subgraph_outgoing_edges( SubParallelComputationGraph const &, - std::unordered_set const &); + std::set const &); -std::unordered_set +std::set get_open_parallel_tensor_uses(SubParallelComputationGraph const &, open_parallel_tensor_guid_t const &); diff --git a/lib/substitutions/include/substitutions/sub_parallel_computation_graph_data.dtg.toml b/lib/substitutions/include/substitutions/sub_parallel_computation_graph_data.dtg.toml index 8836d050b5..622f39be9f 100644 --- a/lib/substitutions/include/substitutions/sub_parallel_computation_graph_data.dtg.toml +++ b/lib/substitutions/include/substitutions/sub_parallel_computation_graph_data.dtg.toml @@ -14,29 +14,29 @@ includes = [ "substitutions/open_parallel_tensor_guid_t.dtg.h", "substitutions/input_parallel_tensor_guid_t.dtg.h", "substitutions/sub_parallel_computation_graph_edge.dtg.h", - "", - "", + "", + "", ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_map.h", - "utils/fmt/unordered_set.h", + "utils/hash/map.h", + "utils/hash/set.h", + "utils/fmt/map.h", + "utils/fmt/set.h", ] [[fields]] name = "node_data" -type = "std::unordered_map<::FlexFlow::parallel_layer_guid_t, ::FlexFlow::ParallelLayerAttrs>" +type = "std::map<::FlexFlow::parallel_layer_guid_t, ::FlexFlow::ParallelLayerAttrs>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::SubParallelComputationGraphEdge>" +type = "std::set<::FlexFlow::SubParallelComputationGraphEdge>" [[fields]] name = "inputs" -type = "std::unordered_set<::FlexFlow::input_parallel_tensor_guid_t>" +type = "std::set<::FlexFlow::input_parallel_tensor_guid_t>" [[fields]] name = "value_data" -type = "std::unordered_map<::FlexFlow::open_parallel_tensor_guid_t, ::FlexFlow::ParallelTensorAttrs>" +type = "std::map<::FlexFlow::open_parallel_tensor_guid_t, ::FlexFlow::ParallelTensorAttrs>" diff --git a/lib/substitutions/include/substitutions/substitution_builder.h b/lib/substitutions/include/substitutions/substitution_builder.h index 248c23ecc8..4f180f6eb4 100644 --- a/lib/substitutions/include/substitutions/substitution_builder.h +++ b/lib/substitutions/include/substitutions/substitution_builder.h @@ -17,19 +17,19 @@ struct SubstitutionBuilder { std::optional const &name = std::nullopt); void equate_outputs(PatternValue const &, OutputGraphExprValue const &); - std::unordered_map add_pattern_node( + std::map add_pattern_node( OperatorAttributePattern const &node_pattern, - std::unordered_map const &inputs, - std::unordered_map const + std::map const &inputs, + std::map const &output_patterns, std::optional const &name = std::nullopt); - std::unordered_map + std::map add_output_graph_node( OutputOperatorAttrsAssignment const &node_expr, - std::unordered_map const + std::map const &inputs, - std::unordered_set const &output_slots); + std::set const &output_slots); PatternNode pattern_node_named(std::string const &) const; PatternInput pattern_input_named(std::string const &) const; diff --git a/lib/substitutions/include/substitutions/tensor_pattern/tensor_attribute_pattern.dtg.toml b/lib/substitutions/include/substitutions/tensor_pattern/tensor_attribute_pattern.dtg.toml index c81d28c8a0..b8f06a0304 100644 --- a/lib/substitutions/include/substitutions/tensor_pattern/tensor_attribute_pattern.dtg.toml +++ b/lib/substitutions/include/substitutions/tensor_pattern/tensor_attribute_pattern.dtg.toml @@ -11,12 +11,12 @@ features = [ ] includes = [ - "", + "", "substitutions/tensor_pattern/tensor_attribute_constraint.dtg.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "attribute_constraints" -type = "std::unordered_set<::FlexFlow::TensorAttributeConstraint>" +type = "std::set<::FlexFlow::TensorAttributeConstraint>" diff --git a/lib/substitutions/include/substitutions/unlabelled/pattern_edge.h b/lib/substitutions/include/substitutions/unlabelled/pattern_edge.h index 13c6e36bc8..a07f72f33b 100644 --- a/lib/substitutions/include/substitutions/unlabelled/pattern_edge.h +++ b/lib/substitutions/include/substitutions/unlabelled/pattern_edge.h @@ -7,13 +7,13 @@ #include "substitutions/unlabelled/standard_pattern_edge.dtg.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge.dtg.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge.dtg.h" -#include +#include namespace FlexFlow { PatternNode get_dst_node(PatternEdge const &); -std::unordered_set get_nodes(PatternEdge const &); +std::set get_nodes(PatternEdge const &); bool is_input_edge(PatternEdge const &); bool is_standard_edge(PatternEdge const &); diff --git a/lib/substitutions/include/substitutions/unlabelled/pattern_split.dtg.toml b/lib/substitutions/include/substitutions/unlabelled/pattern_split.dtg.toml index 3358f320e6..3be30c071d 100644 --- a/lib/substitutions/include/substitutions/unlabelled/pattern_split.dtg.toml +++ b/lib/substitutions/include/substitutions/unlabelled/pattern_split.dtg.toml @@ -10,16 +10,16 @@ features = [ ] includes = [ - "", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "", + "utils/hash/set.h", + "utils/fmt/set.h", "substitutions/unlabelled/pattern_node.dtg.h", ] [[fields]] name = "first" -type = "std::unordered_set<::FlexFlow::PatternNode>" +type = "std::set<::FlexFlow::PatternNode>" [[fields]] name = "second" -type = "std::unordered_set<::FlexFlow::PatternNode>" +type = "std::set<::FlexFlow::PatternNode>" diff --git a/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h b/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h index 716714e2d9..06a76f320f 100644 --- a/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h +++ b/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h @@ -12,29 +12,29 @@ namespace FlexFlow { size_t num_nodes(UnlabelledGraphPattern const &); bool is_singleton_pattern(UnlabelledGraphPattern const &); -std::unordered_set +std::set get_pattern_nodes(UnlabelledGraphPattern const &); -std::unordered_set +std::set get_pattern_values(UnlabelledGraphPattern const &); std::vector get_topological_ordering(UnlabelledGraphPattern const &); -std::unordered_set +std::set get_pattern_inputs(UnlabelledGraphPattern const &); -std::unordered_set +std::set get_pattern_edges(UnlabelledGraphPattern const &); -std::unordered_map +std::map get_inputs_to_pattern_node(UnlabelledGraphPattern const &, PatternNode const &); -std::unordered_map +std::map get_outputs_from_pattern_node(UnlabelledGraphPattern const &, PatternNode const &); UnlabelledGraphPatternSubgraphResult get_pattern_subgraph(UnlabelledGraphPattern const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.dtg.toml b/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.dtg.toml index c7f3b20394..3c5c21eb8e 100644 --- a/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.dtg.toml +++ b/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.dtg.toml @@ -14,13 +14,13 @@ includes = [ "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_value.dtg.h", "substitutions/unlabelled/pattern_input.dtg.h", "substitutions/unlabelled/pattern_node.dtg.h", - "", + "", "op-attrs/tensor_slot_name.dtg.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] @@ -29,4 +29,4 @@ type = "::FlexFlow::bidict<::FlexFlow::PatternNode, ::FlexFlow::Node>" [[fields]] name = "input_assignment" -type = "std::unordered_map<::FlexFlow::PatternInput, ::FlexFlow::OpenKwargDataflowValue>" +type = "std::map<::FlexFlow::PatternInput, ::FlexFlow::OpenKwargDataflowValue>" diff --git a/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h b/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h index d175747852..6311a5a40e 100644 --- a/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h +++ b/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h @@ -6,12 +6,12 @@ #include "substitutions/unlabelled/pattern_value.dtg.h" #include "substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.dtg.h" #include -#include +#include namespace FlexFlow { UnlabelledKwargDataflowGraphPatternMatch empty_unlabelled_pattern_match(); -std::unordered_set +std::set matched_nodes(UnlabelledKwargDataflowGraphPatternMatch const &); std::optional merge_unlabelled_dataflow_graph_pattern_matches( @@ -22,7 +22,7 @@ std::optional bidict const &merged_graph_values_to_inputs_of_2); -std::unordered_map, PatternValue> +std::map, PatternValue> get_output_assignment(SubParallelComputationGraph const &, PCGPattern const &, UnlabelledKwargDataflowGraphPatternMatch const &); diff --git a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc index b8140440b7..6699870669 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc @@ -9,11 +9,11 @@ #include "substitutions/sub_parallel_computation_graph_data.dtg.h" #include "substitutions/sub_parallel_computation_graph_data.h" #include "substitutions/sub_parallel_computation_graph_edge.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" #include "utils/containers/values.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -49,30 +49,30 @@ SubParallelComputationGraph apply_substitution_from_output_result( SubParallelComputationGraphData pre_data = get_sub_pcg_data(spcg); require_sub_parallel_computation_graph_data_is_valid(pre_data); - std::unordered_set pre_nodes = - unordered_keys(pre_data.node_data); - std::unordered_set matched_nodes = - unordered_set_of(values(match.node_assignment)); - std::unordered_set post_nodes_from_original_graph = + std::set pre_nodes = + keys(pre_data.node_data); + std::set matched_nodes = + set_of(values(match.node_assignment)); + std::set post_nodes_from_original_graph = set_minus(pre_nodes, matched_nodes); - std::unordered_map post_node_data = + std::map post_node_data = [&] { - std::unordered_map + std::map post_node_data_from_orig = restrict_keys( pre_data.node_data, post_nodes_from_original_graph); - std::unordered_map + std::map post_node_data_from_sub = output_graph_data.node_data; - return binary_merge_disjoint_unordered_maps(post_node_data_from_orig, + return binary_merge_disjoint_maps(post_node_data_from_orig, post_node_data_from_sub); }(); - std::unordered_set post_inputs = + std::set post_inputs = pre_data.inputs; - std::unordered_set post_edges = [&] { - std::unordered_set post_edges_from_orig = + std::set post_edges = [&] { + std::set post_edges_from_orig = filter(pre_data.edges, [&](SubParallelComputationGraphEdge const &e) { if (e.raw_edge.is_input_edge()) { return true; @@ -86,7 +86,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( } }); - std::unordered_set post_edges_from_sub = + std::set post_edges_from_sub = filter(output_graph_data.edges, [&](SubParallelComputationGraphEdge const &e) { return e.raw_edge.is_internal_edge(); @@ -101,7 +101,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( sub.output_graph_expr, substitution_output_graph); - std::unordered_set incoming_to_sub_edges; + std::set incoming_to_sub_edges; for (auto const &[pattern_input, base_graph_tensor] : match.input_assignment) { OutputGraphExprInput output_expr_input = @@ -109,7 +109,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( input_parallel_tensor_guid_t output_graph_input = output_expr_to_result_sub_pcg_mapping.input_mapping.at_r( output_expr_input); - std::unordered_set uses = + std::set uses = get_open_parallel_tensor_uses( substitution_output_graph, open_parallel_tensor_guid_from_input(output_graph_input)); @@ -120,7 +120,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( } } - std::unordered_set outgoing_from_sub_edges; + std::set outgoing_from_sub_edges; for (ParallelComputationGraphEdge const &outgoing_edge : get_subgraph_outgoing_edges(spcg, matched_nodes)) { parallel_tensor_guid_t original_tensor = @@ -148,9 +148,9 @@ SubParallelComputationGraph apply_substitution_from_output_result( }); }(); - std::unordered_map + std::map post_value_data = [&] { - std::unordered_map + std::map post_value_data_from_orig = filter_keys( pre_data.value_data, [&](open_parallel_tensor_guid_t const &t) { return visit_open_parallel_tensor_guid( @@ -166,9 +166,9 @@ SubParallelComputationGraph apply_substitution_from_output_result( }); }); - std::unordered_map + std::map post_value_data_from_sub = output_graph_data.value_data; - return binary_merge_disjoint_unordered_maps(post_value_data_from_orig, + return binary_merge_disjoint_maps(post_value_data_from_orig, post_value_data_from_sub); }(); diff --git a/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc b/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc index 57a93daefc..28bfac0f69 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc @@ -24,8 +24,8 @@ std::pair evaluate_substitution_output(SubParallelComputationGraph const &spcg, Substitution const &sub, PCGPatternMatch const &match) { - std::unordered_map node_match = - map_values(match.node_assignment.as_unordered_map(), + std::map node_match = + map_values(match.node_assignment.as_map(), [&](parallel_layer_guid_t const &n) { return get_operator_attrs(spcg, n); }); @@ -86,7 +86,7 @@ std::pair [](Node const &n) { return OutputGraphExprNode{n}; }), [](NewNode const &n) { return parallel_layer_guid_t{n.raw_node}; }); - std::unordered_map, ParallelTensorShape> + std::map, ParallelTensorShape> input_shapes = map_values( map_keys(match.input_assignment, [&](PatternInput const &i) { diff --git a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc index 4374a951f8..11aca1bdad 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc @@ -16,9 +16,9 @@ bidict bidict result; for (auto const &[parallel_layer, output_graph_expr_node] : m.node_mapping) { - std::unordered_map layer_outputs = + std::map layer_outputs = get_outgoing_tensors(spcg, parallel_layer); - std::unordered_map + std::map output_graph_expr_outputs = get_node_outputs(output_graph_expr, output_graph_expr_node); diff --git a/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc b/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc index d3ad4ca246..62a67e1082 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/perform_shape_inference.cc @@ -19,7 +19,7 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.h" #include "utils/nonnegative_int/num_elements.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -32,10 +32,10 @@ LabelledOpenKwargDataflowGraphView const &g, - std::unordered_map, + std::map, ParallelTensorShape> const &input_shapes) { - std::unordered_map, + std::map, ParallelTensorShape> inferred = map_keys(input_shapes, @@ -45,7 +45,7 @@ LabelledOpenKwargDataflowGraphView incoming_shapes = + std::map incoming_shapes = map_values(get_incoming_open_kwarg_dataflow_values_for_node(g, n), [&](OpenKwargDataflowValue const &v) { return inferred.at(v); @@ -53,38 +53,38 @@ LabelledOpenKwargDataflowGraphView + std::map incoming_tensor_roles = get_incoming_tensor_roles(n_attrs.op_attrs); - ASSERT(is_subseteq_of(unordered_keys(incoming_shapes), unordered_keys(incoming_tensor_roles))); + ASSERT(is_subseteq_of(keys(incoming_shapes), keys(incoming_tensor_roles))); auto incoming_shapes_with_role = [&](IncomingTensorRole role) - -> std::unordered_map { - std::unordered_set slots_with_desired_role = - unordered_keys(filter_values(incoming_tensor_roles, + -> std::map { + std::set slots_with_desired_role = + keys(filter_values(incoming_tensor_roles, [&](IncomingTensorRole r) { return r == role; })); return restrict_keys(incoming_shapes, slots_with_desired_role); }; - std::unordered_map input_shapes = + std::map input_shapes = incoming_shapes_with_role(IncomingTensorRole::INPUT); - std::unordered_map weight_shapes = + std::map weight_shapes = incoming_shapes_with_role(IncomingTensorRole::WEIGHT); - ASSERT(binary_merge_disjoint_unordered_maps(input_shapes, weight_shapes) == + ASSERT(binary_merge_disjoint_maps(input_shapes, weight_shapes) == incoming_shapes); - std::unordered_map + std::map inferred_weight_shapes = get_weight_shapes(n_attrs.op_attrs, input_shapes); ASSERT(weight_shapes == inferred_weight_shapes); - std::unordered_map output_shapes = + std::map output_shapes = get_output_shapes(n_attrs.op_attrs, input_shapes); - std::unordered_map> + std::map> outputs = get_outgoing_kwarg_dataflow_outputs_for_node(g, n); for (auto const &[output, shape] : diff --git a/lib/substitutions/src/substitutions/operator_pattern/get_attribute_map.cc b/lib/substitutions/src/substitutions/operator_pattern/get_attribute_map.cc index f1b7440aed..b69aaf142e 100644 --- a/lib/substitutions/src/substitutions/operator_pattern/get_attribute_map.cc +++ b/lib/substitutions/src/substitutions/operator_pattern/get_attribute_map.cc @@ -6,9 +6,9 @@ namespace FlexFlow { -std::unordered_map +std::map get_attribute_map(PCGOperatorAttrs const &op_attrs) { - std::unordered_map result; + std::map result; for (OperatorAttributeKey const &attr_key : all_operator_attribute_keys()) { std::optional attr_value = diff --git a/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc b/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc index ce5094190c..529b6d908b 100644 --- a/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc +++ b/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc @@ -1,16 +1,16 @@ #include "substitutions/output_graph/materialize_operator_from_attrs_map.h" #include "utils/containers/contains_key.h" -#include "utils/fmt/unordered_map.h" +#include "utils/fmt/map.h" #include namespace FlexFlow { struct Accessor { Accessor( - std::unordered_map const &m) + std::map const &m) : m(m) {} - std::unordered_map const &m; + std::map const &m; template T get(OperatorAttributeKey k) const { @@ -29,7 +29,7 @@ struct Accessor { }; PCGOperatorAttrs materialize_operator_from_attrs_map( - std::unordered_map const + std::map const &attrs) { OperatorType op_type = attrs.at(OperatorAttributeKey::OP_TYPE).get(); diff --git a/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc b/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc index 3bc1d04abc..1a2a1034a9 100644 --- a/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc +++ b/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc @@ -7,16 +7,16 @@ namespace FlexFlow { -std::unordered_set get_nodes(OutputGraphExpr const &g) { - std::unordered_set raw_nodes = get_nodes(g.raw_graph); +std::set get_nodes(OutputGraphExpr const &g) { + std::set raw_nodes = get_nodes(g.raw_graph); return transform(raw_nodes, [](Node const &n) { return OutputGraphExprNode{n}; }); } -std::unordered_map +std::map get_node_outputs(OutputGraphExpr const &g, OutputGraphExprNode const &n) { - std::unordered_map> + std::map> raw_outputs = get_outgoing_kwarg_dataflow_outputs_for_node( g.raw_graph, n.raw_graph_node); @@ -26,8 +26,8 @@ std::unordered_map }); } -std::unordered_set get_inputs(OutputGraphExpr const &g) { - std::unordered_set> raw_inputs = +std::set get_inputs(OutputGraphExpr const &g) { + std::set> raw_inputs = get_all_kwarg_dataflow_graph_inputs(g.raw_graph); return transform(raw_inputs, [](KwargDataflowGraphInput const &i) { diff --git a/lib/substitutions/src/substitutions/output_graph/output_operator_attribute_expr.cc b/lib/substitutions/src/substitutions/output_graph/output_operator_attribute_expr.cc index e7cfcf232c..00b0329230 100644 --- a/lib/substitutions/src/substitutions/output_graph/output_operator_attribute_expr.cc +++ b/lib/substitutions/src/substitutions/output_graph/output_operator_attribute_expr.cc @@ -6,7 +6,7 @@ namespace FlexFlow { OperatorAttributeValue evaluate_output_operator_attribute_expr( OutputOperatorAttributeExpr const &expr, - std::unordered_map const &node_match) { + std::map const &node_match) { return expr.visit(overload{ [&](OutputOperatorAttrAccess const &a) { return evaluate_attribute_expr(a.attr_expr, node_match.at(a.node)) diff --git a/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc b/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc index 2755716e44..9298c5b35c 100644 --- a/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc +++ b/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc @@ -4,7 +4,7 @@ #include "substitutions/output_graph/output_operator_attribute_expr.h" #include "utils/containers/map_values.h" #include "utils/exception.h" -#include "utils/containers/binary_merge_unordered_maps_with_right_dominating.h" +#include "utils/containers/binary_merge_maps_with_right_dominating.h" namespace FlexFlow { @@ -14,11 +14,11 @@ OutputOperatorAttrsAssignment output_operator_clone_node(PatternNode const &) { PCGOperatorAttrs materialize_output_operator_from_attrs_assignment( OutputOperatorAttrsAssignment const &attrs_assignment, - std::unordered_map const &node_match) { + std::map const &node_match) { - std::unordered_map + std::map template_attrs_map = [&]() - -> std::unordered_map { + -> std::map { if (attrs_assignment.template_operator.has_value()) { PatternNode template_node = attrs_assignment.template_operator.value(); PCGOperatorAttrs template_op_attrs = node_match.at(template_node); @@ -28,15 +28,15 @@ PCGOperatorAttrs materialize_output_operator_from_attrs_assignment( } }(); - std::unordered_map + std::map assignments_attrs_map = map_values( attrs_assignment.assignments, [&](OutputOperatorAttributeExpr const &expr) { return evaluate_output_operator_attribute_expr(expr, node_match); }); - std::unordered_map - joined_attrs_map = binary_merge_unordered_maps_with_right_dominating( + std::map + joined_attrs_map = binary_merge_maps_with_right_dominating( template_attrs_map, assignments_attrs_map); return materialize_operator_from_attrs_map(joined_attrs_map); diff --git a/lib/substitutions/src/substitutions/pcg_pattern.cc b/lib/substitutions/src/substitutions/pcg_pattern.cc index b578383352..d55bc2f9b5 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern.cc @@ -15,8 +15,8 @@ namespace FlexFlow { -std::unordered_set get_nodes(PCGPattern const &p) { - std::unordered_set raw_nodes = get_nodes(p.raw_graph); +std::set get_nodes(PCGPattern const &p) { + std::set raw_nodes = get_nodes(p.raw_graph); return transform(raw_nodes, [](Node const &n) { return PatternNode{n}; }); } @@ -88,8 +88,8 @@ OperatorAttributePattern get_operator_pattern(PCGPattern const &p, return p.raw_graph.at(n.raw_node); } -std::unordered_set get_inputs(PCGPattern const &p) { - std::unordered_set> raw_inputs = +std::set get_inputs(PCGPattern const &p) { + std::set> raw_inputs = get_all_kwarg_dataflow_graph_inputs(p.raw_graph); return transform(raw_inputs, [](KwargDataflowGraphInput const &i) { @@ -97,10 +97,10 @@ std::unordered_set get_inputs(PCGPattern const &p) { }); } -std::unordered_map +std::map get_pattern_node_outputs(PCGPattern const &pattern, PatternNode const &node) { - std::unordered_map> + std::map> raw_outputs = get_outgoing_kwarg_dataflow_outputs_for_node( pattern.raw_graph, node.raw_node); diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index 8a71fe2ad5..9f4b207b30 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -57,26 +57,26 @@ void assert_pcg_pattern_match_is_valid_for_pattern_and_subpcg( PCGPatternMatch const &match, PCGPattern const &pattern, SubParallelComputationGraph const &spcg) { - std::unordered_set spcg_nodes = - get_parallel_layers(spcg); - std::unordered_set match_nodes = + std::set spcg_nodes = + spcg_get_parallel_layers(spcg); + std::set match_nodes = match.node_assignment.right_values(); ASSERT(is_subseteq_of(match_nodes, spcg_nodes)); - std::unordered_set spcg_values = + std::set spcg_values = get_parallel_tensors(spcg); - std::unordered_set match_values = - unordered_set_of(values(match.input_assignment)); + std::set match_values = + set_of(values(match.input_assignment)); ASSERT(is_subseteq_of(match_values, spcg_values)); - std::unordered_set pattern_nodes = get_nodes(pattern); - std::unordered_set match_pattern_nodes = + std::set pattern_nodes = get_nodes(pattern); + std::set match_pattern_nodes = match.node_assignment.left_values(); ASSERT(match_pattern_nodes == pattern_nodes); - std::unordered_set pattern_inputs = get_inputs(pattern); - std::unordered_set match_pattern_inputs = - unordered_keys(match.input_assignment); + std::set pattern_inputs = get_inputs(pattern); + std::set match_pattern_inputs = + keys(match.input_assignment); ASSERT(pattern_inputs == match_pattern_inputs); } diff --git a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc index c0c05ad5b1..427ac6747d 100644 --- a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc +++ b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc @@ -17,13 +17,13 @@ namespace FlexFlow { -std::unordered_set - get_parallel_layers(SubParallelComputationGraph const &sub_pcg) { +std::set + spcg_get_parallel_layers(SubParallelComputationGraph const &sub_pcg) { return transform(get_nodes(sub_pcg.raw_graph), [](Node const &n) { return parallel_layer_guid_t{n}; }); } -std::unordered_set +std::set get_parallel_tensors(SubParallelComputationGraph const &sub_pcg) { return transform(get_all_open_kwarg_dataflow_values(sub_pcg.raw_graph), [](OpenKwargDataflowValue const &v) @@ -80,7 +80,7 @@ parallel_layer_guid_t name); } -std::unordered_map +std::map get_layer_inputs(SubParallelComputationGraph const &pcg, parallel_layer_guid_t const &layer) { return map_values(get_incoming_open_kwarg_dataflow_values_for_node( @@ -90,7 +90,7 @@ std::unordered_map }); } -std::unordered_map +std::map get_outgoing_tensors(SubParallelComputationGraph const &pcg, parallel_layer_guid_t const &layer) { return map_values(get_outgoing_kwarg_dataflow_outputs_for_node( @@ -100,10 +100,10 @@ std::unordered_map }); } -std::unordered_set get_subgraph_outgoing_edges( +std::set get_subgraph_outgoing_edges( SubParallelComputationGraph const &spcg, - std::unordered_set const &layers) { - std::unordered_set> raw_edges = + std::set const &layers) { + std::set> raw_edges = get_kwarg_dataflow_subgraph_outgoing_edges( spcg.raw_graph, transform(layers, [](parallel_layer_guid_t const &l) { return l.raw_graph_node; @@ -113,14 +113,14 @@ std::unordered_set get_subgraph_outgoing_edges( }); } -std::unordered_set get_subgraph_incoming_edges( +std::set get_subgraph_incoming_edges( SubParallelComputationGraph const &spcg, - std::unordered_set const &subgraph) { - std::unordered_set raw_subgraph = + std::set const &subgraph) { + std::set raw_subgraph = transform(subgraph, [](parallel_layer_guid_t const &l) { return l.raw_graph_node; }); - std::unordered_set> + std::set> raw_incoming_edges = get_open_kwarg_dataflow_subgraph_incoming_edges( spcg.raw_graph, raw_subgraph); @@ -130,10 +130,10 @@ std::unordered_set get_subgraph_incoming_edges( }); } -std::unordered_set +std::set get_open_parallel_tensor_uses(SubParallelComputationGraph const &spcg, open_parallel_tensor_guid_t const &t) { - std::unordered_set> raw_uses = + std::set> raw_uses = get_open_kwarg_dataflow_value_uses(spcg.raw_graph, t.raw_open_dataflow_value); return transform(raw_uses, [](KwargDataflowInput const &i) { diff --git a/lib/substitutions/src/substitutions/substitution.cc b/lib/substitutions/src/substitutions/substitution.cc index 9bcfd054a7..d19a2329e8 100644 --- a/lib/substitutions/src/substitutions/substitution.cc +++ b/lib/substitutions/src/substitutions/substitution.cc @@ -36,7 +36,7 @@ bool is_isomorphic_to(Substitution const &l, Substitution const &r) { [&](OutputOperatorAttrsAssignment const &r_attrs) { std::optional l_template_operator = transform(r_attrs.template_operator, l_from_r_pattern_node); - std::unordered_map + std::map l_assignments = map_values( r_attrs.assignments, [&](OutputOperatorAttributeExpr const &r_expr) { @@ -132,9 +132,9 @@ bool is_isomorphic_to(Substitution const &l, Substitution const &r) { bool is_valid_substitution(Substitution const &sub) { { - std::unordered_set pattern_inputs = + std::set pattern_inputs = get_inputs(sub.pcg_pattern); - std::unordered_set mapped_inputs = + std::set mapped_inputs = left_entries(sub.inputs_mapping); if (pattern_inputs != mapped_inputs) { @@ -143,9 +143,9 @@ bool is_valid_substitution(Substitution const &sub) { } { - std::unordered_set output_graph_inputs = + std::set output_graph_inputs = get_inputs(sub.output_graph_expr); - std::unordered_set mapped_inputs = + std::set mapped_inputs = right_entries(sub.inputs_mapping); if (output_graph_inputs != mapped_inputs) { diff --git a/lib/substitutions/src/substitutions/substitution_builder.cc b/lib/substitutions/src/substitutions/substitution_builder.cc index ffda2291b5..0bfcac3c08 100644 --- a/lib/substitutions/src/substitutions/substitution_builder.cc +++ b/lib/substitutions/src/substitutions/substitution_builder.cc @@ -54,11 +54,11 @@ std::pair SubstitutionBuilder::add_input( }; } -std::unordered_map +std::map SubstitutionBuilder::add_pattern_node( OperatorAttributePattern const &node_pattern, - std::unordered_map const &inputs, - std::unordered_map const + std::map const &inputs, + std::map const &output_patterns, std::optional const &maybe_name) { KwargNodeAddedResult node_added = this->pattern_g.add_node( @@ -85,16 +85,16 @@ std::unordered_map }); } -std::unordered_map +std::map SubstitutionBuilder::add_output_graph_node( OutputOperatorAttrsAssignment const &node_expr, - std::unordered_map const &inputs, - std::unordered_set const &output_slots) { + std::map const &inputs, + std::set const &output_slots) { KwargNodeAddedResult node_added = this->output_g.add_node( node_expr, map_values(inputs, raw_open_kwarg_dataflow_value_from_output_graph_expr_value), - generate_unordered_map(output_slots, + generate_map(output_slots, [](TensorSlotName) { return std::monostate{}; })); return map_values( diff --git a/lib/substitutions/src/substitutions/unity_substitution_set.cc b/lib/substitutions/src/substitutions/unity_substitution_set.cc index 25d714e825..551d1101a6 100644 --- a/lib/substitutions/src/substitutions/unity_substitution_set.cc +++ b/lib/substitutions/src/substitutions/unity_substitution_set.cc @@ -75,7 +75,7 @@ std::vector static PatternValue insert_single_output_pattern( SubstitutionBuilder &b, OperatorAttributePattern const &attribute_pattern, - std::unordered_map const &inputs, + std::map const &inputs, TensorAttributePattern const &output_pattern, std::string const &name) { return require_only_key(b.add_pattern_node(attribute_pattern, @@ -94,7 +94,7 @@ static PatternValue insert_single_output_pattern( static OutputGraphExprValue insert_single_output_op( SubstitutionBuilder &b, OutputOperatorAttrsAssignment const &expr, - std::unordered_map const &inputs) { + std::map const &inputs) { return require_only_key( b.add_output_graph_node(expr, inputs, {TensorSlotName::OUTPUT}), TensorSlotName::OUTPUT); @@ -190,7 +190,7 @@ Substitution create_replicate_linear_combine(positive_int num_dims, auto [p_input, o_input] = b.add_input(tensor_attribute_pattern_match_all()); auto [p_weight, o_weight] = b.add_input(tensor_attribute_pattern_match_all()); - std::unordered_map p_inputs = { + std::map p_inputs = { {TensorSlotName::INPUT, p_input}, {TensorSlotName::WEIGHT, p_weight}, }; @@ -227,7 +227,7 @@ Substitution create_replicate_linear_combine(positive_int num_dims, OutputGraphExprValue o_partition_weights_output = insert_partition(b, degree, ff_dim_t{1_n}, o_weight); - std::unordered_map o_linear_inputs = { + std::map o_linear_inputs = { { TensorSlotName::INPUT, o_replicate_input_output, @@ -273,7 +273,7 @@ Substitution create_partition_linear_combine(positive_int num_dims, auto [p_input, o_input] = b.add_input(tensor_attribute_pattern_match_all()); auto [p_weight, o_weight] = b.add_input(tensor_attribute_pattern_match_all()); - std::unordered_map p_inputs = { + std::map p_inputs = { { TensorSlotName::INPUT, p_input, @@ -316,7 +316,7 @@ Substitution create_partition_linear_combine(positive_int num_dims, OutputGraphExprValue o_replicate_weights_output = insert_replicate(b, degree, o_weight); - std::unordered_map o_linear_inputs = { + std::map o_linear_inputs = { { TensorSlotName::INPUT, o_partition_input_output, @@ -364,7 +364,7 @@ Substitution create_partition_conv2d_combine(positive_int num_dims, auto [p_input, o_input] = b.add_input(tensor_attribute_pattern_match_all()); auto [p_weight, o_weight] = b.add_input(tensor_attribute_pattern_match_all()); - std::unordered_map p_inputs = { + std::map p_inputs = { { TensorSlotName::INPUT, p_input, @@ -394,7 +394,7 @@ Substitution create_partition_conv2d_combine(positive_int num_dims, OutputGraphExprValue o_replicate_weights_output = insert_replicate(b, degree, o_weight); - std::unordered_map o_conv2d_inputs = { + std::map o_conv2d_inputs = { { TensorSlotName::INPUT, o_partition_input_output, @@ -430,7 +430,7 @@ Substitution create_partition_attention_combine(positive_int num_heads, b.add_input(tensor_attribute_pattern_match_all()); auto [p_weights, o_weights] = b.add_input(tensor_attribute_pattern_match_all()); - std::unordered_map p_inputs = { + std::map p_inputs = { { TensorSlotName::QUERY, p_query_input, @@ -475,7 +475,7 @@ Substitution create_partition_attention_combine(positive_int num_heads, OutputGraphExprValue o_replicate_weight_output = insert_replicate(b, degree, o_weights); - std::unordered_map o_attention_inputs = + std::map o_attention_inputs = { { TensorSlotName::QUERY, @@ -524,7 +524,7 @@ Substitution create_replicate_attention_reduce(positive_int num_heads, auto [p_weights, o_weights] = b.add_input(tensor_attribute_pattern_match_all()); - std::unordered_map p_inputs = { + std::map p_inputs = { { TensorSlotName::QUERY, p_query_input, @@ -569,7 +569,7 @@ Substitution create_replicate_attention_reduce(positive_int num_heads, OutputGraphExprValue o_partition_weight_output = insert_partition(b, degree, ff_dim_t{1_n}, o_weights); - std::unordered_map o_attention_inputs = + std::map o_attention_inputs = { { TensorSlotName::QUERY, @@ -612,7 +612,7 @@ Substitution create_partition_softmax_combine(ff_dim_t softmax_dim, SubstitutionBuilder b; auto [p_input, o_input] = b.add_input(tensor_attribute_pattern_match_all()); - std::unordered_map p_inputs = { + std::map p_inputs = { { TensorSlotName::INPUT, p_input, @@ -637,7 +637,7 @@ Substitution create_partition_softmax_combine(ff_dim_t softmax_dim, OutputGraphExprValue o_partition_input_output = insert_partition(b, degree, partition_dim, o_input); - std::unordered_map o_softmax_inputs = { + std::map o_softmax_inputs = { { TensorSlotName::INPUT, o_partition_input_output, @@ -666,7 +666,7 @@ Substitution create_partition_add_combine(ff_dim_t parallel_dim, auto [p_input1, o_input1] = b.add_input(tensor_attribute_pattern_match_all()); auto [p_input2, o_input2] = b.add_input(tensor_attribute_pattern_match_all()); - std::unordered_map p_inputs = { + std::map p_inputs = { { TensorSlotName::LHS_INPUT, p_input1, @@ -695,7 +695,7 @@ Substitution create_partition_add_combine(ff_dim_t parallel_dim, OutputGraphExprValue o_partition_input2_output = insert_partition(b, degree, parallel_dim, o_input2); - std::unordered_map o_add_inputs = { + std::map o_add_inputs = { { TensorSlotName::LHS_INPUT, o_partition_input1_output, diff --git a/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc b/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc index a982277d22..3d3c2f8211 100644 --- a/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc +++ b/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc @@ -33,9 +33,9 @@ static std::optional empty_unlabelled_pattern_match(); match.node_assignment.equate(pattern_node, graph_node); - std::unordered_map pattern_outputs = + std::map pattern_outputs = get_outputs_from_pattern_node(pattern, pattern_node); - std::unordered_map> graph_outputs = map_values( get_outgoing_kwarg_dataflow_outputs_for_node(graph, graph_node), @@ -43,25 +43,25 @@ static std::optional return OpenKwargDataflowValue{o}; }); - if (unordered_keys(pattern_outputs) != unordered_keys(graph_outputs)) { + if (keys(pattern_outputs) != keys(graph_outputs)) { return std::nullopt; } - std::unordered_map pattern_node_inputs = + std::map pattern_node_inputs = get_inputs_to_pattern_node(pattern, pattern_node); - std::unordered_set pattern_graph_inputs = + std::set pattern_graph_inputs = get_pattern_inputs(pattern); - ASSERT(unordered_set_of(values(pattern_node_inputs)) == + ASSERT(set_of(values(pattern_node_inputs)) == transform(pattern_graph_inputs, [](PatternInput const &i) { return PatternValue{i}; })); - std::unordered_map> graph_node_inputs = get_incoming_open_kwarg_dataflow_values_for_node(graph, graph_node); - if (unordered_keys(graph_node_inputs) != unordered_keys(pattern_node_inputs)) { + if (keys(graph_node_inputs) != keys(pattern_node_inputs)) { return std::nullopt; } diff --git a/lib/substitutions/src/substitutions/unlabelled/pattern_edge.cc b/lib/substitutions/src/substitutions/unlabelled/pattern_edge.cc index f70d20c7a3..f1bc7ed83e 100644 --- a/lib/substitutions/src/substitutions/unlabelled/pattern_edge.cc +++ b/lib/substitutions/src/substitutions/unlabelled/pattern_edge.cc @@ -6,13 +6,13 @@ namespace FlexFlow { -std::unordered_set get_nodes(PatternEdge const &e) { - return e.visit>(overload{ +std::set get_nodes(PatternEdge const &e) { + return e.visit>(overload{ [](InputPatternEdge const &ee) { - return std::unordered_set{get_dst_node(ee)}; + return std::set{get_dst_node(ee)}; }, [](StandardPatternEdge const &ee) { - return std::unordered_set{ + return std::set{ get_src_node(ee), get_dst_node(ee), }; diff --git a/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc b/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc index 25d505b1fa..931ec43af3 100644 --- a/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc +++ b/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc @@ -26,7 +26,7 @@ namespace FlexFlow { OpenKwargDataflowSubgraphResult subgraph_matched(OpenKwargDataflowGraphView const &g, UnlabelledKwargDataflowGraphPatternMatch const &match) { - std::unordered_set matched_nodes = right_entries(match.node_assignment); + std::set matched_nodes = right_entries(match.node_assignment); return get_open_kwarg_dataflow_graph_subgraph( g, matched_nodes, make_counter_func()); } @@ -124,8 +124,8 @@ bool pattern_matches_subgraph_under( SubgraphConcreteFromPattern concrete_from_pattern{ match, full_graph_values_to_subgraph_inputs}; - std::unordered_set concrete_nodes = get_nodes(subgraph); - std::unordered_set concrete_nodes_from_match = + std::set concrete_nodes = get_nodes(subgraph); + std::set concrete_nodes_from_match = transform(get_pattern_nodes(pattern), concrete_from_pattern); if (concrete_nodes != concrete_nodes_from_match) { @@ -139,9 +139,9 @@ bool pattern_matches_subgraph_under( } } - std::unordered_set> + std::set> concrete_edges = get_all_open_kwarg_dataflow_edges(subgraph); - std::unordered_set> + std::set> concrete_edge_from_match = transform(get_pattern_edges(pattern), [&](PatternEdge const &e) @@ -153,9 +153,9 @@ bool pattern_matches_subgraph_under( return false; } - std::unordered_set> + std::set> concrete_values = get_all_open_kwarg_dataflow_values(subgraph); - std::unordered_set> + std::set> concrete_values_from_match = transform(get_pattern_values(pattern), [&](PatternValue const &v) @@ -183,14 +183,14 @@ bool unlabelled_pattern_does_match( UnlabelledKwargDataflowGraphPatternMatch const &match, MatchAdditionalCriterion const &additional_criterion) { - std::unordered_set> + std::set> matched_by_pattern_inputs = - unordered_set_of(values(match.input_assignment)); + set_of(values(match.input_assignment)); ASSERT(left_entries(match.node_assignment) == get_pattern_nodes(pattern)); ASSERT( is_subseteq_of(right_entries(match.node_assignment), get_nodes(graph))); - ASSERT(unordered_keys(match.input_assignment) == get_pattern_inputs(pattern)); + ASSERT(keys(match.input_assignment) == get_pattern_inputs(pattern)); ASSERT(is_subseteq_of(matched_by_pattern_inputs, get_all_open_kwarg_dataflow_values(graph))); @@ -199,7 +199,7 @@ bool unlabelled_pattern_does_match( OpenKwargDataflowGraphView matched_subgraph = subgraph_result.graph; - std::unordered_set> + std::set> full_values_split_by_subgraph = left_entries(subgraph_result.full_graph_values_to_subgraph_inputs); diff --git a/lib/substitutions/src/substitutions/unlabelled/pattern_split.cc b/lib/substitutions/src/substitutions/unlabelled/pattern_split.cc index 2c46944c99..48e4574b08 100644 --- a/lib/substitutions/src/substitutions/unlabelled/pattern_split.cc +++ b/lib/substitutions/src/substitutions/unlabelled/pattern_split.cc @@ -13,8 +13,8 @@ PatternSplit find_even_split(UnlabelledGraphPattern const &pattern) { int split_point = topological_ordering.size() / 2; auto split = vector_split(topological_ordering, split_point); - std::unordered_set prefix = unordered_set_of(split.first); - std::unordered_set postfix = unordered_set_of(split.second); + std::set prefix = set_of(split.first); + std::set postfix = set_of(split.second); return PatternSplit{prefix, postfix}; } diff --git a/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc b/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc index 432d41ef1d..1b6741ef9e 100644 --- a/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc +++ b/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc @@ -23,26 +23,26 @@ bool is_singleton_pattern(UnlabelledGraphPattern const &pattern) { return num_nodes(pattern) == 1; } -std::unordered_set +std::set get_pattern_nodes(UnlabelledGraphPattern const &p) { return transform(get_nodes(p.raw_graph), [](Node const &n) { return PatternNode{n}; }); } -std::unordered_set +std::set get_pattern_values(UnlabelledGraphPattern const &p) { return transform(get_all_open_kwarg_dataflow_values(p.raw_graph), pattern_value_from_raw_open_kwarg_dataflow_value); } -std::unordered_set +std::set get_pattern_inputs(UnlabelledGraphPattern const &p) { return transform( get_all_kwarg_dataflow_graph_inputs(p.raw_graph), [](KwargDataflowGraphInput const &i) { return PatternInput{i}; }); } -std::unordered_set +std::set get_pattern_edges(UnlabelledGraphPattern const &p) { return transform(get_all_open_kwarg_dataflow_edges(p.raw_graph), pattern_edge_from_raw_open_dataflow_edge); @@ -54,7 +54,7 @@ std::vector [](Node const &n) { return PatternNode{n}; }); } -std::unordered_map +std::map get_inputs_to_pattern_node(UnlabelledGraphPattern const &p, PatternNode const &n) { return map_values( @@ -64,7 +64,7 @@ std::unordered_map }); } -std::unordered_map +std::map get_outputs_from_pattern_node(UnlabelledGraphPattern const &p, PatternNode const &n) { return map_values( @@ -77,7 +77,7 @@ std::unordered_map UnlabelledGraphPatternSubgraphResult get_pattern_subgraph(UnlabelledGraphPattern const &p, - std::unordered_set const &n) { + std::set const &n) { OpenKwargDataflowSubgraphResult raw_result = get_open_kwarg_dataflow_graph_subgraph( p.raw_graph, diff --git a/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc b/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc index 8252f5ef02..ce46faa2f0 100644 --- a/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc +++ b/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc @@ -2,7 +2,7 @@ #include "utils/bidict/try_merge_nondisjoint_bidicts.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/map_keys.h" -#include "utils/containers/try_merge_nondisjoint_unordered_maps.h" +#include "utils/containers/try_merge_nondisjoint_maps.h" namespace FlexFlow { @@ -31,24 +31,24 @@ std::optional result.value(); }); - std::unordered_map> + std::map> merged_input_assignment = ({ - std::unordered_map> lifted_input_assignment_1 = map_keys( subpattern_1.input_assignment, [&](PatternInput const &pi1) { return merged_graph_values_to_inputs_of_1.at_r(pi1); }); - std::unordered_map> lifted_input_assignment_2 = map_keys( subpattern_2.input_assignment, [&](PatternInput const &pi2) { return merged_graph_values_to_inputs_of_2.at_r(pi2); }); std::optional< - std::unordered_map>> - merged = try_merge_nondisjoint_unordered_maps( + merged = try_merge_nondisjoint_maps( lifted_input_assignment_1, lifted_input_assignment_2); if (!merged.has_value()) { return std::nullopt; diff --git a/lib/substitutions/test/src/substitutions/apply_substitution/apply_substitution.cc b/lib/substitutions/test/src/substitutions/apply_substitution/apply_substitution.cc index 89b8eb820e..67bcb6ab6e 100644 --- a/lib/substitutions/test/src/substitutions/apply_substitution/apply_substitution.cc +++ b/lib/substitutions/test/src/substitutions/apply_substitution/apply_substitution.cc @@ -173,7 +173,7 @@ TEST_SUITE(FF_TEST_SUITE) { {b.pattern_node_named("mm"), mm_match_layer}, {b.pattern_node_named("relu"), relu_match_layer}, }, - std::unordered_map{ + std::map{ { b.pattern_input_named("input"), mm_match_layer_input_activations, diff --git a/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc b/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc index cf70b0cf53..dbdb5cb5ed 100644 --- a/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc +++ b/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc @@ -233,7 +233,7 @@ TEST_SUITE(FF_TEST_SUITE) { {pattern_mm_node, mm_match_layer}, {pattern_relu_node, relu_match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{pattern_i_activation}, mm_match_layer_input_activations, @@ -285,14 +285,14 @@ TEST_SUITE(FF_TEST_SUITE) { SubParallelComputationGraphData correct_graph_data = SubParallelComputationGraphData{ - std::unordered_map{{ + std::map{{ result_fused_mm_relu_node, ParallelLayerAttrs{ PCGOperatorAttrs{correct_result_fused_mm_relu_attrs}, /*name=*/std::nullopt, }, }}, - std::unordered_set{ + std::set{ SubParallelComputationGraphEdge{ OpenKwargDataflowEdge{ KwargDataflowInputEdge{ @@ -316,11 +316,11 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, }, - std::unordered_set{ + std::set{ result_i_activation, result_i_weights, }, - std::unordered_map{ { open_parallel_tensor_guid_from_input(result_i_activation), diff --git a/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc b/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc index 1c0f46bb3f..46efd88cc9 100644 --- a/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc +++ b/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc @@ -174,7 +174,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargDataflowOutput o2 = require_only_key(n2_added_result.outputs, TensorSlotName::OUTPUT); - std::unordered_map, ParallelTensorShape> + std::map, ParallelTensorShape> input_shapes = { {i0, i0_shape}, }; diff --git a/lib/substitutions/test/src/substitutions/pcg_pattern.cc b/lib/substitutions/test/src/substitutions/pcg_pattern.cc index b36a5f1d82..ad0e96c143 100644 --- a/lib/substitutions/test/src/substitutions/pcg_pattern.cc +++ b/lib/substitutions/test/src/substitutions/pcg_pattern.cc @@ -65,12 +65,12 @@ TEST_SUITE(FF_TEST_SUITE) { get_parallel_layer_by_name(pcg, x_matmul_name); parallel_layer_guid_t y_matmul = get_parallel_layer_by_name(pcg, y_matmul_name); - std::unordered_map x_incoming = + std::map x_incoming = get_incoming_tensors(pcg, x_matmul); REQUIRE(x_incoming.size() == 2); parallel_tensor_guid_t x_weights = x_incoming.at(TensorSlotName::WEIGHT); - std::unordered_map y_incoming = + std::map y_incoming = get_incoming_tensors(pcg, y_matmul); REQUIRE(y_incoming.size() == 2); parallel_tensor_guid_t y_weights = y_incoming.at(TensorSlotName::WEIGHT); @@ -166,7 +166,7 @@ TEST_SUITE(FF_TEST_SUITE) { PCGPattern pattern = PCGPattern{g}; - std::unordered_set result = unordered_set_of( + std::set result = set_of( find_pattern_matches(pattern, sub_pcg_from_full_pcg(pcg))); PCGPatternMatch match1 = PCGPatternMatch{ @@ -197,7 +197,7 @@ TEST_SUITE(FF_TEST_SUITE) { open_parallel_tensor_guid_from_closed(x_weights)}, }}; - std::unordered_set correct = {match1, match2}; + std::set correct = {match1, match2}; CHECK(result == correct); } @@ -350,7 +350,7 @@ TEST_SUITE(FF_TEST_SUITE) { PCGPattern pattern = PCGPattern{g}; - std::unordered_set result = unordered_set_of( + std::set result = set_of( find_pattern_matches(pattern, sub_pcg_from_full_pcg(pcg))); CHECK(result.size() == 3); diff --git a/lib/substitutions/test/src/substitutions/substitution_builder.cc b/lib/substitutions/test/src/substitutions/substitution_builder.cc index 10e08b09e7..a1c503b1c0 100644 --- a/lib/substitutions/test/src/substitutions/substitution_builder.cc +++ b/lib/substitutions/test/src/substitutions/substitution_builder.cc @@ -25,7 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { OperatorAttributeValue{std::optional{std::nullopt}}), }}; - std::unordered_map + std::map fused_mm_relu_attr_assignments = { set_attr_to_constant(OperatorAttributeKey::ACTIVATION, OperatorAttributeValue{Activation::RELU}), diff --git a/lib/substitutions/test/src/substitutions/unity_substitution_set.cc b/lib/substitutions/test/src/substitutions/unity_substitution_set.cc index df7f28538e..301c1363de 100644 --- a/lib/substitutions/test/src/substitutions/unity_substitution_set.cc +++ b/lib/substitutions/test/src/substitutions/unity_substitution_set.cc @@ -41,9 +41,9 @@ parallel_tensor_guid_t parallel_tensor_guid_t add_single_output_layer( ParallelComputationGraph &pcg, ParallelLayerAttrs const &layer_attrs, - std::unordered_map const &inputs, - std::unordered_map const &weights, - std::optional> const + std::map const &inputs, + std::map const &weights, + std::optional> const &outputs = std::nullopt) { return get_single_output( @@ -141,7 +141,7 @@ parallel_tensor_guid_t add_linear_layer( ASSERT(t_bias.has_value() == linear_attrs.use_bias); - std::unordered_map weights = { + std::map weights = { {TensorSlotName::WEIGHT, t_weight}, }; @@ -184,7 +184,7 @@ parallel_tensor_guid_t add_conv2d_layer( ASSERT(bias.has_value() == conv2d_attrs.use_bias); - std::unordered_map weights = { + std::map weights = { {TensorSlotName::FILTER, t_filter}, }; @@ -282,7 +282,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, match_layer_input_activations, @@ -390,7 +390,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, match_layer_input_activations, @@ -502,7 +502,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, match_layer_input_activations, @@ -612,7 +612,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, match_layer_input_activations, @@ -737,7 +737,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, match_layer_input_activations, @@ -853,7 +853,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, match_layer_query, @@ -974,7 +974,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, match_layer_query, @@ -1069,7 +1069,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{{ + std::map{{ PatternInput{KwargDataflowGraphInput{0}}, match_layer_input, }}, @@ -1158,7 +1158,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, add_match_layer_lhs, @@ -1245,7 +1245,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict{ {PatternNode{Node{0}}, match_layer}, }, - std::unordered_map{{ + std::map{{ PatternInput{KwargDataflowGraphInput{0}}, match_layer_input, }}, @@ -1324,7 +1324,7 @@ TEST_SUITE(FF_TEST_SUITE) { {PatternNode{Node{0}}, mm_match_layer}, {PatternNode{Node{1}}, relu_match_layer}, }, - std::unordered_map{ + std::map{ { PatternInput{KwargDataflowGraphInput{0}}, mm_match_layer_input_activations, diff --git a/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc b/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc index 9cb16f2923..60f7d2929a 100644 --- a/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc +++ b/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc @@ -136,7 +136,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict>{}}; - std::unordered_map> n1_incoming = { { @@ -152,20 +152,20 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_incoming_edges") { SUBCASE("n0") { - std::unordered_map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n0); - std::unordered_map> correct = {}; CHECK(result == correct); } SUBCASE("n1") { - std::unordered_map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n1); - std::unordered_map> correct = n1_incoming; CHECK(result == correct); @@ -173,9 +173,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_open_kwarg_dataflow_subgraph_inputs") { - std::unordered_set> result = + std::set> result = get_open_kwarg_dataflow_subgraph_inputs(graph, {n0, n1}); - std::unordered_set> correct = + std::set> correct = {}; CHECK(result == correct); } @@ -188,20 +188,20 @@ TEST_SUITE(FF_TEST_SUITE) { .graph; SUBCASE("nodes") { - std::unordered_set result = get_nodes(g); - std::unordered_set correct = {n0, n1}; + std::set result = get_nodes(g); + std::set correct = {n0, n1}; CHECK(result == correct); } SUBCASE("inputs") { - std::unordered_set> result = + std::set> result = g.get_inputs(); - std::unordered_set> correct = {}; + std::set> correct = {}; CHECK(result == correct); } SUBCASE("get_all_open_kwarg_dataflow_values") { - std::unordered_set> values = + std::set> values = get_all_open_kwarg_dataflow_values(g); CHECK(values.size() == 2); } @@ -210,8 +210,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("subgraph_matched") { OpenKwargDataflowGraphView result = subgraph_matched(graph, match).graph; - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = {n0, n1}; + std::set result_nodes = get_nodes(result); + std::set correct_nodes = {n0, n1}; CHECK(result_nodes == correct_nodes); } diff --git a/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc b/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc index ef4650c8d5..4866389a20 100644 --- a/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc +++ b/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc @@ -171,7 +171,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict>{}}; - std::unordered_map> n1_incoming = { { @@ -187,20 +187,20 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_incoming_open_kwarg_dataflow_edges_for_node") { SUBCASE("n0") { - std::unordered_map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n0); - std::unordered_map> correct = {}; CHECK(result == correct); } SUBCASE("n1") { - std::unordered_map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n1); - std::unordered_map> correct = n1_incoming; CHECK(result == correct); @@ -208,9 +208,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_open_kwarg_dataflow_subgraph_inputs") { - std::unordered_set> result = + std::set> result = get_open_kwarg_dataflow_subgraph_inputs(graph, {n0, n1}); - std::unordered_set> correct = + std::set> correct = {}; CHECK(result == correct); } @@ -222,20 +222,20 @@ TEST_SUITE(FF_TEST_SUITE) { .graph; SUBCASE("nodes") { - std::unordered_set result = get_nodes(g); - std::unordered_set correct = {n0, n1}; + std::set result = get_nodes(g); + std::set correct = {n0, n1}; CHECK(result == correct); } SUBCASE("inputs") { - std::unordered_set> result = + std::set> result = g.get_inputs(); - std::unordered_set> correct = {}; + std::set> correct = {}; CHECK(result == correct); } SUBCASE("get_all_open_kwarg_dataflow_values") { - std::unordered_set> values = + std::set> values = get_all_open_kwarg_dataflow_values(g); CHECK(values.size() == 2); } @@ -244,8 +244,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("subgraph_matched") { OpenKwargDataflowGraphView result = subgraph_matched(graph, match).graph; - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = {n0, n1}; + std::set result_nodes = get_nodes(result); + std::set correct_nodes = {n0, n1}; CHECK(result_nodes == correct_nodes); } diff --git a/lib/substitutions/test/src/substitutions/unlabelled/pattern_split.cc b/lib/substitutions/test/src/substitutions/unlabelled/pattern_split.cc index 458ab8a811..42456cb89a 100644 --- a/lib/substitutions/test/src/substitutions/unlabelled/pattern_split.cc +++ b/lib/substitutions/test/src/substitutions/unlabelled/pattern_split.cc @@ -35,8 +35,8 @@ TEST_SUITE(FF_TEST_SUITE) { PatternValue pv1 = pattern_value_from_raw_open_kwarg_dataflow_value(v1); PatternSplit even_split = PatternSplit{ - std::unordered_set{p0}, - std::unordered_set{p1}, + std::set{p0}, + std::set{p1}, }; SUBCASE("find_even_split") { @@ -48,15 +48,15 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("apply_split") { PatternSplitResult split_result = apply_split(pattern, even_split); SUBCASE("subpattern_1") { - std::unordered_set result = + std::set result = get_pattern_nodes(split_result.subpattern_1); - std::unordered_set correct = even_split.first; + std::set correct = even_split.first; CHECK(result == correct); } SUBCASE("subpattern_2") { - std::unordered_set result = + std::set result = get_pattern_nodes(split_result.subpattern_2); - std::unordered_set correct = even_split.second; + std::set correct = even_split.second; CHECK(result == correct); } SUBCASE("full_pattern_values_to_subpattern_1_inputs") { @@ -113,22 +113,22 @@ TEST_SUITE(FF_TEST_SUITE) { PatternValue pv1 = pattern_value_from_raw_open_kwarg_dataflow_value(v1); PatternSplit even_split = PatternSplit{ - std::unordered_set{p0}, - std::unordered_set{p1}, + std::set{p0}, + std::set{p1}, }; SUBCASE("apply_split") { PatternSplitResult split_result = apply_split(pattern, even_split); SUBCASE("subpattern_1") { - std::unordered_set result = + std::set result = get_pattern_nodes(split_result.subpattern_1); - std::unordered_set correct = even_split.first; + std::set correct = even_split.first; CHECK(result == correct); } SUBCASE("subpattern_2") { - std::unordered_set result = + std::set result = get_pattern_nodes(split_result.subpattern_2); - std::unordered_set correct = even_split.second; + std::set correct = even_split.second; CHECK(result == correct); } SUBCASE("full_pattern_values_to_subpattern_1_inputs") { diff --git a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h index 7a383ee8eb..da51750cd7 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h @@ -13,14 +13,14 @@ bool value_is_mapped(DynamicValueAttrs const &); bool no_part_of_graph_is_copy_inserted(DynamicOpenDataflowGraph const &); bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &); -std::unordered_set copies_for_invocation_inputs( +std::set copies_for_invocation_inputs( DynamicNodeInvocation const &i, - std::unordered_map const + std::map const &unmapped_value_to_mapped_source_value); -std::unordered_set perform_copy_insertion_for_invocation( +std::set perform_copy_insertion_for_invocation( DynamicNodeInvocation const &i, - std::unordered_map const + std::map const &unmapped_value_to_mapped_source_value); DynamicOpenDataflowGraph diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml index 4d6d27444a..d59b0d9060 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.dtg.toml @@ -9,7 +9,7 @@ features = [ ] includes = [ - "", + "", "task-spec/dynamic_graph/dynamic_node_attrs.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h", "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h", diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml index cac13465a0..8f5876beba 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml @@ -8,13 +8,13 @@ features = [ includes = [ "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", - "", + "", ] src_includes = [ - "utils/fmt/unordered_set.h", + "utils/fmt/set.h", ] [[fields]] name = "invocations" -type = "std::unordered_set<::FlexFlow::DynamicNodeInvocation>" +type = "std::set<::FlexFlow::DynamicNodeInvocation>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h index 1aba00a675..3c7dccbbdf 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h @@ -29,13 +29,13 @@ void require_full_dynamic_graph_satisfies( std::function const &, std::function const &); -std::unordered_multiset +std::multiset get_dynamic_nodes(DynamicOpenDataflowGraph const &); -std::unordered_multiset +std::multiset get_dynamic_values(DynamicOpenDataflowGraph const &); -std::unordered_multiset +std::multiset get_dynamic_tensor_slots(DynamicOpenDataflowGraph const &); -std::unordered_set +std::set get_dynamic_invocation_set(DynamicOpenDataflowGraph const &); std::optional @@ -50,11 +50,11 @@ DynamicOpenDataflowGraph transform_dynamic_invocation_set( DynamicOpenDataflowGraph flatmap_dynamic_invocation_set( DynamicOpenDataflowGraph const &, - std::function( + std::function( DynamicNodeInvocation const &)> const &); DynamicOpenDataflowGraph dynamic_open_dataflow_graph_from_invocation_set( - std::unordered_set const &); + std::set const &); std::pair +std::set perform_machine_slicing_for_invocation(DynamicNodeInvocation const &, MachineSpaceCoordinate const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml index 6051a49876..17a48fb6fe 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.toml @@ -10,7 +10,7 @@ features = [ ] includes = [ - "", + "", "task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h", "task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.h", diff --git a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h index 50c713ee3b..4a93575b51 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h @@ -33,10 +33,10 @@ namespace FlexFlow { DynamicNodeInvocation const &, DynamicNodeInvocationShardingInfo const &); -[[nodiscard]] std::unordered_set +[[nodiscard]] std::set generate_shard_expansion_for_invocation(DynamicNodeInvocation const &); -[[nodiscard]] std::unordered_set +[[nodiscard]] std::set perform_shard_expansion_for_invocation(DynamicNodeInvocation const &); [[nodiscard]] DynamicOpenDataflowGraph diff --git a/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h index 9818152b34..f38518e927 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h @@ -15,7 +15,7 @@ bool value_is_ready_for_update_insertion(DynamicValueAttrs const &); bool no_part_of_graph_has_had_update_insertion_performed(DynamicOpenDataflowGraph const &); bool graph_is_ready_for_update_insertion(DynamicOpenDataflowGraph const &); -std::unordered_set +std::set perform_update_insertion_for_invocation(DynamicNodeInvocation const &, OptimizerAttrs const &); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 05e7a0fa24..cf2ee441d4 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -92,16 +92,16 @@ static DynamicValueAttrs map_dynamic_value_attrs_for_task_group( static std::pair filter_mapping_to_avoid_degenerate_copies(DynamicValueAttrs const &input, DynamicValueAttrs const &output) { - std::unordered_set< + std::set< std::pair> input_mapping = unstructured_relation_from_bidict(assert_unwrap(input.mapping)); - std::unordered_set< + std::set< std::pair> output_mapping = unstructured_relation_from_bidict(assert_unwrap(output.mapping)); // Exclude the point shared between the input and output mappings, because // those will not result in actual copies once shard expansion is performed - std::unordered_set< + std::set< std::pair> remove = set_intersection(input_mapping, output_mapping); @@ -116,9 +116,9 @@ static std::pair return std::pair{filtered_input, filtered_output}; } -std::unordered_set copies_for_invocation_inputs( +std::set copies_for_invocation_inputs( DynamicNodeInvocation const &i, - std::unordered_map const &unmapped_value_to_src_mapped_value) + std::map const &unmapped_value_to_src_mapped_value) { if (training_op_attrs_has_op_type(assert_unwrap(i.node_attrs.op_attrs), OperatorType::REPLICATE)) { // copies should not be inserted before a replicate, as the replicate @@ -136,7 +136,7 @@ std::unordered_set copies_for_invocation_inputs( std::map mapped_inputs = map_values2(i.inputs, map_tensor); - std::unordered_set result; + std::set result; for (auto const &[slot, input] : i.inputs) { if (!contains_key(unmapped_value_to_src_mapped_value, input)) { @@ -189,9 +189,9 @@ std::unordered_set copies_for_invocation_inputs( return result; } -std::unordered_set perform_copy_insertion_for_invocation( +std::set perform_copy_insertion_for_invocation( DynamicNodeInvocation const &i, - std::unordered_map const + std::map const &unmapped_value_to_mapped_source_value) { MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); @@ -213,9 +213,9 @@ std::unordered_set perform_copy_insertion_for_invocation( return r; }(); - std::unordered_set result = set_union( + std::set result = set_union( copies_for_invocation_inputs(i, unmapped_value_to_mapped_source_value), - std::unordered_set{ + std::set{ mapped_i, }); @@ -228,7 +228,7 @@ DynamicOpenDataflowGraph ASSERT(no_part_of_graph_is_copy_inserted(g)); require_graph_is_ready_for_copy_insertion(g); - std::unordered_map + std::map unmapped_value_to_mapped_source_value; for (DynamicNodeInvocation const &i : g.invocations) { for (auto const &[slot, value] : i.outputs) { diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index 38d49ef183..2cb908aa8c 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -20,14 +20,13 @@ #include "utils/graph/open_kwarg_dataflow_graph/kwarg_dataflow_graph_input.dtg.h" #include "utils/many_to_one/many_to_one.h" #include "utils/containers/require_all_of.h" -#include "utils/containers/unordered_map_from_map.h" -#include "utils/containers/map_from_unordered.h" +#include "utils/containers/multiset_of.h" namespace FlexFlow { DynamicOpenDataflowGraph make_empty_dynamic_open_dataflow_graph() { return DynamicOpenDataflowGraph{ - std::unordered_set{}, + std::set{}, }; } @@ -71,34 +70,34 @@ void require_full_dynamic_graph_satisfies( } -std::unordered_multiset +std::multiset get_dynamic_nodes(DynamicOpenDataflowGraph const &g) { - return transform(unordered_multiset_of(g.invocations), + return transform(multiset_of(g.invocations), [&](DynamicNodeInvocation const &i) -> DynamicNodeAttrs { return i.node_attrs; }); } -std::unordered_multiset +std::multiset get_dynamic_values(DynamicOpenDataflowGraph const &g) { - return flatmap(unordered_multiset_of(g.invocations), + return flatmap(multiset_of(g.invocations), [&](DynamicNodeInvocation const &i) - -> std::unordered_multiset { + -> std::multiset { return multiset_union(values(i.inputs), values(i.outputs)); }); } -std::unordered_multiset +std::multiset get_dynamic_tensor_slots(DynamicOpenDataflowGraph const &g) { - return flatmap(unordered_multiset_of(g.invocations), + return flatmap(multiset_of(g.invocations), [&](DynamicNodeInvocation const &i) - -> std::unordered_multiset { - return unordered_multiset_of( + -> std::multiset { + return multiset_of( set_union(keys(i.inputs), keys(i.outputs))); }); } -std::unordered_set +std::set get_dynamic_invocation_set(DynamicOpenDataflowGraph const &g) { return g.invocations; } @@ -121,9 +120,9 @@ DynamicOpenDataflowGraph transform_dynamic_invocation_set( DynamicOpenDataflowGraph const &g, std::function const &f) { - std::unordered_set current_invocation_set = + std::set current_invocation_set = get_dynamic_invocation_set(g); - std::unordered_set new_invocation_set = + std::set new_invocation_set = transform(current_invocation_set, f); return dynamic_open_dataflow_graph_from_invocation_set(new_invocation_set); @@ -131,10 +130,10 @@ DynamicOpenDataflowGraph transform_dynamic_invocation_set( DynamicOpenDataflowGraph flatmap_dynamic_invocation_set( DynamicOpenDataflowGraph const &g, - std::function( + std::function( DynamicNodeInvocation const &)> const &f) { - std::unordered_set current_invocation_set = + std::set current_invocation_set = get_dynamic_invocation_set(g); std::vector new_invocation_set = flatmap(vector_of(current_invocation_set), f); @@ -142,11 +141,11 @@ DynamicOpenDataflowGraph flatmap_dynamic_invocation_set( ASSERT(!contains_duplicates(new_invocation_set)); return dynamic_open_dataflow_graph_from_invocation_set( - unordered_set_of(new_invocation_set)); + set_of(new_invocation_set)); } DynamicOpenDataflowGraph dynamic_open_dataflow_graph_from_invocation_set( - std::unordered_set const &invocation_set) { + std::set const &invocation_set) { return DynamicOpenDataflowGraph{ invocation_set, @@ -161,8 +160,8 @@ std::pair all_values = - unordered_set_of(get_dynamic_values(g)); + std::set all_values = + set_of(get_dynamic_values(g)); ManyToOne value_to_producer; for (DynamicNodeInvocation const &invocation : @@ -172,7 +171,7 @@ std::pair graph_inputs = + std::set graph_inputs = filter(all_values, [&](DynamicValueAttrs const &v) -> bool { return !value_to_producer.contains_l(v); }); @@ -212,22 +211,22 @@ std::pair node_map; - std::unordered_set to_add = g.invocations; + std::set to_add = g.invocations; auto add_invocation_to_graph = [&](DynamicNodeInvocation const &invocation) -> void { KwargNodeAddedResult added = result.add_node( invocation.node_attrs, - map_values(unordered_map_from_map(invocation.inputs), + map_values(invocation.inputs, [&](DynamicValueAttrs const &input) -> OpenKwargDataflowValue { return value_map.at_r(input); }), - unordered_map_from_map(invocation.outputs)); + invocation.outputs); node_map.equate(added.node, invocation); for (auto const &[k, v] : - zip_values_strict(invocation.outputs, map_from_unordered(added.outputs))) { + zip_values_strict(invocation.outputs, added.outputs)) { DynamicValueAttrs invocation_output = v.first; KwargDataflowOutput graph_output = v.second; value_map.equate( @@ -344,8 +343,8 @@ std::string }; std::function( - std::unordered_set const &)> - order_slots = [](std::unordered_set const &slot_names) + std::set const &)> + order_slots = [](std::set const &slot_names) -> std::vector { return sorted(slot_names); }; return labelled_open_kwarg_dataflow_graph_view_as_dot(labelled_g, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc index 6dd73bed7d..1fc987926d 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set +std::set perform_machine_slicing_for_invocation( DynamicNodeInvocation const &invocation, MachineSpaceCoordinate const &device_coord) { @@ -23,7 +23,7 @@ DynamicOpenDataflowGraph DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [&](DynamicNodeInvocation const &invocation) - -> std::unordered_set { + -> std::set { return perform_machine_slicing_for_invocation(invocation, device_coord); }); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc index 740a93415e..5a7f21a26b 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc @@ -7,9 +7,9 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include -#include +#include #include #include "utils/containers/map_from_unordered.h" @@ -31,7 +31,7 @@ DynamicOpenDataflowGraph /*per_device_op_state=*/std::nullopt, }; - std::unordered_map result_inputs = + std::map result_inputs = transform( get_incoming_tensors(cg, layer), [&](TensorSlotName const &slot_name, tensor_guid_t const &tensor) { @@ -53,7 +53,7 @@ DynamicOpenDataflowGraph }; }); - std::unordered_map result_outputs = + std::map result_outputs = transform( get_outgoing_tensors(cg, layer), [&](TensorSlotName const &slot_name, tensor_guid_t const &tensor) { @@ -75,7 +75,7 @@ DynamicOpenDataflowGraph }; }); - result.invocations.emplace(map_from_unordered(result_inputs), result_attrs, map_from_unordered(result_outputs)); + result.invocations.emplace(result_inputs, result_attrs, result_outputs); } return result; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 9c7638440f..285b02d58d 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -15,9 +15,8 @@ #include "utils/containers/require_only_key.h" #include "utils/containers/transform_pairs.h" #include -#include +#include #include -#include "utils/containers/unordered_map_from_map.h" #include "utils/bidict/algorithms/bidict_unordered_set_of.h" namespace FlexFlow { @@ -75,7 +74,7 @@ DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( MappedParallelComputationGraph const &mpcg) { return dynamic_open_dataflow_graph_from_invocation_set( - transform(unordered_set_of(mpcg_get_invocation_set(mpcg)), make_dynamic_node_invocation_from_mapped)); + transform(set_of(mpcg_get_invocation_set(mpcg)), make_dynamic_node_invocation_from_mapped)); } } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index a348ba77da..28342ec2e9 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -161,11 +161,11 @@ DynamicOpenDataflowGraph DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [](DynamicNodeInvocation const &invocation) { if (invocation.inputs.empty()) { - return std::unordered_set{ + return std::set{ perform_fwd_pass_expansion_for_invocation(invocation), }; } else { - return std::unordered_set{ + return std::set{ perform_fwd_pass_expansion_for_invocation(invocation), perform_bwd_pass_expansion_for_invocation(invocation), }; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index 440668ea95..0146e80de9 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -417,10 +417,10 @@ static std::set return transform(input_grad_tensor_shards, invocation_sharding_info_for_input_grad_tensor_shard); } -std::unordered_set +std::set perform_shard_expansion_for_invocation(DynamicNodeInvocation const &i) { - std::unordered_set + std::set shard_expansion_info = generate_shard_expansion_for_invocation(i); return transform( @@ -509,22 +509,22 @@ DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( return result; } -std::unordered_set +std::set generate_shard_expansion_for_invocation(DynamicNodeInvocation const &i) { require_invocation_is_ready_for_shard_expansion(i); if (i.node_attrs.op_attrs.value().is_copy()) { - return unordered_set_of(generate_shard_expansion_for_copy(i)); + return set_of(generate_shard_expansion_for_copy(i)); } if (training_op_attrs_has_op_type(i.node_attrs.op_attrs.value(), OperatorType::REPLICATE)) { DynamicTaskType task_type = assert_unwrap(i.node_attrs.task_type); switch (task_type) { case DynamicTaskType::FWD: - return unordered_set_of(generate_shard_expansion_for_fwd_replicate(i)); + return set_of(generate_shard_expansion_for_fwd_replicate(i)); case DynamicTaskType::BWD: - return unordered_set_of(generate_shard_expansion_for_bwd_replicate(i)); + return set_of(generate_shard_expansion_for_bwd_replicate(i)); default: PANIC("Unexpected task type for Replicate: {}", task_type); } @@ -532,7 +532,7 @@ std::unordered_set MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); - std::unordered_set shard_machine_coords = + std::set shard_machine_coords = mapping.get_shard_bindings().left_values(); return transform( diff --git a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc index 51a79cff59..36a7c58432 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc @@ -67,8 +67,8 @@ static DynamicNodeInvocation get_update_invocation_for_invocation( }; }; - std::unordered_set tensor_roles = set_union( - std::unordered_set{ + std::set tensor_roles = set_union( + std::set{ mk_dynamic_tensor_role_fwd(), mk_dynamic_tensor_role_bwd(), }, @@ -83,7 +83,7 @@ static DynamicNodeInvocation get_update_invocation_for_invocation( }; } -std::unordered_set +std::set perform_update_insertion_for_invocation( DynamicNodeInvocation const &invocation, OptimizerAttrs const &optimizer_attrs) { @@ -91,12 +91,12 @@ std::unordered_set if (invocation.node_attrs.task_type.value() == DynamicTaskType::FWD && invocation.node_attrs.op_attrs.value().is_pcg_op() && invocation.node_attrs.op_attrs.value().require_pcg_op().is_weight()) { - return std::unordered_set{ + return std::set{ invocation, get_update_invocation_for_invocation(invocation, optimizer_attrs), }; } else { - return std::unordered_set{ + return std::set{ invocation, }; }; diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index 9bb0e7a7c4..33e240dde6 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -4,7 +4,7 @@ #include "task-spec/dynamic_graph/dynamic_task_type.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include #include "task-spec/dynamic_graph/dynamic_value_attrs.h" #include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" @@ -292,15 +292,15 @@ TEST_SUITE(FF_TEST_SUITE) { }; SUBCASE("same mapping, no copies") { - std::unordered_map sources_same{ + std::map sources_same{ {graph_input1, graph_input1_src_same}, {graph_input2, graph_input2_src_same}, }; - std::unordered_set result = + std::set result = copies_for_invocation_inputs(input, sources_same); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result.size() == correct.size()); CHECK(result == correct); @@ -340,14 +340,14 @@ TEST_SUITE(FF_TEST_SUITE) { graph_input1, get_tensor_bindings_for_slot_name(input_mapping_copy1_diff_vs_use, TensorSlotName::OUTPUT)); - std::unordered_map sources_copy1{ + std::map sources_copy1{ {graph_input1, graph_input1_src_copy1}, {graph_input2, graph_input2_src_same}}; - std::unordered_set result = + std::set result = copies_for_invocation_inputs(input, sources_copy1); - std::unordered_set correct = { + std::set correct = { mk_copy(graph_input1_src_copy1_diff_vs_use, graph_input1_use_diff_vs_copy1), }; @@ -391,14 +391,14 @@ TEST_SUITE(FF_TEST_SUITE) { graph_input2, get_tensor_bindings_for_slot_name(weight_mapping_copy2, TensorSlotName::OUTPUT)); - std::unordered_map sources_copy2{ + std::map sources_copy2{ {graph_input1, graph_input1_src_copy2}, {graph_input2, graph_input2_src_copy2}}; - std::unordered_set result = + std::set result = copies_for_invocation_inputs(input, sources_copy2); - std::unordered_set correct = { + std::set correct = { mk_copy(graph_input1_src_copy2, graph_input1_use), mk_copy(graph_input2_src_copy2, graph_input2_use), }; @@ -493,7 +493,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_map unmapped_to_mapped_source_value = { + std::map unmapped_to_mapped_source_value = { { graph_input_unmapped, decide_dynamic_value_attrs_mapping( @@ -507,10 +507,10 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set result = copies_for_invocation_inputs( + std::set result = copies_for_invocation_inputs( input, unmapped_to_mapped_source_value); - std::unordered_set correct = {}; + std::set correct = {}; nlohmann::json result_j = transform(result, dynamic_node_invocation_to_serializable); nlohmann::json correct_j = transform(correct, dynamic_node_invocation_to_serializable); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index 96523b6c31..7c1219c5b3 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -129,7 +129,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*outputs=*/std::map{}, }; - std::unordered_set invocation_set = { + std::set invocation_set = { invocation_1, invocation_2, invocation_3, diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc index 789d89a676..19fb423350 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc @@ -230,7 +230,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicOpenDataflowGraph correct = dynamic_open_dataflow_graph_from_invocation_set( - std::unordered_set{}); + std::set{}); CHECK(dynamic_open_dataflow_graphs_are_isomorphic(result, correct)); } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc index 90fbdec5f7..1a90aa3f40 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc @@ -352,7 +352,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); - std::unordered_set invocation_set = { + std::set invocation_set = { DynamicNodeInvocation{ /*inputs=*/std::map{}, /*node_attrs=*/n1, @@ -418,7 +418,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicValueAttrs v2_gradient = mk_value_attrs(1, mk_dynamic_tensor_role_bwd()); - std::unordered_set invocation_set = { + std::set invocation_set = { DynamicNodeInvocation{ /*inputs=*/std::map{}, /*node_attrs=*/n1_fwd, diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index 60dbfea9aa..cfbb8aa309 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -3,7 +3,7 @@ #include "task-spec/dynamic_graph/copy_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_copy_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "op-attrs/ops/element_unary.h" @@ -234,7 +234,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set result = + std::set result = generate_shard_expansion_for_invocation(input); auto mk_invocation_shard = @@ -255,7 +255,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::unordered_set correct = { + std::set correct = { mk_invocation_shard(mc1, mc1_input_coord, mc1_weight_coord, @@ -320,7 +320,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set result = + std::set result = generate_shard_expansion_for_invocation(input); auto mk_invocation_shard = @@ -349,7 +349,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::unordered_set correct = { + std::set correct = { mk_invocation_shard(mc1, pt1), mk_invocation_shard(mc2, pt2), }; @@ -459,7 +459,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set result = + std::set result = generate_shard_expansion_for_invocation(input); @@ -482,7 +482,7 @@ TEST_SUITE(FF_TEST_SUITE) { auto mk_invocation_shard = [&](nonempty_set const &device_coords, ParallelTensorSpaceCoordinate const &input_shard_coord, - std::unordered_set const &output_task_shards) + std::set const &output_task_shards) -> DynamicNodeInvocationShardingInfo { return DynamicNodeInvocationShardingInfo{ @@ -506,7 +506,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::unordered_set correct = { + std::set correct = { mk_invocation_shard(nonempty_set{mc1, mc2}, pt1, {mc1, mc2}), mk_invocation_shard(nonempty_set{mc3, mc4}, pt2, {mc3, mc4}), }; @@ -567,7 +567,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set result = + std::set result = generate_shard_expansion_for_invocation(input); auto mk_output_grad_binding = [&](MachineSpaceCoordinate const &mc) @@ -588,7 +588,7 @@ TEST_SUITE(FF_TEST_SUITE) { auto mk_invocation_shard = [&](nonempty_set const &device_coords, - std::unordered_set const &output_grad_task_shards, + std::set const &output_grad_task_shards, ParallelTensorSpaceCoordinate const &input_grad_shard_coord) -> DynamicNodeInvocationShardingInfo { @@ -613,7 +613,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::unordered_set correct = { + std::set correct = { mk_invocation_shard(nonempty_set{mc1, mc2}, {mc1, mc2}, pt1), mk_invocation_shard(nonempty_set{mc3, mc4}, {mc3, mc4}, pt2), }; diff --git a/lib/utils/benchmark/src/internal/random_dag.cc b/lib/utils/benchmark/src/internal/random_dag.cc index 186d070be7..8ded12030d 100644 --- a/lib/utils/benchmark/src/internal/random_dag.cc +++ b/lib/utils/benchmark/src/internal/random_dag.cc @@ -26,7 +26,7 @@ DiGraphView random_dag(nonnegative_int num_nodes, float edges_fraction) { DiGraph g = DiGraph::create(); std::vector n = add_nodes(g, num_nodes.unwrap_nonnegative()); - std::unordered_set edges; + std::set edges; while (edges.size() < num_edges) { Node n1 = select_random(n); Node n2 = select_random(n); diff --git a/lib/utils/include/utils/bidict/algorithms/bidict_from_map.h b/lib/utils/include/utils/bidict/algorithms/bidict_from_map.h index b4d74df97c..aca6bc4d5e 100644 --- a/lib/utils/include/utils/bidict/algorithms/bidict_from_map.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_from_map.h @@ -6,7 +6,7 @@ namespace FlexFlow { template -bidict bidict_from_map(std::unordered_map const &m) { +bidict bidict_from_map(std::map const &m) { bidict result; for (auto const &[k, v] : m) { ASSERT(!result.contains_r(v)); @@ -16,7 +16,7 @@ bidict bidict_from_map(std::unordered_map const &m) { } template -bidict bidict_from_map(std::map const &m) { +bidict bidict_from_map(std::unordered_map const &m) { bidict result; for (auto const &[k, v] : m) { ASSERT(!result.contains_r(v)); diff --git a/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h b/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h index 1f6d7893b6..3f7a3e4ff2 100644 --- a/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h @@ -7,7 +7,7 @@ namespace FlexFlow { template bidict bidict_from_unstructured_relation( - std::unordered_set> const &relation) { + std::set> const &relation) { bidict result; for (auto const &lr : relation) { result.equate_strict(lr); diff --git a/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h b/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h new file mode 100644 index 0000000000..25731f3d89 --- /dev/null +++ b/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h @@ -0,0 +1,24 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_TRANSFORM_KEYS_AND_VALUES_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_TRANSFORM_KEYS_AND_VALUES_H + +#include "utils/bidict/bidict.h" + +namespace FlexFlow { + +template , + typename V2 = std::invoke_result_t> +bidict bidict_transform_keys_and_values(bidict const &m, KF &&kf, VF &&vf) { + bidict result; + for (auto const &kv : m) { + result.equate_strict(kf(kv.first), vf(kv.second)); + } + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/bidict/algorithms/left_entries.h b/lib/utils/include/utils/bidict/algorithms/left_entries.h index a3fab172b1..f98dae6e99 100644 --- a/lib/utils/include/utils/bidict/algorithms/left_entries.h +++ b/lib/utils/include/utils/bidict/algorithms/left_entries.h @@ -2,13 +2,13 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_LEFT_ENTRIES_H #include "utils/bidict/bidict.h" -#include +#include namespace FlexFlow { template -std::unordered_set left_entries(bidict const &b) { - std::unordered_set result; +std::set left_entries(bidict const &b) { + std::set result; for (auto const &[l, _] : b) { result.insert(l); } diff --git a/lib/utils/include/utils/bidict/algorithms/right_entries.h b/lib/utils/include/utils/bidict/algorithms/right_entries.h index ec0e822c74..ea157fdbf6 100644 --- a/lib/utils/include/utils/bidict/algorithms/right_entries.h +++ b/lib/utils/include/utils/bidict/algorithms/right_entries.h @@ -2,13 +2,13 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_RIGHT_ENTRIES_H #include "utils/bidict/bidict.h" -#include +#include namespace FlexFlow { template -std::unordered_set right_entries(bidict const &b) { - std::unordered_set result; +std::set right_entries(bidict const &b) { + std::set result; for (auto const &[_, r] : b) { result.insert(r); } diff --git a/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h b/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h index 63d25332d3..d5c09ce039 100644 --- a/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h +++ b/lib/utils/include/utils/bidict/algorithms/unstructured_relation_from_bidict.h @@ -1,15 +1,21 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_UNSTRUCTURED_RELATION_FROM_BIDICT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_UNSTRUCTURED_RELATION_FROM_BIDICT_H -#include "utils/bidict/algorithms/bidict_unordered_set_of.h" #include "utils/bidict/bidict.h" namespace FlexFlow { template -std::unordered_set> +std::set> unstructured_relation_from_bidict(bidict const &b) { - return bidict_unordered_set_of(b); + + std::set> result; + + for (auto const &lr : b) { + result.insert(lr); + } + + return result; } } // namespace FlexFlow diff --git a/lib/utils/include/utils/bidict/bidict.h b/lib/utils/include/utils/bidict/bidict.h index 7fcc59f116..79c88d6bd9 100644 --- a/lib/utils/include/utils/bidict/bidict.h +++ b/lib/utils/include/utils/bidict/bidict.h @@ -1,22 +1,22 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_BIDICT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_BIDICT_H -#include "utils/containers/unordered_keys.h" #include "utils/containers/map_from_keys_and_values.h" -#include "utils/fmt/unordered_map.h" -#include "utils/hash/unordered_map.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" -#include "utils/ord/unordered_map.h" #include #include #include #include -#include #include "utils/containers/require_same.h" #include "utils/containers/values.h" -#include "utils/containers/unordered_set_of.h" #include "utils/containers/contains_key.h" +#include "utils/containers/unordered_map_from_map.h" +#include "utils/check_fmtable.h" +#include "utils/containers/keys.h" +#include "utils/containers/set_of.h" +#include "utils/hash/map.h" +#include "utils/fmt/map.h" namespace FlexFlow { @@ -118,12 +118,12 @@ struct bidict { return bwd_map.at(r); } - std::unordered_set left_values() const { - return unordered_keys(this->fwd_map); + std::set left_values() const { + return keys(this->fwd_map); } - std::unordered_set right_values() const { - return unordered_keys(this->bwd_map); + std::set right_values() const { + return keys(this->bwd_map); } std::size_t size() const { @@ -135,7 +135,7 @@ struct bidict { return this->size() == 0; } - using const_iterator = typename std::unordered_map::const_iterator; + using const_iterator = typename std::map::const_iterator; using value_type = std::pair; using reference = value_type &; using const_reference = value_type const &; @@ -148,7 +148,7 @@ struct bidict { /* using pointer = std::pair const *; */ /* using reference = std::pair const &; */ - /* explicit const_iterator(typename std::unordered_map, + /* explicit const_iterator(typename std::map, * tl::optional>::const_iterator); */ /* reference operator*() const { */ @@ -177,7 +177,7 @@ struct bidict { /* } */ /* private: */ /* mutable tl::optional> current; */ - /* typename std::unordered_map, + /* typename std::map, * tl::optional>::const_iterator it; */ /* }; */ @@ -217,33 +217,56 @@ struct bidict { return bidict(bwd_map, fwd_map); } - operator std::unordered_map const &() const { + operator std::map const &() const { return this->fwd_map; } - std::unordered_map const &as_unordered_map() const { + operator std::unordered_map () const { + return unordered_map_from_map(this->fwd_map); + } + + std::map const &as_map() const { return this->fwd_map; } - std::unordered_map const &l_to_r() const { + std::unordered_map as_unordered_map() const { + return unordered_map_from_map(this->fwd_map); + } + + std::map const &l_to_r() const { return this->fwd_map; } - std::unordered_map const &r_to_l() const { + std::map const &r_to_l() const { return this->bwd_map; } - bidict(std::unordered_map const &fwd_map, - std::unordered_map const &bwd_map) + bidict(std::map const &fwd_map, + std::map const &bwd_map) : fwd_map(fwd_map), bwd_map(bwd_map) {} + bool operator<(bidict const &other) const { + return this->fwd_map < other.fwd_map; + } + + bool operator<=(bidict const &other) const { + return this->fwd_map <= other.fwd_map; + } + + bool operator>(bidict const &other) const { + return this->fwd_map > other.fwd_map; + } + + bool operator>=(bidict const &other) const { + return this->fwd_map >= other.fwd_map; + } private: void check_invariants() const { - std::unordered_set fwd_l_vals = unordered_keys(this->fwd_map); - std::unordered_set bwd_l_vals = unordered_set_of(values(this->bwd_map)); + std::set fwd_l_vals = keys(this->fwd_map); + std::set bwd_l_vals = set_of(values(this->bwd_map)); - std::unordered_set bwd_r_vals = unordered_keys(this->bwd_map); - std::unordered_set fwd_r_vals = unordered_set_of(values(this->fwd_map)); + std::set bwd_r_vals = keys(this->bwd_map); + std::set fwd_r_vals = set_of(values(this->fwd_map)); ASSERT(fwd_l_vals == bwd_l_vals); ASSERT(fwd_r_vals == bwd_r_vals); @@ -255,18 +278,12 @@ struct bidict { friend struct bidict; - std::unordered_map fwd_map; - std::unordered_map bwd_map; + std::map fwd_map; + std::map bwd_map; }; template -std::enable_if_t && is_lt_comparable_v, bool> - operator<(bidict const &lhs, bidict const &rhs) { - return lhs.as_unordered_map() < rhs.as_unordered_map(); -} - -template -std::unordered_map format_as(bidict const &b) { +std::map format_as(bidict const &b) { return b; } @@ -288,7 +305,7 @@ struct adl_serializer<::FlexFlow::bidict> { CHECK_IS_JSON_DESERIALIZABLE(L); CHECK_IS_JSON_DESERIALIZABLE(R); - std::unordered_map m = j; + std::map m = j; ::FlexFlow::bidict b{m.cbegin(), m.cend()}; @@ -298,7 +315,7 @@ struct adl_serializer<::FlexFlow::bidict> { CHECK_IS_JSON_SERIALIZABLE(L); CHECK_IS_JSON_SERIALIZABLE(R); - j = b.as_unordered_map(); + j = b.as_map(); } }; @@ -310,16 +327,16 @@ template struct Arbitrary<::FlexFlow::bidict> { static Gen<::FlexFlow::bidict> arbitrary() { return gen::map( - gen::withSize([](int size) -> Gen> { + gen::withSize([](int size) -> Gen> { return gen::apply( [](std::vector const &keys, - std::vector const &values) -> std::unordered_map { + std::vector const &values) -> std::map { return ::FlexFlow::map_from_keys_and_values(keys, values); }, gen::unique>(size, gen::arbitrary()), gen::unique>(size, gen::arbitrary())); }), - [](std::unordered_map const &m) { + [](std::map const &m) { return ::FlexFlow::bidict{m.cbegin(), m.cend()}; }); } @@ -332,7 +349,7 @@ namespace std { template struct hash<::FlexFlow::bidict> { size_t operator()(::FlexFlow::bidict const &b) const { - return hash>{}(b); + return hash>{}(b.as_map()); } }; diff --git a/lib/utils/include/utils/cli/cli_argument_key.dtg.toml b/lib/utils/include/utils/cli/cli_argument_key.dtg.toml index bea9ed3eb8..e2432e8c57 100644 --- a/lib/utils/include/utils/cli/cli_argument_key.dtg.toml +++ b/lib/utils/include/utils/cli/cli_argument_key.dtg.toml @@ -3,6 +3,7 @@ name = "CLIArgumentKey" type = "variant" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/cli/cli_flag_spec.dtg.toml b/lib/utils/include/utils/cli/cli_flag_spec.dtg.toml index 1bb29953e6..73a4e48be4 100644 --- a/lib/utils/include/utils/cli/cli_flag_spec.dtg.toml +++ b/lib/utils/include/utils/cli/cli_flag_spec.dtg.toml @@ -3,6 +3,7 @@ name = "CLIFlagSpec" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/cli/cli_parse_result.dtg.toml b/lib/utils/include/utils/cli/cli_parse_result.dtg.toml index fc52613e51..dbba820186 100644 --- a/lib/utils/include/utils/cli/cli_parse_result.dtg.toml +++ b/lib/utils/include/utils/cli/cli_parse_result.dtg.toml @@ -3,26 +3,27 @@ name = "CLIParseResult" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] includes = [ - "", + "", "", "utils/cli/cli_flag_key.dtg.h", "utils/cli/cli_positional_argument_key.dtg.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "flags" -type = "std::unordered_map<::FlexFlow::CLIFlagKey, bool>" +type = "std::map<::FlexFlow::CLIFlagKey, bool>" [[fields]] name = "positional_arguments" -type = "std::unordered_map<::FlexFlow::CLIPositionalArgumentKey, std::string>" +type = "std::map<::FlexFlow::CLIPositionalArgumentKey, std::string>" diff --git a/lib/utils/include/utils/cli/cli_positional_argument_key.dtg.toml b/lib/utils/include/utils/cli/cli_positional_argument_key.dtg.toml index 2ed6eed5b4..e36f4422b9 100644 --- a/lib/utils/include/utils/cli/cli_positional_argument_key.dtg.toml +++ b/lib/utils/include/utils/cli/cli_positional_argument_key.dtg.toml @@ -3,6 +3,7 @@ name = "CLIPositionalArgumentKey" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/cli/cli_positional_argument_spec.dtg.toml b/lib/utils/include/utils/cli/cli_positional_argument_spec.dtg.toml index 34312f6a04..9b640b39a5 100644 --- a/lib/utils/include/utils/cli/cli_positional_argument_spec.dtg.toml +++ b/lib/utils/include/utils/cli/cli_positional_argument_spec.dtg.toml @@ -3,13 +3,14 @@ name = "CLIPositionalArgumentSpec" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] includes = [ "", - "", + "", "", ] diff --git a/lib/utils/include/utils/cli/cli_spec.dtg.toml b/lib/utils/include/utils/cli/cli_spec.dtg.toml index 39ff325eac..4de17f8952 100644 --- a/lib/utils/include/utils/cli/cli_spec.dtg.toml +++ b/lib/utils/include/utils/cli/cli_spec.dtg.toml @@ -3,20 +3,21 @@ name = "CLISpec" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] includes = [ - "", + "", "utils/cli/cli_flag_spec.dtg.h", "utils/cli/cli_positional_argument_spec.dtg.h", "", ] src_includes = [ - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h", + "utils/fmt/set.h", + "utils/hash/set.h", "utils/fmt/vector.h", "utils/hash/vector.h", ] diff --git a/lib/utils/include/utils/cli/cli_spec.h b/lib/utils/include/utils/cli/cli_spec.h index 2c0df08c55..f88c9cffb7 100644 --- a/lib/utils/include/utils/cli/cli_spec.h +++ b/lib/utils/include/utils/cli/cli_spec.h @@ -4,7 +4,7 @@ #include "utils/cli/cli_argument_key.dtg.h" #include "utils/cli/cli_flag_spec.dtg.h" #include "utils/cli/cli_spec.dtg.h" -#include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/commutative_pair.h b/lib/utils/include/utils/commutative_pair.h index 1a4009fedb..6ff28e689a 100644 --- a/lib/utils/include/utils/commutative_pair.h +++ b/lib/utils/include/utils/commutative_pair.h @@ -98,8 +98,8 @@ template struct hash<::FlexFlow::commutative_pair> { size_t operator()(::FlexFlow::commutative_pair const &p) { size_t result = 0; - ::FlexFlow::unordered_hash_combine(result, p.first); - ::FlexFlow::unordered_hash_combine(result, p.second); + ::FlexFlow::hash_combine(result, p.first); + ::FlexFlow::hash_combine(result, p.second); return result; } }; diff --git a/lib/utils/include/utils/containers/are_all_distinct.h b/lib/utils/include/utils/containers/are_all_distinct.h index d02845ba16..097406d885 100644 --- a/lib/utils/include/utils/containers/are_all_distinct.h +++ b/lib/utils/include/utils/containers/are_all_distinct.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_UTILS_INCLUDE_UTILS_CONTAINERS_ARE_ALL_DISTINCT_H #define _FLEXFLOW_UTILS_INCLUDE_UTILS_CONTAINERS_ARE_ALL_DISTINCT_H -#include "utils/containers/unordered_multiset_of.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/multiset_of.h" +#include "utils/containers/set_of.h" namespace FlexFlow { template bool are_all_distinct(C const &c) { - return unordered_set_of(c).size() == unordered_multiset_of(c).size(); + return set_of(c).size() == multiset_of(c).size(); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/containers/are_disjoint.h b/lib/utils/include/utils/containers/are_disjoint.h index 0d5882889c..3be7712dd2 100644 --- a/lib/utils/include/utils/containers/are_disjoint.h +++ b/lib/utils/include/utils/containers/are_disjoint.h @@ -11,6 +11,12 @@ bool are_disjoint(std::unordered_set const &l, return set_intersection(l, r).empty(); } +template +bool are_disjoint(std::set const &l, + std::set const &r) { + return set_intersection(l, r).empty(); +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/binary_cartesian_product.h b/lib/utils/include/utils/containers/binary_cartesian_product.h index 1e9f5febbf..471f99cc6e 100644 --- a/lib/utils/include/utils/containers/binary_cartesian_product.h +++ b/lib/utils/include/utils/containers/binary_cartesian_product.h @@ -1,16 +1,15 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_CARTESIAN_PRODUCT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_CARTESIAN_PRODUCT_H -#include "utils/hash/pair.h" -#include +#include namespace FlexFlow { template -std::unordered_set> - binary_cartesian_product(std::unordered_set const &lhs, - std::unordered_set const &rhs) { - std::unordered_set> result; +std::set> + binary_cartesian_product(std::set const &lhs, + std::set const &rhs) { + std::set> result; for (A const &a : lhs) { for (B const &b : rhs) { diff --git a/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h b/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h index 536d402be0..761465532e 100644 --- a/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h +++ b/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h @@ -4,7 +4,7 @@ #include #include "utils/containers/binary_merge_unordered_maps_with.h" #include "utils/containers/unordered_keys.h" -#include "utils/containers/intersection.h" +#include "utils/containers/set_intersection.h" namespace FlexFlow { @@ -16,7 +16,7 @@ std::unordered_map std::unordered_set lhs_keys = unordered_keys(lhs); std::unordered_set rhs_keys = unordered_keys(rhs); - std::unordered_set shared_keys = intersection(lhs_keys, rhs_keys); + std::unordered_set shared_keys = set_intersection(lhs_keys, rhs_keys); ASSERT(shared_keys.empty()); return binary_merge_unordered_maps_with( diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with.h b/lib/utils/include/utils/containers/binary_merge_maps_with.h index 243bbe6e06..eb404f02bf 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with.h @@ -7,7 +7,7 @@ #include "utils/containers/restrict_keys.h" #include "utils/containers/set_intersection.h" #include "utils/containers/set_minus.h" -#include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h index 2f8be45802..ef4ccebea3 100644 --- a/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h +++ b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_H #include "utils/containers/generate_unordered_map.h" -#include "utils/containers/intersection.h" +#include "utils/containers/set_intersection.h" #include "utils/containers/unordered_keys.h" #include "utils/containers/merge_unordered_maps_with_right_dominating.h" #include "utils/containers/restrict_keys.h" @@ -22,7 +22,7 @@ std::unordered_map std::unordered_set l_only_keys = set_minus(l_keys, r_keys); std::unordered_set r_only_keys = set_minus(r_keys, l_keys); - std::unordered_set both_keys = intersection(r_keys, l_keys); + std::unordered_set both_keys = set_intersection(r_keys, l_keys); std::unordered_map l_only = restrict_keys(lhs, l_only_keys); std::unordered_map r_only = restrict_keys(rhs, r_only_keys); diff --git a/lib/utils/include/utils/containers/cartesian_product.h b/lib/utils/include/utils/containers/cartesian_product.h index 28d0fb118c..be19bebae1 100644 --- a/lib/utils/include/utils/containers/cartesian_product.h +++ b/lib/utils/include/utils/containers/cartesian_product.h @@ -3,15 +3,15 @@ #include "utils/hash/vector.h" #include -#include +#include #include namespace FlexFlow { template -std::unordered_multiset> +std::multiset> cartesian_product(std::vector const &containers) { - std::unordered_multiset> result; + std::multiset> result; std::function &, size_t)> recurse = [&](std::vector ¤t, size_t depth) { diff --git a/lib/utils/include/utils/containers/enumerate.h b/lib/utils/include/utils/containers/enumerate.h index 1e8bc1f3dc..3952d497b2 100644 --- a/lib/utils/include/utils/containers/enumerate.h +++ b/lib/utils/include/utils/containers/enumerate.h @@ -3,7 +3,7 @@ #include "utils/containers/enumerate_vector.h" #include -#include +#include #include namespace FlexFlow { @@ -34,7 +34,7 @@ std::map enumerate(std::vector const &c) { * the result of this function should still function as expected. */ template -std::map enumerate(std::unordered_set const &c) { +std::map enumerate(std::set const &c) { std::map result; nonnegative_int idx = 0_n; for (auto const &v : c) { diff --git a/lib/utils/include/utils/containers/filter_values.h b/lib/utils/include/utils/containers/filter_values.h index 7636962e39..6d4be1254f 100644 --- a/lib/utils/include/utils/containers/filter_values.h +++ b/lib/utils/include/utils/containers/filter_values.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_FILTER_VALUES_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_FILTER_VALUES_H -#include +#include namespace FlexFlow { template -std::unordered_map filter_values(std::unordered_map const &m, +std::map filter_values(std::map const &m, F const &f) { - std::unordered_map result; + std::map result; for (auto const &kv : m) { if (f(kv.second)) { result.insert(kv); diff --git a/lib/utils/include/utils/containers/find.h b/lib/utils/include/utils/containers/find.h index 7b103fed16..68226479df 100644 --- a/lib/utils/include/utils/containers/find.h +++ b/lib/utils/include/utils/containers/find.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_FIND_H #include -#include +#include namespace FlexFlow { @@ -13,8 +13,8 @@ typename Container::const_iterator } template -typename std::unordered_set::const_iterator - find(std::unordered_set const &c, V const &e) { +typename std::set::const_iterator + find(std::set const &c, V const &e) { return c.find(e); } diff --git a/lib/utils/include/utils/containers/flatmap.h b/lib/utils/include/utils/containers/flatmap.h index e0440cc791..bd8836d596 100644 --- a/lib/utils/include/utils/containers/flatmap.h +++ b/lib/utils/include/utils/containers/flatmap.h @@ -5,7 +5,8 @@ #include "utils/containers/get_element_type.h" #include #include -#include +#include +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -83,6 +84,23 @@ std::unordered_map flatmap(std::unordered_map const &m, return result; } +template < + typename InK, + typename InV, + typename F, + typename OutK = typename std::invoke_result_t::key_type, + typename OutV = typename std::invoke_result_t::mapped_type> +std::map flatmap(std::map const &m, + F &&f) { + std::map result; + + for (auto const &[k, v] : m) { + result = binary_merge_disjoint_maps(result, f(k, v)); + } + + return result; +} + template ::value_type> diff --git a/lib/utils/include/utils/containers/get_all_assignments.h b/lib/utils/include/utils/containers/get_all_assignments.h index 8f77ffbc24..d1667700aa 100644 --- a/lib/utils/include/utils/containers/get_all_assignments.h +++ b/lib/utils/include/utils/containers/get_all_assignments.h @@ -2,16 +2,22 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GET_ALL_ASSIGNMENTS_H #include "utils/containers/cartesian_product.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_map_from_pairs.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/map_from_pairs.h" +#include "utils/containers/set_of.h" #include "utils/containers/vector_of.h" #include "utils/containers/zip.h" #include "utils/hash/unordered_map.h" -#include -#include +#include +#include #include +#include "utils/containers/keys.h" +#include "utils/containers/map_from_pairs.h" +#include "utils/containers/set_of.h" +#include "utils/containers/unordered_keys.h" +#include "utils/containers/unordered_set_of.h" +#include "utils/containers/unordered_map_from_pairs.h" namespace FlexFlow { @@ -39,6 +45,31 @@ std::unordered_set> get_all_assignments( return result; } +/** + * @note If \p options_per_key is empty, an set containing a single empty + * assignment is returned + */ +template +std::set> get_all_assignments( + std::map> const &options_per_key) { + if (options_per_key.empty()) { + return {{}}; + } + + std::vector ordered_keys = vector_of(keys(options_per_key)); + std::vector> ordered_value_option_sets = transform( + ordered_keys, [&](K const &k) { return options_per_key.at(k); }); + + std::set> result = transform( + set_of(cartesian_product(ordered_value_option_sets)), + [&](std::vector const &chosen_values) { + return map_from_pairs(zip(ordered_keys, chosen_values)); + }); + + return result; +} + + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/get_all_permutations_with_repetition.h b/lib/utils/include/utils/containers/get_all_permutations_with_repetition.h index d8160e2f6a..a446e998cb 100644 --- a/lib/utils/include/utils/containers/get_all_permutations_with_repetition.h +++ b/lib/utils/include/utils/containers/get_all_permutations_with_repetition.h @@ -2,9 +2,8 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GET_ALL_PERMUTATIONS_WITH_REPETITION_H #include "utils/containers/transform.h" -#include "utils/hash/vector.h" #include "utils/nonnegative_int/nonnegative_int.h" -#include +#include #include namespace FlexFlow { @@ -16,10 +15,10 @@ namespace FlexFlow { * https://en.wikipedia.org/wiki/Permutation#Permutations_with_repetition **/ template -std::unordered_multiset> +std::multiset> get_all_permutations_with_repetition(C const &container, nonnegative_int n) { - std::unordered_multiset> result; + std::multiset> result; if (container.empty() || n == 0) { return {{}}; diff --git a/lib/utils/include/utils/containers/get_element_counts.h b/lib/utils/include/utils/containers/get_element_counts.h index 58cb436040..e32ca7b552 100644 --- a/lib/utils/include/utils/containers/get_element_counts.h +++ b/lib/utils/include/utils/containers/get_element_counts.h @@ -3,14 +3,14 @@ #include "utils/containers/contains_key.h" #include -#include +#include #include namespace FlexFlow { template -std::unordered_map get_element_counts(std::vector const &v) { - std::unordered_map counts; +std::map get_element_counts(std::vector const &v) { + std::map counts; for (T const &t : v) { if (!contains_key(counts, t)) { counts[t] = 0; @@ -20,7 +20,7 @@ std::unordered_map get_element_counts(std::vector const &v) { return counts; } -std::unordered_map get_element_counts(std::string const &); +std::map get_element_counts(std::string const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/containers/get_one_of.h b/lib/utils/include/utils/containers/get_one_of.h index 47c46fb1d6..5bada9dfa6 100644 --- a/lib/utils/include/utils/containers/get_one_of.h +++ b/lib/utils/include/utils/containers/get_one_of.h @@ -2,12 +2,12 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GET_ONE_OF_H #include "utils/exception.h" -#include "utils/fmt/unordered_set.h" -#include +#include "utils/fmt/set.h" +#include namespace FlexFlow { template -T get_one_of(std::unordered_set const &s) { +T get_one_of(std::set const &s) { if (s.empty()) { throw mk_runtime_error(fmt::format( "get_one_of expected non-empty container but receieved {}", s)); diff --git a/lib/utils/include/utils/containers/get_only.h b/lib/utils/include/utils/containers/get_only.h index d5e2faf67a..79b799edde 100644 --- a/lib/utils/include/utils/containers/get_only.h +++ b/lib/utils/include/utils/containers/get_only.h @@ -20,6 +20,13 @@ std::pair get_only(std::unordered_map const &m) { return *m.cbegin(); } +template +std::pair get_only(std::map const &m) { + ASSERT(m.size() == 1); + + return *m.cbegin(); +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/group_by.h b/lib/utils/include/utils/containers/group_by.h index e5269e6958..057467bdfd 100644 --- a/lib/utils/include/utils/containers/group_by.h +++ b/lib/utils/include/utils/containers/group_by.h @@ -29,9 +29,9 @@ OneToMany group_by(std::set const &vs, F &&f) { } template > -std::unordered_map> group_by(std::vector const &vs, +std::map> group_by(std::vector const &vs, F &&f) { - std::unordered_map> result; + std::map> result; for (V const &v : vs) { result[f(v)].push_back(v); } diff --git a/lib/utils/include/utils/containers/index.dox b/lib/utils/include/utils/containers/index.dox index 2ffda1cfdf..9b3865dd78 100644 --- a/lib/utils/include/utils/containers/index.dox +++ b/lib/utils/include/utils/containers/index.dox @@ -9,7 +9,7 @@ Some of the most commonly-used functions are listed below, but you should ideall - \ref containers/transform.h - \ref containers/filter.h - \ref containers/contains.h -- \ref containers/generate_unordered_map.h +- \ref containers/generate_map.h - \ref containers/get_only.h - \ref containers/slice.h - \ref containers/merge_disjoint_maps.h diff --git a/lib/utils/include/utils/containers/invert_map.h b/lib/utils/include/utils/containers/invert_map.h index 6f0c04a189..c1ada072dd 100644 --- a/lib/utils/include/utils/containers/invert_map.h +++ b/lib/utils/include/utils/containers/invert_map.h @@ -2,20 +2,23 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_INVERT_MAP_H #include +#include +#include #include #include namespace FlexFlow { template -std::unordered_map> - invert_map(std::unordered_map const &m) { - std::unordered_map> result; +std::map> + invert_map(std::map const &m) { + std::map> result; for (auto const &[key, value] : m) { result[value].insert(key); } return result; } + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/invert_unordered_map.h b/lib/utils/include/utils/containers/invert_unordered_map.h new file mode 100644 index 0000000000..9efc092227 --- /dev/null +++ b/lib/utils/include/utils/containers/invert_unordered_map.h @@ -0,0 +1,22 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_INVERT_UNORDERED_MAP_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_INVERT_UNORDERED_MAP_H + +#include +#include + +namespace FlexFlow { + +template +std::unordered_map> + invert_unordered_map(std::unordered_map const &m) { + std::unordered_map> result; + for (auto const &[key, value] : m) { + result[value].insert(key); + } + return result; +} + + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/is_submapeq_of.h b/lib/utils/include/utils/containers/is_submapeq_of.h index e50e85e745..a9d52c871b 100644 --- a/lib/utils/include/utils/containers/is_submapeq_of.h +++ b/lib/utils/include/utils/containers/is_submapeq_of.h @@ -1,16 +1,16 @@ #ifndef _FLEXFLOW_UTILS_INCLUDE_UTILS_CONTAINERS_IS_SUBMAP_H #define _FLEXFLOW_UTILS_INCLUDE_UTILS_CONTAINERS_IS_SUBMAP_H -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/containers/restrict_keys.h" -#include +#include namespace FlexFlow { template -bool is_submapeq_of(std::unordered_map const &sub, - std::unordered_map const &m) { - return restrict_keys(m, unordered_keys(sub)) == sub; +bool is_submapeq_of(std::map const &sub, + std::map const &m) { + return restrict_keys(m, keys(sub)) == sub; } } // namespace FlexFlow diff --git a/lib/utils/include/utils/containers/is_superseteq_of.h b/lib/utils/include/utils/containers/is_superseteq_of.h index 23b16d92f9..7f1580e874 100644 --- a/lib/utils/include/utils/containers/is_superseteq_of.h +++ b/lib/utils/include/utils/containers/is_superseteq_of.h @@ -2,13 +2,13 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_IS_SUPERSETEQ_OF_H #include "utils/containers/is_subseteq_of.h" -#include +#include namespace FlexFlow { template -bool is_superseteq_of(std::unordered_set const &super, - std::unordered_set const &sub) { +bool is_superseteq_of(std::set const &super, + std::set const &sub) { return is_subseteq_of(sub, super); } diff --git a/lib/utils/include/utils/containers/keys.h b/lib/utils/include/utils/containers/keys.h index bd080b7087..e9421a6188 100644 --- a/lib/utils/include/utils/containers/keys.h +++ b/lib/utils/include/utils/containers/keys.h @@ -2,20 +2,10 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_KEYS_H #include -#include #include namespace FlexFlow { -template -std::set keys(std::unordered_map const &c) { - std::set result; - for (auto const &kv : c) { - result.insert(kv.first); - } - return result; -} - template std::set keys(std::map const &c) { std::set result; diff --git a/lib/utils/include/utils/containers/lift_optional_through_map.h b/lib/utils/include/utils/containers/lift_optional_through_map.h index 189d1b0519..03aee73200 100644 --- a/lib/utils/include/utils/containers/lift_optional_through_map.h +++ b/lib/utils/include/utils/containers/lift_optional_through_map.h @@ -6,16 +6,16 @@ #include "utils/containers/values.h" #include #include -#include +#include namespace FlexFlow { template -static std::optional> lift_optional_through_map( - std::unordered_map> const &m) { +static std::optional> lift_optional_through_map( + std::map> const &m) { ASSERT(!m.empty()); - std::unordered_multiset> m_values = values(m); + std::multiset> m_values = values(m); bool has_all_values = all_of(m_values, [](std::optional const &t) -> bool { return t.has_value(); diff --git a/lib/utils/include/utils/containers/lookup_in_map.h b/lib/utils/include/utils/containers/lookup_in_map.h index 339b13f042..81161df9e6 100644 --- a/lib/utils/include/utils/containers/lookup_in_map.h +++ b/lib/utils/include/utils/containers/lookup_in_map.h @@ -1,17 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_LOOKUP_IN_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_LOOKUP_IN_MAP_H -#include "utils/fmt/unordered_map.h" +#include "utils/fmt/map.h" #include "utils/containers/contains_key.h" #include #include -#include +#include #include namespace FlexFlow { template -std::function lookup_in_map(std::unordered_map const &m) { +std::function lookup_in_map(std::map const &m) { return [m](K const &key) -> V { if (!contains_key(m, key)) { PANIC("Key {} is not present in the underlying map {}", key, m); diff --git a/lib/utils/include/utils/containers/map_from_keys_and_values.h b/lib/utils/include/utils/containers/map_from_keys_and_values.h index 8a9e36ff4e..7a9c8f450d 100644 --- a/lib/utils/include/utils/containers/map_from_keys_and_values.h +++ b/lib/utils/include/utils/containers/map_from_keys_and_values.h @@ -4,17 +4,17 @@ #include "utils/containers/zip.h" #include #include -#include +#include namespace FlexFlow { template -std::unordered_map +std::map map_from_keys_and_values(std::vector const &keys, std::vector const &values) { ASSERT(keys.size() == values.size()); - std::unordered_map result; + std::map result; for (auto const &[k, v] : zip(keys, values)) { result.insert({k, v}); } diff --git a/lib/utils/include/utils/containers/map_keys2.h b/lib/utils/include/utils/containers/map_keys2.h index da68fe05b4..90320f848a 100644 --- a/lib/utils/include/utils/containers/map_keys2.h +++ b/lib/utils/include/utils/containers/map_keys2.h @@ -3,7 +3,7 @@ #include "utils/containers/keys.h" #include -#include +#include namespace FlexFlow { @@ -11,10 +11,10 @@ template > -std::unordered_map map_keys2(std::unordered_map const &m, +std::map map_keys2(std::map const &m, F const &f) { - std::unordered_map result; + std::map result; for (auto const &kv : m) { result.insert({f(kv.first, kv.second), kv.second}); } diff --git a/lib/utils/include/utils/containers/map_keys_with_value_merging.h b/lib/utils/include/utils/containers/map_keys_with_value_merging.h index 93c046f017..34efff1de5 100644 --- a/lib/utils/include/utils/containers/map_keys_with_value_merging.h +++ b/lib/utils/include/utils/containers/map_keys_with_value_merging.h @@ -3,6 +3,7 @@ #include "utils/containers/contains_key.h" #include +#include namespace FlexFlow { @@ -32,6 +33,32 @@ std::unordered_map map_keys_with_value_merging( return result; } +template > +std::map map_keys_with_value_merging( + std::map const &m, F &&key_func, MergeF &&merge_values) { + + std::map result; + + for (auto const &kv : m) { + K k = kv.first; + V v = kv.second; + + K2 k2 = key_func(k); + + if (contains_key(result, k2)) { + result.at(k2) = merge_values(result.at(k2), v); + } else { + result.insert({k2, v}); + } + } + + return result; +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/merge_maps_with.h b/lib/utils/include/utils/containers/merge_maps_with.h index eef6c9af67..31658056ce 100644 --- a/lib/utils/include/utils/containers/merge_maps_with.h +++ b/lib/utils/include/utils/containers/merge_maps_with.h @@ -3,7 +3,7 @@ #include "utils/containers/binary_merge_maps_with.h" #include "utils/containers/foldl.h" -#include +#include #include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/minimum.h b/lib/utils/include/utils/containers/minimum.h index 8bdd6ea985..bd17b50e74 100644 --- a/lib/utils/include/utils/containers/minimum.h +++ b/lib/utils/include/utils/containers/minimum.h @@ -1,17 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MINIMUM_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MINIMUM_H -#include "utils/exception.h" #include +#include namespace FlexFlow { template typename C::value_type minimum(C const &c) { - if (c.empty()) { - throw mk_runtime_error( - fmt::format("minimum expected non-empty container but received {}", c)); - } + ASSERT( + c.size() > 0, + "minimum expected non-empty container but received {}", c + ); return *std::min_element(c.begin(), c.end()); } diff --git a/lib/utils/include/utils/containers/multiset_union.h b/lib/utils/include/utils/containers/multiset_union.h index 6f2b2a7889..a2d39ce876 100644 --- a/lib/utils/include/utils/containers/multiset_union.h +++ b/lib/utils/include/utils/containers/multiset_union.h @@ -32,8 +32,8 @@ std::multiset multiset_union(std::multiset const &lhs, } template -std::unordered_multiset multiset_union(C const &c) { - std::unordered_multiset result; +std::multiset multiset_union(C const &c) { + std::multiset result; for (auto const &s : c) { for (T const &element : s) { result.insert(element); diff --git a/lib/utils/include/utils/containers/require_two_keys.h b/lib/utils/include/utils/containers/require_two_keys.h index 8da683c2f0..044e6ff21a 100644 --- a/lib/utils/include/utils/containers/require_two_keys.h +++ b/lib/utils/include/utils/containers/require_two_keys.h @@ -2,12 +2,12 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REQUIRE_TWO_KEYS_H #include -#include +#include namespace FlexFlow { template -std::pair require_two_keys(std::unordered_map const &m, +std::pair require_two_keys(std::map const &m, K const &k1, K const &k2) { ASSERT(k1 != k2); diff --git a/lib/utils/include/utils/containers/set_difference.h b/lib/utils/include/utils/containers/set_difference.h index d9c36df755..b4250b8c32 100644 --- a/lib/utils/include/utils/containers/set_difference.h +++ b/lib/utils/include/utils/containers/set_difference.h @@ -3,13 +3,13 @@ #include "utils/containers/contains.h" #include "utils/containers/filter.h" -#include +#include namespace FlexFlow { template -std::unordered_set set_difference(std::unordered_set const &l, - std::unordered_set const &r) { +std::set set_difference(std::set const &l, + std::set const &r) { return filter(l, [&](T const &element) { return !contains(r, element); }); } diff --git a/lib/utils/include/utils/containers/set_of.h b/lib/utils/include/utils/containers/set_of.h index 14658209aa..14cb0f8ee0 100644 --- a/lib/utils/include/utils/containers/set_of.h +++ b/lib/utils/include/utils/containers/set_of.h @@ -2,6 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_SET_OF_H #include +#include namespace FlexFlow { @@ -14,6 +15,16 @@ std::set set_of(C const &c) { return result; } +template +std::set> + set_of(std::map const &m) { + std::set> result; + for (auto const &[k, v] : m) { + result.insert({k, v}); + } + return result; +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/set_union.h b/lib/utils/include/utils/containers/set_union.h index cd29b1e02e..277b72a119 100644 --- a/lib/utils/include/utils/containers/set_union.h +++ b/lib/utils/include/utils/containers/set_union.h @@ -22,8 +22,8 @@ std::set set_union(std::set const &l, std::set const &r) { } template -std::unordered_set set_union(C const &sets) { - std::unordered_set result; +std::set set_union(C const &sets) { + std::set result; for (auto const &s : sets) { for (T const &element : s) { result.insert(element); diff --git a/lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h b/lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h new file mode 100644 index 0000000000..dbf380eb41 --- /dev/null +++ b/lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h @@ -0,0 +1,40 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_TRY_MERGE_NONDISJOINT_MAPS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_TRY_MERGE_NONDISJOINT_MAPS_H + +#include "utils/containers/contains_key.h" +#include +#include + +namespace FlexFlow { + +template +std::optional> + try_merge_nondisjoint_maps(std::map const &m1, + std::map const &m2) { + std::map result; + auto try_insert = [&](K const &k, V const &v) { + if (contains_key(result, k) && result.at(k) != v) { + return false; + } + result.insert({k, v}); + return true; + }; + + for (auto const &[k, v] : m1) { + if (!try_insert(k, v)) { + return std::nullopt; + } + } + + for (auto const &[k, v] : m2) { + if (!try_insert(k, v)) { + return std::nullopt; + } + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h b/lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h new file mode 100644 index 0000000000..ff916f3704 --- /dev/null +++ b/lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h @@ -0,0 +1,26 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_MAP_FROM_KEYS_AND_VALUES_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_MAP_FROM_KEYS_AND_VALUES_H + +#include "utils/containers/zip.h" +#include +#include +#include + +namespace FlexFlow { + +template +std::unordered_map + unordered_map_from_keys_and_values(std::vector const &keys, + std::vector const &values) { + ASSERT(keys.size() == values.size()); + + std::unordered_map result; + for (auto const &[k, v] : zip(keys, values)) { + result.insert({k, v}); + } + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/unstructured_exhaustive_relational_join.h b/lib/utils/include/utils/containers/unstructured_exhaustive_relational_join.h index 3894f4099d..635bdb0169 100644 --- a/lib/utils/include/utils/containers/unstructured_exhaustive_relational_join.h +++ b/lib/utils/include/utils/containers/unstructured_exhaustive_relational_join.h @@ -4,29 +4,29 @@ #include "utils/containers/transform.h" #include "utils/hash/pair.h" #include -#include +#include namespace FlexFlow { template -std::unordered_set> unstructured_exhaustive_relational_join( - std::unordered_set> const &lhs, - std::unordered_set> const &rhs) { - std::unordered_set> result; +std::set> unstructured_exhaustive_relational_join( + std::set> const &lhs, + std::set> const &rhs) { + std::set> result; - std::unordered_set lhs_ls = + std::set lhs_ls = transform(lhs, [](std::pair const &lc) { return lc.first; }); - std::unordered_set lhs_cs = + std::set lhs_cs = transform(lhs, [](std::pair const &lc) { return lc.second; }); - std::unordered_set rhs_cs = + std::set rhs_cs = transform(rhs, [](std::pair const &cr) { return cr.first; }); - std::unordered_set rhs_rs = + std::set rhs_rs = transform(rhs, [](std::pair const &cr) { return cr.second; }); ASSERT(lhs_cs == rhs_cs); - std::unordered_set result_ls; - std::unordered_set result_rs; + std::set result_ls; + std::set result_rs; for (auto const &[l, c1] : lhs) { for (auto const &[c2, r] : rhs) { diff --git a/lib/utils/include/utils/containers/value_all.h b/lib/utils/include/utils/containers/value_all.h index 5727bd8396..28a68b665d 100644 --- a/lib/utils/include/utils/containers/value_all.h +++ b/lib/utils/include/utils/containers/value_all.h @@ -19,7 +19,7 @@ std::vector value_all(std::vector> const &v) { } template -std::unordered_set value_all(std::unordered_set> const &v) { +std::set value_all(std::set> const &v) { return transform(v, [&](std::optional const &element) { return unwrap(element, [&] { throw mk_runtime_error(fmt::format( diff --git a/lib/utils/include/utils/containers/values.h b/lib/utils/include/utils/containers/values.h index 2a730ccc42..ff5c164243 100644 --- a/lib/utils/include/utils/containers/values.h +++ b/lib/utils/include/utils/containers/values.h @@ -1,13 +1,13 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_VALUES_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_VALUES_H -#include +#include namespace FlexFlow { template -std::unordered_multiset values(C const &c) { - std::unordered_multiset result; +std::multiset values(C const &c) { + std::multiset result; for (auto const &kv : c) { result.insert(kv.second); } diff --git a/lib/utils/include/utils/containers/vector_from_idx_map.h b/lib/utils/include/utils/containers/vector_from_idx_map.h index cd32385d1f..fae100dea9 100644 --- a/lib/utils/include/utils/containers/vector_from_idx_map.h +++ b/lib/utils/include/utils/containers/vector_from_idx_map.h @@ -4,7 +4,7 @@ #include "utils/containers/contains_key.h" #include "utils/nonnegative_int/nonnegative_int.h" #include -#include +#include #include namespace FlexFlow { @@ -24,6 +24,21 @@ std::optional> return result; } +template +std::optional> + vector_from_idx_map(std::map const &m) { + std::vector result; + + for (nonnegative_int i = 0_n; i < m.size(); i++) { + if (!contains_key(m, i)) { + return std::nullopt; + } + result.push_back(m.at(i)); + } + + return result; +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/without_nullopts.h b/lib/utils/include/utils/containers/without_nullopts.h index faf05090e0..3c6a2e74d8 100644 --- a/lib/utils/include/utils/containers/without_nullopts.h +++ b/lib/utils/include/utils/containers/without_nullopts.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_WITHOUT_NULLOPTS_H #include -#include +#include #include namespace FlexFlow { @@ -19,9 +19,9 @@ std::vector without_nullopts(std::vector> const &v) { } template -std::unordered_set - without_nullopts(std::unordered_set> const &s) { - std::unordered_set result; +std::set + without_nullopts(std::set> const &s) { + std::set result; for (std::optional const &t : s) { if (t.has_value()) { result.insert(t.value()); diff --git a/lib/utils/include/utils/containers/zip_values_strict_with.h b/lib/utils/include/utils/containers/zip_values_strict_with.h index 5fc4bb7f5b..1cd0165e1b 100644 --- a/lib/utils/include/utils/containers/zip_values_strict_with.h +++ b/lib/utils/include/utils/containers/zip_values_strict_with.h @@ -1,11 +1,11 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_WITH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_WITH_H -#include "utils/containers/generate_unordered_map.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/generate_map.h" +#include "utils/containers/keys.h" #include "utils/containers/require_same.h" #include -#include +#include namespace FlexFlow { @@ -14,14 +14,14 @@ template > -std::unordered_map - zip_values_strict_with(std::unordered_map const &m1, - std::unordered_map const &m2, +std::map + zip_values_strict_with(std::map const &m1, + std::map const &m2, F &&f) { - ASSERT(unordered_keys(m1) == unordered_keys(m2)); + ASSERT(keys(m1) == keys(m2)); - return generate_unordered_map(require_same(unordered_keys(m1), unordered_keys(m2)), + return generate_map(require_same(keys(m1), keys(m2)), [&](K const &k) -> Out { return f(m1.at(k), m2.at(k)); }); } diff --git a/lib/utils/include/utils/deduplicated_priority_queue.h b/lib/utils/include/utils/deduplicated_priority_queue.h index afad3f5889..b4a59f69a2 100644 --- a/lib/utils/include/utils/deduplicated_priority_queue.h +++ b/lib/utils/include/utils/deduplicated_priority_queue.h @@ -4,15 +4,14 @@ #include "utils/containers/contains.h" #include #include -#include +#include #include namespace FlexFlow { template , - typename Compare = std::less, - typename Hash = std::hash> + typename Compare = std::less> class DeduplicatedPriorityQueue { public: Elem const &top() const { @@ -28,14 +27,14 @@ class DeduplicatedPriorityQueue { } void push(Elem const &e) { - if (!contains(hashmap, e)) { + if (!contains(seen, e)) { impl.push(e); - hashmap.insert(e); + seen.insert(e); } } void pop() { - hashmap.erase(impl.top()); + seen.erase(impl.top()); impl.pop(); } @@ -51,7 +50,7 @@ class DeduplicatedPriorityQueue { private: std::priority_queue impl; - std::unordered_set hashmap; + std::set seen; }; } // namespace FlexFlow diff --git a/lib/utils/include/utils/disjoint_set.h b/lib/utils/include/utils/disjoint_set.h index 4810e5b29e..d8eea1d7a8 100644 --- a/lib/utils/include/utils/disjoint_set.h +++ b/lib/utils/include/utils/disjoint_set.h @@ -5,7 +5,7 @@ #include #include #include -#include +#include namespace FlexFlow { @@ -38,7 +38,7 @@ class m_disjoint_set { mapping[t] = std::nullopt; } } - mutable std::unordered_map, std::optional> mapping; + mutable std::map, std::optional> mapping; }; // Custom comparator for optional diff --git a/lib/utils/include/utils/dot/dot_file.h b/lib/utils/include/utils/dot/dot_file.h index eaef0832ae..e427beda22 100644 --- a/lib/utils/include/utils/dot/dot_file.h +++ b/lib/utils/include/utils/dot/dot_file.h @@ -11,8 +11,8 @@ #include #include #include -#include -#include +#include +#include #include namespace FlexFlow { @@ -28,9 +28,9 @@ class DotFile { size_t node_id = 0; size_t subgraph_id = 0; std::map node_ids; - std::unordered_map> subgraphs; - std::unordered_map> subgraph_children; - std::unordered_map> subgraph_parents; + std::map> subgraphs; + std::map> subgraph_children; + std::map> subgraph_parents; std::optional owned_fstream = std::nullopt; std::ostream *out = nullptr; diff --git a/lib/utils/include/utils/fmt.h b/lib/utils/include/utils/fmt.h index 378b2d07b9..f40212e618 100644 --- a/lib/utils/include/utils/fmt.h +++ b/lib/utils/include/utils/fmt.h @@ -5,7 +5,7 @@ #include "utils/type_traits_core.h" #include #include -#include +#include #include #include diff --git a/lib/utils/include/utils/fmt/unordered_map.h b/lib/utils/include/utils/fmt/unordered_map.h index 12faa64e32..ecd0257f4e 100644 --- a/lib/utils/include/utils/fmt/unordered_map.h +++ b/lib/utils/include/utils/fmt/unordered_map.h @@ -6,19 +6,19 @@ #include "utils/join_strings.h" #include #include -#include +#include #include namespace fmt { template struct formatter< - ::std::unordered_map, + ::std::map, Char, - std::enable_if_t>::value>> + std::enable_if_t>::value>> : formatter<::std::string> { template - auto format(::std::unordered_map const &m, FormatContext &ctx) const + auto format(::std::map const &m, FormatContext &ctx) const -> decltype(ctx.out()) { CHECK_FMTABLE(K); CHECK_FMTABLE(V); @@ -38,7 +38,7 @@ struct formatter< namespace FlexFlow { template -std::ostream &operator<<(std::ostream &s, std::unordered_map const &m) { +std::ostream &operator<<(std::ostream &s, std::map const &m) { CHECK_FMTABLE(K); CHECK_FMTABLE(V); diff --git a/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h b/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h index 9cf5d63210..9b6d981d30 100644 --- a/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h +++ b/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h @@ -6,20 +6,20 @@ #include "utils/full_binary_tree/binary_tree_path.dtg.h" #include "utils/full_binary_tree/binary_tree_path.h" #include "utils/full_binary_tree/visit.h" -#include +#include namespace FlexFlow { template -std::unordered_set find_paths_to_leaf( +std::set find_paths_to_leaf( Tree const &tree, FullBinaryTreeImplementation const &impl, Leaf const &needle) { - auto visitor = FullBinaryTreeVisitor, + auto visitor = FullBinaryTreeVisitor, Tree, Parent, Leaf>{ - [&](Parent const &parent) -> std::unordered_set { + [&](Parent const &parent) -> std::set { return set_union( transform( find_paths_to_leaf(impl.get_left_child(parent), impl, needle), @@ -32,7 +32,7 @@ std::unordered_set find_paths_to_leaf( return nest_inside_right_child(path); })); }, - [&](Leaf const &leaf) -> std::unordered_set { + [&](Leaf const &leaf) -> std::set { if (leaf == needle) { return {binary_tree_root_path()}; } else { diff --git a/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h b/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h index 822acfe9ee..f3796d3f8e 100644 --- a/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h +++ b/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h @@ -7,19 +7,19 @@ #include "utils/full_binary_tree/binary_tree_path.h" #include "utils/full_binary_tree/visit.h" #include "utils/overload.h" -#include +#include namespace FlexFlow { template -std::unordered_set get_all_leaf_paths( +std::set get_all_leaf_paths( Tree const &tree, FullBinaryTreeImplementation const &impl) { - auto visitor = FullBinaryTreeVisitor, + auto visitor = FullBinaryTreeVisitor, Tree, Parent, Leaf>{ - [&](Parent const &parent) -> std::unordered_set { + [&](Parent const &parent) -> std::set { return set_union( transform(get_all_leaf_paths(impl.get_left_child(parent), impl), [](BinaryTreePath const &path) { @@ -30,7 +30,7 @@ std::unordered_set get_all_leaf_paths( return nest_inside_right_child(path); })); }, - [&](Leaf const &leaf) -> std::unordered_set { + [&](Leaf const &leaf) -> std::set { return {binary_tree_root_path()}; }, }; diff --git a/lib/utils/include/utils/full_binary_tree/get_leaves.h b/lib/utils/include/utils/full_binary_tree/get_leaves.h index 8f9d8e919f..929b5327f4 100644 --- a/lib/utils/include/utils/full_binary_tree/get_leaves.h +++ b/lib/utils/include/utils/full_binary_tree/get_leaves.h @@ -4,23 +4,23 @@ #include "utils/containers/multiset_union.h" #include "utils/full_binary_tree/full_binary_tree_visitor.dtg.h" #include "utils/full_binary_tree/visit.h" -#include +#include namespace FlexFlow { template -std::unordered_multiset +std::multiset get_leaves(Tree const &tree, FullBinaryTreeImplementation const &impl) { auto visitor = - FullBinaryTreeVisitor, Tree, Parent, Leaf>{ - [&](Parent const &parent) -> std::unordered_multiset { + FullBinaryTreeVisitor, Tree, Parent, Leaf>{ + [&](Parent const &parent) -> std::multiset { return multiset_union( get_leaves(impl.get_left_child(parent), impl), get_leaves(impl.get_right_child(parent), impl)); }, - [](Leaf const &leaf) -> std::unordered_multiset { + [](Leaf const &leaf) -> std::multiset { return {leaf}; }, }; diff --git a/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h b/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h index e3947605c0..15df552bb8 100644 --- a/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h +++ b/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h @@ -1,39 +1,39 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_FULL_BINARY_TREE_GET_PATH_TO_LEAF_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_FULL_BINARY_TREE_GET_PATH_TO_LEAF_MAP_H -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/map_keys.h" #include "utils/containers/multiset_union.h" #include "utils/full_binary_tree/binary_tree_path.dtg.h" #include "utils/full_binary_tree/binary_tree_path.h" #include "utils/full_binary_tree/full_binary_tree_visitor.dtg.h" #include "utils/full_binary_tree/visit.h" -#include +#include namespace FlexFlow { template -std::unordered_map get_path_to_leaf_map( +std::map get_path_to_leaf_map( Tree const &tree, FullBinaryTreeImplementation const &impl) { - auto visitor = FullBinaryTreeVisitor, + auto visitor = FullBinaryTreeVisitor, Tree, Parent, Leaf>{ - [&](Parent const &parent) -> std::unordered_map { - std::unordered_map left_map = map_keys( + [&](Parent const &parent) -> std::map { + std::map left_map = map_keys( get_path_to_leaf_map(impl.get_left_child(parent), impl), [](BinaryTreePath const &p) { return nest_inside_left_child(p); }); - std::unordered_map right_map = map_keys( + std::map right_map = map_keys( get_path_to_leaf_map(impl.get_right_child(parent), impl), [](BinaryTreePath const &p) { return nest_inside_right_child(p); }); - return binary_merge_disjoint_unordered_maps(left_map, right_map); + return binary_merge_disjoint_maps(left_map, right_map); }, - [](Leaf const &leaf) -> std::unordered_map { - return std::unordered_map{ + [](Leaf const &leaf) -> std::map { + return std::map{ {binary_tree_root_path(), leaf}, }; }, diff --git a/lib/utils/include/utils/graph/algorithms.h b/lib/utils/include/utils/graph/algorithms.h index f78a0225d5..c54a8418cd 100644 --- a/lib/utils/include/utils/graph/algorithms.h +++ b/lib/utils/include/utils/graph/algorithms.h @@ -14,8 +14,8 @@ std::vector add_nodes(Graph &, int); std::vector add_nodes(UndirectedGraph &, int); std::vector add_nodes(DiGraph &, int); -std::unordered_set query_nodes(GraphView const &, - std::unordered_set const &); +std::set query_nodes(GraphView const &, + std::set const &); void remove_node(DiGraph &, Node const &); void remove_node(UndirectedGraph &, Node const &); @@ -29,51 +29,51 @@ void add_edges(DiGraph &, std::vector const &); void add_edges(UndirectedGraph &, std::vector const &); void add_edges(DiGraph &, std::initializer_list); void add_edges(UndirectedGraph &, std::initializer_list); -void add_edges(DiGraph &, std::unordered_set const &); -void add_edges(UndirectedGraph &, std::unordered_set const &); +void add_edges(DiGraph &, std::set const &); +void add_edges(UndirectedGraph &, std::set const &); bool contains_node(GraphView const &, Node const &); bool contains_edge(DiGraphView const &, DirectedEdge const &); bool contains_edge(UndirectedGraphView const &, UndirectedEdge const &); -void remove_edges(DiGraph &, std::unordered_set const &); +void remove_edges(DiGraph &, std::set const &); void remove_edges(UndirectedGraph &, - std::unordered_set const &); + std::set const &); -std::unordered_set get_edges(UndirectedGraphView const &); +std::set get_edges(UndirectedGraphView const &); -std::unordered_set get_node_edges(UndirectedGraphView const &, +std::set get_node_edges(UndirectedGraphView const &, Node const &); -std::unordered_set get_node_edges(UndirectedGraphView const &, +std::set get_node_edges(UndirectedGraphView const &, Node const &); -std::unordered_set +std::set get_node_edges(UndirectedGraphView const &, - std::unordered_set const &); + std::set const &); -std::unordered_set get_neighbors(UndirectedGraphView const &, +std::set get_neighbors(UndirectedGraphView const &, Node const &); -std::unordered_set get_neighbors(DiGraphView const &, Node const &); +std::set get_neighbors(DiGraphView const &, Node const &); std::vector get_dfs_ordering(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); std::vector get_unchecked_dfs_ordering(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); std::vector get_bfs_ordering(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); std::vector get_unchecked_topological_ordering(DiGraphView const &); -std::unordered_set +std::set get_transitive_reduction_delta(DiGraphView const &); UndirectedGraphView get_subgraph(UndirectedGraphView const &, - std::unordered_set const &); -DiGraphView get_subgraph(DiGraphView const &, std::unordered_set const &); + std::set const &); +DiGraphView get_subgraph(DiGraphView const &, std::set const &); DiGraphView join(DiGraphView const &lhs, DiGraphView const &rhs); UndirectedGraphView join(UndirectedGraphView const &lhs, @@ -82,7 +82,7 @@ UndirectedGraphView join(UndirectedGraphView const &lhs, DiGraphView flipped(DiGraphView const &); DiGraphView with_added_edges(DiGraphView const &, - std::unordered_set const &); + std::set const &); UndirectedGraphView as_undirected(DiGraphView const &); DiGraphView as_digraph(UndirectedGraphView const &); diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms.h index d50facee57..58de430af0 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms.h @@ -6,14 +6,14 @@ namespace FlexFlow { -std::unordered_set get_edges(DataflowGraphView const &); +std::set get_edges(DataflowGraphView const &); std::vector get_input_values(DataflowGraphView const &, Node const &); std::vector get_dataflow_inputs(DataflowGraphView const &, Node const &); std::vector get_outputs(DataflowGraphView const &, Node const &); -std::unordered_set +std::set get_all_dataflow_outputs(DataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_data.dtg.toml b/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_data.dtg.toml index a957ad82c7..d45e590245 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_data.dtg.toml +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_data.dtg.toml @@ -3,6 +3,7 @@ name = "DataflowGraphData" type = "struct" features = [ "eq", + "ord", "hash", "fmt", "json", @@ -12,22 +13,22 @@ includes = [ "utils/graph/node/node.dtg.h", "utils/graph/dataflow_graph/dataflow_edge.dtg.h", "utils/graph/dataflow_graph/dataflow_output.dtg.h", - "", + "", ] src_includes = [ - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "nodes" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::DataflowEdge>" +type = "std::set<::FlexFlow::DataflowEdge>" [[fields]] name = "outputs" -type = "std::unordered_set<::FlexFlow::DataflowOutput>" +type = "std::set<::FlexFlow::DataflowOutput>" diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_isomorphism.dtg.toml b/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_isomorphism.dtg.toml index 78efd9fba2..24fc984fa5 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_isomorphism.dtg.toml +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/dataflow_graph_isomorphism.dtg.toml @@ -3,6 +3,7 @@ name = "DataflowGraphIsomorphism" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h index dda69ea69a..effce9f3cd 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h @@ -6,7 +6,7 @@ namespace FlexFlow { -std::unordered_set +std::set find_isomorphisms(DataflowGraphView const &, DataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.h index de7ead8fb6..2f3dfc2787 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_dataflow_edges_from_node_to_node( +std::set get_dataflow_edges_from_node_to_node( DataflowGraphView const &g, Node const &src, Node const &dst); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h index a4cd27bf9d..aa4408bc4f 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h @@ -7,9 +7,9 @@ namespace FlexFlow { std::vector get_incoming_edges(DataflowGraphView const &, Node const &); -std::unordered_set +std::set get_incoming_edges(DataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h index a8b5efe66e..d0f03751f6 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h @@ -5,11 +5,11 @@ namespace FlexFlow { -std::unordered_set get_outgoing_edges(DataflowGraphView const &, +std::set get_outgoing_edges(DataflowGraphView const &, Node const &); -std::unordered_set +std::set get_outgoing_edges(DataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h index 2ed0bc02be..0ec7d26796 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h @@ -5,9 +5,9 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_incoming_edges(DataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h index f26ea20473..6a4898c341 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h @@ -5,9 +5,9 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_outgoing_edges(DataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.h index 09135acb51..d1c00e4085 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.h @@ -6,7 +6,7 @@ namespace FlexFlow { -std::unordered_set get_transitive_reduced_edges_across_split( +std::set get_transitive_reduced_edges_across_split( TransitiveReducedDataflowGraphView const &, BinarySeriesSplit const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.h index 00b213845d..1494fc2f90 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.h @@ -6,7 +6,7 @@ namespace FlexFlow { -std::unordered_set get_transitive_reduced_outputs_across_split( +std::set get_transitive_reduced_outputs_across_split( TransitiveReducedDataflowGraphView const &, BinarySeriesSplit const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/split_boundary_nodes.dtg.toml b/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/split_boundary_nodes.dtg.toml index 6b23df3f4f..1c0443e20c 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/split_boundary_nodes.dtg.toml +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/split_boundary_nodes.dtg.toml @@ -9,18 +9,18 @@ features = [ includes = [ "utils/graph/node/node.dtg.h", - "", + "", ] src_includes = [ - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "pre_split_boundary" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "post_split_boundary" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.h index b12e20124f..75901b5d79 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.h @@ -12,11 +12,11 @@ struct ViewDataflowGraphAsOpenDataflowGraph final ViewDataflowGraphAsOpenDataflowGraph() = delete; ViewDataflowGraphAsOpenDataflowGraph(DataflowGraphView const &); - std::unordered_set query_nodes(NodeQuery const &) const override; - std::unordered_set + std::set query_nodes(NodeQuery const &) const override; + std::set query_outputs(DataflowOutputQuery const &) const override; - std::unordered_set get_inputs() const override; - std::unordered_set + std::set get_inputs() const override; + std::set query_edges(OpenDataflowEdgeQuery const &) const override; ViewDataflowGraphAsOpenDataflowGraph *clone() const override; diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.h index 39a38e88aa..e4aaba5710 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.h @@ -12,10 +12,10 @@ struct ViewFromDataflowGraphData final : virtual public IDataflowGraphView { public: explicit ViewFromDataflowGraphData(DataflowGraphData const &); - std::unordered_set query_nodes(NodeQuery const &query) const override; - std::unordered_set + std::set query_nodes(NodeQuery const &query) const override; + std::set query_edges(DataflowEdgeQuery const &query) const override; - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &query) const override; ViewFromDataflowGraphData *clone() const override; diff --git a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h index 043187208c..5f3890e42d 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h +++ b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h @@ -17,9 +17,9 @@ struct DataflowGraph : virtual public DataflowGraphView { std::vector const &inputs, std::vector const &outputs); - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(DataflowEdgeQuery const &) const; - std::unordered_set + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(DataflowEdgeQuery const &) const; + std::set query_outputs(DataflowOutputQuery const &) const; template diff --git a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h index 61b914c6e7..dce09008d9 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h @@ -12,9 +12,9 @@ struct DataflowGraphView : virtual public DiGraphView { DataflowGraphView(DataflowGraphView const &) = default; DataflowGraphView &operator=(DataflowGraphView const &) = default; - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(DataflowEdgeQuery const &) const; - std::unordered_set + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(DataflowEdgeQuery const &) const; + std::set query_outputs(DataflowOutputQuery const &) const; template diff --git a/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.dtg.toml b/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.dtg.toml index 2bd1068f4f..c8afe605e1 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.dtg.toml +++ b/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.dtg.toml @@ -15,7 +15,7 @@ includes = [ ] src_includes = [ - "utils/fmt/unordered_set.h", + "utils/fmt/set.h", ] [[fields]] diff --git a/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.h b/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.h index fc1a222f1e..ab0c323945 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.h +++ b/lib/utils/include/utils/graph/dataflow_graph/dataflow_output_query.h @@ -11,9 +11,9 @@ DataflowOutputQuery dataflow_output_query_none(); bool dataflow_output_query_includes_dataflow_output(DataflowOutputQuery const &, DataflowOutput const &); DataflowOutputQuery dataflow_output_query_for_output(DataflowOutput const &); -std::unordered_set +std::set apply_dataflow_output_query(DataflowOutputQuery const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/i_dataflow_graph_view.h b/lib/utils/include/utils/graph/dataflow_graph/i_dataflow_graph_view.h index 9166beab01..5c8a833967 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/i_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/dataflow_graph/i_dataflow_graph_view.h @@ -10,12 +10,12 @@ namespace FlexFlow { struct IDataflowGraphView : virtual public IDiGraphView { - virtual std::unordered_set + virtual std::set query_edges(DataflowEdgeQuery const &) const = 0; - virtual std::unordered_set + virtual std::set query_outputs(DataflowOutputQuery const &) const = 0; - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override final; virtual ~IDataflowGraphView() = default; diff --git a/lib/utils/include/utils/graph/digraph/algorithms/apply_contraction.h b/lib/utils/include/utils/graph/digraph/algorithms/apply_contraction.h index 792a8376c3..4a28689706 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/apply_contraction.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/apply_contraction.h @@ -6,7 +6,7 @@ namespace FlexFlow { DiGraphView apply_contraction(DiGraphView const &g, - std::unordered_map const &nodes); + std::map const &nodes); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/calculate_topo_rank.h b/lib/utils/include/utils/graph/digraph/algorithms/calculate_topo_rank.h index d19e1b4b48..05312fb237 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/calculate_topo_rank.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/calculate_topo_rank.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map calculate_topo_rank(DiGraphView const &); +std::map calculate_topo_rank(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/bipartite_component.dtg.toml b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/bipartite_component.dtg.toml index be7c58a9f6..cd17b9c51f 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/bipartite_component.dtg.toml +++ b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/bipartite_component.dtg.toml @@ -3,24 +3,25 @@ name = "BipartiteComponent" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] includes = [ - "", + "", "utils/graph/node/node.dtg.h", ] src_includes = [ - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h", + "utils/fmt/set.h", + "utils/hash/set.h", ] [[fields]] name = "head_nodes" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "tail_nodes" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" diff --git a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.dtg.toml b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.dtg.toml index 386466d62e..ec4cb9bee9 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.dtg.toml +++ b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.dtg.toml @@ -12,10 +12,10 @@ includes = [ ] src_includes = [ - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h", + "utils/fmt/set.h", + "utils/hash/set.h", ] [[fields]] name = "subgraphs" -type = "std::unordered_set<::FlexFlow::BipartiteComponent>" +type = "std::set<::FlexFlow::BipartiteComponent>" diff --git a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.h b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.h index 475ad0b125..0bb71d9666 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.h @@ -10,9 +10,9 @@ std::optional get_component_containing_node_in_head( CompleteBipartiteCompositeDecomposition const &, Node const &); std::optional get_component_containing_node_in_tail( CompleteBipartiteCompositeDecomposition const &, Node const &); -std::unordered_set> +std::set> get_head_subcomponents(CompleteBipartiteCompositeDecomposition const &); -std::unordered_set> +std::set> get_tail_subcomponents(CompleteBipartiteCompositeDecomposition const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.h b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.h index 3066886e37..5735edb975 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.h @@ -7,7 +7,7 @@ namespace FlexFlow { bool is_complete_bipartite_digraph(DiGraphView const &); bool is_complete_bipartite_digraph(DiGraphView const &, - std::unordered_set const &srcs); + std::set const &srcs); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h b/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h index a5de0ca61f..9fe8dc6fb7 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h @@ -12,9 +12,9 @@ struct ContractNodeView : public IDiGraphView { Node const &into) : g(g), from(removed), to(into) {} - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; ContractNodeView *clone() const override; diff --git a/lib/utils/include/utils/graph/digraph/algorithms/flipped.h b/lib/utils/include/utils/graph/digraph/algorithms/flipped.h index de13b125a3..a19dac4908 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/flipped.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/flipped.h @@ -10,9 +10,9 @@ struct FlippedView : public IDiGraphView { FlippedView() = delete; explicit FlippedView(DiGraphView const &); - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; FlippedView *clone() const override; diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_ancestors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_ancestors.h index f8d2b84845..33ef434d5a 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_ancestors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_ancestors.h @@ -12,7 +12,7 @@ namespace FlexFlow { * @note `n` is not considered to be its own ancestor, and is thus not * included in the returned set. **/ -std::unordered_set get_ancestors(DiGraphView const &g, Node const &n); +std::set get_ancestors(DiGraphView const &g, Node const &n); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h b/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h index 2e4c9eb5a3..6831f4d321 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h @@ -12,7 +12,7 @@ namespace FlexFlow { * @note `starting_node` is not considered to be its own descendant, and is thus * not included in the returned set. **/ -std::unordered_set get_descendants(DiGraphView const &g, +std::set get_descendants(DiGraphView const &g, Node const &starting_node); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h index 96e8864bc1..6f5bc6fcef 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h @@ -12,7 +12,7 @@ namespace FlexFlow { * dominates itself. * */ -std::unordered_set get_dominators(DiGraphView const &, Node const &); +std::set get_dominators(DiGraphView const &, Node const &); /** * @brief Returns the intersection of the dominators of the given set of nodes. @@ -21,8 +21,8 @@ std::unordered_set get_dominators(DiGraphView const &, Node const &); * that all edges belonging to the set of nodes now pass through a single * unified node). */ -std::unordered_set get_dominators(DiGraphView const &, - std::unordered_set const &); +std::set get_dominators(DiGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h index 51737834a8..b6845f3041 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_edges.h b/lib/utils/include/utils/graph/digraph/algorithms/get_edges.h index d4340f6a0b..91e0f6ac33 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_edges.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_edges.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_edges(DiGraphView const &); +std::set get_edges(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h b/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h index 240fc66426..06ab2953da 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h @@ -4,10 +4,10 @@ #include "utils/graph/digraph/digraph_view.h" namespace FlexFlow { -std::unordered_set +std::set get_edges_from_subgraph_to_subgraph(DiGraphView const &, - std::unordered_set const &, - std::unordered_set const &); + std::set const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h index c6adc83470..60fcbecc36 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_imm_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominator.h b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominator.h index 704f7025f4..99a3d2329e 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominator.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominator.h @@ -7,7 +7,7 @@ namespace FlexFlow { std::optional get_imm_post_dominator(DiGraphView const &, Node const &); std::optional get_imm_post_dominator(DiGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominators_map.h index 5a49113f89..186271daa1 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_post_dominators_map.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_imm_post_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h index 26b5a4f371..53367b1743 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h @@ -5,10 +5,10 @@ namespace FlexFlow { -std::unordered_set get_incoming_edges(DiGraphView const &, +std::set get_incoming_edges(DiGraphView const &, Node const &); -std::unordered_map> - get_incoming_edges(DiGraphView const &, std::unordered_set const &); +std::map> + get_incoming_edges(DiGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_initial_nodes.h b/lib/utils/include/utils/graph/digraph/algorithms/get_initial_nodes.h index bf907bac52..f2ff7bf1f3 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_initial_nodes.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_initial_nodes.h @@ -8,7 +8,7 @@ namespace FlexFlow { /** * @brief Returns the set of nodes in the graph with no incoming edges. */ -std::unordered_set get_initial_nodes(DiGraphView const &graph); +std::set get_initial_nodes(DiGraphView const &graph); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h b/lib/utils/include/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h index 16ab2798a5..ec98df0cea 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h @@ -3,7 +3,7 @@ #include "utils/graph/digraph/digraph_view.h" #include "utils/nonnegative_int/nonnegative_int.h" -#include +#include namespace FlexFlow { @@ -11,26 +11,26 @@ namespace FlexFlow { * @brief Computes the longest path lengths from the root in directed acyclic * graph. * - * @return std::unordered_map For each node n, returns the length + * @return std::map For each node n, returns the length * (i.e. number of nodes) of the longest path from the root to n. * * @note The root has a path length of 1. g must be acyclic. */ -std::unordered_map +std::map get_longest_path_lengths_from_root(DiGraphView const &g); /** * @brief Computes the weighted longest path lengths from the root in a directed * acyclic graph. * - * @return std::unordered_map For each node n, returns the length + * @return std::map For each node n, returns the length * (i.e. the sum of the weights of all the nodes) of the longest path from the * root to n. * * @note The root has a path length equal to its weight. g must be acyclic. */ -std::unordered_map get_weighted_longest_path_lengths_from_root( - DiGraphView const &g, std::unordered_map const &node_costs); +std::map get_weighted_longest_path_lengths_from_root( + DiGraphView const &g, std::map const &node_costs); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_lowest_common_ancestors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_lowest_common_ancestors.h index 60b0d32ae2..266aab5274 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_lowest_common_ancestors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_lowest_common_ancestors.h @@ -32,9 +32,9 @@ namespace FlexFlow { * In a Directed Acyclic Graph, a set of nodes can have no LCA, a unique node as * LCA, or a set of nodes as LCA. */ -std::optional> +std::optional> get_lowest_common_ancestors(DiGraphView const &g, - std::unordered_set const &nodes); + std::set const &nodes); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.h b/lib/utils/include/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.h index 0a701e13a1..c8752fcdbd 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.h @@ -5,7 +5,7 @@ namespace FlexFlow { -Node get_node_with_greatest_topo_rank(std::unordered_set const &, +Node get_node_with_greatest_topo_rank(std::set const &, DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h b/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h index 34ca643517..b3880d9322 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h @@ -5,10 +5,10 @@ namespace FlexFlow { -std::unordered_set get_outgoing_edges(DiGraphView const &, +std::set get_outgoing_edges(DiGraphView const &, Node const &); -std::unordered_map> - get_outgoing_edges(DiGraphView const &, std::unordered_set const &); +std::map> + get_outgoing_edges(DiGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators.h b/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators.h index d1a93ed834..6b76b76367 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_post_dominators(DiGraphView const &, Node const &); +std::set get_post_dominators(DiGraphView const &, Node const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h index e4f310db22..034294df8a 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_post_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h index b8e83268c7..4f11595430 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h @@ -5,11 +5,11 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_predecessors(DiGraphView const &); -std::unordered_set get_predecessors(DiGraphView const &, Node const &); -std::unordered_map> - get_predecessors(DiGraphView const &, std::unordered_set const &); +std::set get_predecessors(DiGraphView const &, Node const &); +std::map> + get_predecessors(DiGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h index aa94a44daf..0e957d2df4 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_strict_dominators(DiGraphView const &, +std::set get_strict_dominators(DiGraphView const &, Node const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h index 1e8b4d3b4f..7fb71b9c8d 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_strict_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h index 6d98c5c20d..74fbf51ea0 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h @@ -5,9 +5,9 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_outgoing_edges(DiGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h index 2c48d327c4..90372193fc 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h @@ -5,9 +5,9 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_successors(DiGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h index 0fa6e44c3d..23195bb256 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h @@ -5,11 +5,11 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_successors(DiGraphView const &); -std::unordered_set get_successors(DiGraphView const &, Node const &); -std::unordered_map> - get_successors(DiGraphView const &, std::unordered_set const &); +std::set get_successors(DiGraphView const &, Node const &); +std::map> + get_successors(DiGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_terminal_nodes.h b/lib/utils/include/utils/graph/digraph/algorithms/get_terminal_nodes.h index 3c8620e134..2b9f4787eb 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_terminal_nodes.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_terminal_nodes.h @@ -8,7 +8,7 @@ namespace FlexFlow { /** * @brief Returns the set of nodes in the graph with no outgoing edges. */ -std::unordered_set get_terminal_nodes(DiGraphView const &graph); +std::set get_terminal_nodes(DiGraphView const &graph); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h b/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h index 4d0e9a51d8..9d2343d923 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set> +std::set> get_weakly_connected_components(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h b/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h index ad11c6388c..5d71737e1f 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h @@ -9,17 +9,17 @@ namespace FlexFlow { struct DirectedEdgeMaskView final : public IDiGraphView { DirectedEdgeMaskView() = delete; explicit DirectedEdgeMaskView(DiGraphView const &, - std::unordered_set const &); + std::set const &); - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; DirectedEdgeMaskView *clone() const override; private: DiGraphView g; - std::unordered_set edge_mask; + std::set edge_mask; }; DiGraph transitive_reduction(DiGraphView const &); diff --git a/lib/utils/include/utils/graph/digraph/digraph.h b/lib/utils/include/utils/graph/digraph/digraph.h index 3d320b1c06..3c33d24399 100644 --- a/lib/utils/include/utils/graph/digraph/digraph.h +++ b/lib/utils/include/utils/graph/digraph/digraph.h @@ -24,8 +24,8 @@ struct DiGraph : virtual DiGraphView { void add_edge(Edge const &); void remove_edge(Edge const &); - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(EdgeQuery const &) const; + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(EdgeQuery const &) const; template static typename std::enable_if::value, diff --git a/lib/utils/include/utils/graph/digraph/digraph_view.h b/lib/utils/include/utils/graph/digraph/digraph_view.h index 0380751c55..1dc34846ce 100644 --- a/lib/utils/include/utils/graph/digraph/digraph_view.h +++ b/lib/utils/include/utils/graph/digraph/digraph_view.h @@ -16,8 +16,8 @@ struct DiGraphView : virtual public GraphView { DiGraphView(DiGraphView const &) = default; DiGraphView &operator=(DiGraphView const &) = default; - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(EdgeQuery const &) const; + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(EdgeQuery const &) const; template static typename std::enable_if::value, diff --git a/lib/utils/include/utils/graph/digraph/i_digraph_view.h b/lib/utils/include/utils/graph/digraph/i_digraph_view.h index f626a93805..baa3386a78 100644 --- a/lib/utils/include/utils/graph/digraph/i_digraph_view.h +++ b/lib/utils/include/utils/graph/digraph/i_digraph_view.h @@ -18,7 +18,7 @@ struct IDiGraphView : virtual public IGraphView { IDiGraphView(IDiGraphView const &) = delete; IDiGraphView &operator=(IDiGraphView const &) = delete; - virtual std::unordered_set query_edges(EdgeQuery const &) const = 0; + virtual std::set query_edges(EdgeQuery const &) const = 0; virtual ~IDiGraphView() = default; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(IDiGraphView); diff --git a/lib/utils/include/utils/graph/graph_split.dtg.toml b/lib/utils/include/utils/graph/graph_split.dtg.toml index 05624318c9..f653edef05 100644 --- a/lib/utils/include/utils/graph/graph_split.dtg.toml +++ b/lib/utils/include/utils/graph/graph_split.dtg.toml @@ -8,16 +8,16 @@ features = [ ] includes = [ - "", + "", "utils/graph/node/node.dtg.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "first" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "second" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" diff --git a/lib/utils/include/utils/graph/index.dox b/lib/utils/include/utils/graph/index.dox index 75793b2ed4..ee4e851027 100644 --- a/lib/utils/include/utils/graph/index.dox +++ b/lib/utils/include/utils/graph/index.dox @@ -118,7 +118,7 @@ These GraphView objects represent read-only (i.e., immutable) graphs. Similar to C++'s \c const semantics, Graphs can be coerced \ref graph-footnote-2 "[2]" to GraphViews, but not the other way around. To transform a GraphView (e.g., \ref DiGraphView) to a Graph (e.g., \ref DiGraph), we can perform an explicit copy with a materialize function (e.g., \ref materialize_digraph_view). Both Graph and GraphView types follow normal value semantics. -This may seem wasteful (oftentimes graphs are large objects that are passed around via reference to avoid making additional copies), but the Graph and GraphView types internally implement copy-on-write optimizations to only perform the minimum number of actual copies while maintaining immutability and lifetime safety (if you allocate a \ref DiGraph use for example \ref "get_subgraph(DiGraphView const &, std::unordered_set const &)" "get_subgraph" to get a \ref DiGraphView representing a part of this graph, modifications to the underlying \ref DiGraph will not be mirrored in the \ref DiGraphView and the \ref DiGraphView will remain valid even after the base \ref DiGraph leaves scope. +This may seem wasteful (oftentimes graphs are large objects that are passed around via reference to avoid making additional copies), but the Graph and GraphView types internally implement copy-on-write optimizations to only perform the minimum number of actual copies while maintaining immutability and lifetime safety (if you allocate a \ref DiGraph use for example \ref "get_subgraph(DiGraphView const &, std::set const &)" "get_subgraph" to get a \ref DiGraphView representing a part of this graph, modifications to the underlying \ref DiGraph will not be mirrored in the \ref DiGraphView and the \ref DiGraphView will remain valid even after the base \ref DiGraph leaves scope. At this point, however, we still have not discussed how to create a graph. The user-facing graph interface is intentionally separated from the underlying graph representations, so representations can be changed without requiring any user-side code modifications besides the choice of which implementation to use. @@ -209,7 +209,7 @@ While the interfaces of these graphs differ slightly from the core graph variant Note that all of the labelled graph types require that each element of the labelled types have a label, which is enforced via the interfaces they provide. Partial labelling can be implement via wrapping the label type in \c std::optional. Interacting with \c Node and \c Edge objects is still necessary to use the labelled graph types: intuitively the labelled graph types can be thought of as a pair of a core graph variant and a hash map the maps nodes/edges to labels. -As such, the labelled graph types provide the typical \ref LabelledDataflowGraph::at method (as on \c std::unordered_map \ref graph-footnote-3 "[3]") and can be coerced to their underlying core graph variants. +As such, the labelled graph types provide the typical \ref LabelledDataflowGraph::at method (as on \c std::map \ref graph-footnote-3 "[3]") and can be coerced to their underlying core graph variants. \section graph-internals Internals diff --git a/lib/utils/include/utils/graph/instances/adjacency_digraph.h b/lib/utils/include/utils/graph/instances/adjacency_digraph.h index b6ae76fd55..783d6e381c 100644 --- a/lib/utils/include/utils/graph/instances/adjacency_digraph.h +++ b/lib/utils/include/utils/graph/instances/adjacency_digraph.h @@ -3,8 +3,8 @@ #include "utils/graph/digraph/digraph.h" #include "utils/graph/node/node_source.h" -#include -#include +#include +#include namespace FlexFlow { @@ -17,19 +17,19 @@ class AdjacencyDiGraph : public IDiGraph { void remove_node_unsafe(Node const &) override; void add_edge(Edge const &) override; void remove_edge(Edge const &) override; - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; AdjacencyDiGraph *clone() const override; private: AdjacencyDiGraph( NodeSource const &node_source, - std::unordered_map> const &adjacency); + std::map> const &adjacency); NodeSource node_source; - std::unordered_map> adjacency; + std::map> adjacency; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(AdjacencyDiGraph); diff --git a/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h b/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h index 8b2a0431be..49ea58a659 100644 --- a/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h +++ b/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h @@ -15,8 +15,8 @@ struct AdjacencyMultiDiGraph final : public IMultiDiGraph { MultiDiEdge add_edge(Node const &, Node const &) override; void remove_node(Node const &) override; void remove_edge(MultiDiEdge const &) override; - std::unordered_set query_nodes(NodeQuery const &) const override; - std::unordered_set + std::set query_nodes(NodeQuery const &) const override; + std::set query_edges(MultiDiEdgeQuery const &) const override; Node get_multidiedge_src(MultiDiEdge const &) const override; Node get_multidiedge_dst(MultiDiEdge const &) const override; @@ -28,18 +28,18 @@ struct AdjacencyMultiDiGraph final : public IMultiDiGraph { AdjacencyMultiDiGraph( NodeSource const &, MultiDiEdgeSource const &, - std::unordered_map< + std::map< Node, - std::unordered_map>> const &, - std::unordered_map> const &); + std::map>> const &, + std::map> const &); private: NodeSource node_source; MultiDiEdgeSource edge_source; - std::unordered_map>> + std::map>> adjacency; - std::unordered_map> edge_nodes; + std::map> edge_nodes; }; } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h b/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h index 8630277fe8..189c19d5f8 100644 --- a/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h +++ b/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h @@ -14,9 +14,9 @@ class HashmapUndirectedGraph : public IUndirectedGraph { void remove_node_unsafe(Node const &) override; void add_edge(Edge const &) override; void remove_edge(Edge const &) override; - std::unordered_set + std::set query_edges(UndirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; friend bool operator==(HashmapUndirectedGraph const &, HashmapUndirectedGraph const &); @@ -28,7 +28,7 @@ class HashmapUndirectedGraph : public IUndirectedGraph { } private: - using ContentsType = std::unordered_map>; + using ContentsType = std::map>; HashmapUndirectedGraph(std::size_t next_node_idx, ContentsType adjacency) : next_node_idx(next_node_idx), adjacency(adjacency) {} diff --git a/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h index ecba7921af..c99bb47a56 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h @@ -19,12 +19,12 @@ struct UnorderedSetDataflowGraph final : virtual public IDataflowGraph, nonnegative_int num_outputs) override; DataflowGraphInput add_input() override; - std::unordered_set query_nodes(NodeQuery const &) const override; - std::unordered_set + std::set query_nodes(NodeQuery const &) const override; + std::set query_edges(OpenDataflowEdgeQuery const &) const override; - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &) const override; - std::unordered_set get_inputs() const override; + std::set get_inputs() const override; void add_node_unsafe(Node const &node, std::vector const &inputs, @@ -42,18 +42,18 @@ struct UnorderedSetDataflowGraph final : virtual public IDataflowGraph, UnorderedSetDataflowGraph( NodeSource const &node_source, DataflowGraphInputSource const &graph_input_source, - std::unordered_set const &nodes, - std::unordered_set const &edges, - std::unordered_set const &outputs, - std::unordered_set const &graph_inputs); + std::set const &nodes, + std::set const &edges, + std::set const &outputs, + std::set const &graph_inputs); private: NodeSource node_source; DataflowGraphInputSource graph_input_source; - std::unordered_set nodes; - std::unordered_set edges; - std::unordered_set outputs; - std::unordered_set graph_inputs; + std::set nodes; + std::set edges; + std::set outputs; + std::set graph_inputs; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(UnorderedSetDataflowGraph); diff --git a/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h index 4bb8373666..d23231bd43 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_KWARG_DATAFLOW_GRAPH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_KWARG_DATAFLOW_GRAPH_H -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/set_union.h" #include "utils/containers/values.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" @@ -20,13 +20,13 @@ struct UnorderedSetKwargDataflowGraph final UnorderedSetKwargDataflowGraph() = default; KwargNodeAddedResult add_node( - std::unordered_map> const &inputs, - std::unordered_set const &output_slots) override { + std::map> const &inputs, + std::set const &output_slots) override { Node new_node = this->node_source.new_node(); - std::unordered_map> outputs = - generate_unordered_map( + std::map> outputs = + generate_map( output_slots, [&](SlotName const &output_slot) -> KwargDataflowOutput { KwargDataflowOutput output = @@ -50,8 +50,8 @@ struct UnorderedSetKwargDataflowGraph final void add_node_unsafe( Node const &node, - std::unordered_map> const &inputs, - std::unordered_map> const + std::map> const &inputs, + std::map> const &outputs) override { this->nodes.insert(node); @@ -69,22 +69,22 @@ struct UnorderedSetKwargDataflowGraph final this->edges.insert(in_edge); } - this->outputs = set_union(this->outputs, unordered_set_of(values(outputs))); + this->outputs = set_union(this->outputs, set_of(values(outputs))); } - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return filter(this->nodes, [&](Node const &n) { return includes(q.nodes, n); }); } - std::unordered_set> + std::set> query_edges(KwargDataflowEdgeQuery const &q) const override { return filter(this->edges, [&](KwargDataflowEdge const &e) { return kwarg_dataflow_edge_query_includes(q, e); }); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &q) const override { return filter(this->outputs, [&](KwargDataflowOutput const &output) { @@ -111,18 +111,18 @@ struct UnorderedSetKwargDataflowGraph final private: UnorderedSetKwargDataflowGraph( NodeSource const &node_source, - std::unordered_set const &nodes, - std::unordered_set> const &edges, - std::unordered_set> const &outputs) + std::set const &nodes, + std::set> const &edges, + std::set> const &outputs) : node_source(node_source), nodes(nodes), edges(edges), outputs(outputs) { } private: NodeSource node_source; - std::unordered_set nodes; - std::unordered_set> edges; - std::unordered_set> outputs; + std::set nodes; + std::set> edges; + std::set> outputs; }; } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h index f05e6cb58a..00266ca1ca 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h @@ -4,8 +4,8 @@ #include "utils/containers/count.h" #include "utils/containers/enumerate_vector.h" #include "utils/containers/filter.h" -#include "utils/containers/generate_unordered_map.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/generate_map.h" +#include "utils/containers/keys.h" #include "utils/containers/map_keys.h" #include "utils/containers/transform.h" #include "utils/containers/without_nullopts.h" @@ -23,7 +23,7 @@ #include "utils/graph/open_dataflow_graph/dataflow_graph_input_source.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge_query.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" namespace FlexFlow { @@ -81,22 +81,22 @@ struct UnorderedSetLabelledOpenDataflowGraph final return new_input; } - std::unordered_set query_nodes(NodeQuery const &q) const override { - return filter(unordered_keys(this->nodes), + std::set query_nodes(NodeQuery const &q) const override { + return filter(keys(this->nodes), [&](Node const &n) { return includes(q.nodes, n); }); } - std::unordered_set + std::set query_edges(OpenDataflowEdgeQuery const &q) const override { return filter(this->edges, [&](OpenDataflowEdge const &e) { return open_dataflow_edge_query_includes(q, e); }); } - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &q) const override { return without_nullopts(transform( - unordered_keys(this->values), + keys(this->values), [&](OpenDataflowValue const &v) -> std::optional { if (!v.has()) { return std::nullopt; @@ -111,7 +111,7 @@ struct UnorderedSetLabelledOpenDataflowGraph final })); } - std::unordered_set get_inputs() const override { + std::set get_inputs() const override { return this->inputs; } @@ -125,16 +125,16 @@ struct UnorderedSetLabelledOpenDataflowGraph final virtual void inplace_materialize_from( LabelledDataflowGraphView const &view) override { - std::unordered_set nodes = get_nodes(view); - std::unordered_set outputs = get_all_dataflow_outputs(view); - std::unordered_set edges = get_edges(view); - std::unordered_map labelled_outputs = - generate_unordered_map(outputs, + std::set nodes = get_nodes(view); + std::set outputs = get_all_dataflow_outputs(view); + std::set edges = get_edges(view); + std::map labelled_outputs = + generate_map(outputs, [&](DataflowOutput const &o) { return view.at(o); }); this->inputs.clear(); this->nodes = - generate_unordered_map(nodes, [&](Node const &n) { return view.at(n); }); + generate_map(nodes, [&](Node const &n) { return view.at(n); }); this->edges = transform( edges, [](DataflowEdge const &e) { return OpenDataflowEdge{e}; }); this->values = map_keys(labelled_outputs, [](DataflowOutput const &o) { @@ -146,14 +146,14 @@ struct UnorderedSetLabelledOpenDataflowGraph final LabelledOpenDataflowGraphView const &view) override { - std::unordered_map nodes = generate_unordered_map( + std::map nodes = generate_map( get_nodes(view), [&](Node const &n) { return view.at(n); }); - std::unordered_set edges = get_edges(view); - std::unordered_set inputs = + std::set edges = get_edges(view); + std::set inputs = ::FlexFlow::get_open_dataflow_graph_inputs(view); - std::unordered_map values = - generate_unordered_map(get_open_dataflow_values(view), + std::map values = + generate_map(get_open_dataflow_values(view), [&](OpenDataflowValue const &v) { return view.at(v); }); this->inputs = inputs; @@ -177,20 +177,20 @@ struct UnorderedSetLabelledOpenDataflowGraph final UnorderedSetLabelledOpenDataflowGraph( NodeSource const &node_source, DataflowGraphInputSource const &input_source, - std::unordered_set const &inputs, - std::unordered_map const &nodes, - std::unordered_set const &edges, - std::unordered_map const &values) + std::set const &inputs, + std::map const &nodes, + std::set const &edges, + std::map const &values) : node_source(node_source), input_source(input_source), inputs(inputs), nodes(nodes), edges(edges), values(values) {} private: NodeSource node_source; DataflowGraphInputSource input_source; - std::unordered_set inputs; - std::unordered_map nodes; - std::unordered_set edges; - std::unordered_map values; + std::set inputs; + std::map nodes; + std::set edges; + std::map values; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT( UnorderedSetLabelledOpenDataflowGraph); diff --git a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h index ab60e3d364..b27f5cfd56 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h @@ -4,7 +4,7 @@ #include "utils/containers/contains_key.h" #include "utils/containers/enumerate.h" #include "utils/containers/extend.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/map_values.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" @@ -17,7 +17,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge.h" #include "utils/overload.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" namespace FlexFlow { @@ -36,8 +36,8 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::unordered_map> const &inputs, - std::unordered_map const &output_labels) override { + std::map> const &inputs, + std::map const &output_labels) override { return this->add_node( node_label, map_values(inputs, @@ -49,10 +49,10 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::unordered_map> const &inputs, - std::unordered_map const &output_labels) override { + std::map const &output_labels) override { Node new_node = this->node_source.new_node(); this->nodes.insert({new_node, node_label}); @@ -68,9 +68,9 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final this->edges.insert(in_edge); } - std::unordered_map> outputs = - generate_unordered_map( - unordered_keys(output_labels), + std::map> outputs = + generate_map( + keys(output_labels), [&](SlotName const &output_slot) -> KwargDataflowOutput { ValueLabel value_label = output_labels.at(output_slot); @@ -106,12 +106,12 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final return input; } - std::unordered_set query_nodes(NodeQuery const &q) const override { - return filter(unordered_keys(this->nodes), + std::set query_nodes(NodeQuery const &q) const override { + return filter(keys(this->nodes), [&](Node const &n) { return includes(q.nodes, n); }); } - std::unordered_set> + std::set> query_edges(OpenKwargDataflowEdgeQuery const &q) const override { return filter( @@ -121,17 +121,17 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final }); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &q) const override { - return filter(unordered_keys(this->outputs), + return filter(keys(this->outputs), [&](KwargDataflowOutput const &output) { return kwarg_dataflow_output_query_includes(q, output); }); } - std::unordered_set> + std::set> get_inputs() const override { - return unordered_keys(this->graph_inputs); + return keys(this->graph_inputs); } NodeLabel at(Node const &n) const override { @@ -152,15 +152,15 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final void inplace_materialize_from( LabelledKwargDataflowGraphView const &view) override { - std::unordered_set view_nodes = get_nodes(view); - std::unordered_set> view_edges = + std::set view_nodes = get_nodes(view); + std::set> view_edges = get_all_kwarg_dataflow_edges(view); - std::unordered_set> view_outputs = + std::set> view_outputs = get_all_kwarg_dataflow_outputs(view); this->graph_inputs.clear(); this->nodes = - generate_unordered_map(view_nodes, [&](Node const &n) { return view.at(n); }); + generate_map(view_nodes, [&](Node const &n) { return view.at(n); }); this->edges = transform(view_edges, @@ -169,7 +169,7 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final return OpenKwargDataflowEdge{e}; }); this->outputs = - generate_unordered_map(view_outputs, [&](KwargDataflowOutput const &o) { + generate_map(view_outputs, [&](KwargDataflowOutput const &o) { return view.at(o); }); } @@ -179,24 +179,24 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final ValueLabel, GraphInputName, SlotName> const &view) override { - std::unordered_set> view_inputs = + std::set> view_inputs = get_all_kwarg_dataflow_graph_inputs(view); - std::unordered_set view_nodes = get_nodes(view); - std::unordered_set> + std::set view_nodes = get_nodes(view); + std::set> view_edges = get_all_open_kwarg_dataflow_edges(view); - std::unordered_set> view_outputs = + std::set> view_outputs = get_all_kwarg_dataflow_outputs(view); - this->graph_inputs = generate_unordered_map( + this->graph_inputs = generate_map( view_inputs, [&](KwargDataflowGraphInput const &i) { return view.at(OpenKwargDataflowValue{i}); }); this->nodes = - generate_unordered_map(view_nodes, [&](Node const &n) { return view.at(n); }); + generate_map(view_nodes, [&](Node const &n) { return view.at(n); }); this->edges = view_edges; this->outputs = - generate_unordered_map(view_outputs, [&](KwargDataflowOutput const &o) { + generate_map(view_outputs, [&](KwargDataflowOutput const &o) { return view.at(OpenKwargDataflowValue{o}); }); } @@ -214,12 +214,12 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final private: UnorderedSetLabelledOpenKwargDataflowGraph( NodeSource const &node_source, - std::unordered_map, + std::map, ValueLabel> const &graph_inputs, - std::unordered_map const &nodes, - std::unordered_set> const + std::map const &nodes, + std::set> const &edges, - std::unordered_map, ValueLabel> const + std::map, ValueLabel> const &outputs) : node_source(node_source), graph_inputs(graph_inputs), nodes(nodes), edges(edges), outputs(outputs) {} @@ -227,11 +227,11 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final private: NodeSource node_source; - std::unordered_map, ValueLabel> + std::map, ValueLabel> graph_inputs; - std::unordered_map nodes; - std::unordered_set> edges; - std::unordered_map, ValueLabel> outputs; + std::map nodes; + std::set> edges; + std::map, ValueLabel> outputs; }; } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h index 9746533c17..1013f69691 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_OPEN_KWARG_DATAFLOW_GRAPH_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_INSTANCES_UNORDERED_SET_OPEN_KWARG_DATAFLOW_GRAPH_H -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.h" #include "utils/graph/node/node_source.h" #include "utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h" @@ -16,10 +16,10 @@ struct UnorderedSetOpenKwargDataflowGraph final UnorderedSetOpenKwargDataflowGraph() = default; KwargNodeAddedResult add_node( - std::unordered_map> const &inputs, - std::unordered_set const &output_slots) override { + std::set const &output_slots) override { Node new_node = this->node_source.new_node(); this->nodes.insert(new_node); @@ -35,8 +35,8 @@ struct UnorderedSetOpenKwargDataflowGraph final this->edges.insert(in_edge); } - std::unordered_map> outputs = - generate_unordered_map( + std::map> outputs = + generate_map( output_slots, [&](SlotName const &output_slot) -> KwargDataflowOutput { KwargDataflowOutput output = @@ -66,12 +66,12 @@ struct UnorderedSetOpenKwargDataflowGraph final return input; } - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return filter(this->nodes, [&](Node const &n) { return includes(q.nodes, n); }); } - std::unordered_set> + std::set> query_edges(OpenKwargDataflowEdgeQuery const &q) const override { return filter( @@ -81,7 +81,7 @@ struct UnorderedSetOpenKwargDataflowGraph final }); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &q) const override { return filter(this->outputs, [&](KwargDataflowOutput const &output) { @@ -89,7 +89,7 @@ struct UnorderedSetOpenKwargDataflowGraph final }); } - std::unordered_set> + std::set> get_inputs() const override { return this->graph_inputs; } @@ -107,22 +107,22 @@ struct UnorderedSetOpenKwargDataflowGraph final private: UnorderedSetOpenKwargDataflowGraph( NodeSource const &node_source, - std::unordered_set> const + std::set> const &graph_inputs, - std::unordered_set const &nodes, - std::unordered_set> const + std::set const &nodes, + std::set> const &edges, - std::unordered_set> const &outputs) + std::set> const &outputs) : node_source(node_source), graph_inputs(graph_inputs), nodes(nodes), edges(edges), outputs(outputs) {} private: NodeSource node_source; - std::unordered_set> graph_inputs; - std::unordered_set nodes; - std::unordered_set> edges; - std::unordered_set> outputs; + std::set> graph_inputs; + std::set nodes; + std::set> edges; + std::set> outputs; }; } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/instances/unordered_set_undirected_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_undirected_graph.h index db2526f973..438ffa9a33 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_undirected_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_undirected_graph.h @@ -16,20 +16,20 @@ struct UnorderedSetUndirectedGraph final : public IUndirectedGraph { void add_edge(UndirectedEdge const &) override; void remove_edge(UndirectedEdge const &) override; - std::unordered_set query_nodes(NodeQuery const &) const override; - std::unordered_set + std::set query_nodes(NodeQuery const &) const override; + std::set query_edges(UndirectedEdgeQuery const &) const override; UnorderedSetUndirectedGraph *clone() const override; private: UnorderedSetUndirectedGraph(NodeSource const &, - std::unordered_set const &, - std::unordered_set const &); + std::set const &, + std::set const &); NodeSource node_source; - std::unordered_set nodes; - std::unordered_set edges; + std::set nodes; + std::set edges; }; } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h index 154072eebf..06b19b6986 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h @@ -7,7 +7,7 @@ #include "utils/containers/transform.h" #include "utils/graph/dataflow_graph/algorithms/dataflow_graph_data.dtg.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.dtg.h" -#include "utils/nonempty_unordered_set/nonempty_unordered_set.h" +#include "utils/nonempty_set/nonempty_set.h" #include "utils/one_to_many/one_to_many_transform_values.h" namespace FlexFlow { @@ -16,41 +16,41 @@ template DataflowGraphData dataflow_graph_data_from_kwarg_dataflow_graph_data( KwargDataflowGraphData const &kwarg_data, std::function( - std::unordered_set const &)> const &order_slots) { - std::unordered_set> all_inputs = transform( + std::set const &)> const &order_slots) { + std::set> all_inputs = transform( kwarg_data.edges, [](KwargDataflowEdge const &e) -> KwargDataflowInput { return e.dst; }); - std::unordered_set> all_outputs = + std::set> all_outputs = kwarg_data.outputs; - std::unordered_map> + std::map> incoming_slots_by_node = map_values( group_by(all_inputs, [](KwargDataflowInput const &i) -> Node { return i.node; }) .l_to_r(), - [](nonempty_unordered_set> const &is) - -> std::unordered_set { - return transform(is.unwrap_as_unordered_set(), + [](nonempty_set> const &is) + -> std::set { + return transform(is.unwrap_as_set(), [](KwargDataflowInput const &i) { return i.slot_name; }); }); - std::unordered_map> + std::map> outgoing_slots_by_node = map_values( group_by(all_outputs, [](KwargDataflowOutput const &o) -> Node { return o.node; }) .l_to_r(), - [](nonempty_unordered_set> const &os) - -> std::unordered_set { - return transform(os.unwrap_as_unordered_set(), + [](nonempty_set> const &os) + -> std::set { + return transform(os.unwrap_as_set(), [](KwargDataflowOutput const &o) { return o.slot_name; }); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h index d07250ae5b..329b47b6b8 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h @@ -13,7 +13,7 @@ template DataflowGraphView dataflow_graph_from_kwarg_dataflow_graph( KwargDataflowGraphView const &kwarg_dg, std::function( - std::unordered_set const &)> const &order_slots) { + std::set const &)> const &order_slots) { KwargDataflowGraphData kwarg_data = get_kwarg_dataflow_graph_data(kwarg_dg); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h index 146408123d..ef11d5bc56 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h @@ -15,7 +15,7 @@ std::optional> KwargDataflowGraphView const &lhs, KwargDataflowGraphView const &rhs) { - std::unordered_set> open_isomorphisms = + std::set> open_isomorphisms = find_isomorphisms_between_open_kwarg_dataflow_graphs( view_as_open_kwarg_dataflow_graph(lhs), view_as_open_kwarg_dataflow_graph(rhs)); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h index b881cd9584..1fbdbaa505 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h @@ -7,7 +7,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_all_kwarg_dataflow_edges(KwargDataflowGraphView const &g) { return g.query_edges(kwarg_dataflow_edge_query_all()); } diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.h index 97d13b498f..7973faf469 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.h @@ -7,7 +7,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_all_kwarg_dataflow_inputs(KwargDataflowGraphView const &v) { return transform( v.query_edges(kwarg_dataflow_edge_query_all()), diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h index 2405284fd3..ac563d7363 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h @@ -7,7 +7,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_all_kwarg_dataflow_outputs( KwargDataflowGraphView const &view) { return view.query_outputs(kwarg_dataflow_output_query_all()); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.h index 8e7d111ecf..c9c393d0ac 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.h @@ -1,13 +1,13 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_KWARG_DATAFLOW_EDGES_FOR_NODE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_KWARG_DATAFLOW_EDGES_FOR_NODE_H -#include "utils/containers/unordered_map_from_pairs.h" +#include "utils/containers/map_from_pairs.h" #include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h" namespace FlexFlow { template -std::unordered_map> +std::map> get_incoming_kwarg_dataflow_edges_for_node( KwargDataflowGraphView const &g, Node const &n) { KwargDataflowEdgeQuery query = KwargDataflowEdgeQuery{ @@ -17,7 +17,7 @@ std::unordered_map> /*dst_slots=*/query_set::matchall(), }; - return unordered_map_from_pairs( + return map_from_pairs( transform(g.query_edges(query), [](KwargDataflowEdge const &e) { return std::pair{ e.dst.slot_name, diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.h index b0940570b0..8514c62b93 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.h @@ -8,7 +8,7 @@ namespace FlexFlow { template -std::unordered_map> +std::map> get_incoming_kwarg_dataflow_outputs_for_node( KwargDataflowGraphView const &g, Node const &n) { return map_values(get_incoming_kwarg_dataflow_edges_for_node(g, n), diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h index 87ebab4c2d..160aa23dba 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h @@ -1,16 +1,16 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_SLOTS_FOR_NODE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_SLOTS_FOR_NODE_H -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.h" namespace FlexFlow { template -std::unordered_set +std::set get_incoming_slots_for_node(KwargDataflowGraphView const &g, Node n) { - return unordered_keys(get_incoming_kwarg_dataflow_edges_for_node(g, n)); + return keys(get_incoming_kwarg_dataflow_edges_for_node(g, n)); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.h index 45a9fddc5b..51643deea0 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.h @@ -6,7 +6,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_kwarg_dataflow_edges_from_node_to_node( KwargDataflowGraphView const &g, Node const &src, diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_data.h index 365c09486c..7cb6c3312c 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_data.h @@ -12,9 +12,9 @@ template KwargDataflowGraphData get_kwarg_dataflow_graph_data(KwargDataflowGraphView const &g) { return KwargDataflowGraphData{ - /*nodes=*/get_nodes(g), - /*edges=*/get_all_kwarg_dataflow_edges(g), - /*outputs=*/get_all_kwarg_dataflow_outputs(g), + /*nodes=*/set_of(get_nodes(g)), + /*edges=*/set_of(get_all_kwarg_dataflow_edges(g)), + /*outputs=*/set_of(get_all_kwarg_dataflow_outputs(g)), }; } diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h index dba05a6170..ea5448c96e 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h @@ -11,19 +11,19 @@ namespace FlexFlow { template KwargDataflowGraphView get_kwarg_dataflow_graph_subgraph( KwargDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { KwargDataflowGraphData g_data = get_kwarg_dataflow_graph_data(g); - std::unordered_set nodes = - set_intersection(g_data.nodes, subgraph_nodes); + std::set nodes = + set_intersection(g_data.nodes, set_of(subgraph_nodes)); - std::unordered_set> edges = + std::set> edges = filter(g_data.edges, [&](KwargDataflowEdge const &e) -> bool { return contains(subgraph_nodes, e.src.node) && contains(subgraph_nodes, e.dst.node); }); - std::unordered_set> outputs = filter( + std::set> outputs = filter( g_data.outputs, [&](KwargDataflowOutput const &o) -> bool { return contains(subgraph_nodes, o.node); }); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.h index 908b805e58..d8f7dfc9c3 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.h @@ -9,11 +9,11 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_kwarg_dataflow_subgraph_incoming_edges( KwargDataflowGraphView const &g, - std::unordered_set const &subgraph) { - std::unordered_set all_nodes = get_nodes(g); + std::set const &subgraph) { + std::set all_nodes = get_nodes(g); query_set src_query = query_set::match_values_in(set_of(set_minus(all_nodes, subgraph))); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.h index 5b86f9492f..e7b77ad243 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.h @@ -9,11 +9,11 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_kwarg_dataflow_subgraph_outgoing_edges( KwargDataflowGraphView const &g, - std::unordered_set const &subgraph) { - std::unordered_set all_nodes = get_nodes(g); + std::set const &subgraph) { + std::set all_nodes = get_nodes(g); query_set dst_query = query_set::match_values_in(set_of(set_minus(all_nodes, subgraph))); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h index 52c225d157..fead19c1b8 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.h @@ -6,7 +6,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_kwarg_dataflow_value_uses(KwargDataflowGraphView const &g, KwargDataflowOutput const &v) { @@ -17,7 +17,7 @@ std::unordered_set> /*dst_slots=*/query_set::matchall(), }; - std::unordered_set> edges = g.query_edges(query); + std::set> edges = g.query_edges(query); return transform(edges, [&](KwargDataflowEdge const &e) { return e.dst; }); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h index 8b70dd80ff..77ddd6b820 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h @@ -6,7 +6,7 @@ namespace FlexFlow { template -std::unordered_map> +std::map> get_outgoing_kwarg_dataflow_outputs_for_node( KwargDataflowGraphView const &g, Node const &n) { KwargDataflowOutputQuery query = KwargDataflowOutputQuery{ @@ -14,7 +14,7 @@ std::unordered_map> /*output_idxs=*/query_set::matchall(), }; - std::unordered_map> result; + std::map> result; for (KwargDataflowOutput const &output : g.query_outputs(query)) { result.insert({output.slot_name, output}); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h index 6cf2b4b8a4..888da4b173 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h @@ -1,16 +1,16 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_OUTGOING_SLOTS_FOR_NODE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_OUTGOING_SLOTS_FOR_NODE_H -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h" namespace FlexFlow { template -std::unordered_set +std::set get_outgoing_slots_for_node(KwargDataflowGraphView const &g, Node n) { - return unordered_keys(get_outgoing_kwarg_dataflow_outputs_for_node(g, n)); + return keys(get_outgoing_kwarg_dataflow_outputs_for_node(g, n)); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h index aacdadee4f..855333e500 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h @@ -6,6 +6,7 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h" #include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h" +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -17,11 +18,11 @@ std::string kwarg_dataflow_graph_as_dot( &render_value, std::function const &render_slot_name, std::function( - std::unordered_set const &)> const &order_slots) { + std::set const &)> const &order_slots) { std::function get_input_label = [&](DataflowInput const &i) -> nlohmann::json { std::vector slot_ordering = - order_slots(get_incoming_slots_for_node(g, i.node)); + order_slots(set_of(get_incoming_slots_for_node(g, i.node))); SlotName slot_name = slot_ordering.at(i.idx.unwrap_nonnegative()); @@ -31,7 +32,7 @@ std::string kwarg_dataflow_graph_as_dot( std::function get_output_label = [&](DataflowOutput const &o) -> nlohmann::json { std::vector slot_ordering = - order_slots(get_outgoing_slots_for_node(g, o.node)); + order_slots(set_of(get_outgoing_slots_for_node(g, o.node))); SlotName slot_name = slot_ordering.at(o.idx.unwrap_nonnegative()); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.dtg.toml b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.dtg.toml index e3429f16d0..8b447d0c78 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.dtg.toml +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.dtg.toml @@ -3,6 +3,7 @@ name = "KwargDataflowGraphData" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -15,22 +16,22 @@ includes = [ "utils/graph/node/node.dtg.h", "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_edge.dtg.h", "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.h", - "", + "", ] src_includes = [ - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "nodes" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::KwargDataflowEdge>" +type = "std::set<::FlexFlow::KwargDataflowEdge>" [[fields]] name = "outputs" -type = "std::unordered_set<::FlexFlow::KwargDataflowOutput>" +type = "std::set<::FlexFlow::KwargDataflowOutput>" diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h index 62644a5c6b..18ea5baa51 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h @@ -13,10 +13,10 @@ template void require_kwarg_dataflow_graph_data_is_valid( KwargDataflowGraphData const &data) { - std::unordered_set nodes_from_edges = flatmap( + std::set nodes_from_edges = flatmap( data.edges, - [](KwargDataflowEdge const &e) -> std::unordered_set { - return std::unordered_set{ + [](KwargDataflowEdge const &e) -> std::set { + return std::set{ e.src.node, e.dst.node, }; @@ -24,13 +24,13 @@ void require_kwarg_dataflow_graph_data_is_valid( ASSERT(is_subseteq_of(nodes_from_edges, data.nodes)); - std::unordered_set nodes_from_outputs = transform( + std::set nodes_from_outputs = transform( data.outputs, [](KwargDataflowOutput const &o) -> Node { return o.node; }); ASSERT(is_subseteq_of(nodes_from_outputs, data.nodes)); - std::unordered_set> outputs_from_edges = + std::set> outputs_from_edges = transform(data.edges, [](KwargDataflowEdge const &e) -> KwargDataflowOutput { return e.src; }); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_boundary_nodes_for_kwarg_dataflow_graph_split.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_boundary_nodes_for_kwarg_dataflow_graph_split.h index af7b288aee..8f952c4626 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_boundary_nodes_for_kwarg_dataflow_graph_split.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_boundary_nodes_for_kwarg_dataflow_graph_split.h @@ -14,13 +14,13 @@ SplitBoundaryNodes TransitiveReducedKwargDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { - std::unordered_set> edges = + std::set> edges = get_transitive_reduced_kwarg_dataflow_edges_across_split(tr_g, split); - std::unordered_set src_boundary_nodes = transform( + std::set src_boundary_nodes = transform( edges, [](KwargDataflowEdge const &e) { return e.src.node; }); - std::unordered_set dst_boundary_nodes = transform( + std::set dst_boundary_nodes = transform( edges, [](KwargDataflowEdge const &e) { return e.dst.node; }); return SplitBoundaryNodes{ diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h index 56a90c833f..0f8940abea 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h @@ -11,17 +11,17 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_transitive_reduced_kwarg_dataflow_edges_across_split( TransitiveReducedKwargDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { - std::unordered_set src_subgraph = - unordered_set_of(get_leaves(split.get_left_child())); - std::unordered_set dst_subgraph = - unordered_set_of(get_leaves(split.get_right_child())); + std::set src_subgraph = + set_of(get_leaves(split.get_left_child())); + std::set dst_subgraph = + set_of(get_leaves(split.get_right_child())); - std::unordered_set raw_edges = + std::set raw_edges = get_edges_from_subgraph_to_subgraph( tr_g.transitive_reduction, src_subgraph, dst_subgraph); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.h index 82804acc04..46f06746cb 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.h @@ -9,7 +9,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_transitive_reduced_kwarg_dataflow_outputs_across_split( TransitiveReducedKwargDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_as_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_as_open_kwarg_dataflow_graph.h index e1ceddc466..04fe95fa55 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_as_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_as_open_kwarg_dataflow_graph.h @@ -13,11 +13,11 @@ struct KwargDataflowGraphAsOpenView final KwargDataflowGraphAsOpenView(KwargDataflowGraphView const &g) : g(g) {} - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return this->g.query_nodes(q); } - std::unordered_set> + std::set> query_edges(OpenKwargDataflowEdgeQuery const &q) const override { return transform(this->g.query_edges(q.standard_edge_query), @@ -27,12 +27,12 @@ struct KwargDataflowGraphAsOpenView final }); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &q) const override { return this->g.query_outputs(q); } - std::unordered_set> + std::set> get_inputs() const override { return {}; } diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h index 8e6daadd3c..a0c5b9b21b 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h @@ -15,20 +15,20 @@ struct ViewFromKwargDataflowGraphData final KwargDataflowGraphData const &data) : data(data) {} - std::unordered_set query_nodes(NodeQuery const &query) const override { - return apply_node_query(query, this->data.nodes); + std::set query_nodes(NodeQuery const &query) const override { + return apply_node_query(query, set_of(this->data.nodes)); } - std::unordered_set> query_edges( + std::set> query_edges( KwargDataflowEdgeQuery const &query) const override { - return filter(this->data.edges, [&](KwargDataflowEdge const &e) { + return filter(set_of(this->data.edges), [&](KwargDataflowEdge const &e) { return kwarg_dataflow_edge_query_includes(query, e); }); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &query) const override { - return filter(this->data.outputs, + return filter(set_of(this->data.outputs), [&](KwargDataflowOutput const &o) { return kwarg_dataflow_output_query_includes(query, o); }); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h index 36f36b3fc8..2ff3ddb40e 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h @@ -10,13 +10,13 @@ namespace FlexFlow { template struct IKwargDataflowGraph : virtual public IKwargDataflowGraphView { virtual KwargNodeAddedResult add_node( - std::unordered_map> const &inputs, - std::unordered_set const &outputs) = 0; + std::map> const &inputs, + std::set const &outputs) = 0; virtual void add_node_unsafe( Node const &node, - std::unordered_map> const &inputs, - std::unordered_map> const + std::map> const &inputs, + std::map> const &outputs) = 0; virtual void diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph_view.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph_view.h index 6b85d9f470..652d43ea96 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph_view.h @@ -12,12 +12,12 @@ namespace FlexFlow { template struct IKwargDataflowGraphView : virtual public IDiGraphView { - virtual std::unordered_set> + virtual std::set> query_edges(KwargDataflowEdgeQuery const &) const = 0; - virtual std::unordered_set> + virtual std::set> query_outputs(KwargDataflowOutputQuery const &) const = 0; - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &q) const override final { KwargDataflowEdgeQuery dataflow_query = KwargDataflowEdgeQuery{ q.srcs, @@ -25,7 +25,7 @@ struct IKwargDataflowGraphView : virtual public IDiGraphView { q.dsts, matchall(), }; - std::unordered_set> dataflow_edges = + std::set> dataflow_edges = this->query_edges(dataflow_query); return transform(dataflow_edges, [](KwargDataflowEdge const &e) { diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h index 2b6fd22f86..d199a64fb4 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h @@ -11,29 +11,29 @@ template struct KwargDataflowGraph : virtual public KwargDataflowGraphView { public: KwargNodeAddedResult add_node( - std::unordered_map> const &inputs, - std::unordered_set const &outputs) { + std::map> const &inputs, + std::set const &outputs) { return this->get_interface().add_node(inputs, outputs); } void add_node_unsafe( Node const &node, - std::unordered_map> const &inputs, - std::unordered_map> const + std::map> const &inputs, + std::map> const &outputs) { return this->get_interface().add_node_unsafe(node, inputs, outputs); } - std::unordered_set query_nodes(NodeQuery const &q) const { + std::set query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } - std::unordered_set> + std::set> query_edges(KwargDataflowEdgeQuery const &q) const { return this->get_interface().query_edges(q); } - std::unordered_set> + std::set> query_outputs(KwargDataflowOutputQuery const &q) const { return this->get_interface().query_outputs(q); } diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h index 70edc3d9dd..85b1c38ea7 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h @@ -13,16 +13,16 @@ struct KwargDataflowGraphView : virtual public DiGraphView { KwargDataflowGraphView(KwargDataflowGraphView const &) = default; KwargDataflowGraphView &operator=(KwargDataflowGraphView const &) = default; - std::unordered_set query_nodes(NodeQuery const &q) const { + std::set query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } - std::unordered_set> + std::set> query_edges(KwargDataflowEdgeQuery const &q) const { return this->get_interface().query_edges(q); } - std::unordered_set> + std::set> query_outputs(KwargDataflowOutputQuery const &q) const { return this->get_interface().query_outputs(q); } diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.dtg.toml b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.dtg.toml index 8b5de44cc3..23fc460605 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.dtg.toml +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.dtg.toml @@ -19,7 +19,7 @@ includes = [ ] src_includes = [ - "utils/fmt/unordered_set.h", + "utils/fmt/set.h", ] [[fields]] diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_node_added_result.dtg.toml b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_node_added_result.dtg.toml index 5686368f66..66fda1f3c7 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_node_added_result.dtg.toml +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_node_added_result.dtg.toml @@ -12,13 +12,13 @@ template_params = [ ] includes = [ - "", + "", "utils/graph/node/node.dtg.h", "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.h", ] src_includes = [ - "utils/fmt/unordered_map.h", + "utils/fmt/map.h", ] [[fields]] @@ -27,4 +27,4 @@ type = "::FlexFlow::Node" [[fields]] name = "outputs" -type = "std::unordered_map>" +type = "std::map>" diff --git a/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/create_lazy_copy_of_labelled_dataflow_graph_view.h b/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/create_lazy_copy_of_labelled_dataflow_graph_view.h index b9894fbac3..c4d7498c31 100644 --- a/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/create_lazy_copy_of_labelled_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/create_lazy_copy_of_labelled_dataflow_graph_view.h @@ -32,16 +32,16 @@ struct LazyLabelledDataflowGraph final node_label, inputs, output_labels); } - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return this->get_view().query_nodes(q); } - std::unordered_set + std::set query_edges(DataflowEdgeQuery const &q) const override { return this->get_view().query_edges(q); } - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &q) const override { return this->get_view().query_outputs(q); } diff --git a/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/view_as_labelled_open_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/view_as_labelled_open_dataflow_graph.h index f1cdfd9690..af56f8d200 100644 --- a/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/view_as_labelled_open_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_dataflow_graph/algorithms/view_as_labelled_open_dataflow_graph.h @@ -14,22 +14,22 @@ struct LabelledDataflowGraphAsOpenView final LabelledDataflowGraphView const &g) : g(g) {} - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return this->g.query_nodes(q); } - std::unordered_set + std::set query_edges(OpenDataflowEdgeQuery const &q) const override { return transform(this->g.query_edges(q.standard_edge_query), [](DataflowEdge const &e) { return OpenDataflowEdge{e}; }); } - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &q) const override { return this->g.query_outputs(q); } - std::unordered_set get_inputs() const override { + std::set get_inputs() const override { return {}; } diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h index fabb47a5ab..a981bbd8da 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_NODE_LABEL_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_NODE_LABEL_MAP_H -#include "utils/containers/generate_map.h" #include "utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { template -std::unordered_map +std::map get_labelled_kwarg_dataflow_graph_node_label_map( LabelledKwargDataflowGraphView const &g) { diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h index 45366d4d96..ef88f5e7aa 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_OUTPUT_LABEL_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_OUTPUT_LABEL_MAP_H -#include "utils/containers/generate_map.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" #include "utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph_view.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { template -std::unordered_map, OutputLabel> +std::map, OutputLabel> get_labelled_kwarg_dataflow_graph_output_label_map( LabelledKwargDataflowGraphView const &g) { diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h index 3aa7953446..9ed5b5fe44 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h @@ -17,14 +17,14 @@ LabelledKwargDataflowGraphView get_labelled_kwarg_dataflow_graph_subgraph( LabelledKwargDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { KwargDataflowGraphView unlabelled_subgraph = get_kwarg_dataflow_graph_subgraph(g, subgraph_nodes); - std::unordered_map g_node_labelling = + std::map g_node_labelling = get_labelled_kwarg_dataflow_graph_node_label_map(g); - std::unordered_map, OutputLabel> + std::map, OutputLabel> g_output_labelling = get_labelled_kwarg_dataflow_graph_output_label_map(g); diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h index 782e63889b..e09c83e91e 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h @@ -14,22 +14,22 @@ struct KwargDataflowGraphLabellingWrapper final KwargDataflowGraphLabellingWrapper() = delete; KwargDataflowGraphLabellingWrapper( KwargDataflowGraphView const &unlabelled, - std::unordered_map const &node_labels, - std::unordered_map, OutputLabel> const + std::map const &node_labels, + std::map, OutputLabel> const &output_labels) : unlabelled(unlabelled), node_labels(node_labels), output_labels(output_labels) {} - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return this->unlabelled.query_nodes(q); } - std::unordered_set> + std::set> query_edges(KwargDataflowEdgeQuery const &q) const override { return this->unlabelled.query_edges(q); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &q) const override { return this->unlabelled.query_outputs(q); } @@ -52,16 +52,16 @@ struct KwargDataflowGraphLabellingWrapper final private: KwargDataflowGraphView unlabelled; - std::unordered_map node_labels; - std::unordered_map, OutputLabel> output_labels; + std::map node_labels; + std::map, OutputLabel> output_labels; }; template LabelledKwargDataflowGraphView kwarg_dataflow_graph_view_with_labelling( KwargDataflowGraphView const &g, - std::unordered_map const &node_labels, - std::unordered_map, OutputLabel> const + std::map const &node_labels, + std::map, OutputLabel> const &value_labels) { return LabelledKwargDataflowGraphView:: template create< diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_data.dtg.toml b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_data.dtg.toml index 470f2712f8..af021e5d82 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_data.dtg.toml +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_data.dtg.toml @@ -17,25 +17,25 @@ includes = [ "utils/graph/node/node.dtg.h", "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_edge.dtg.h", "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.h", - "", - "", + "", + "", ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_map.h", - "utils/fmt/unordered_set.h", + "utils/hash/map.h", + "utils/hash/set.h", + "utils/fmt/map.h", + "utils/fmt/set.h", ] [[fields]] name = "node_data" -type = "std::unordered_map<::FlexFlow::Node, NodeLabel>" +type = "std::map<::FlexFlow::Node, NodeLabel>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::KwargDataflowEdge>" +type = "std::set<::FlexFlow::KwargDataflowEdge>" [[fields]] name = "output_data" -type = "std::unordered_map<::FlexFlow::KwargDataflowOutput, ValueLabel>" +type = "std::map<::FlexFlow::KwargDataflowOutput, ValueLabel>" diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h index 2bbd88103e..752a0fbdc2 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h @@ -14,7 +14,7 @@ std::string labelled_kwarg_dataflow_graph_view_as_dot( std::function const &render_value_label, std::function const &render_slot_name, std::function( - std::unordered_set const &)> const &order_slots) { + std::set const &)> const &order_slots) { std::function render_node = [&](Node const &n) -> nlohmann::json { return render_node_label(g.at(n)); diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/view_as_labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/view_as_labelled_open_kwarg_dataflow_graph.h index 58e71a4587..f02f2a2651 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/view_as_labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/view_as_labelled_open_kwarg_dataflow_graph.h @@ -22,11 +22,11 @@ struct LabelledKwargDataflowGraphAsOpenView final LabelledKwargDataflowGraphView const &g) : g(g) {} - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return this->g.query_nodes(q); } - std::unordered_set> + std::set> query_edges(OpenKwargDataflowEdgeQuery const &q) const override { return transform(this->g.query_edges(q.standard_edge_query), @@ -36,12 +36,12 @@ struct LabelledKwargDataflowGraphAsOpenView final }); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &q) const override { return this->g.query_outputs(q); } - std::unordered_set> + std::set> get_inputs() const override { return {}; } diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h index 9bf0e51413..c4f7e22276 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h @@ -15,8 +15,8 @@ struct ILabelledKwargDataflowGraph public: virtual KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::unordered_map> const &inputs, - std::unordered_map const &output_labels) = 0; + std::map> const &inputs, + std::map const &output_labels) = 0; virtual void inplace_materialize_from( LabelledKwargDataflowGraphView const &) = 0; diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h index 3f4b469471..3b38254f01 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h @@ -20,8 +20,8 @@ struct LabelledKwargDataflowGraph KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::unordered_map> const &inputs, - std::unordered_map const &output_labels) { + std::map> const &inputs, + std::map const &output_labels) { return this->get_interface().add_node(node_label, inputs, output_labels); } diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/find_isomorphism.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/find_isomorphism.h index 8306dad1ec..3d6ead0f60 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/find_isomorphism.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/find_isomorphism.h @@ -19,7 +19,7 @@ template std::optional find_isomorphism( LabelledOpenDataflowGraphView const &src, LabelledOpenDataflowGraphView const &dst) { - std::unordered_set unlabelled_isomorphisms = + std::set unlabelled_isomorphisms = find_isomorphisms(static_cast(src), static_cast(dst)); diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h index 106d500464..e12cb1a440 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h @@ -15,8 +15,8 @@ template LabelledOpenDataflowGraphView from_labelled_open_dataflow_graph_data( LabelledOpenDataflowGraphData const &data) { - std::unordered_set values = keys(data.value_data); - std::unordered_set outputs = + std::set values = keys(data.value_data); + std::set outputs = filtrans(values, try_get_dataflow_output); OpenDataflowGraphData unlabelled_data = OpenDataflowGraphData{ diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h index 502eeab73b..1d037d5ac2 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h @@ -13,15 +13,15 @@ template LabelledOpenDataflowGraphData get_graph_data( LabelledOpenDataflowGraphView const &g) { - std::unordered_map node_data = - generate_unordered_map(get_nodes(g), [&](Node const &n) { return g.at(n); }); + std::map node_data = + generate_map(get_nodes(g), [&](Node const &n) { return g.at(n); }); - std::unordered_set edges = get_edges(g); + std::set edges = get_edges(g); - std::unordered_set inputs = g.get_inputs(); + std::set inputs = g.get_inputs(); - std::unordered_map value_data = - generate_unordered_map(get_open_dataflow_values(g), + std::map value_data = + generate_map(get_open_dataflow_values(g), [&](OpenDataflowValue const &v) { return g.at(v); }); return LabelledOpenDataflowGraphData{ diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/labelled_open_dataflow_graph_data.dtg.toml b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/labelled_open_dataflow_graph_data.dtg.toml index 32880d697b..beb8ff8a62 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/labelled_open_dataflow_graph_data.dtg.toml +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/labelled_open_dataflow_graph_data.dtg.toml @@ -14,29 +14,29 @@ includes = [ "utils/graph/open_dataflow_graph/open_dataflow_edge.dtg.h", "utils/graph/open_dataflow_graph/dataflow_graph_input.dtg.h", "utils/graph/open_dataflow_graph/open_dataflow_value.dtg.h", - "", - "", + "", + "", ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_map.h", - "utils/fmt/unordered_set.h", + "utils/hash/map.h", + "utils/hash/set.h", + "utils/fmt/map.h", + "utils/fmt/set.h", ] [[fields]] name = "node_data" -type = "std::unordered_map<::FlexFlow::Node, NodeLabel>" +type = "std::map<::FlexFlow::Node, NodeLabel>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::OpenDataflowEdge>" +type = "std::set<::FlexFlow::OpenDataflowEdge>" [[fields]] name = "inputs" -type = "std::unordered_set<::FlexFlow::DataflowGraphInput>" +type = "std::set<::FlexFlow::DataflowGraphInput>" [[fields]] name = "value_data" -type = "std::unordered_map<::FlexFlow::OpenDataflowValue, ValueLabel>" +type = "std::map<::FlexFlow::OpenDataflowValue, ValueLabel>" diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h index 580a35b3f7..b526fd0169 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h @@ -29,11 +29,11 @@ LabelledOpenDataflowGraphView permute_input_ids( }); }; - std::unordered_map node_labels = - generate_unordered_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); + std::map node_labels = + generate_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); - std::unordered_map value_labels = - generate_unordered_map(get_open_dataflow_values(permuted), + std::map value_labels = + generate_map(get_open_dataflow_values(permuted), [&](OpenDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h index 5119587654..b158068aa4 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_NODE_IDS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_NODE_IDS_H -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h" #include "utils/graph/labelled_open_dataflow_graph/labelled_open_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" @@ -36,13 +36,13 @@ LabelledOpenDataflowGraphView permute_node_ids( }); }; - std::unordered_map node_labels = - generate_unordered_map(get_nodes(permuted), [&](Node const &new_node) { + std::map node_labels = + generate_map(get_nodes(permuted), [&](Node const &new_node) { return g.at(old_node_from_new(new_node)); }); - std::unordered_map value_labels = - generate_unordered_map(get_open_dataflow_values(permuted), + std::map value_labels = + generate_map(get_open_dataflow_values(permuted), [&](OpenDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h index fde90497e7..44b7835f0d 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/rewrite_labels.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELS_H -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h" #include "utils/graph/labelled_open_dataflow_graph/labelled_open_dataflow_graph_view.h" #include "utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h" @@ -26,10 +26,10 @@ LabelledOpenDataflowGraphView rewrite_labels( return f(v, g.at(v)); }; - std::unordered_map node_labels = - generate_unordered_map(get_nodes(g), get_new_node_label); - std::unordered_map value_labels = - generate_unordered_map(get_open_dataflow_values(g), get_new_value_label); + std::map node_labels = + generate_map(get_nodes(g), get_new_node_label); + std::map value_labels = + generate_map(get_open_dataflow_values(g), get_new_value_label); return with_labelling(g, node_labels, value_labels); } diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h index 3697ab0f93..a03673436e 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/with_labelling.h @@ -13,26 +13,26 @@ struct OpenDataflowGraphLabellingWrapper final OpenDataflowGraphLabellingWrapper() = delete; OpenDataflowGraphLabellingWrapper( OpenDataflowGraphView const &unlabelled, - std::unordered_map const &node_labels, - std::unordered_map const &value_labels) + std::map const &node_labels, + std::map const &value_labels) : unlabelled(unlabelled), node_labels(node_labels), value_labels(value_labels) {} - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return this->unlabelled.query_nodes(q); } - std::unordered_set + std::set query_edges(OpenDataflowEdgeQuery const &q) const override { return this->unlabelled.query_edges(q); } - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &q) const override { return this->unlabelled.query_outputs(q); } - std::unordered_set get_inputs() const override { + std::set get_inputs() const override { return this->unlabelled.get_inputs(); } @@ -54,15 +54,15 @@ struct OpenDataflowGraphLabellingWrapper final private: OpenDataflowGraphView unlabelled; - std::unordered_map node_labels; - std::unordered_map value_labels; + std::map node_labels; + std::map value_labels; }; template LabelledOpenDataflowGraphView with_labelling( OpenDataflowGraphView const &g, - std::unordered_map const &node_labels, - std::unordered_map const &value_labels) { + std::map const &node_labels, + std::map const &value_labels) { return LabelledOpenDataflowGraphView::template create< OpenDataflowGraphLabellingWrapper>( g, node_labels, value_labels); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.h index d50670ed41..42cf4d50a7 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.h @@ -22,7 +22,7 @@ std::optional> ValueLabel, GraphInputName, SlotName> const &dst) { - std::unordered_set> + std::set> unlabelled_isomorphisms = find_isomorphisms_between_open_kwarg_dataflow_graphs( static_cast>( diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h index e98b858019..0f66eecc00 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h @@ -1,13 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H -#include "utils/containers/generate_unordered_map.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.h" +#include "utils/containers/set_of.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { @@ -28,12 +29,12 @@ LabelledOpenKwargDataflowGraphData{ - /*nodes=*/generate_unordered_map( + /*nodes=*/generate_map( get_nodes(g), [&](Node const &n) -> NodeLabel { return g.at(n); }), - /*edges=*/get_all_open_kwarg_dataflow_edges(g), - /*inputs=*/get_all_kwarg_dataflow_graph_inputs(g), + /*edges=*/set_of(get_all_open_kwarg_dataflow_edges(g)), + /*inputs=*/set_of(get_all_kwarg_dataflow_graph_inputs(g)), /*outputs=*/ - generate_unordered_map( + generate_map( get_all_open_kwarg_dataflow_values(g), [&](OpenKwargDataflowValue const &v) -> ValueLabel { return g.at(v); }), diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.toml b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.toml index a3ba15c6ff..87d56db913 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.toml +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.toml @@ -8,7 +8,7 @@ features = [ ] template_params = [ - "NodeLabel", + "NodeLabel", "ValueLabel", "GraphInputName", "SlotName", @@ -19,29 +19,29 @@ includes = [ "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge.dtg.h", "utils/graph/open_kwarg_dataflow_graph/kwarg_dataflow_graph_input.dtg.h", "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_value.dtg.h", - "", - "", + "", + "", ] src_includes = [ - "utils/hash/unordered_map.h", - "utils/hash/unordered_set.h", - "utils/fmt/unordered_map.h", - "utils/fmt/unordered_set.h", + "utils/hash/map.h", + "utils/hash/set.h", + "utils/fmt/map.h", + "utils/fmt/set.h", ] [[fields]] name = "node_data" -type = "std::unordered_map<::FlexFlow::Node, NodeLabel>" +type = "std::map<::FlexFlow::Node, NodeLabel>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::OpenKwargDataflowEdge>" +type = "std::set<::FlexFlow::OpenKwargDataflowEdge>" [[fields]] name = "inputs" -type = "std::unordered_set<::FlexFlow::KwargDataflowGraphInput>" +type = "std::set<::FlexFlow::KwargDataflowGraphInput>" [[fields]] name = "value_data" -type = "std::unordered_map<::FlexFlow::OpenKwargDataflowValue, ValueLabel>" +type = "std::map<::FlexFlow::OpenKwargDataflowValue, ValueLabel>" diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h index b6a4366fc2..d06b96e37f 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.h @@ -21,12 +21,12 @@ OpenKwargDataflowGraphData SlotName> const &labelled_data) { OpenKwargDataflowGraphData result = OpenKwargDataflowGraphData{ - /*nodes=*/unordered_keys(labelled_data.node_data), + /*nodes=*/keys(labelled_data.node_data), /*edges=*/labelled_data.edges, /*inputs=*/labelled_data.inputs, /*outputs=*/ filtrans( - unordered_keys(labelled_data.value_data), + keys(labelled_data.value_data), [](OpenKwargDataflowValue const &v) { return v.try_require_internal(); }), diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h index 120833020f..50b59025f2 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h @@ -19,7 +19,7 @@ std::string labelled_open_kwarg_dataflow_graph_view_as_dot( std::function const &render_value_label, std::function const &render_slot_name, std::function( - std::unordered_set const &)> const &order_slots) { + std::set const &)> const &order_slots) { std::function render_node = [&](Node const &n) -> nlohmann::json { return render_node_label(g.at(n)); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h index bee9ad0c37..689a5380c0 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h @@ -20,28 +20,28 @@ struct OpenKwargDataflowGraphLabellingWrapper final OpenKwargDataflowGraphLabellingWrapper() = delete; OpenKwargDataflowGraphLabellingWrapper( OpenKwargDataflowGraphView const &unlabelled, - std::unordered_map const &node_labels, - std::unordered_map, + std::map const &node_labels, + std::map, ValueLabel> const &value_labels) : unlabelled(unlabelled), node_labels(node_labels), value_labels(value_labels) {} - std::unordered_set query_nodes(NodeQuery const &q) const override { + std::set query_nodes(NodeQuery const &q) const override { return this->unlabelled.query_nodes(q); } - std::unordered_set> + std::set> query_edges(OpenKwargDataflowEdgeQuery const &q) const override { return this->unlabelled.query_edges(q); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &q) const override { return this->unlabelled.query_outputs(q); } - std::unordered_set> + std::set> get_inputs() const override { return this->unlabelled.get_inputs(); } @@ -65,8 +65,8 @@ struct OpenKwargDataflowGraphLabellingWrapper final private: OpenKwargDataflowGraphView unlabelled; - std::unordered_map node_labels; - std::unordered_map, + std::map node_labels; + std::map, ValueLabel> value_labels; }; @@ -81,8 +81,8 @@ LabelledOpenKwargDataflowGraphView open_kwarg_dataflow_graph_view_with_labelling( OpenKwargDataflowGraphView const &g, - std::unordered_map const &node_labels, - std::unordered_map, + std::map const &node_labels, + std::map, ValueLabel> const &value_labels) { return LabelledOpenKwargDataflowGraphView node_labels = - generate_unordered_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); + std::map node_labels = + generate_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); - std::unordered_map, + std::map, ValueLabel> - value_labels = generate_unordered_map( + value_labels = generate_map( get_all_open_kwarg_dataflow_values(permuted), [&](OpenKwargDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h index d89c935d52..29a8235900 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_NODE_IDS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_PERMUTE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_NODE_IDS_H -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms/new_node.dtg.h" @@ -55,14 +55,14 @@ LabelledOpenKwargDataflowGraphView node_labels = - generate_unordered_map(get_nodes(permuted), [&](Node const &new_node) { + std::map node_labels = + generate_map(get_nodes(permuted), [&](Node const &new_node) { return g.at(old_node_from_new(new_node)); }); - std::unordered_map, + std::map, ValueLabel> - value_labels = generate_unordered_map( + value_labels = generate_map( get_all_open_kwarg_dataflow_values(permuted), [&](OpenKwargDataflowValue const &new_value) { return g.at(old_value_from_new(new_value)); }); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h index d5d1435462..61bd3957b5 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_LABELS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_REWRITE_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_LABELS_H -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h" #include "utils/graph/node/algorithms.h" @@ -38,11 +38,11 @@ LabelledOpenKwargDataflowGraphView const &v) -> NewValueLabel { return f(v, g.at(v)); }; - std::unordered_map node_labels = - generate_unordered_map(get_nodes(g), get_new_node_label); - std::unordered_map, + std::map node_labels = + generate_map(get_nodes(g), get_new_node_label); + std::map, NewValueLabel> - value_labels = generate_unordered_map(get_all_open_kwarg_dataflow_values(g), + value_labels = generate_map(get_all_open_kwarg_dataflow_values(g), get_new_value_label); return open_kwarg_dataflow_graph_view_with_labelling( g, node_labels, value_labels); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h index bec1c540ea..45f82ad823 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h @@ -22,10 +22,10 @@ struct ILabelledOpenKwargDataflowGraph SlotName> { virtual KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::unordered_map> const &inputs, - std::unordered_map const &output_labels) = 0; + std::map const &output_labels) = 0; virtual KwargDataflowGraphInput add_input(GraphInputName const &name, ValueLabel const &value_label) = 0; diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h index 1f03f4d341..5f92a86b48 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h @@ -29,10 +29,10 @@ struct LabelledOpenKwargDataflowGraph KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::unordered_map> const &inputs, - std::unordered_map const &output_labels) { + std::map const &output_labels) { return this->get_interface().add_node(node_label, inputs, output_labels); } diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h index d6c1ffd95c..fa18a86a8d 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map get_edge_counts(MultiDiGraphView const &); +std::map get_edge_counts(MultiDiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_edges.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_edges.h index bde2193241..7eb6b133f9 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_edges.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_edges.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_edges(MultiDiGraphView const &); +std::set get_edges(MultiDiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h index 471a12a44b..a8683d1ac4 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h @@ -5,12 +5,12 @@ namespace FlexFlow { -std::unordered_set get_incoming_edges(MultiDiGraphView const &, +std::set get_incoming_edges(MultiDiGraphView const &, Node const &); -std::unordered_map> +std::map> get_incoming_edges(MultiDiGraphView const &g, - std::unordered_set const &nodes); + std::set const &nodes); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.h index 967184e397..916bd4b116 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_map +std::map get_multidiedge_to_diedge_map(MultiDiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h index bd8c364f7e..b4d1254512 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h @@ -2,16 +2,16 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_MULTIDIGRAPH_ALGORITHMS_GET_OUTGOING_EDGES_H #include "utils/graph/multidigraph/multidigraph_view.h" -#include +#include namespace FlexFlow { -std::unordered_set get_outgoing_edges(MultiDiGraphView const &, +std::set get_outgoing_edges(MultiDiGraphView const &, Node const &); -std::unordered_map> +std::map> get_outgoing_edges(MultiDiGraphView const &g, - std::unordered_set const &ns); + std::set const &ns); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/multidigraph/i_multidigraph_view.h b/lib/utils/include/utils/graph/multidigraph/i_multidigraph_view.h index 4c92880067..2e2359d4db 100644 --- a/lib/utils/include/utils/graph/multidigraph/i_multidigraph_view.h +++ b/lib/utils/include/utils/graph/multidigraph/i_multidigraph_view.h @@ -17,13 +17,13 @@ struct IMultiDiGraphView : virtual public IDiGraphView { IMultiDiGraphView(IMultiDiGraphView const &) = delete; IMultiDiGraphView &operator=(IMultiDiGraphView const &) = delete; - virtual std::unordered_set query_edges(EdgeQuery const &) const = 0; + virtual std::set query_edges(EdgeQuery const &) const = 0; virtual Node get_multidiedge_src(MultiDiEdge const &) const = 0; virtual Node get_multidiedge_dst(MultiDiEdge const &) const = 0; virtual ~IMultiDiGraphView() = default; - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override final; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(IMultiDiGraphView); diff --git a/lib/utils/include/utils/graph/multidigraph/multidigraph.h b/lib/utils/include/utils/graph/multidigraph/multidigraph.h index 692ee33783..bab0aebd5d 100644 --- a/lib/utils/include/utils/graph/multidigraph/multidigraph.h +++ b/lib/utils/include/utils/graph/multidigraph/multidigraph.h @@ -13,8 +13,8 @@ struct MultiDiGraph : virtual public MultiDiGraphView { void remove_node(Node const &); void remove_edge(MultiDiEdge const &); - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(MultiDiEdgeQuery const &) const; + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(MultiDiEdgeQuery const &) const; Node get_multidiedge_src(MultiDiEdge const &) const; Node get_multidiedge_dst(MultiDiEdge const &) const; diff --git a/lib/utils/include/utils/graph/multidigraph/multidigraph_view.h b/lib/utils/include/utils/graph/multidigraph/multidigraph_view.h index 229c859338..c5ef74a842 100644 --- a/lib/utils/include/utils/graph/multidigraph/multidigraph_view.h +++ b/lib/utils/include/utils/graph/multidigraph/multidigraph_view.h @@ -10,8 +10,8 @@ struct MultiDiGraphView : virtual public DiGraphView { MultiDiGraphView(MultiDiGraphView const &) = default; MultiDiGraphView &operator=(MultiDiGraphView const &) = default; - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(MultiDiEdgeQuery const &) const; + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(MultiDiEdgeQuery const &) const; Node get_multidiedge_src(MultiDiEdge const &) const; Node get_multidiedge_dst(MultiDiEdge const &) const; diff --git a/lib/utils/include/utils/graph/node/algorithms.h b/lib/utils/include/utils/graph/node/algorithms.h index 5c11a0cd96..3eb8b7d36b 100644 --- a/lib/utils/include/utils/graph/node/algorithms.h +++ b/lib/utils/include/utils/graph/node/algorithms.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_nodes(GraphView const &); +std::set get_nodes(GraphView const &); bool has_node(GraphView const &, Node const &); size_t num_nodes(GraphView const &); bool empty(GraphView const &); diff --git a/lib/utils/include/utils/graph/node/graph.h b/lib/utils/include/utils/graph/node/graph.h index 1d94d1a65e..3a7856432a 100644 --- a/lib/utils/include/utils/graph/node/graph.h +++ b/lib/utils/include/utils/graph/node/graph.h @@ -18,7 +18,7 @@ struct Graph : virtual GraphView { void add_node_unsafe(Node const &); void remove_node_unsafe(Node const &); - std::unordered_set query_nodes(NodeQuery const &) const; + std::set query_nodes(NodeQuery const &) const; template static typename std::enable_if::value, Graph>::type diff --git a/lib/utils/include/utils/graph/node/graph_view.h b/lib/utils/include/utils/graph/node/graph_view.h index 8d904e05f2..4556d80448 100644 --- a/lib/utils/include/utils/graph/node/graph_view.h +++ b/lib/utils/include/utils/graph/node/graph_view.h @@ -8,7 +8,7 @@ namespace FlexFlow { struct GraphView { - std::unordered_set query_nodes(NodeQuery const &) const; + std::set query_nodes(NodeQuery const &) const; friend bool is_ptr_equal(GraphView const &, GraphView const &); template diff --git a/lib/utils/include/utils/graph/node/i_graph_view.h b/lib/utils/include/utils/graph/node/i_graph_view.h index be5b07a685..010586f518 100644 --- a/lib/utils/include/utils/graph/node/i_graph_view.h +++ b/lib/utils/include/utils/graph/node/i_graph_view.h @@ -13,7 +13,7 @@ struct IGraphView { virtual IGraphView *clone() const = 0; - virtual std::unordered_set query_nodes(NodeQuery const &) const = 0; + virtual std::set query_nodes(NodeQuery const &) const = 0; virtual ~IGraphView(){}; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(IGraphView); diff --git a/lib/utils/include/utils/graph/node/node_query.h b/lib/utils/include/utils/graph/node/node_query.h index 2ec8958083..9ab9bde9fe 100644 --- a/lib/utils/include/utils/graph/node/node_query.h +++ b/lib/utils/include/utils/graph/node/node_query.h @@ -8,8 +8,8 @@ namespace FlexFlow { NodeQuery node_query_all(); NodeQuery query_intersection(NodeQuery const &, NodeQuery const &); NodeQuery query_union(NodeQuery const &, NodeQuery const &); -std::unordered_set apply_node_query(NodeQuery const &, - std::unordered_set const &); +std::set apply_node_query(NodeQuery const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.h index 022fc5b9fd..873d445167 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.h @@ -6,7 +6,7 @@ namespace FlexFlow { -std::unordered_set +std::set find_isomorphisms(OpenDataflowGraphView const &, OpenDataflowGraphView const &); diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.h index 1fbbea21b0..b210390386 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.h @@ -10,12 +10,12 @@ struct FromOpenDataflowGraphDataView final : virtual public IOpenDataflowGraphView { FromOpenDataflowGraphDataView(OpenDataflowGraphData const &); - std::unordered_set query_nodes(NodeQuery const &) const override; - std::unordered_set + std::set query_nodes(NodeQuery const &) const override; + std::set query_edges(OpenDataflowEdgeQuery const &) const override; - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &) const override; - std::unordered_set get_inputs() const override; + std::set get_inputs() const override; FromOpenDataflowGraphDataView *clone() const override; diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_edges.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_edges.h index 0710b3d970..15717f4afa 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_edges.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_edges.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_edges(OpenDataflowGraphView const &); +std::set get_edges(OpenDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h index 22d66a0c0f..4266c66e18 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h @@ -5,13 +5,13 @@ namespace FlexFlow { -std::unordered_set +std::set get_incoming_edges(OpenDataflowGraphView const &); std::vector get_incoming_edges(OpenDataflowGraphView const &, Node const &); -std::unordered_map> +std::map> get_incoming_edges(OpenDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.h index 98231c8f8c..956d4b237d 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set +std::set get_open_dataflow_graph_inputs(OpenDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.h index bd7749a172..e6dab6b92a 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.h @@ -6,7 +6,7 @@ namespace FlexFlow { -std::unordered_set +std::set get_open_dataflow_value_uses(OpenDataflowGraphView const &view, OpenDataflowValue const &value); diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h index 5d8f58540e..5ee559fb17 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h @@ -6,7 +6,7 @@ namespace FlexFlow { -std::unordered_set +std::set get_open_dataflow_values(OpenDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.h index a89b4e1bc1..e3bc7ee237 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_source_nodes(OpenDataflowGraphView const &); +std::set get_source_nodes(OpenDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h index f5bbbc228d..425ae32d44 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h @@ -9,16 +9,16 @@ namespace FlexFlow { OpenDataflowSubgraphResult get_subgraph(OpenDataflowGraphView const &, - std::unordered_set const &); + std::set const &); bidict get_full_graph_values_to_subgraph_inputs( OpenDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes); + std::set const &subgraph_nodes); OpenDataflowGraphData get_subgraph_data(OpenDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes, + std::set const &subgraph_nodes, bidict const &full_graph_values_to_subgraph_inputs); diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.h index 0df5f8458c..1416d5de77 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.h @@ -5,9 +5,9 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_incoming_edges(OpenDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h index 136ac071b5..017bac26b9 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h @@ -6,9 +6,9 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_inputs(OpenDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.h index 2325dcfbda..15076aa4e7 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set +std::set get_unused_open_dataflow_graph_inputs(OpenDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_data.dtg.toml b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_data.dtg.toml index 98285329fd..50bebc614f 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_data.dtg.toml +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_data.dtg.toml @@ -12,26 +12,26 @@ includes = [ "utils/graph/open_dataflow_graph/open_dataflow_edge.dtg.h", "utils/graph/open_dataflow_graph/dataflow_graph_input.dtg.h", "utils/graph/dataflow_graph/dataflow_output.dtg.h", - "", + "", ] src_includes = [ - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "nodes" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::OpenDataflowEdge>" +type = "std::set<::FlexFlow::OpenDataflowEdge>" [[fields]] name = "inputs" -type = "std::unordered_set<::FlexFlow::DataflowGraphInput>" +type = "std::set<::FlexFlow::DataflowGraphInput>" [[fields]] name = "outputs" -type = "std::unordered_set<::FlexFlow::DataflowOutput>" +type = "std::set<::FlexFlow::DataflowOutput>" diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_isomorphism.dtg.toml b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_isomorphism.dtg.toml index d85a29403f..8bc92ba2be 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_isomorphism.dtg.toml +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/open_dataflow_graph_isomorphism.dtg.toml @@ -3,6 +3,7 @@ name = "OpenDataflowGraphIsomorphism" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.h b/lib/utils/include/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.h index b47b3814fc..50755d2279 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.h @@ -9,11 +9,11 @@ namespace FlexFlow { struct IOpenDataflowGraphView : virtual public IDataflowGraphView { - virtual std::unordered_set get_inputs() const = 0; - virtual std::unordered_set + virtual std::set get_inputs() const = 0; + virtual std::set query_edges(OpenDataflowEdgeQuery const &) const = 0; - std::unordered_set + std::set query_edges(DataflowEdgeQuery const &) const override final; virtual ~IOpenDataflowGraphView() = default; diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h index ae6e30549b..6c0e31f4cc 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h @@ -15,9 +15,9 @@ OpenDataflowEdgeQuery open_dataflow_edge_query_all_outgoing_from(OpenDataflowValue const &); OpenDataflowEdgeQuery open_dataflow_edge_query_all_incoming_to(DataflowInput const &); -std::unordered_set apply_open_dataflow_edge_query( +std::set apply_open_dataflow_edge_query( OpenDataflowEdgeQuery const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h index e1bbc231c2..f29db36920 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h @@ -11,8 +11,8 @@ struct OpenDataflowGraphView : virtual public DataflowGraphView { OpenDataflowGraphView(OpenDataflowGraphView const &) = default; OpenDataflowGraphView &operator=(OpenDataflowGraphView const &) = default; - std::unordered_set get_inputs() const; - std::unordered_set + std::set get_inputs() const; + std::set query_edges(OpenDataflowEdgeQuery const &) const; template diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.h b/lib/utils/include/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.h index f3d54e4329..ca0da509e2 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.h @@ -14,12 +14,12 @@ struct UnorderedSetOpenDataflowGraph : public IOpenDataflowGraph { NodeAddedResult add_node(std::vector const &inputs, nonnegative_int num_outputs) override; - std::unordered_set query_nodes(NodeQuery const &) const override; - std::unordered_set + std::set query_nodes(NodeQuery const &) const override; + std::set query_edges(OpenDataflowEdgeQuery const &) const override; - std::unordered_set + std::set query_outputs(DataflowOutputQuery const &) const override; - std::unordered_set get_inputs() const override; + std::set get_inputs() const override; DataflowGraphInput add_input() override; UnorderedSetOpenDataflowGraph *clone() const override; @@ -28,20 +28,20 @@ struct UnorderedSetOpenDataflowGraph : public IOpenDataflowGraph { UnorderedSetOpenDataflowGraph( NodeSource const &node_source, DataflowGraphInputSource const &input_source, - std::unordered_set const &nodes, - std::unordered_set const &standard_edges, - std::unordered_set const &input_edges, - std::unordered_set const &outputs, - std::unordered_set const &graph_inputs); + std::set const &nodes, + std::set const &standard_edges, + std::set const &input_edges, + std::set const &outputs, + std::set const &graph_inputs); private: NodeSource node_source; DataflowGraphInputSource input_source; - std::unordered_set nodes; - std::unordered_set standard_edges; - std::unordered_set input_edges; - std::unordered_set outputs; - std::unordered_set graph_inputs; + std::set nodes; + std::set standard_edges; + std::set input_edges; + std::set outputs; + std::set graph_inputs; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(UnorderedSetOpenDataflowGraph); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h index 32710e75bf..cfa7c7c7de 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h @@ -30,34 +30,34 @@ std::optional> KwargDataflowGraphInput> const &unused_graph_inputs_mapping) { { - std::unordered_set already_mapped_src_nodes = + std::set already_mapped_src_nodes = left_entries(sink_node_mapping); - std::unordered_set src_g_sink_nodes = get_terminal_nodes(src_g); + std::set src_g_sink_nodes = set_of(get_terminal_nodes(src_g)); ASSERT(already_mapped_src_nodes == src_g_sink_nodes); } { - std::unordered_set already_mapped_dst_nodes = + std::set already_mapped_dst_nodes = right_entries(sink_node_mapping); - std::unordered_set dst_g_sink_nodes = get_terminal_nodes(dst_g); + std::set dst_g_sink_nodes = set_of(get_terminal_nodes(dst_g)); ASSERT(already_mapped_dst_nodes == dst_g_sink_nodes); } { - std::unordered_set> + std::set> already_mapped_src_inputs = left_entries(unused_graph_inputs_mapping); - std::unordered_set> + std::set> src_g_unused_inputs = - get_unused_open_kwarg_dataflow_graph_inputs(src_g); + set_of(get_unused_open_kwarg_dataflow_graph_inputs(src_g)); ASSERT(already_mapped_src_inputs == src_g_unused_inputs); } { - std::unordered_set> + std::set> already_mapped_dst_inputs = right_entries(unused_graph_inputs_mapping); - std::unordered_set> + std::set> dst_g_unused_inputs = - get_unused_open_kwarg_dataflow_graph_inputs(dst_g); + set_of(get_unused_open_kwarg_dataflow_graph_inputs(dst_g)); ASSERT(already_mapped_dst_inputs == dst_g_unused_inputs); } @@ -178,11 +178,11 @@ std::optional> result->node_mapping.equate(src_node, dst_node); - std::unordered_map> src_incoming_edges = get_incoming_open_kwarg_dataflow_edges_for_node(src_g, src_node); - std::unordered_map> dst_incoming_edges = get_incoming_open_kwarg_dataflow_edges_for_node(dst_g, dst_node); @@ -206,14 +206,14 @@ std::optional> } template -std::unordered_set> +std::set> find_isomorphisms_between_open_kwarg_dataflow_graphs( OpenKwargDataflowGraphView const &src, OpenKwargDataflowGraphView const &dst) { - std::unordered_set> result; + std::set> result; std::vector src_sink_nodes = vector_of(get_terminal_nodes(src)); - std::unordered_set dst_sink_nodes = get_terminal_nodes(dst); + std::set dst_sink_nodes = get_terminal_nodes(dst); if (src_sink_nodes.size() != dst_sink_nodes.size()) { return {}; @@ -221,7 +221,7 @@ std::unordered_set> std::vector> src_unused_graph_inputs = vector_of(get_unused_open_kwarg_dataflow_graph_inputs(src)); - std::unordered_set> + std::set> dst_unused_graph_inputs = get_unused_open_kwarg_dataflow_graph_inputs(dst); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.h index 060463d6be..885259b77e 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.h @@ -14,7 +14,7 @@ bidict, generate_new_kwarg_dataflow_graph_input_id_permutation( OpenKwargDataflowGraphView const &g, std::function const &input_id_source) { - std::unordered_set> old_graph_inputs = + std::set> old_graph_inputs = get_all_kwarg_dataflow_graph_inputs(g); auto fresh_input_id = [&]() -> GraphInputName { diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.h index b14deaebf4..ae9f9728be 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.h @@ -6,7 +6,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_all_kwarg_dataflow_graph_inputs( OpenKwargDataflowGraphView const &view) { return view.get_inputs(); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h index 7459dac065..dbb7eff9ed 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h @@ -8,7 +8,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_all_open_kwarg_dataflow_edges( OpenKwargDataflowGraphView const &view) { return view.query_edges( diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.h index 73b0e9c29d..ca8a26b218 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.h @@ -9,13 +9,13 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_all_open_kwarg_dataflow_values( OpenKwargDataflowGraphView const &g) { - std::unordered_set> internal_values = + std::set> internal_values = get_all_kwarg_dataflow_outputs(g); - std::unordered_set> external_values = + std::set> external_values = get_all_kwarg_dataflow_graph_inputs(g); return set_union( diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.h index 20e078fc52..2e442db240 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_OPEN_KWARG_DATAFLOW_EDGES_FOR_NODE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_INCOMING_OPEN_KWARG_DATAFLOW_EDGES_FOR_NODE_H -#include "utils/containers/unordered_map_from_pairs.h" +#include "utils/containers/map_from_pairs.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" namespace FlexFlow { template -std::unordered_map> +std::map> get_incoming_open_kwarg_dataflow_edges_for_node( OpenKwargDataflowGraphView const &g, Node const &n) { @@ -29,7 +29,7 @@ std::unordered_map> }, }; - return unordered_map_from_pairs( + return map_from_pairs( transform(g.query_edges(query), [](OpenKwargDataflowEdge const &e) { return std::pair{ diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.h index e30d554e89..482274cd46 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.h @@ -9,7 +9,7 @@ namespace FlexFlow { template -std::unordered_map> +std::map> get_incoming_open_kwarg_dataflow_values_for_node( OpenKwargDataflowGraphView const &g, Node const &n) { diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h index b27ad13cc9..10f10cc58c 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h @@ -7,6 +7,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -15,10 +16,10 @@ OpenKwargDataflowGraphData get_open_kwarg_dataflow_graph_data( OpenKwargDataflowGraphView const &g) { return OpenKwargDataflowGraphData{ - /*nodes=*/get_nodes(g), - /*edges=*/get_all_open_kwarg_dataflow_edges(g), - /*inputs=*/get_all_kwarg_dataflow_graph_inputs(g), - /*outputs=*/get_all_kwarg_dataflow_outputs(g), + /*nodes=*/set_of(get_nodes(g)), + /*edges=*/set_of(get_all_open_kwarg_dataflow_edges(g)), + /*inputs=*/set_of(get_all_kwarg_dataflow_graph_inputs(g)), + /*outputs=*/set_of(get_all_kwarg_dataflow_outputs(g)), }; } diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h index a6a5391ed6..2c8d94be20 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h @@ -11,6 +11,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_subgraph_result.dtg.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h" #include "utils/overload.h" +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -18,7 +19,7 @@ template OpenKwargDataflowSubgraphResult get_open_kwarg_dataflow_graph_subgraph( OpenKwargDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes, + std::set const &subgraph_nodes, std::function const &input_source) { bidict, KwargDataflowGraphInput> @@ -39,7 +40,7 @@ bidict, KwargDataflowGraphInput> get_full_kwarg_dataflow_graph_values_to_subgraph_inputs( OpenKwargDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes, + std::set const &subgraph_nodes, std::function const &input_source) { return generate_bidict( get_open_kwarg_dataflow_subgraph_inputs(g, subgraph_nodes), @@ -63,13 +64,14 @@ template OpenKwargDataflowGraphData get_open_kwarg_dataflow_subgraph_data( OpenKwargDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes, + std::set const &subgraph_nodes, bidict, KwargDataflowGraphInput> const &full_graph_values_to_subgraph_inputs) { - std::unordered_set> + + std::set> subgraph_input_edges = transform( - get_open_kwarg_dataflow_subgraph_incoming_edges(g, subgraph_nodes), + set_of(get_open_kwarg_dataflow_subgraph_incoming_edges(g, set_of(subgraph_nodes))), [&](OpenKwargDataflowEdge const &edge) { return edge.template visit< OpenKwargDataflowEdge>(overload{ @@ -112,23 +114,23 @@ OpenKwargDataflowGraphData }, }; - std::unordered_set> - subgraph_interior_edges = g.query_edges(subgraph_interior_edges_query); + std::set> + subgraph_interior_edges = set_of(g.query_edges(subgraph_interior_edges_query)); - std::unordered_set> subgraph_inputs = - unordered_set_of(values(full_graph_values_to_subgraph_inputs)); + std::set> subgraph_inputs = + set_of(values(full_graph_values_to_subgraph_inputs)); - std::unordered_set> subgraph_outputs = + std::set> subgraph_outputs = filter(g.query_outputs(kwarg_dataflow_output_query_all()), [&](KwargDataflowOutput const &o) { return contains(subgraph_nodes, o.node); }); return OpenKwargDataflowGraphData{ - subgraph_nodes, - set_union(subgraph_input_edges, subgraph_interior_edges), - subgraph_inputs, - subgraph_outputs, + /*nodes=*/set_of(subgraph_nodes), + /*edges=*/set_union(subgraph_input_edges, subgraph_interior_edges), + /*inputs=*/subgraph_inputs, + /*outputs=*/subgraph_outputs, }; } diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.h index 975180a12e..e31157d46b 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.h @@ -9,11 +9,11 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_open_kwarg_dataflow_subgraph_incoming_edges( OpenKwargDataflowGraphView const &g, - std::unordered_set const &subgraph) { - std::unordered_set all_nodes = get_nodes(g); + std::set const &subgraph) { + std::set all_nodes = get_nodes(g); query_set src_query = query_set::match_values_in(set_of(set_minus(all_nodes, subgraph))); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.h index 2e724bc217..8de5546fb4 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.h @@ -8,10 +8,10 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_open_kwarg_dataflow_subgraph_inputs( OpenKwargDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { return transform( get_open_kwarg_dataflow_subgraph_incoming_edges(g, subgraph_nodes), diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h index 2c80a5ab8d..77cac17d1c 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h @@ -11,7 +11,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_open_kwarg_dataflow_value_uses( OpenKwargDataflowGraphView const &g, OpenKwargDataflowValue const &v) { @@ -41,7 +41,7 @@ std::unordered_set> }; }}); - std::unordered_set> edges = + std::set> edges = g.query_edges(query); return transform( diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.h index b6880abecd..42b04b0580 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.h @@ -8,7 +8,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_unused_open_kwarg_dataflow_graph_inputs( OpenKwargDataflowGraphView const &g) { return filter( diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h index 423f7a9a2c..48b5729ab0 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h @@ -7,6 +7,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_value.dtg.h" +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -32,8 +33,8 @@ std::string open_kwarg_dataflow_graph_as_dot( return j; }; - std::function(std::unordered_set const &)> - order_slots = [](std::unordered_set const &unordered) { + std::function(std::set const &)> + order_slots = [](std::set const &unordered) { return sorted(unordered); }; @@ -50,7 +51,7 @@ std::string open_kwarg_dataflow_graph_as_dot( &render_value, std::function const &render_slot_name, std::function( - std::unordered_set const &)> const &order_slots) { + std::set const &)> const &order_slots) { std::pair>, bidict, Node>> closed_g_and_mapping = @@ -104,11 +105,11 @@ std::string open_kwarg_dataflow_graph_as_dot( }; std::function>( - std::unordered_set> const &)> + std::set> const &)> closed_order_slots = - [&](std::unordered_set> const &unsorted) + [&](std::set> const &unsorted) -> std::vector> { - std::unordered_set not_nullopt = filtrans( + std::set not_nullopt = filtrans( unsorted, [](std::optional const &s) -> std::optional { return s; diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.toml b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.toml index 49c139cd0a..20a05fc206 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.toml +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.toml @@ -3,6 +3,7 @@ name = "OpenKwargDataflowGraphData" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -17,26 +18,26 @@ includes = [ "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge.dtg.h", "utils/graph/open_kwarg_dataflow_graph/kwarg_dataflow_graph_input.dtg.h", "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.h", - "", + "", ] src_includes = [ - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "nodes" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::OpenKwargDataflowEdge>" +type = "std::set<::FlexFlow::OpenKwargDataflowEdge>" [[fields]] name = "inputs" -type = "std::unordered_set<::FlexFlow::KwargDataflowGraphInput>" +type = "std::set<::FlexFlow::KwargDataflowGraphInput>" [[fields]] name = "outputs" -type = "std::unordered_set<::FlexFlow::KwargDataflowOutput>" +type = "std::set<::FlexFlow::KwargDataflowOutput>" diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h index c7328a8f2a..b72bf95701 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h @@ -13,7 +13,7 @@ namespace FlexFlow { template void require_open_kwarg_dataflow_graph_data_is_valid( OpenKwargDataflowGraphData const &data) { - std::unordered_set> + std::set> inputs_from_edges = filtrans( data.edges, [](OpenKwargDataflowEdge const &e) diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_isomorphism.dtg.toml b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_isomorphism.dtg.toml index f7b08229d5..c9b6eecd98 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_isomorphism.dtg.toml +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_isomorphism.dtg.toml @@ -3,6 +3,7 @@ name = "OpenKwargDataflowGraphIsomorphism" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.h index 38c9fe89d2..93db7cdc42 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.h @@ -16,7 +16,7 @@ OpenKwargDataflowGraphView bidict, KwargDataflowGraphInput> const &new_input_to_old_input) { - std::unordered_set> g_inputs = + std::set> g_inputs = get_all_kwarg_dataflow_graph_inputs(g); ASSERT(g_inputs == new_input_to_old_input.right_values()); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h index fd644c355a..d3db64c5b3 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h @@ -12,7 +12,7 @@ std::optional> try_find_isomorphism_between_open_kwarg_dataflow_graphs( OpenKwargDataflowGraphView const &src, OpenKwargDataflowGraphView const &dst) { - std::unordered_set> + std::set> isomorphisms = find_isomorphisms_between_open_kwarg_dataflow_graphs(src, dst); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h index 1626c5833b..7e1fe0e60c 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h @@ -13,6 +13,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" #include "utils/overload.h" +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -88,14 +89,17 @@ std::pair>, KwargDataflowGraphData> closed_g_data = KwargDataflowGraphData>{ - /*nodes=*/set_union(open_g_data.nodes, - right_entries(graph_input_nodes)), - /*edges=*/transform(open_g_data.edges, convert_edge), + /*nodes=*/set_of( + set_union( + open_g_data.nodes, + right_entries(graph_input_nodes))), + /*edges=*/set_of(transform(open_g_data.edges, convert_edge)), /*outputs=*/ - set_union( - transform(open_g_data.outputs, convert_kwarg_dataflow_output), - transform(open_g_data.inputs, - kwarg_dataflow_output_for_graph_input)), + set_of( + set_union( + transform(open_g_data.outputs, convert_kwarg_dataflow_output), + transform(open_g_data.inputs, + kwarg_dataflow_output_for_graph_input))), }; ASSERT(closed_g_data.edges.size() == open_g_data.edges.size()); @@ -103,7 +107,7 @@ std::pair>, KwargDataflowGraphView> closed_g = view_from_kwarg_dataflow_graph_data(closed_g_data); - ASSERT(closed_g_data.edges == get_all_kwarg_dataflow_edges(closed_g)); + ASSERT(closed_g_data.edges == set_of(get_all_kwarg_dataflow_edges(closed_g))); return std::pair{ closed_g, diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h index 178d64933a..56e43a73bd 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h @@ -7,6 +7,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge_query.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" +#include "utils/containers/set_of.h" namespace FlexFlow { @@ -17,28 +18,28 @@ struct ViewFromOpenKwargDataflowGraphData final OpenKwargDataflowGraphData const &data) : data(data) {} - std::unordered_set query_nodes(NodeQuery const &query) const override { - return apply_node_query(query, this->data.nodes); + std::set query_nodes(NodeQuery const &query) const override { + return apply_node_query(query, set_of(this->data.nodes)); } - std::unordered_set> + std::set> get_inputs() const override { - return this->data.inputs; + return set_of(this->data.inputs); } - std::unordered_set> + std::set> query_edges(OpenKwargDataflowEdgeQuery const &query) const override { return filter( - this->data.edges, + set_of(this->data.edges), [&](OpenKwargDataflowEdge const &e) { return open_kwarg_dataflow_edge_query_includes(query, e); }); } - std::unordered_set> query_outputs( + std::set> query_outputs( KwargDataflowOutputQuery const &query) const override { - return filter(this->data.outputs, + return filter(set_of(this->data.outputs), [&](KwargDataflowOutput const &o) { return kwarg_dataflow_output_query_includes(query, o); }); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h index 7f59628705..0be5f46e37 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h @@ -11,10 +11,10 @@ template struct IOpenKwargDataflowGraph : virtual public IOpenKwargDataflowGraphView { virtual KwargNodeAddedResult add_node( - std::unordered_map> const &inputs, - std::unordered_set const &outputs) = 0; + std::set const &outputs) = 0; virtual KwargDataflowGraphInput add_input(GraphInputName const &name) = 0; virtual IOpenKwargDataflowGraph *clone() const = 0; diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h index aee878a268..162224853d 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h @@ -12,13 +12,13 @@ namespace FlexFlow { template struct IOpenKwargDataflowGraphView : virtual public IKwargDataflowGraphView { - virtual std::unordered_set> + virtual std::set> get_inputs() const = 0; - virtual std::unordered_set> + virtual std::set> query_edges(OpenKwargDataflowEdgeQuery const &) const = 0; - std::unordered_set> query_edges( + std::set> query_edges( KwargDataflowEdgeQuery const &query) const override final { OpenKwargDataflowEdgeQuery open_query = OpenKwargDataflowEdgeQuery{ @@ -28,7 +28,7 @@ struct IOpenKwargDataflowGraphView /*standard_edge_query=*/query, }; - std::unordered_set> + std::set> open_edges = this->query_edges(open_query); return transform( diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h index 1c903b1dab..335ec6a67d 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h @@ -13,10 +13,10 @@ struct OpenKwargDataflowGraph : virtual public OpenKwargDataflowGraphView { public: KwargNodeAddedResult add_node( - std::unordered_map> const &inputs, - std::unordered_set const &outputs) { + std::set const &outputs) { return this->get_interface().add_node(inputs, outputs); } diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h index 3153d1cbe9..989634bcf5 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h @@ -14,12 +14,12 @@ struct OpenKwargDataflowGraphView OpenKwargDataflowGraphView & operator=(OpenKwargDataflowGraphView const &) = default; - std::unordered_set> + std::set> get_inputs() const { return this->get_interface().get_inputs(); } - std::unordered_set> + std::set> query_edges( OpenKwargDataflowEdgeQuery const &q) const { return this->get_interface().query_edges(q); diff --git a/lib/utils/include/utils/graph/query_set.h b/lib/utils/include/utils/graph/query_set.h index c59268d6eb..a36327875f 100644 --- a/lib/utils/include/utils/graph/query_set.h +++ b/lib/utils/include/utils/graph/query_set.h @@ -9,15 +9,15 @@ #include "utils/containers/set_of.h" #include "utils/containers/set_union.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/exception.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/hash-utils.h" #include "utils/hash/set.h" #include "utils/optional.h" #include #include -#include +#include namespace FlexFlow { @@ -68,10 +68,10 @@ struct query_set { return !q.query.has_value(); } - friend std::unordered_set allowed_values(query_set const &q) { + friend std::set allowed_values(query_set const &q) { assert(!is_matchall(q)); std::set query_value = q.query.value(); - return std::unordered_set{query_value.begin(), query_value.end()}; + return std::set{query_value.begin(), query_value.end()}; } std::optional> const &value() const { @@ -108,19 +108,19 @@ bool includes(query_set const &q, T const &v) { } template -std::unordered_set apply_query(query_set const &q, C const &c) { +std::set apply_query(query_set const &q, C const &c) { if (is_matchall(q)) { - return unordered_set_of(c); + return set_of(c); } - return filter(unordered_set_of(c), + return filter(set_of(c), [&](T const &t) { return includes(q, t); }); } template -std::unordered_map query_keys(query_set const &q, C const &m) { +std::map query_keys(query_set const &q, C const &m) { if (is_matchall(q)) { return m; } @@ -130,7 +130,7 @@ std::unordered_map query_keys(query_set const &q, C const &m) { template -std::unordered_map query_values(query_set const &q, C const &m) { +std::map query_values(query_set const &q, C const &m) { if (is_matchall(q)) { return m; } diff --git a/lib/utils/include/utils/graph/render_dot.h b/lib/utils/include/utils/graph/render_dot.h index 632ba736ea..e29b07b2b3 100644 --- a/lib/utils/include/utils/graph/render_dot.h +++ b/lib/utils/include/utils/graph/render_dot.h @@ -3,15 +3,15 @@ #include "utils/graph/labelled_open_dataflow_graph/labelled_open_dataflow_graph_view.h" #include -#include +#include namespace FlexFlow { std::string escape_dot_string(std::string const &); std::string render_dot_node_attrs( - std::unordered_map const &attrs); + std::map const &attrs); std::string render_dot( - LabelledDataflowGraphView, + LabelledDataflowGraphView, std::string> const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_parallel_split.dtg.toml b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_parallel_split.dtg.toml index bb607264d8..4cb03dcf50 100644 --- a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_parallel_split.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_parallel_split.dtg.toml @@ -3,6 +3,7 @@ name = "BinaryParallelSplit" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_series_split.dtg.toml b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_series_split.dtg.toml index 0e4cfaadcd..d96296744d 100644 --- a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_series_split.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_series_split.dtg.toml @@ -3,6 +3,7 @@ name = "BinarySeriesSplit" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.dtg.toml b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.dtg.toml index faaf38626e..6c44bcad74 100644 --- a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.dtg.toml @@ -3,6 +3,7 @@ name = "BinarySPDecompositionTree" type = "variant" features = [ "eq", + "ord", "hash", "fmt", ] diff --git a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h index 34b77f4d37..2772d518ae 100644 --- a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h +++ b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h @@ -9,7 +9,7 @@ #include "utils/graph/series_parallel/sp_decomposition_tree_node_type.dtg.h" #include "utils/nonnegative_int/nonnegative_int.h" #include -#include +#include namespace FlexFlow { @@ -22,7 +22,7 @@ GenericBinarySPDecompositionTreeImplementation get_leaves(BinarySPDecompositionTree const &); +std::multiset get_leaves(BinarySPDecompositionTree const &); SPDecompositionTreeNodeType get_node_type(BinarySPDecompositionTree const &); diff --git a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.h b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.h index 105f5490a4..6db143c044 100644 --- a/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.h +++ b/lib/utils/include/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.h @@ -7,7 +7,7 @@ namespace FlexFlow { template -std::unordered_set find_paths_to_leaf( +std::set find_paths_to_leaf( Tree const &tree, GenericBinarySPDecompositionTreeImplementation -std::unordered_set get_all_leaf_paths( +std::set get_all_leaf_paths( Tree const &tree, GenericBinarySPDecompositionTreeImplementation -std::unordered_multiset get_leaves( +std::multiset get_leaves( Tree const &tree, GenericBinarySPDecompositionTreeImplementation -std::unordered_map get_path_to_leaf_map( +std::map get_path_to_leaf_map( Tree const &tree, GenericBinarySPDecompositionTreeImplementation parallel_extend(DiGraph &g, +std::map parallel_extend(DiGraph &g, DiGraphView const &ext); -std::unordered_map serial_extend(DiGraph &g, +std::map serial_extend(DiGraph &g, DiGraphView const &ext); DiGraph series_composition(DiGraphView const &g1, DiGraphView const &g2); DiGraph parallel_composition(DiGraphView const &g1, DiGraphView const &g2); diff --git a/lib/utils/include/utils/graph/series_parallel/extended_parallel_reduction.dtg.toml b/lib/utils/include/utils/graph/series_parallel/extended_parallel_reduction.dtg.toml index 331b0eed8a..f4d1772be6 100644 --- a/lib/utils/include/utils/graph/series_parallel/extended_parallel_reduction.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/extended_parallel_reduction.dtg.toml @@ -3,6 +3,7 @@ name = "ExtendedParallelReduction" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] @@ -14,14 +15,14 @@ docstring = """\ includes = [ "utils/graph/multidigraph/multidiedge.dtg.h", - "" + "" ] -src_includes = [ - "utils/hash/unordered_set.h", - "utils/fmt/unordered_set.h", +src_includes = [ + "utils/hash/set.h", + "utils/fmt/set.h", ] [[fields]] name = "edges" -type = "std::unordered_set<::FlexFlow::MultiDiEdge>" +type = "std::set<::FlexFlow::MultiDiEdge>" diff --git a/lib/utils/include/utils/graph/series_parallel/extended_series_reduction.dtg.toml b/lib/utils/include/utils/graph/series_parallel/extended_series_reduction.dtg.toml index 166cb71b46..32c929ce23 100644 --- a/lib/utils/include/utils/graph/series_parallel/extended_series_reduction.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/extended_series_reduction.dtg.toml @@ -13,6 +13,7 @@ docstring = """\ features = [ "eq", + "ord", "hash", "fmt", ] @@ -22,7 +23,7 @@ includes = [ "" ] -src_includes = [ +src_includes = [ "utils/hash/vector.h", "utils/fmt/vector.h", ] diff --git a/lib/utils/include/utils/graph/series_parallel/get_ancestors.h b/lib/utils/include/utils/graph/series_parallel/get_ancestors.h index e658c04060..15e3993442 100644 --- a/lib/utils/include/utils/graph/series_parallel/get_ancestors.h +++ b/lib/utils/include/utils/graph/series_parallel/get_ancestors.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_SERIAL_PARALLLEL_GET_ANCESTORS_H #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" -#include +#include namespace FlexFlow { @@ -42,7 +42,7 @@ namespace FlexFlow { * n5 | {n0, n1, n2, n3, n4} * */ -std::unordered_set get_ancestors(SeriesParallelDecomposition const &sp, +std::set get_ancestors(SeriesParallelDecomposition const &sp, Node const &node); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.h b/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.h index eeb4590c79..436c48d2ac 100644 --- a/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.h +++ b/lib/utils/include/utils/graph/series_parallel/non_normal_sp_decomposition.h @@ -3,7 +3,7 @@ #include "utils/graph/series_parallel/non_normal_sp_decomposition.dtg.h" #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" -#include +#include #include namespace FlexFlow { @@ -14,7 +14,7 @@ NonNormalSPDecomposition non_normal_series_composition( std::vector const &sp_compositions); NonNormalSPDecomposition non_normal_parallel_composition( - std::unordered_multiset const &sp_compositions); + std::multiset const &sp_compositions); NonNormalSPDecomposition as_non_normal(SeriesParallelDecomposition const &sp); diff --git a/lib/utils/include/utils/graph/series_parallel/parallel_reduction.h b/lib/utils/include/utils/graph/series_parallel/parallel_reduction.h index 598548bec1..2322b1e096 100644 --- a/lib/utils/include/utils/graph/series_parallel/parallel_reduction.h +++ b/lib/utils/include/utils/graph/series_parallel/parallel_reduction.h @@ -5,7 +5,7 @@ #include "utils/graph/series_parallel/extended_parallel_reduction.dtg.h" #include "utils/graph/series_parallel/parallel_reduction.dtg.h" #include -#include +#include namespace FlexFlow { @@ -18,7 +18,7 @@ std::optional /** * @brief Finds all ExtendedParallelReduction for a given MultiDiGraph */ -std::unordered_set +std::set find_all_extended_parallel_reductions(MultiDiGraphView const &); MultiDiEdge apply_parallel_reduction(MultiDiGraph &, ParallelReduction const &); diff --git a/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h b/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h index 06db05b8aa..9e448d0db6 100644 --- a/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h +++ b/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h @@ -13,10 +13,10 @@ std::variant internal_to_final_ast( SeriesParallelDecomposition to_final_ast(std::variant const &); -std::unordered_multiset get_nodes(SeriesParallelDecomposition const &sp); -std::unordered_multiset get_nodes(SeriesSplit const &); -std::unordered_multiset get_nodes(ParallelSplit const &); -std::unordered_multiset get_nodes(Node const &); +std::multiset get_nodes(SeriesParallelDecomposition const &sp); +std::multiset get_nodes(SeriesSplit const &); +std::multiset get_nodes(ParallelSplit const &); +std::multiset get_nodes(Node const &); bool has_no_duplicate_nodes(SeriesParallelDecomposition const &sp); @@ -30,7 +30,7 @@ nonnegative_int num_nodes(SeriesParallelDecomposition const &sp); SeriesParallelDecomposition series_composition( std::vector const &sp_compositions); SeriesParallelDecomposition parallel_composition( - std::unordered_multiset const + std::multiset const &sp_compositions); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h b/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h index 935b8a9e52..ca9e555803 100644 --- a/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h +++ b/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h @@ -4,7 +4,7 @@ #include "utils/graph/digraph/digraph_view.h" #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" #include "utils/nonnegative_int/nonnegative_int.h" -#include +#include namespace FlexFlow { @@ -12,7 +12,7 @@ namespace FlexFlow { * @brief Maps each node to the number of times it appears in the decomposition. * */ -std::unordered_map +std::map get_num_occurrences_of_nodes(SeriesParallelDecomposition const &sp); /** @@ -21,10 +21,10 @@ std::unordered_map * */ float work_cost(SeriesParallelDecomposition const &sp, - std::unordered_map cost_map); + std::map cost_map); float work_cost(DiGraphView const &g, - std::unordered_map const &cost_map); + std::map const &cost_map); /** * @brief Computes the total number of edges the decomposition has when viewed @@ -36,10 +36,10 @@ nonnegative_int num_dependencies(SeriesParallelDecomposition const &sp); nonnegative_int num_dependencies(DiGraphView const &g); float critical_path_cost(SeriesParallelDecomposition const &sp, - std::unordered_map const &cost_map); + std::map const &cost_map); float critical_path_cost(DiGraphView const &g, - std::unordered_map const &cost_map); + std::map const &cost_map); /** * @brief Calculates the relative increase in total work cost between the @@ -48,7 +48,7 @@ float critical_path_cost(DiGraphView const &g, */ float relative_work_increase(DiGraphView const &g, SeriesParallelDecomposition const &sp, - std::unordered_map const &cost_map); + std::map const &cost_map); /** * @brief Calculates the relative increase in critical path cost between the @@ -58,7 +58,7 @@ float relative_work_increase(DiGraphView const &g, float relative_critical_path_cost_increase( DiGraphView const &g, SeriesParallelDecomposition const &sp, - std::unordered_map const &cost_map); + std::map const &cost_map); /** * @brief Calculates the relative increase in the number of dependencies between diff --git a/lib/utils/include/utils/graph/series_parallel/series_reduction.h b/lib/utils/include/utils/graph/series_parallel/series_reduction.h index 9d11e2bdfb..b71a8c8bc6 100644 --- a/lib/utils/include/utils/graph/series_parallel/series_reduction.h +++ b/lib/utils/include/utils/graph/series_parallel/series_reduction.h @@ -28,7 +28,7 @@ std::optional find_series_reduction(MultiDiGraphView const &); * We have that [(A,B), (B,D), (D,E)] and [(A,C), (C,E)] both constitute * `ExtendedSeriesReduction`. */ -std::unordered_set +std::set find_all_extended_series_reductions(MultiDiGraphView const &g); MultiDiEdge apply_series_reduction(MultiDiGraph &, SeriesReduction const &); diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h index a71bacc05a..9e409dab39 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h @@ -3,7 +3,7 @@ #include "utils/graph/digraph/digraph_view.h" #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" -#include +#include namespace FlexFlow { /** diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h index 60d3aa6aa9..4cee222872 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h @@ -5,17 +5,17 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" #include "utils/graph/series_parallel/sp_ization/node_role.dtg.h" #include "utils/nonnegative_int/nonnegative_int.h" -#include +#include namespace FlexFlow { DiGraph add_dummy_nodes(DiGraph g, - std::unordered_map &node_roles); + std::map &node_roles); -std::unordered_set +std::set get_component(DiGraph const &g, Node const &node, - std::unordered_map const &depth_map, - std::unordered_map const &node_roles); + std::map const &depth_map, + std::map const &node_roles); /** * \brief See \ref spization-escribano. diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/flexible_algo.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/flexible_algo.h index 93a4e29fa2..4a1449154c 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/flexible_algo.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/flexible_algo.h @@ -7,8 +7,8 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" #include "utils/graph/series_parallel/sp_ization/node_role.dtg.h" #include "utils/graph/series_parallel/sp_ization/up_down_partition.dtg.h" -#include -#include +#include +#include namespace FlexFlow { @@ -17,7 +17,7 @@ namespace FlexFlow { */ SeriesParallelDecomposition flexible_sp_ization(DiGraphView const &g, - std::unordered_map const &cost_map); + std::map const &cost_map); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h index 3de889a140..d3f87c8e15 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h @@ -4,11 +4,11 @@ #include "utils/graph/digraph/digraph.h" #include "utils/graph/digraph/digraph_view.h" #include "utils/graph/series_parallel/sp_ization/node_role.dtg.h" -#include +#include namespace FlexFlow { -std::unordered_map +std::map get_initial_node_role_map(DiGraphView const &g); /** @@ -20,7 +20,7 @@ std::unordered_map DiGraph contract_out_nodes_of_given_role( DiGraph g, NodeRole const &role, - std::unordered_map const &node_roles); + std::map const &node_roles); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/sp_ization_combined_benchmark_result.dtg.toml b/lib/utils/include/utils/graph/series_parallel/sp_ization/sp_ization_combined_benchmark_result.dtg.toml index 1f7244ddcd..ecda844cda 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/sp_ization_combined_benchmark_result.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/sp_ization_combined_benchmark_result.dtg.toml @@ -5,10 +5,10 @@ features = [] includes = [ "", - "", + "", "utils/graph/series_parallel/sp_ization/sp_ization_benchmark_result.dtg.h", ] [[fields]] name = "by_technique" -type = "std::unordered_map" +type = "std::map" diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.dtg.toml b/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.dtg.toml index 8eeb7a7d40..158a384f7f 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.dtg.toml +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.dtg.toml @@ -3,25 +3,26 @@ name = "UpDownPartition" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] includes = [ - "", + "", "utils/graph/node/node.dtg.h", ] src_includes = [ - "utils/fmt/unordered_set.h", - "utils/hash/unordered_set.h", + "utils/fmt/set.h", + "utils/hash/set.h", ] [[fields]] name = "up" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" [[fields]] name = "down" -type = "std::unordered_set<::FlexFlow::Node>" +type = "std::set<::FlexFlow::Node>" diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h index 39db2984c1..5cc3b9bd6f 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h @@ -3,7 +3,7 @@ #include "utils/graph/digraph/digraph.h" #include "utils/graph/series_parallel/sp_ization/up_down_partition.dtg.h" -#include +#include namespace FlexFlow { @@ -11,14 +11,14 @@ namespace FlexFlow { * @brief Returns the nodes n in the up set such that in the up subgraph, there * is no outgoing edge from n. */ -std::unordered_set get_up_frontier(DiGraph const &sp, +std::set get_up_frontier(DiGraph const &sp, UpDownPartition const &partition); /** * @brief Returns the nodes n in the down set such that in the down subgraph, * there is no incoming edge to n. */ -std::unordered_set get_down_frontier(DiGraph const &sp, +std::set get_down_frontier(DiGraph const &sp, UpDownPartition const &partition); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.h index c6dd87d2a0..9d371f29e0 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.h @@ -3,7 +3,7 @@ #include "utils/graph/digraph/digraph_view.h" #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" -#include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/graph/traversal.h b/lib/utils/include/utils/graph/traversal.h index 44ddc39eb8..8fe4040377 100644 --- a/lib/utils/include/utils/graph/traversal.h +++ b/lib/utils/include/utils/graph/traversal.h @@ -18,7 +18,7 @@ struct unchecked_dfs_iterator { unchecked_dfs_iterator(DiGraphView const &g, std::vector const &); unchecked_dfs_iterator(DiGraphView const &g, - std::unordered_set const &); + std::set const &); reference operator*() const; pointer operator->(); @@ -50,9 +50,9 @@ struct checked_dfs_iterator { checked_dfs_iterator(DiGraphView const &g, std::vector const &, - std::unordered_set const &); + std::set const &); checked_dfs_iterator(DiGraphView const &g, - std::unordered_set const &starting_points); + std::set const &starting_points); reference operator*() const; pointer operator->(); @@ -64,7 +64,7 @@ struct checked_dfs_iterator { private: unchecked_dfs_iterator iter; - std::unordered_set seen; + std::set seen; }; struct bfs_iterator { @@ -76,9 +76,9 @@ struct bfs_iterator { bfs_iterator(DiGraphView const &, std::queue const &, - std::optional> const &); + std::optional> const &); bfs_iterator(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); reference operator*() const; pointer operator->(); @@ -91,13 +91,13 @@ struct bfs_iterator { private: DiGraphView graph; std::queue q; - std::optional> seen; + std::optional> seen; }; struct CheckedDFSView { CheckedDFSView() = delete; explicit CheckedDFSView(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); checked_dfs_iterator begin() const; checked_dfs_iterator end() const; @@ -106,13 +106,13 @@ struct CheckedDFSView { private: DiGraphView graph; - std::unordered_set starting_points; + std::set starting_points; }; struct UncheckedDFSView { UncheckedDFSView() = delete; explicit UncheckedDFSView(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); unchecked_dfs_iterator begin() const; unchecked_dfs_iterator end() const; @@ -121,13 +121,13 @@ struct UncheckedDFSView { private: DiGraphView graph; - std::unordered_set starting_points; + std::set starting_points; }; struct BFSView { BFSView() = delete; explicit BFSView(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); bfs_iterator begin() const; bfs_iterator end() const; @@ -136,7 +136,7 @@ struct BFSView { private: DiGraphView graph; - std::unordered_set starting_points; + std::set starting_points; }; /* struct BoundaryDFSView { */ @@ -151,7 +151,7 @@ struct BFSView { /* using reference = Node const &; */ /* boundary_dfs_iterator(IDiGraphView const &g, std::vector const &, - * std::unordered_set const &); */ + * std::set const &); */ /* reference operator*() const; */ /* pointer operator->(); */ @@ -172,13 +172,13 @@ struct BFSView { /* }; */ UncheckedDFSView unchecked_dfs(DiGraphView const &, - std::unordered_set const &starting_points); -/* BoundaryDFSView boundary_dfs(IDiGraphView const &, std::unordered_set + std::set const &starting_points); +/* BoundaryDFSView boundary_dfs(IDiGraphView const &, std::set * const &starting_points); */ CheckedDFSView dfs(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); BFSView bfs(DiGraphView const &, - std::unordered_set const &starting_points); + std::set const &starting_points); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h b/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h index d595d2baab..3cd3f5abf6 100644 --- a/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h +++ b/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set> +std::set> get_connected_components(UndirectedGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/algorithms/get_edges.h b/lib/utils/include/utils/graph/undirected/algorithms/get_edges.h index 3e951b1db1..b416f1d38b 100644 --- a/lib/utils/include/utils/graph/undirected/algorithms/get_edges.h +++ b/lib/utils/include/utils/graph/undirected/algorithms/get_edges.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_edges(UndirectedGraphView const &); +std::set get_edges(UndirectedGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h b/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h index bc605360d2..f98ae4b51d 100644 --- a/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h +++ b/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h @@ -5,7 +5,7 @@ namespace FlexFlow { -std::unordered_set get_neighboring_nodes(UndirectedGraphView const &, +std::set get_neighboring_nodes(UndirectedGraphView const &, Node const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/i_undirected_graph.h b/lib/utils/include/utils/graph/undirected/i_undirected_graph.h index 4761275031..24527ae98a 100644 --- a/lib/utils/include/utils/graph/undirected/i_undirected_graph.h +++ b/lib/utils/include/utils/graph/undirected/i_undirected_graph.h @@ -12,7 +12,7 @@ struct IUndirectedGraph : public IUndirectedGraphView { virtual void add_edge(UndirectedEdge const &) = 0; virtual void remove_edge(UndirectedEdge const &) = 0; - virtual std::unordered_set + virtual std::set query_nodes(NodeQuery const &query) const = 0; virtual IUndirectedGraph *clone() const = 0; diff --git a/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h b/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h index 2ffe061dbe..3e5c82b519 100644 --- a/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h +++ b/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h @@ -14,7 +14,7 @@ struct IUndirectedGraphView : public IGraphView { IUndirectedGraphView(IUndirectedGraphView const &) = delete; IUndirectedGraphView &operator=(IUndirectedGraphView const &) = delete; - virtual std::unordered_set + virtual std::set query_edges(UndirectedEdgeQuery const &) const = 0; virtual ~IUndirectedGraphView() = default; diff --git a/lib/utils/include/utils/graph/undirected/undirected_edge.h b/lib/utils/include/utils/graph/undirected/undirected_edge.h index 1eeea7b3c2..fcfaaa2175 100644 --- a/lib/utils/include/utils/graph/undirected/undirected_edge.h +++ b/lib/utils/include/utils/graph/undirected/undirected_edge.h @@ -8,7 +8,7 @@ namespace FlexFlow { bool is_connected_to(UndirectedEdge const &e, Node const &n); -std::unordered_set get_endpoints(UndirectedEdge const &); +std::set get_endpoints(UndirectedEdge const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/undirected_graph.h b/lib/utils/include/utils/graph/undirected/undirected_graph.h index 09b6495699..dd4dba8329 100644 --- a/lib/utils/include/utils/graph/undirected/undirected_graph.h +++ b/lib/utils/include/utils/graph/undirected/undirected_graph.h @@ -22,8 +22,8 @@ struct UndirectedGraph : virtual UndirectedGraphView { void add_edge(Edge const &); void remove_edge(Edge const &); - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(EdgeQuery const &) const; + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(EdgeQuery const &) const; template static typename std::enable_if::value, diff --git a/lib/utils/include/utils/graph/undirected/undirected_graph_view.h b/lib/utils/include/utils/graph/undirected/undirected_graph_view.h index 90dd5dd5d8..0bb4547738 100644 --- a/lib/utils/include/utils/graph/undirected/undirected_graph_view.h +++ b/lib/utils/include/utils/graph/undirected/undirected_graph_view.h @@ -16,8 +16,8 @@ struct UndirectedGraphView : virtual GraphView { UndirectedGraphView(UndirectedGraphView const &) = default; UndirectedGraphView &operator=(UndirectedGraphView const &) = default; - std::unordered_set query_nodes(NodeQuery const &) const; - std::unordered_set query_edges(EdgeQuery const &query) const; + std::set query_nodes(NodeQuery const &) const; + std::set query_edges(EdgeQuery const &query) const; template static diff --git a/lib/utils/include/utils/graph/views/views.h b/lib/utils/include/utils/graph/views/views.h index 5e0109ed5b..639c8750b7 100644 --- a/lib/utils/include/utils/graph/views/views.h +++ b/lib/utils/include/utils/graph/views/views.h @@ -12,56 +12,56 @@ struct UndirectedSubgraphView : public IUndirectedGraphView { public: UndirectedSubgraphView() = delete; UndirectedSubgraphView(UndirectedGraphView const &, - std::unordered_set const &); + std::set const &); - std::unordered_set + std::set query_edges(UndirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; UndirectedSubgraphView *clone() const override; private: UndirectedGraphView g; - std::unordered_set subgraph_nodes; + std::set subgraph_nodes; }; struct DiSubgraphView : public IDiGraphView { public: DiSubgraphView() = delete; - DiSubgraphView(DiGraphView const &, std::unordered_set const &); + DiSubgraphView(DiGraphView const &, std::set const &); - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; DiSubgraphView *clone() const override; private: DiGraphView g; - std::unordered_set subgraph_nodes; + std::set subgraph_nodes; }; UndirectedGraphView view_subgraph(UndirectedGraphView const &, - std::unordered_set const &); + std::set const &); DiGraphView view_subgraph(DiGraphView const &, - std::unordered_set const &); + std::set const &); UndirectedEdge to_undirected_edge(DirectedEdge const &); -std::unordered_set - to_undirected_edges(std::unordered_set const &); +std::set + to_undirected_edges(std::set const &); -std::unordered_set to_directed_edges(UndirectedEdge const &); -std::unordered_set - to_directed_edges(std::unordered_set const &); +std::set to_directed_edges(UndirectedEdge const &); +std::set + to_directed_edges(std::set const &); struct ViewDiGraphAsUndirectedGraph : public IUndirectedGraphView { public: explicit ViewDiGraphAsUndirectedGraph(DiGraphView const &); - std::unordered_set + std::set query_edges(UndirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; ViewDiGraphAsUndirectedGraph *clone() const override; @@ -73,9 +73,9 @@ struct ViewUndirectedGraphAsDiGraph : public IDiGraphView { public: explicit ViewUndirectedGraphAsDiGraph(UndirectedGraphView const &); - std::unordered_set + std::set query_edges(DirectedEdgeQuery const &) const override; - std::unordered_set query_nodes(NodeQuery const &) const override; + std::set query_nodes(NodeQuery const &) const override; ViewUndirectedGraphAsDiGraph *clone() const override; diff --git a/lib/utils/include/utils/many_to_one/many_to_one.h b/lib/utils/include/utils/many_to_one/many_to_one.h index a501a0672c..9a01d77e5c 100644 --- a/lib/utils/include/utils/many_to_one/many_to_one.h +++ b/lib/utils/include/utils/many_to_one/many_to_one.h @@ -3,24 +3,21 @@ #include "utils/containers/require_same.h" #include "utils/containers/try_at.h" -#include "utils/containers/unordered_set_of.h" #include "utils/containers/values.h" #include "utils/exception.h" -#include "utils/fmt/unordered_map.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/map.h" +#include "utils/fmt/set.h" #include "utils/hash-utils.h" #include "utils/hash/tuple.h" -#include "utils/hash/unordered_map.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/map.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" #include #include #include -#include -#include #include "utils/containers/set_of.h" -#include "utils/containers/unordered_keys.h" +#include "utils/nonempty_set/nonempty_set.h" +#include "utils/containers/keys.h" namespace FlexFlow { @@ -59,16 +56,21 @@ struct ManyToOne { if (!found_r.has_value()) { this->m_l_to_r.insert({l, r}); - this->m_r_to_l[r].insert(l); + + if (contains_key(this->m_r_to_l, r)) { + this->m_r_to_l.at(r).insert(l); + } else { + this->m_r_to_l.insert({r, nonempty_set{{l}}}); + } } else if (found_r.value() == r) { return; } else { - PANIC(fmt::format( + PANIC( "Existing mapping found for left value {}: tried to map to right " "value {}, but is already bound to right value {}", l, r, - found_r.value())); + found_r.value()); } } @@ -84,27 +86,27 @@ struct ManyToOne { return this->m_l_to_r.at(l); } - std::unordered_set const &at_r(R const &r) const { + nonempty_set const &at_r(R const &r) const { return this->m_r_to_l.at(r); } - std::unordered_set left_values() const { - return unordered_keys(this->m_l_to_r); + std::set left_values() const { + return keys(this->m_l_to_r); } - std::unordered_set> left_groups() const { - return unordered_set_of(values(this->m_r_to_l)); + std::set> left_groups() const { + return set_of(values(this->m_r_to_l)); } - std::unordered_set right_values() const { - return unordered_keys(this->m_r_to_l); + std::set right_values() const { + return keys(this->m_r_to_l); } - std::unordered_map const &l_to_r() const { + std::map const &l_to_r() const { return this->m_l_to_r; } - std::unordered_map> const &r_to_l() const { + std::map> const &r_to_l() const { return this->m_r_to_l; } @@ -113,8 +115,8 @@ struct ManyToOne { } private: - std::unordered_map m_l_to_r; - std::unordered_map> m_r_to_l; + std::map m_l_to_r; + std::map> m_r_to_l; private: std::tuple @@ -126,9 +128,9 @@ struct ManyToOne { }; template -std::unordered_map, R> +std::map, R> format_as(ManyToOne const &m) { - std::unordered_map, R> result; + std::map, R> result; for (R const &r : m.right_values()) { result.insert({m.at_r(r), r}); @@ -143,14 +145,14 @@ std::ostream &operator<<(std::ostream &s, ManyToOne const &m) { } template -std::unordered_set> +std::set> unstructured_relation_from_many_to_one(ManyToOne const &many_to_one) { - return unordered_set_of(many_to_one.l_to_r()); + return set_of(many_to_one.l_to_r()); } template ManyToOne many_to_one_from_unstructured_relation( - std::unordered_set> const &relation) { + std::set> const &relation) { ManyToOne result; for (auto const &lr : relation) { result.insert(lr); @@ -168,7 +170,7 @@ struct adl_serializer<::FlexFlow::ManyToOne> { CHECK_IS_JSON_DESERIALIZABLE(L); CHECK_IS_JSON_DESERIALIZABLE(R); - std::unordered_set> s = j; + std::set> s = j; return ::FlexFlow::many_to_one_from_unstructured_relation(s); } diff --git a/lib/utils/include/utils/many_to_one/many_to_one_from_map.h b/lib/utils/include/utils/many_to_one/many_to_one_from_map.h index e0484d2131..8e1f6e4ad1 100644 --- a/lib/utils/include/utils/many_to_one/many_to_one_from_map.h +++ b/lib/utils/include/utils/many_to_one/many_to_one_from_map.h @@ -6,7 +6,7 @@ namespace FlexFlow { template -ManyToOne many_to_one_from_map(std::unordered_map const &m) { +ManyToOne many_to_one_from_map(std::map const &m) { ManyToOne result; for (auto const &[l, r] : m) { @@ -17,7 +17,7 @@ ManyToOne many_to_one_from_map(std::unordered_map const &m) { } template -ManyToOne many_to_one_from_map(std::map const &m) { +ManyToOne many_to_one_from_map(std::unordered_map const &m) { ManyToOne result; for (auto const &[l, r] : m) { diff --git a/lib/utils/include/utils/nonempty_set/nonempty_set.h b/lib/utils/include/utils/nonempty_set/nonempty_set.h index fe4b152bd5..93d2b37def 100644 --- a/lib/utils/include/utils/nonempty_set/nonempty_set.h +++ b/lib/utils/include/utils/nonempty_set/nonempty_set.h @@ -7,9 +7,10 @@ #include "utils/hash/set.h" #include "utils/fmt/set.h" #include "utils/positive_int/positive_int.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" +#include "utils/containers/unordered_set_of.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/nonempty_unordered_set/nonempty_unordered_set.h b/lib/utils/include/utils/nonempty_unordered_set/nonempty_unordered_set.h index 2bc070fb5e..67f8e0afde 100644 --- a/lib/utils/include/utils/nonempty_unordered_set/nonempty_unordered_set.h +++ b/lib/utils/include/utils/nonempty_unordered_set/nonempty_unordered_set.h @@ -6,7 +6,7 @@ #include "utils/hash/unordered_set.h" #include "utils/positive_int/positive_int.h" #include -#include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/one_to_many/one_to_many.h b/lib/utils/include/utils/one_to_many/one_to_many.h index d57622b950..f46afaae42 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many.h +++ b/lib/utils/include/utils/one_to_many/one_to_many.h @@ -87,12 +87,12 @@ struct OneToMany { } else if (found_l.value() == l) { return; } else { - throw mk_runtime_error( - fmt::format("Existing mapping found for right value {}: tried to map " - "to left value {}, but is already bound to left value {}", - r, - l, - found_l.value())); + PANIC( + "Existing mapping found for right value {}: tried to map " + "to left value {}, but is already bound to left value {}", + r, + l, + found_l.value()); } } @@ -160,9 +160,9 @@ std::ostream &operator<<(std::ostream &s, OneToMany const &m) { } template -std::unordered_set> +std::set> unstructured_relation_from_one_to_many(OneToMany const &one_to_many) { - return transform(unordered_set_of(one_to_many.r_to_l()), + return transform(set_of(one_to_many.r_to_l()), [](std::pair const &rl) -> std::pair { return std::pair{rl.second, rl.first}; }); @@ -170,7 +170,7 @@ std::unordered_set> template OneToMany one_to_many_from_unstructured_relation( - std::unordered_set> const &rel) { + std::set> const &rel) { OneToMany result; for (auto const &lr : rel) { result.insert(lr); @@ -188,7 +188,7 @@ struct adl_serializer<::FlexFlow::OneToMany> { CHECK_IS_JSON_DESERIALIZABLE(L); CHECK_IS_JSON_DESERIALIZABLE(R); - std::unordered_set> s = j; + std::set> s = j; return ::FlexFlow::one_to_many_from_unstructured_relation(s); } diff --git a/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h b/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h index bed62caaf6..eae63e6ade 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h +++ b/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h @@ -8,7 +8,7 @@ namespace FlexFlow { template OneToMany one_to_many_from_l_to_r_mapping( - std::unordered_map> const &m) { + std::map> const &m) { OneToMany result; for (auto const &[l, rs] : m) { diff --git a/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h b/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h index 050f6ad7dd..9d5fbb3665 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h +++ b/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h @@ -14,7 +14,7 @@ template one_to_many_transform_values(OneToMany const &input, F f) { return one_to_many_from_unstructured_relation(transform( - unordered_set_of(input.relation()), + set_of(input.relation()), [&](std::pair const &p) -> std::pair { return {p.first, f(p.second)}; })); diff --git a/lib/utils/include/utils/ord/unordered_map.h b/lib/utils/include/utils/ord/unordered_map.h index 1cfbdb27b6..072386b3fc 100644 --- a/lib/utils/include/utils/ord/unordered_map.h +++ b/lib/utils/include/utils/ord/unordered_map.h @@ -3,8 +3,8 @@ #include "utils/type_traits_core.h" #include -#include #include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/orthotope/dim_coord.dtg.toml b/lib/utils/include/utils/orthotope/dim_coord.dtg.toml index 3ac729fc7e..5ff32e75af 100644 --- a/lib/utils/include/utils/orthotope/dim_coord.dtg.toml +++ b/lib/utils/include/utils/orthotope/dim_coord.dtg.toml @@ -13,15 +13,15 @@ template_params = [ ] includes = [ - "", + "", "utils/nonnegative_int/nonnegative_int.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "raw" -type = "std::unordered_map" +type = "std::map" diff --git a/lib/utils/include/utils/orthotope/dim_coord.h b/lib/utils/include/utils/orthotope/dim_coord.h index b9a10d6750..ccb2cb320b 100644 --- a/lib/utils/include/utils/orthotope/dim_coord.h +++ b/lib/utils/include/utils/orthotope/dim_coord.h @@ -3,10 +3,9 @@ #include "utils/containers/all_of.h" #include "utils/containers/contains_key.h" -#include "utils/containers/generate_unordered_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/is_subseteq_of.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/map_values.h" #include "utils/containers/product.h" @@ -15,7 +14,7 @@ #include "utils/containers/scanr.h" #include "utils/containers/sorted_by.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/zip_with_strict.h" #include "utils/exception.h" #include "utils/nonnegative_int/nonnegative_range.h" @@ -24,17 +23,19 @@ #include "utils/orthotope/dim_domain.h" #include "utils/orthotope/minimal_dim_domain.h" #include "utils/orthotope/orthotope.h" +#include "utils/containers/set_of.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { template -std::unordered_set get_coord_dims(DimCoord const &coord) { - return unordered_keys(coord.raw); +std::set get_coord_dims(DimCoord const &coord) { + return keys(coord.raw); } template DimCoord restrict_coord_to_dims(DimCoord const &coord, - std::unordered_set const &dims) { + std::set const &dims) { return DimCoord{ restrict_keys(coord.raw, dims), }; @@ -52,7 +53,7 @@ OrthotopeCoord template DimCoord dim_coord_from_orthotope_coord(OrthotopeCoord const &coord, - std::unordered_set const &dims, + std::set const &dims, DimOrdering const &dim_ordering) { return DimCoord{ map_from_keys_and_values(sorted_by(dims, dim_ordering.lt), coord.raw), @@ -61,11 +62,11 @@ DimCoord dim_coord_from_orthotope_coord(OrthotopeCoord const &coord, template DimCoord lift_dim_coord(DimCoord const &coord, - std::unordered_set const &lifted_dims) { + std::set const &lifted_dims) { ASSERT(is_subseteq_of(get_coord_dims(coord), lifted_dims)); return DimCoord{ - generate_unordered_map(lifted_dims, + generate_map(lifted_dims, [&](T const &dim) { if (contains_key(coord.raw, dim)) { return coord.raw.at(dim); @@ -77,27 +78,27 @@ DimCoord lift_dim_coord(DimCoord const &coord, } template -std::unordered_set> +std::set> get_coords_in_dim_domain(DimDomain const &dim_domain) { - std::unordered_map> + std::map> component_possible_values = map_values( dim_domain.dims, [](positive_int component_size) - -> std::unordered_set { - return unordered_set_of(nonnegative_range(component_size)); + -> std::set { + return set_of(nonnegative_range(component_size)); }); - return transform( + return set_of(transform( get_all_assignments(component_possible_values), - [](std::unordered_map const &assignment) { + [](std::map const &assignment) { return DimCoord{ assignment, }; - }); + })); } template -std::unordered_set> get_coords_in_minimal_dim_domain( +std::set> get_coords_in_minimal_dim_domain( MinimalDimDomain const &minimal_dim_domain) { return get_coords_in_dim_domain(lift_minimal_dim_domain(minimal_dim_domain)); } @@ -127,7 +128,7 @@ bool dim_domain_contains_coord(DimDomain const &domain, DimCoord const &coord) { ASSERT(get_domain_dims(domain) == get_coord_dims(coord)); - std::unordered_set dims = + std::set dims = require_same(get_domain_dims(domain), get_coord_dims(coord)); return all_of(dims, [&](T const &dim) { return coord.raw.at(dim) < domain.dims.at(dim); diff --git a/lib/utils/include/utils/orthotope/dim_domain.dtg.toml b/lib/utils/include/utils/orthotope/dim_domain.dtg.toml index ccad639aac..5f3b4b7234 100644 --- a/lib/utils/include/utils/orthotope/dim_domain.dtg.toml +++ b/lib/utils/include/utils/orthotope/dim_domain.dtg.toml @@ -14,15 +14,15 @@ template_params = [ ] includes = [ - "", + "", "utils/positive_int/positive_int.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "dims" -type = "std::unordered_map" +type = "std::map" diff --git a/lib/utils/include/utils/orthotope/dim_domain.h b/lib/utils/include/utils/orthotope/dim_domain.h index 6bd63faeae..7c6abf509a 100644 --- a/lib/utils/include/utils/orthotope/dim_domain.h +++ b/lib/utils/include/utils/orthotope/dim_domain.h @@ -2,7 +2,6 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ORTHOTOPE_DIM_DOMAIN_H #include "utils/containers/filter.h" -#include "utils/containers/unordered_keys.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" @@ -12,6 +11,7 @@ #include "utils/orthotope/dim_domain.dtg.h" #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/orthotope.dtg.h" +#include "utils/containers/keys.h" namespace FlexFlow { @@ -26,24 +26,24 @@ nonnegative_int dim_domain_num_dims(DimDomain const &domain) { } template -std::unordered_set get_domain_dims(DimDomain const &domain) { - return unordered_keys(domain.dims); +std::set get_domain_dims(DimDomain const &domain) { + return keys(domain.dims); } template -std::unordered_set get_trivial_domain_dims(DimDomain const &domain) { +std::set get_trivial_domain_dims(DimDomain const &domain) { return filter(get_domain_dims(domain), [&](T const &idx) { return domain.dims.at(idx) == 1; }); } template -std::unordered_set get_nontrivial_domain_dims(DimDomain const &domain) { +std::set get_nontrivial_domain_dims(DimDomain const &domain) { return set_minus(get_domain_dims(domain), get_trivial_domain_dims(domain)); } template DimDomain restrict_domain_to_dims(DimDomain const &domain, - std::unordered_set const &allowed) { + std::set const &allowed) { return DimDomain{restrict_keys(domain.dims, allowed)}; } @@ -58,7 +58,7 @@ Orthotope orthotope_from_dim_domain(DimDomain const &domain, template DimDomain dim_domain_from_orthotope(Orthotope const &orthotope, - std::unordered_set const &dims, + std::set const &dims, DimOrdering const &dim_ordering) { return DimDomain{ map_from_keys_and_values(sorted_by(dims, dim_ordering.lt), diff --git a/lib/utils/include/utils/orthotope/dim_projection.h b/lib/utils/include/utils/orthotope/dim_projection.h index fa47edd897..44b7a7cebb 100644 --- a/lib/utils/include/utils/orthotope/dim_projection.h +++ b/lib/utils/include/utils/orthotope/dim_projection.h @@ -32,9 +32,9 @@ DimProjection } template -std::unordered_set +std::set input_dims_of_projection(DimProjection const &projection) { - return projection.template visit>(overload{ + return projection.template visit>(overload{ [](UpProjection const &p) { return input_dims_of_up_projection(p); }, @@ -48,9 +48,9 @@ std::unordered_set } template -std::unordered_set +std::set output_dims_of_projection(DimProjection const &projection) { - return projection.template visit>(overload{ + return projection.template visit>(overload{ [](UpProjection const &p) { return output_dims_of_up_projection(p); }, @@ -100,11 +100,11 @@ DimCoord compute_dim_projection(DimProjection const &projection, input_coord); { - std::unordered_set nontrivial_input_domain_dims = + std::set nontrivial_input_domain_dims = get_nontrivial_domain_dims(input_domain); - std::unordered_set projection_input_dims = + std::set projection_input_dims = input_dims_of_projection(projection); - std::unordered_set all_input_domain_dims = get_domain_dims(input_domain); + std::set all_input_domain_dims = get_domain_dims(input_domain); ASSERT(is_subseteq_of(nontrivial_input_domain_dims, projection_input_dims), nontrivial_input_domain_dims, @@ -115,11 +115,11 @@ DimCoord compute_dim_projection(DimProjection const &projection, } { - std::unordered_set nontrivial_output_domain_dims = + std::set nontrivial_output_domain_dims = get_nontrivial_domain_dims(output_domain); - std::unordered_set projection_output_dims = + std::set projection_output_dims = output_dims_of_projection(projection); - std::unordered_set all_output_domain_dims = + std::set all_output_domain_dims = get_domain_dims(output_domain); ASSERT( @@ -146,7 +146,7 @@ DimCoord compute_dim_projection(DimProjection const &projection, }); DimCoord lifted_output_coord = - lift_dim_coord(output_coord, get_domain_dims(output_domain)); + lift_dim_coord(output_coord, set_of(get_domain_dims(output_domain))); ASSERT(dim_domain_contains_coord(output_domain, lifted_output_coord), output_domain, diff --git a/lib/utils/include/utils/orthotope/down_projection.h b/lib/utils/include/utils/orthotope/down_projection.h index 8fb1487ad3..2dda306371 100644 --- a/lib/utils/include/utils/orthotope/down_projection.h +++ b/lib/utils/include/utils/orthotope/down_projection.h @@ -22,13 +22,13 @@ DownProjection make_empty_down_projection() { } template -std::unordered_set +std::set input_dims_of_down_projection(DownProjection const &projection) { return projection.dim_mapping.left_values(); } template -std::unordered_set +std::set output_dims_of_down_projection(DownProjection const &projection) { return projection.dim_mapping.right_values(); } @@ -38,21 +38,21 @@ DimCoord compute_down_projection(DownProjection const &projection, DimCoord const &coord, DimDomain const &input_domain, DimOrdering const &input_dim_ordering) { - std::unordered_set input_dims = input_dims_of_down_projection(projection); - std::unordered_set coord_dims = get_coord_dims(coord); + std::set input_dims = input_dims_of_down_projection(projection); + std::set coord_dims = get_coord_dims(coord); ASSERT(input_dims == coord_dims, "compute_down_projection expected coord dimensions to match " "projection input dimensions"); - std::unordered_set output_dims = + std::set output_dims = output_dims_of_down_projection(projection); return DimCoord{ - generate_unordered_map( + generate_map( output_dims, - [&](R const &output_dim) { - std::unordered_set src_dims = - projection.dim_mapping.at_r(output_dim); + [&](R const &output_dim) -> nonnegative_int { + std::set src_dims = + projection.dim_mapping.at_r(output_dim).unwrap_as_set(); DimCoord src_coord = restrict_coord_to_dims(coord, src_dims); DimDomain src_domain = @@ -65,7 +65,7 @@ DimCoord compute_down_projection(DownProjection const &projection, template void project_dims(DownProjection &proj, - std::unordered_set const &from, + std::set const &from, R const &onto) { ASSERT(from.size() > 0); diff --git a/lib/utils/include/utils/orthotope/eq_projection.h b/lib/utils/include/utils/orthotope/eq_projection.h index 6b394ac04b..12a0046583 100644 --- a/lib/utils/include/utils/orthotope/eq_projection.h +++ b/lib/utils/include/utils/orthotope/eq_projection.h @@ -15,13 +15,13 @@ EqProjection make_empty_eq_projection() { } template -std::unordered_set +std::set input_dims_of_eq_projection(EqProjection const &projection) { return projection.dim_mapping.left_values(); } template -std::unordered_set +std::set output_dims_of_eq_projection(EqProjection const &projection) { return projection.dim_mapping.right_values(); } diff --git a/lib/utils/include/utils/orthotope/minimal_dim_domain.dtg.toml b/lib/utils/include/utils/orthotope/minimal_dim_domain.dtg.toml index 18d2aad26f..66bf4b33f8 100644 --- a/lib/utils/include/utils/orthotope/minimal_dim_domain.dtg.toml +++ b/lib/utils/include/utils/orthotope/minimal_dim_domain.dtg.toml @@ -14,15 +14,15 @@ template_params = [ ] includes = [ - "", + "", "utils/int_ge_two/int_ge_two.h", ] src_includes = [ - "utils/fmt/unordered_map.h", - "utils/hash/unordered_map.h", + "utils/fmt/map.h", + "utils/hash/map.h", ] [[fields]] name = "dims" -type = "std::unordered_map" +type = "std::map" diff --git a/lib/utils/include/utils/orthotope/minimal_dim_domain.h b/lib/utils/include/utils/orthotope/minimal_dim_domain.h index f9bf5fa979..0720d2e8ec 100644 --- a/lib/utils/include/utils/orthotope/minimal_dim_domain.h +++ b/lib/utils/include/utils/orthotope/minimal_dim_domain.h @@ -3,7 +3,7 @@ #include "utils/containers/are_disjoint.h" #include "utils/containers/filtermap_values.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/map_values.h" #include "utils/containers/restrict_keys.h" @@ -14,8 +14,8 @@ #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/minimal_dim_domain.dtg.h" #include "utils/orthotope/minimal_orthotope.dtg.h" -#include "utils/containers/unordered_keys.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/keys.h" +#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -59,31 +59,31 @@ MinimalDimDomain template DimDomain dim_domain_from_minimal_dim_domain( MinimalDimDomain const &minimal_dim_domain, - std::unordered_set const &trivial_dims) { - std::unordered_set nontrivial_dims = + std::set const &trivial_dims) { + std::set nontrivial_dims = get_minimal_domain_dims(minimal_dim_domain); ASSERT(are_disjoint(nontrivial_dims, trivial_dims)); return DimDomain{ - /*dims=*/binary_merge_disjoint_unordered_maps( + /*dims=*/binary_merge_disjoint_maps( map_values( minimal_dim_domain.dims, [](int_ge_two x) { return x.positive_int_from_int_ge_two(); }), - generate_unordered_map(trivial_dims, [](T const &) { return 1_p; })), + generate_map(trivial_dims, [](T const &) { return 1_p; })), }; } template -std::unordered_set +std::set get_minimal_domain_dims(MinimalDimDomain const &domain) { - return unordered_keys(domain.dims); + return keys(domain.dims); } template MinimalDimDomain restrict_minimal_domain_to_dims(MinimalDimDomain const &domain, - std::unordered_set const &allowed) { + std::set const &allowed) { return MinimalDimDomain{restrict_keys(domain.dims, allowed)}; } @@ -100,7 +100,7 @@ MinimalOrthotope minimal_orthotope_from_minimal_dim_domain( template MinimalDimDomain minimal_dim_domain_from_minimal_orthotope( MinimalOrthotope const &orthotope, - std::unordered_set const &dims, + std::set const &dims, DimOrdering const &dim_ordering) { return MinimalDimDomain{ diff --git a/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h b/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h index 1ebff61701..c15152ceef 100644 --- a/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h +++ b/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h @@ -4,8 +4,6 @@ #include "utils/bidict/algorithms/exhaustive_relational_join.h" #include "utils/bidict/algorithms/left_entries.h" #include "utils/bidict/algorithms/right_entries.h" -#include "utils/bidict/algorithms/transform_keys.h" -#include "utils/bidict/algorithms/transform_values.h" #include "utils/bidict/bidict.h" #include "utils/bidict/generate_bidict.h" #include "utils/hash/tuple.h" @@ -15,6 +13,7 @@ #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/dim_projection.h" #include "utils/orthotope/minimal_dim_domain.dtg.h" +#include "utils/bidict/algorithms/bidict_transform_keys_and_values.h" namespace FlexFlow { @@ -93,23 +92,24 @@ template MinimalDimDomainMapping minimal_mapping_from_dim_domain_mapping(DimDomainMapping const &m) { - std::unordered_set l_nontrivial_dims = + std::set l_nontrivial_dims = get_nontrivial_domain_dims(m.l_domain); - std::unordered_set r_nontrivial_dims = + std::set r_nontrivial_dims = get_nontrivial_domain_dims(m.r_domain); return MinimalDimDomainMapping{ /*coord_mapping=*/ - transform_keys(transform_values(m.coord_mapping, - [&](DimCoord const &r_coord) { - return restrict_coord_to_dims( - r_coord, r_nontrivial_dims); - }), - [&](DimCoord const &l_coord) { - return restrict_coord_to_dims(l_coord, - l_nontrivial_dims); - }), + bidict_transform_keys_and_values( + m.coord_mapping, + [&](DimCoord const &l_coord) { + return restrict_coord_to_dims(l_coord, + l_nontrivial_dims); + }, + [&](DimCoord const &r_coord) { + return restrict_coord_to_dims( + r_coord, r_nontrivial_dims); + }), /*l_domain=*/minimal_dim_domain_from_dim_domain(m.l_domain), /*r_domain=*/minimal_dim_domain_from_dim_domain(m.r_domain), }; @@ -118,27 +118,28 @@ MinimalDimDomainMapping template DimDomainMapping dim_domain_mapping_from_minimal_dim_domain( MinimalDimDomainMapping const &m, - std::unordered_set const &l_trivial_dims, - std::unordered_set const &r_trivial_dims) { + std::set const &l_trivial_dims, + std::set const &r_trivial_dims) { DimDomain l_domain = dim_domain_from_minimal_dim_domain(m.l_domain, l_trivial_dims); DimDomain r_domain = dim_domain_from_minimal_dim_domain(m.r_domain, r_trivial_dims); - std::unordered_set all_l_dims = get_domain_dims(l_domain); - std::unordered_set all_r_dims = get_domain_dims(r_domain); + std::set all_l_dims = get_domain_dims(l_domain); + std::set all_r_dims = get_domain_dims(r_domain); return DimDomainMapping{ /*coord_mapping=*/ - transform_keys(transform_values(m.coord_mapping, - [&](DimCoord const &r_coord) { - return lift_dim_coord(r_coord, - all_r_dims); - }), - [&](DimCoord const &l_coord) { - return lift_dim_coord(l_coord, all_l_dims); - }), + bidict_transform_keys_and_values( + m.coord_mapping, + [&](DimCoord const &l_coord) { + return lift_dim_coord(l_coord, all_l_dims); + }, + [&](DimCoord const &r_coord) { + return lift_dim_coord(r_coord, + all_r_dims); + }), /*l_domain=*/l_domain, /*r_domain=*/r_domain, }; @@ -206,13 +207,13 @@ DimDomainMapping compose_dim_domain_mappings_through_minimal( MinimalDimDomainMapping minimal_lhs = minimal_mapping_from_dim_domain_mapping(lhs); - std::unordered_set t1_trivial_dims = + std::set t1_trivial_dims = get_trivial_domain_dims(lhs.l_domain); MinimalDimDomainMapping minimal_rhs = minimal_mapping_from_dim_domain_mapping(rhs); - std::unordered_set t3_trivial_dims = + std::set t3_trivial_dims = get_trivial_domain_dims(rhs.r_domain); return dim_domain_mapping_from_minimal_dim_domain( diff --git a/lib/utils/include/utils/orthotope/orthotope.h b/lib/utils/include/utils/orthotope/orthotope.h index 509497ff00..ef8ad75846 100644 --- a/lib/utils/include/utils/orthotope/orthotope.h +++ b/lib/utils/include/utils/orthotope/orthotope.h @@ -10,7 +10,7 @@ nonnegative_int orthotope_get_num_dims(Orthotope const &); positive_int orthotope_get_volume(Orthotope const &); -std::unordered_set +std::set get_all_coords_in_orthotope(Orthotope const &); bool orthotope_contains_coord(Orthotope const &, OrthotopeCoord const &); diff --git a/lib/utils/include/utils/orthotope/up_projection.h b/lib/utils/include/utils/orthotope/up_projection.h index 1e241108e2..19b753c69f 100644 --- a/lib/utils/include/utils/orthotope/up_projection.h +++ b/lib/utils/include/utils/orthotope/up_projection.h @@ -22,15 +22,15 @@ UpProjection make_empty_up_projection() { } template -std::unordered_set +std::set input_dims_of_up_projection(UpProjection const &projection) { - return unordered_set_of(projection.dim_mapping.left_values()); + return projection.dim_mapping.left_values(); } template -std::unordered_set +std::set output_dims_of_up_projection(UpProjection const &projection) { - return unordered_set_of(projection.dim_mapping.right_values()); + return projection.dim_mapping.right_values(); } template @@ -38,22 +38,24 @@ DimCoord compute_up_projection(UpProjection const &projection, DimCoord const &coord, DimDomain const &output_domain, DimOrdering const &output_dim_ordering) { - std::unordered_set input_dims = input_dims_of_up_projection(projection); - std::unordered_set coord_dims = get_coord_dims(coord); + std::set input_dims = input_dims_of_up_projection(projection); + std::set coord_dims = get_coord_dims(coord); ASSERT(input_dims == coord_dims, "compute_up_projection expected coord dimensions to match projection " "input dimensions"); - std::unordered_set output_dims = output_dims_of_up_projection(projection); - std::unordered_set output_domain_dims = get_domain_dims(output_domain); + std::set output_dims = output_dims_of_up_projection(projection); + std::set output_domain_dims = get_domain_dims(output_domain); ASSERT(is_subseteq_of(output_dims, output_domain_dims)); DimCoord unlifted = DimCoord{ flatmap(coord.raw, - [&](L const &input_dim, nonnegative_int input_dim_val) { - std::unordered_set dst_dims = + [&](L const &input_dim, nonnegative_int input_dim_val) + -> std::map + { + std::set dst_dims = projection.dim_mapping.at_l(input_dim) - .unwrap_as_unordered_set(); + .unwrap_as_set(); DimDomain dst_domain = restrict_domain_to_dims(output_domain, dst_dims); @@ -71,7 +73,7 @@ DimCoord compute_up_projection(UpProjection const &projection, template void project_dims(UpProjection &proj, L const &onto, - std::unordered_set const &from) { + std::set const &from) { ASSERT(from.size() > 0); for (R const &r : from) { diff --git a/lib/utils/include/utils/record_formatter.h b/lib/utils/include/utils/record_formatter.h index 9d2527ab3f..89d28b9aea 100644 --- a/lib/utils/include/utils/record_formatter.h +++ b/lib/utils/include/utils/record_formatter.h @@ -63,7 +63,7 @@ RecordFormatter mk_kv_record(std::string const &k, std::optional const &v) { } template -RecordFormatter mk_record_for_map(std::unordered_map const &m) { +RecordFormatter mk_record_for_map(std::map const &m) { RecordFormatter result = mk_empty_record(Orientation::VERTICAL); for (K const &k : sorted(keys(m))) { diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_filter_keys.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filter_keys.cc index 5c84ec85b1..453ebb4f86 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_filter_keys.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filter_keys.cc @@ -1,10 +1,10 @@ #include "utils/bidict/algorithms/bidict_filter_keys.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; using F = std::function; template bidict bidict_filter_keys(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_filter_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filter_values.cc index 7adf808f85..54be4f72a1 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_filter_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filter_values.cc @@ -1,10 +1,10 @@ #include "utils/bidict/algorithms/bidict_filter_values.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; using F = std::function; template bidict bidict_filter_values(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_keys.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_keys.cc index 954558a66a..18f676f9c9 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_keys.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_keys.cc @@ -1,11 +1,11 @@ #include "utils/bidict/algorithms/bidict_filtrans_keys.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; -using K2 = value_type<2>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; +using K2 = ordered_value_type<2>; using F = std::function(K)>; template bidict bidict_filtrans_keys(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_values.cc index 303d40575a..401d620abf 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_filtrans_values.cc @@ -1,11 +1,11 @@ #include "utils/bidict/algorithms/bidict_filtrans_values.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; -using V2 = value_type<2>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; +using V2 = ordered_value_type<2>; using F = std::function(V)>; template bidict bidict_filtrans_values(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_enumerating.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_enumerating.cc index 7a56bb34bc..65d982d8e6 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_enumerating.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_enumerating.cc @@ -1,9 +1,9 @@ #include "utils/bidict/algorithms/bidict_from_enumerating.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T = value_type<0>; +using T = ordered_value_type<0>; template bidict bidict_from_enumerating(std::vector const &); diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc index 34562f40c1..e52c8703d8 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc @@ -1 +1,14 @@ #include "utils/bidict/algorithms/bidict_from_keys_and_values.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template + bidict bidict_from_keys_and_values( + std::vector const &, + std::vector const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_map.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_map.cc index 64e850db62..2a10d642ae 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_map.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_map.cc @@ -4,10 +4,11 @@ namespace FlexFlow { -template bidict, value_type<1>> - bidict_from_map(std::unordered_map, value_type<1>> const &); +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; -template bidict, value_type<1>> - bidict_from_map(std::map, value_type<1>> const &); +template bidict bidict_from_map(std::map const &); + +template bidict bidict_from_map(std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc index c8a27b8143..271fc35582 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc @@ -1 +1,12 @@ #include "utils/bidict/algorithms/bidict_from_pairs.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template + bidict bidict_from_pairs(std::vector> const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc index de226c5bcc..7ff4b06576 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc @@ -1,12 +1,12 @@ #include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template bidict bidict_from_unstructured_relation( - std::unordered_set> const &); + std::set> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc new file mode 100644 index 0000000000..8e836c2da6 --- /dev/null +++ b/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc @@ -0,0 +1,15 @@ +#include "utils/bidict/algorithms/bidict_transform_keys_and_values.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; +using K2 = ordered_value_type<2>; +using V2 = ordered_value_type<3>; +using KF = std::function; +using VF = std::function; + +template bidict bidict_transform_keys_and_values(bidict const &, KF &&, VF &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc b/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc index 425c1d099c..339754a8f2 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc @@ -1,10 +1,10 @@ #include "utils/bidict/algorithms/bidict_unordered_set_of.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; std::unordered_set> bidict_unordered_set_of(bidict const &); diff --git a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc index 13a1bcd968..d0d2aa3b8f 100644 --- a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc +++ b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc @@ -1,10 +1,10 @@ #include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; template bidict binary_merge_disjoint_bidicts(bidict const &, bidict const &); diff --git a/lib/utils/src/utils/bidict/algorithms/exhaustive_relational_join.cc b/lib/utils/src/utils/bidict/algorithms/exhaustive_relational_join.cc index 589a864f98..07a6be6e4c 100644 --- a/lib/utils/src/utils/bidict/algorithms/exhaustive_relational_join.cc +++ b/lib/utils/src/utils/bidict/algorithms/exhaustive_relational_join.cc @@ -1,11 +1,11 @@ #include "utils/bidict/algorithms/exhaustive_relational_join.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T1 = value_type<0>; -using T2 = value_type<1>; -using T3 = value_type<2>; +using T1 = ordered_value_type<0>; +using T2 = ordered_value_type<1>; +using T3 = ordered_value_type<2>; template bidict exhaustive_relational_join(bidict const &, bidict const &); diff --git a/lib/utils/src/utils/bidict/algorithms/filter_bidict.cc b/lib/utils/src/utils/bidict/algorithms/filter_bidict.cc index 53cee0548e..fac8e88159 100644 --- a/lib/utils/src/utils/bidict/algorithms/filter_bidict.cc +++ b/lib/utils/src/utils/bidict/algorithms/filter_bidict.cc @@ -1,10 +1,10 @@ #include "utils/bidict/algorithms/filter_bidict.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; using F = std::function; template bidict filter_bidict(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/left_entries.cc b/lib/utils/src/utils/bidict/algorithms/left_entries.cc index a2c19de124..99c7d39a7f 100644 --- a/lib/utils/src/utils/bidict/algorithms/left_entries.cc +++ b/lib/utils/src/utils/bidict/algorithms/left_entries.cc @@ -1 +1,11 @@ #include "utils/bidict/algorithms/left_entries.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template std::set left_entries(bidict const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc index 2c27821d3b..741e1be6db 100644 --- a/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc +++ b/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc @@ -1,10 +1,10 @@ #include "utils/bidict/algorithms/merge_disjoint_bidicts.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; template bidict merge_disjoint_bidicts(std::vector> const &); diff --git a/lib/utils/src/utils/bidict/algorithms/right_entries.cc b/lib/utils/src/utils/bidict/algorithms/right_entries.cc index 2f517a0af6..0e657e8d66 100644 --- a/lib/utils/src/utils/bidict/algorithms/right_entries.cc +++ b/lib/utils/src/utils/bidict/algorithms/right_entries.cc @@ -1 +1,11 @@ #include "utils/bidict/algorithms/right_entries.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template std::set right_entries(bidict const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/transform.cc b/lib/utils/src/utils/bidict/algorithms/transform.cc index e6f48f2f24..f6cb4c79f0 100644 --- a/lib/utils/src/utils/bidict/algorithms/transform.cc +++ b/lib/utils/src/utils/bidict/algorithms/transform.cc @@ -1,12 +1,12 @@ #include "utils/bidict/algorithms/transform.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; -using K2 = value_type<2>; -using V2 = value_type<3>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; +using K2 = ordered_value_type<2>; +using V2 = ordered_value_type<3>; using F = std::function(K, V)>; template bidict transform(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/transform_keys.cc b/lib/utils/src/utils/bidict/algorithms/transform_keys.cc index e0ca5c1f0a..a96ec0487b 100644 --- a/lib/utils/src/utils/bidict/algorithms/transform_keys.cc +++ b/lib/utils/src/utils/bidict/algorithms/transform_keys.cc @@ -1,11 +1,11 @@ #include "utils/bidict/algorithms/transform_keys.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; -using K2 = value_type<2>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; +using K2 = ordered_value_type<2>; using F = std::function; template bidict transform_keys(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/transform_values.cc b/lib/utils/src/utils/bidict/algorithms/transform_values.cc index 55337d029c..d6e3d57c13 100644 --- a/lib/utils/src/utils/bidict/algorithms/transform_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/transform_values.cc @@ -1,11 +1,11 @@ #include "utils/bidict/algorithms/transform_values.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; -using V2 = value_type<2>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; +using V2 = ordered_value_type<2>; using F = std::function; template bidict transform_values(bidict const &, F &&); diff --git a/lib/utils/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc b/lib/utils/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc index 472d197371..c08afcc4e3 100644 --- a/lib/utils/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc +++ b/lib/utils/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc @@ -1,12 +1,12 @@ #include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; -template std::unordered_set> +template std::set> unstructured_relation_from_bidict(bidict const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/bidict.cc b/lib/utils/src/utils/bidict/bidict.cc index e3e0b8e9ae..534b99f039 100644 --- a/lib/utils/src/utils/bidict/bidict.cc +++ b/lib/utils/src/utils/bidict/bidict.cc @@ -2,31 +2,25 @@ #include "utils/archetypes/jsonable_value_type.h" #include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/rapidcheckable_value_type.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template struct bidict; -template std::unordered_map format_as(bidict const &); +template std::map format_as(bidict const &); template std::ostream &operator<<(std::ostream &, bidict const &); -using L_Ordered = ordered_value_type<0>; -using R_Ordered = ordered_value_type<1>; - -template bool operator<(bidict const &, - bidict const &); - } // namespace FlexFlow namespace nlohmann { -using L = ::FlexFlow::jsonable_value_type<0>; -using R = ::FlexFlow::jsonable_value_type<1>; +using L = ::FlexFlow::jsonable_ordered_value_type<0>; +using R = ::FlexFlow::jsonable_ordered_value_type<1>; template struct adl_serializer<::FlexFlow::bidict>; @@ -44,8 +38,8 @@ template struct Arbitrary<::FlexFlow::bidict>; namespace std { -using L = ::FlexFlow::value_type<0>; -using R = ::FlexFlow::value_type<1>; +using L = ::FlexFlow::ordered_value_type<0>; +using R = ::FlexFlow::ordered_value_type<1>; template struct hash<::FlexFlow::bidict>; diff --git a/lib/utils/src/utils/cli/cli_parse.cc b/lib/utils/src/utils/cli/cli_parse.cc index 8f5f81324c..36d5837f9c 100644 --- a/lib/utils/src/utils/cli/cli_parse.cc +++ b/lib/utils/src/utils/cli/cli_parse.cc @@ -2,7 +2,7 @@ #include "utils/cli/cli_spec.h" #include "utils/containers/contains.h" #include "utils/containers/enumerate.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" namespace FlexFlow { @@ -27,7 +27,7 @@ tl::expected cli_parse_flag(CLISpec const &cli, tl::expected cli_parse(CLISpec const &cli, std::vector const &args) { CLIParseResult result = CLIParseResult{ - generate_unordered_map(cli_get_flag_keys(cli), + generate_map(cli_get_flag_keys(cli), [](CLIFlagKey const &) { return false; }), {}, }; diff --git a/lib/utils/src/utils/containers/are_disjoint.cc b/lib/utils/src/utils/containers/are_disjoint.cc index 0f0ab8b61c..b71cb3730b 100644 --- a/lib/utils/src/utils/containers/are_disjoint.cc +++ b/lib/utils/src/utils/containers/are_disjoint.cc @@ -1 +1,15 @@ #include "utils/containers/are_disjoint.h" +#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using T = value_type<0>; + +template bool are_disjoint(std::unordered_set const &, std::unordered_set const &); + +using R = ordered_value_type<0>; + +template bool are_disjoint(std::set const &, std::set const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/argmax.cc b/lib/utils/src/utils/containers/argmax.cc index 64b1d6ea4d..f84c33e8d4 100644 --- a/lib/utils/src/utils/containers/argmax.cc +++ b/lib/utils/src/utils/containers/argmax.cc @@ -3,7 +3,7 @@ #include "utils/archetypes/value_type.h" #include #include -#include +#include #include namespace FlexFlow { @@ -13,8 +13,8 @@ using K1 = ordered_value_type<1>; using F1 = std::function; template T1 argmax(std::vector const &, F1 &&); -template T1 argmax(std::unordered_set const &, F1 &&); -template T1 argmax(std::unordered_multiset const &, F1 &&); +template T1 argmax(std::set const &, F1 &&); +template T1 argmax(std::multiset const &, F1 &&); using T2 = ordered_value_type<0>; using K2 = ordered_value_type<1>; diff --git a/lib/utils/src/utils/containers/argmin.cc b/lib/utils/src/utils/containers/argmin.cc index 547f76c117..9f1434861e 100644 --- a/lib/utils/src/utils/containers/argmin.cc +++ b/lib/utils/src/utils/containers/argmin.cc @@ -3,7 +3,7 @@ #include "utils/archetypes/value_type.h" #include #include -#include +#include #include namespace FlexFlow { @@ -13,8 +13,8 @@ using K1 = ordered_value_type<1>; using F1 = std::function; template T1 argmin(std::vector const &, F1 &&); -template T1 argmin(std::unordered_set const &, F1 &&); -template T1 argmin(std::unordered_multiset const &, F1 &&); +template T1 argmin(std::set const &, F1 &&); +template T1 argmin(std::multiset const &, F1 &&); using T2 = ordered_value_type<0>; using K2 = ordered_value_type<1>; diff --git a/lib/utils/src/utils/containers/binary_cartesian_product.cc b/lib/utils/src/utils/containers/binary_cartesian_product.cc index e9cc9c86b6..70195e2a1f 100644 --- a/lib/utils/src/utils/containers/binary_cartesian_product.cc +++ b/lib/utils/src/utils/containers/binary_cartesian_product.cc @@ -1,16 +1,13 @@ #include "utils/containers/binary_cartesian_product.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using A = value_type<0>; -using B = value_type<1>; +using A = ordered_value_type<0>; +using B = ordered_value_type<1>; -template std::unordered_set> - binary_cartesian_product(std::unordered_set const &, - std::unordered_set const &); -template std::unordered_set> - binary_cartesian_product(std::unordered_set const &, - std::unordered_set const &); +template std::set> + binary_cartesian_product(std::set const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc index d777eb0e29..d2f58b90eb 100644 --- a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc @@ -1,9 +1,10 @@ #include "utils/containers/binary_merge_unordered_maps_with_left_dominating.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; template std::unordered_map diff --git a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc index f5586cec6b..62382bfa03 100644 --- a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc @@ -1,13 +1,14 @@ -#include "utils/containers/binary_merge_unordered_maps_with_right_dominating.h" +#include "utils/containers/binary_merge_maps_with_right_dominating.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; template - std::unordered_map binary_merge_unordered_maps_with_right_dominating( - std::unordered_map const &, std::unordered_map const &); + std::map binary_merge_maps_with_right_dominating( + std::map const &, std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/contains_duplicates.cc b/lib/utils/src/utils/containers/contains_duplicates.cc index 3cd07edcf0..757882319e 100644 --- a/lib/utils/src/utils/containers/contains_duplicates.cc +++ b/lib/utils/src/utils/containers/contains_duplicates.cc @@ -1,14 +1,16 @@ #include "utils/containers/contains_duplicates.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { using T = value_type<0>; +using O_T = ordered_value_type<0>; template bool contains_duplicates(std::vector const &); template bool contains_duplicates(std::unordered_multiset const &); -template bool contains_duplicates(std::multiset const &); +template bool contains_duplicates(std::multiset const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/contains_value.cc b/lib/utils/src/utils/containers/contains_value.cc index d9d2118658..4a8332d631 100644 --- a/lib/utils/src/utils/containers/contains_value.cc +++ b/lib/utils/src/utils/containers/contains_value.cc @@ -1,5 +1,6 @@ #include "utils/containers/contains_value.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -8,6 +9,9 @@ using V = value_type<1>; template bool contains_value(std::unordered_map const &, V const &); -template bool contains_value(std::map const &, V const &); +using O_K = ordered_value_type<0>; + +template bool contains_value(std::map const &, V const &); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/enumerate.cc b/lib/utils/src/utils/containers/enumerate.cc index ca5ad6ddc1..9b3da98dee 100644 --- a/lib/utils/src/utils/containers/enumerate.cc +++ b/lib/utils/src/utils/containers/enumerate.cc @@ -7,6 +7,6 @@ using T = value_type<0>; template std::map enumerate(std::vector const &); -template std::map enumerate(std::unordered_set const &); +template std::map enumerate(std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/extend.cc b/lib/utils/src/utils/containers/extend.cc index 1016add5d3..af0ea0e431 100644 --- a/lib/utils/src/utils/containers/extend.cc +++ b/lib/utils/src/utils/containers/extend.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using T = value_type<0>; -using C = std::unordered_multiset; +using C = std::multiset; template void extend(std::vector &, C const &); @@ -14,7 +14,7 @@ template void extend(std::unordered_set &, C const &); template void extend(std::unordered_multiset &, C const &); using T2 = ordered_value_type<0>; -using C2 = std::unordered_multiset; +using C2 = std::multiset; template void extend(std::set &, C2 const &); diff --git a/lib/utils/src/utils/containers/extend_vector.cc b/lib/utils/src/utils/containers/extend_vector.cc index 173e4fe181..b2171d45fd 100644 --- a/lib/utils/src/utils/containers/extend_vector.cc +++ b/lib/utils/src/utils/containers/extend_vector.cc @@ -1,11 +1,11 @@ #include "utils/containers/extend_vector.h" #include "utils/archetypes/value_type.h" -#include +#include namespace FlexFlow { using T = value_type<0>; -template void extend_vector(std::vector &, std::unordered_set const &); +template void extend_vector(std::vector &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/flatmap.cc b/lib/utils/src/utils/containers/flatmap.cc index 676b77542b..2f71264c2e 100644 --- a/lib/utils/src/utils/containers/flatmap.cc +++ b/lib/utils/src/utils/containers/flatmap.cc @@ -32,6 +32,15 @@ using F3 = std::function(InK, InV)>; template std::unordered_map flatmap(std::unordered_map const &, F3 &&); +using O_InK = ordered_value_type<0>; +using O_InV = value_type<1>; +using O_OutK = ordered_value_type<2>; +using O_OutV = value_type<3>; +using O_F3 = std::function(O_InK, O_InV)>; + +template std::map + flatmap(std::map const &, O_F3 &&); + using F4 = std::function(In const &)>; template std::optional flatmap(std::optional const &o, F4 &&); diff --git a/lib/utils/src/utils/containers/generate_unordered_map.cc b/lib/utils/src/utils/containers/generate_unordered_map.cc index 2287e45632..73ecfb48ff 100644 --- a/lib/utils/src/utils/containers/generate_unordered_map.cc +++ b/lib/utils/src/utils/containers/generate_unordered_map.cc @@ -8,6 +8,8 @@ using K = value_type<0>; using V = value_type<1>; using F = std::function; -template std::unordered_map generate_unordered_map(std::unordered_set const &, F &&); +template + std::unordered_map generate_unordered_map( + std::unordered_set const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/get_all_assignments.cc b/lib/utils/src/utils/containers/get_all_assignments.cc index f920ba1c1a..fbe9eaddb0 100644 --- a/lib/utils/src/utils/containers/get_all_assignments.cc +++ b/lib/utils/src/utils/containers/get_all_assignments.cc @@ -1,12 +1,16 @@ #include "utils/containers/get_all_assignments.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/hash/unordered_map.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; template std::unordered_set> get_all_assignments(std::unordered_map> const &); +template std::set> get_all_assignments(std::map> const &); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/get_all_permutations_with_repetition.cc b/lib/utils/src/utils/containers/get_all_permutations_with_repetition.cc index 010bd4562c..b8c41bd099 100644 --- a/lib/utils/src/utils/containers/get_all_permutations_with_repetition.cc +++ b/lib/utils/src/utils/containers/get_all_permutations_with_repetition.cc @@ -1,11 +1,11 @@ #include "utils/containers/get_all_permutations_with_repetition.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T = value_type<0>; +using T = ordered_value_type<0>; -template std::unordered_multiset> +template std::multiset> get_all_permutations_with_repetition(std::vector const &, nonnegative_int n); diff --git a/lib/utils/src/utils/containers/get_element_counts.cc b/lib/utils/src/utils/containers/get_element_counts.cc index ac8e289523..70eda44608 100644 --- a/lib/utils/src/utils/containers/get_element_counts.cc +++ b/lib/utils/src/utils/containers/get_element_counts.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_map get_element_counts(std::string const &s) { +std::map get_element_counts(std::string const &s) { return get_element_counts(vector_of(s)); } diff --git a/lib/utils/src/utils/containers/get_only.cc b/lib/utils/src/utils/containers/get_only.cc index 58cbe085bc..8c24aa77e5 100644 --- a/lib/utils/src/utils/containers/get_only.cc +++ b/lib/utils/src/utils/containers/get_only.cc @@ -1,17 +1,23 @@ #include "utils/containers/get_only.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { using T = value_type<0>; +using O_T = ordered_value_type<0>; template T get_only(std::vector const &); -template T get_only(std::set const &); template T get_only(std::unordered_set const &); +template O_T get_only(std::set const &); using K = value_type<1>; using V = value_type<2>; template std::pair get_only(std::unordered_map const &); +using O_K = ordered_value_type<1>; + +template std::pair get_only(std::map const &); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/group_by.cc b/lib/utils/src/utils/containers/group_by.cc index a41ab4dc62..ee23cab032 100644 --- a/lib/utils/src/utils/containers/group_by.cc +++ b/lib/utils/src/utils/containers/group_by.cc @@ -8,14 +8,9 @@ using K = ordered_value_type<0>; using V = ordered_value_type<1>; using F = std::function; +template OneToMany group_by(std::set const &, F &&); template OneToMany group_by(std::unordered_set const &, F &&); -template std::unordered_map> group_by(std::vector const &, - F &&); - -using V2 = ordered_value_type<1>; -using F2 = std::function; - -template OneToMany group_by(std::set const &, F2 &&); +template std::map> group_by(std::vector const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/invert_map.cc b/lib/utils/src/utils/containers/invert_map.cc index 699503ff96..575f0854d9 100644 --- a/lib/utils/src/utils/containers/invert_map.cc +++ b/lib/utils/src/utils/containers/invert_map.cc @@ -1,11 +1,12 @@ #include "utils/containers/invert_map.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; -template std::unordered_map> - invert_map(std::unordered_map const &); +using O_K = ordered_value_type<0>; +using O_V = ordered_value_type<1>; + +template std::map> + invert_map(std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/invert_unordered_map.cc b/lib/utils/src/utils/containers/invert_unordered_map.cc new file mode 100644 index 0000000000..5cf85480ab --- /dev/null +++ b/lib/utils/src/utils/containers/invert_unordered_map.cc @@ -0,0 +1,13 @@ +#include "utils/containers/invert_unordered_map.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template std::unordered_map> + invert_unordered_map(std::unordered_map const &); + + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/is_submapeq_of.cc b/lib/utils/src/utils/containers/is_submapeq_of.cc index f8fd627b3d..db0c349557 100644 --- a/lib/utils/src/utils/containers/is_submapeq_of.cc +++ b/lib/utils/src/utils/containers/is_submapeq_of.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using K = value_type<0>; using V = value_type<1>; -bool is_submapeq_of(std::unordered_map const &, std::unordered_map const &); +bool is_submapeq_of(std::map const &, std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/is_subseteq_of.cc b/lib/utils/src/utils/containers/is_subseteq_of.cc index 5a71100f92..c3ab1f8ef2 100644 --- a/lib/utils/src/utils/containers/is_subseteq_of.cc +++ b/lib/utils/src/utils/containers/is_subseteq_of.cc @@ -9,7 +9,7 @@ template bool is_subseteq_of(std::unordered_set const &, std::unordered_set const &); using T2 = ordered_value_type<0>; -template bool is_subseteq_of(std::unordered_set const &, - std::unordered_set const &); +template bool is_subseteq_of(std::set const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/items.cc b/lib/utils/src/utils/containers/items.cc index 193cec9ca8..7d4b27e9f7 100644 --- a/lib/utils/src/utils/containers/items.cc +++ b/lib/utils/src/utils/containers/items.cc @@ -1,14 +1,14 @@ #include "utils/containers/items.h" #include "utils/archetypes/ordered_value_type.h" -#include #include +#include namespace FlexFlow { using K = ordered_value_type<0>; using V = ordered_value_type<1>; -template std::set> items(std::unordered_map const &); template std::set> items(std::map const &); +template std::set> items(std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/keys.cc b/lib/utils/src/utils/containers/keys.cc index 96db33f4c7..eb81baa96d 100644 --- a/lib/utils/src/utils/containers/keys.cc +++ b/lib/utils/src/utils/containers/keys.cc @@ -7,7 +7,6 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<1>; -template std::set keys(std::unordered_map const &); template std::set keys(std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/lift_optional_through_map.cc b/lib/utils/src/utils/containers/lift_optional_through_map.cc index 334f9cfdf0..e0100b0020 100644 --- a/lib/utils/src/utils/containers/lift_optional_through_map.cc +++ b/lib/utils/src/utils/containers/lift_optional_through_map.cc @@ -1,12 +1,12 @@ #include "utils/containers/lift_optional_through_map.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; -using V = value_type<1>; +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; -template std::optional> - lift_optional_through_map(std::unordered_map> const &); +template std::optional> + lift_optional_through_map(std::map> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/lookup_in_map.cc b/lib/utils/src/utils/containers/lookup_in_map.cc index a0d7db8e82..c8f3e2feab 100644 --- a/lib/utils/src/utils/containers/lookup_in_map.cc +++ b/lib/utils/src/utils/containers/lookup_in_map.cc @@ -1,12 +1,13 @@ #include "utils/containers/lookup_in_map.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; template std::function - lookup_in_map(std::unordered_map const &map); + lookup_in_map(std::map const &map); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_from_keys_and_values.cc b/lib/utils/src/utils/containers/map_from_keys_and_values.cc index e20b68150a..0c94aace3c 100644 --- a/lib/utils/src/utils/containers/map_from_keys_and_values.cc +++ b/lib/utils/src/utils/containers/map_from_keys_and_values.cc @@ -1,12 +1,13 @@ #include "utils/containers/map_from_keys_and_values.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K1 = value_type<0>; +using K1 = ordered_value_type<0>; using V1 = value_type<1>; -template std::unordered_map +template std::map map_from_keys_and_values(std::vector const &, std::vector const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_from_pairs.cc b/lib/utils/src/utils/containers/map_from_pairs.cc index 8dc8ffa29c..47936dc3da 100644 --- a/lib/utils/src/utils/containers/map_from_pairs.cc +++ b/lib/utils/src/utils/containers/map_from_pairs.cc @@ -1,7 +1,7 @@ #include "utils/containers/map_from_pairs.h" #include "utils/archetypes/ordered_value_type.h" -#include #include +#include #include namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/map_keys.cc b/lib/utils/src/utils/containers/map_keys.cc index daf9ed25d4..f5c07b91ec 100644 --- a/lib/utils/src/utils/containers/map_keys.cc +++ b/lib/utils/src/utils/containers/map_keys.cc @@ -4,19 +4,19 @@ namespace FlexFlow { -using VT0 = value_type<0>; -using VT1 = value_type<1>; -using VT2 = value_type<2>; +using K1 = value_type<0>; +using K2 = value_type<1>; +using V = value_type<2>; template - std::unordered_map map_keys(std::unordered_map const &, - std::function &&); + std::unordered_map map_keys(std::unordered_map const &, + std::function &&); -using OV0 = ordered_value_type<0>; -using OV1 = ordered_value_type<1>; +using O_K1 = ordered_value_type<0>; +using O_K2 = ordered_value_type<1>; template - std::map map_keys(std::map const &m, - std::function &&); + std::map map_keys(std::map const &m, + std::function &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_keys2.cc b/lib/utils/src/utils/containers/map_keys2.cc index 059bac9cf0..2401d64e81 100644 --- a/lib/utils/src/utils/containers/map_keys2.cc +++ b/lib/utils/src/utils/containers/map_keys2.cc @@ -1,14 +1,15 @@ #include "utils/containers/map_keys2.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; -using K2 = value_type<2>; +using K2 = ordered_value_type<2>; using F = std::function; -template std::unordered_map map_keys2(std::unordered_map const &, - F const &); +template std::map map_keys2(std::map const &, + F const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_keys_with_value_merging.cc b/lib/utils/src/utils/containers/map_keys_with_value_merging.cc index 882955a40b..7a6fd94e3f 100644 --- a/lib/utils/src/utils/containers/map_keys_with_value_merging.cc +++ b/lib/utils/src/utils/containers/map_keys_with_value_merging.cc @@ -1,4 +1,5 @@ #include "utils/containers/map_keys_with_value_merging.h" +#include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -13,4 +14,12 @@ using MergeF = std::function; template std::unordered_map map_keys_with_value_merging( std::unordered_map const &, F &&, MergeF &&); +using O_K = ordered_value_type<0>; +using O_K2 = ordered_value_type<2>; + +using O_F = std::function; + +template std::map map_keys_with_value_merging( + std::map const &, O_F &&, MergeF &&); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_values2.cc b/lib/utils/src/utils/containers/map_values2.cc index 8840f15aee..9f95d32324 100644 --- a/lib/utils/src/utils/containers/map_values2.cc +++ b/lib/utils/src/utils/containers/map_values2.cc @@ -4,18 +4,18 @@ namespace FlexFlow { -using VT0 = value_type<0>; -using VT1 = value_type<1>; -using VT2 = value_type<2>; +using K = value_type<0>; +using V1 = value_type<1>; +using V2 = value_type<2>; -template std::unordered_map map_values2( - std::unordered_map const &, - std::function &&); +template std::unordered_map map_values2( + std::unordered_map const &, + std::function &&); -using OT0 = ordered_value_type<0>; +using O_K = ordered_value_type<0>; -template std::map map_values2( - std::map const &, - std::function &&); +template std::map map_values2( + std::map const &, + std::function &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc b/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc index 1ef6af0877..6257356457 100644 --- a/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc +++ b/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc @@ -5,8 +5,9 @@ namespace FlexFlow { using K = value_type<0>; using V = value_type<1>; -using C = std::vector>; -template std::unordered_map merge_disjoint_unordered_maps(C const &); +template + std::unordered_map merge_disjoint_unordered_maps( + std::vector> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_unordered_maps_with.cc b/lib/utils/src/utils/containers/merge_unordered_maps_with.cc index 60218312f3..c9f09e61f0 100644 --- a/lib/utils/src/utils/containers/merge_unordered_maps_with.cc +++ b/lib/utils/src/utils/containers/merge_unordered_maps_with.cc @@ -1,4 +1,4 @@ -#include "utils/containers/merge_unordered_maps_with.h" +#include "utils/containers/merge_maps_with.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -7,8 +7,8 @@ using K = value_type<0>; using V = value_type<1>; using F = std::function; -std::unordered_map - merge_unordered_maps_with(std::vector> const &, +std::map + merge_maps_with(std::vector> const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/multiset_union.cc b/lib/utils/src/utils/containers/multiset_union.cc index a053d05fa6..980f9ea262 100644 --- a/lib/utils/src/utils/containers/multiset_union.cc +++ b/lib/utils/src/utils/containers/multiset_union.cc @@ -1 +1,24 @@ #include "utils/containers/multiset_union.h" +#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using T = value_type<0>; + +template + std::unordered_multiset + multiset_union(std::unordered_multiset const &, + std::unordered_multiset const &); + +using O_T = ordered_value_type<0>; + +template + std::multiset + multiset_union(std::multiset const &, + std::multiset const &); + +template + std::multiset multiset_union(std::vector> const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/require_all_of.cc b/lib/utils/src/utils/containers/require_all_of.cc index 7fbf48c54f..5d85ed3d1c 100644 --- a/lib/utils/src/utils/containers/require_all_of.cc +++ b/lib/utils/src/utils/containers/require_all_of.cc @@ -2,7 +2,7 @@ #include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/value_type.h" #include -#include +#include namespace FlexFlow { @@ -10,8 +10,8 @@ using T1 = value_type<0>; using F1 = std::function; template void require_all_of(std::vector const &, F1 &&); -template void require_all_of(std::unordered_set const &, F1 &&); -template void require_all_of(std::unordered_multiset const &, F1 &&); +template void require_all_of(std::set const &, F1 &&); +template void require_all_of(std::multiset const &, F1 &&); using T2 = ordered_value_type<0>; using F2 = std::function; @@ -23,7 +23,7 @@ using K3 = value_type<0>; using V3 = value_type<1>; using F3 = std::function; -template void require_all_of(std::unordered_map const &, F3 &&); +template void require_all_of(std::map const &, F3 &&); using K4 = ordered_value_type<0>; using V4 = ordered_value_type<1>; diff --git a/lib/utils/src/utils/containers/require_all_same.cc b/lib/utils/src/utils/containers/require_all_same.cc index 87d186caf5..513b681f1a 100644 --- a/lib/utils/src/utils/containers/require_all_same.cc +++ b/lib/utils/src/utils/containers/require_all_same.cc @@ -1,6 +1,6 @@ #include "utils/containers/require_all_same.h" #include "utils/archetypes/value_type.h" -#include +#include namespace FlexFlow { @@ -8,6 +8,6 @@ using T = value_type<0>; template std::optional require_all_same(std::vector const &); -template std::optional require_all_same(std::unordered_set const &); +template std::optional require_all_same(std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/require_all_same1.cc b/lib/utils/src/utils/containers/require_all_same1.cc index 07a9540d97..8f50b48556 100644 --- a/lib/utils/src/utils/containers/require_all_same1.cc +++ b/lib/utils/src/utils/containers/require_all_same1.cc @@ -1,6 +1,6 @@ #include "utils/containers/require_all_same1.h" #include "utils/archetypes/value_type.h" -#include +#include namespace FlexFlow { @@ -8,8 +8,8 @@ using T = value_type<0>; template T require_all_same1(std::vector const &); -template T require_all_same1(std::unordered_set const &); +template T require_all_same1(std::set const &); -template T require_all_same1(std::unordered_multiset const &); +template T require_all_same1(std::multiset const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/require_two_keys.cc b/lib/utils/src/utils/containers/require_two_keys.cc index c1aa4cb3a7..30cc7419a3 100644 --- a/lib/utils/src/utils/containers/require_two_keys.cc +++ b/lib/utils/src/utils/containers/require_two_keys.cc @@ -1,12 +1,13 @@ #include "utils/containers/require_two_keys.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V = value_type<1>; template std::pair - require_two_keys(std::unordered_map const &, K const &, K const &); + require_two_keys(std::map const &, K const &, K const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/restrict_keys.cc b/lib/utils/src/utils/containers/restrict_keys.cc index 13584abec1..7c314733ba 100644 --- a/lib/utils/src/utils/containers/restrict_keys.cc +++ b/lib/utils/src/utils/containers/restrict_keys.cc @@ -8,8 +8,9 @@ using VT0 = value_type<0>; using VT1 = value_type<1>; template - std::unordered_map restrict_keys(std::unordered_map const &, - std::unordered_set const &); + std::unordered_map + restrict_keys(std::unordered_map const &, + std::unordered_set const &); using OV0 = ordered_value_type<0>; diff --git a/lib/utils/src/utils/containers/set_of.cc b/lib/utils/src/utils/containers/set_of.cc index 3a12ee539d..f35f45a0ed 100644 --- a/lib/utils/src/utils/containers/set_of.cc +++ b/lib/utils/src/utils/containers/set_of.cc @@ -1 +1,19 @@ #include "utils/containers/set_of.h" +#include "utils/archetypes/ordered_value_type.h" +#include + +namespace FlexFlow { + +using T = ordered_value_type<0>; + +template std::set set_of(std::vector const &); +template std::set set_of(std::multiset const &); +template std::set set_of(std::unordered_set const &); + +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; + +template std::set> set_of(std::map const &); + + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/set_union.cc b/lib/utils/src/utils/containers/set_union.cc index b993d3a458..bb2ae9be76 100644 --- a/lib/utils/src/utils/containers/set_union.cc +++ b/lib/utils/src/utils/containers/set_union.cc @@ -6,14 +6,13 @@ namespace FlexFlow { using T = value_type<0>; -template std::unordered_set set_union(std::unordered_set const &, - std::unordered_set const &); +template std::unordered_set set_union(std::unordered_set const &, std::unordered_set const &); -using T2 = ordered_value_type<0>; +using O_T = ordered_value_type<0>; -template std::set set_union(std::set const &, std::set const &); +template std::set set_union(std::set const &, std::set const &); -template std::unordered_set - set_union(std::vector> const &); +template std::set + set_union(std::vector> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/transform_pairs.cc b/lib/utils/src/utils/containers/transform_pairs.cc index 4afda936e4..1c6deb907d 100644 --- a/lib/utils/src/utils/containers/transform_pairs.cc +++ b/lib/utils/src/utils/containers/transform_pairs.cc @@ -1,5 +1,6 @@ #include "utils/containers/transform_pairs.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -11,7 +12,12 @@ using F = std::function; template std::vector transform_pairs(std::vector> const &, F &&); -template std::unordered_set - transform_pairs(std::unordered_set> const &, F &&); +using O_L = ordered_value_type<0>; +using O_R = ordered_value_type<1>; +using O_Out = ordered_value_type<2>; +using O_F = std::function; + +template std::set + transform_pairs(std::set> const &, O_F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/try_get_one_of.cc b/lib/utils/src/utils/containers/try_get_one_of.cc index f715d430f2..c39ac76ce3 100644 --- a/lib/utils/src/utils/containers/try_get_one_of.cc +++ b/lib/utils/src/utils/containers/try_get_one_of.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using T = value_type<0>; -template std::optional try_get_one_of(std::unordered_set const &); +template std::optional try_get_one_of(std::set const &); using R = ordered_value_type<0>; diff --git a/lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc b/lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc new file mode 100644 index 0000000000..c97da474e5 --- /dev/null +++ b/lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc @@ -0,0 +1,15 @@ +#include "utils/containers/try_merge_nondisjoint_maps.h" +#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = value_type<1>; + +template + std::optional> + try_merge_nondisjoint_maps(std::map const &, + std::map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc b/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc index e1551068b6..3576b1b7b5 100644 --- a/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc +++ b/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc @@ -1 +1,14 @@ #include "utils/containers/try_merge_nondisjoint_unordered_maps.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template + std::optional> + try_merge_nondisjoint_unordered_maps(std::unordered_map const &, + std::unordered_map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_items.cc b/lib/utils/src/utils/containers/unordered_items.cc index 9b58cfd18e..0ec4d51b15 100644 --- a/lib/utils/src/utils/containers/unordered_items.cc +++ b/lib/utils/src/utils/containers/unordered_items.cc @@ -1,15 +1,14 @@ #include "utils/containers/unordered_items.h" -#include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/value_type.h" -#include #include namespace FlexFlow { -using K = ordered_value_type<0>; +using K = value_type<0>; using V = value_type<1>; -template std::unordered_set> unordered_items(std::unordered_map const &); -template std::unordered_set> unordered_items(std::map const &); +template + std::unordered_set> + unordered_items(std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_keys.cc b/lib/utils/src/utils/containers/unordered_keys.cc index e850b5f460..57ad2efa1c 100644 --- a/lib/utils/src/utils/containers/unordered_keys.cc +++ b/lib/utils/src/utils/containers/unordered_keys.cc @@ -1,4 +1,4 @@ -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/value_type.h" @@ -7,7 +7,7 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<1>; -template std::unordered_set unordered_keys(std::unordered_map const &); -std::unordered_set unordered_keys(std::map const &); +template std::set keys(std::map const &); +std::set keys(std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc b/lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc new file mode 100644 index 0000000000..9a56082145 --- /dev/null +++ b/lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc @@ -0,0 +1,14 @@ +#include "utils/containers/unordered_map_from_keys_and_values.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using K = value_type<0>; +using V = value_type<1>; + +template + std::unordered_map + unordered_map_from_keys_and_values(std::vector const &, + std::vector const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_map_from_pairs.cc b/lib/utils/src/utils/containers/unordered_map_from_pairs.cc index 60cc978be7..eee5f4abaf 100644 --- a/lib/utils/src/utils/containers/unordered_map_from_pairs.cc +++ b/lib/utils/src/utils/containers/unordered_map_from_pairs.cc @@ -1 +1 @@ -#include "utils/containers/unordered_map_from_pairs.h" +#include "utils/containers/map_from_pairs.h" diff --git a/lib/utils/src/utils/containers/unordered_multiset_of.cc b/lib/utils/src/utils/containers/unordered_multiset_of.cc index 5add043c76..3f4e5371b1 100644 --- a/lib/utils/src/utils/containers/unordered_multiset_of.cc +++ b/lib/utils/src/utils/containers/unordered_multiset_of.cc @@ -1 +1 @@ -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" diff --git a/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc b/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc index 75f8d0ccd7..1c975c36e7 100644 --- a/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc +++ b/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc @@ -1,15 +1,15 @@ #include "utils/containers/unstructured_exhaustive_relational_join.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using C = value_type<1>; -using R = value_type<2>; +using L = ordered_value_type<0>; +using C = ordered_value_type<1>; +using R = ordered_value_type<2>; -template std::unordered_set> +template std::set> unstructured_exhaustive_relational_join( - std::unordered_set> const &, - std::unordered_set> const &); + std::set> const &, + std::set> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/vector_from_idx_map.cc b/lib/utils/src/utils/containers/vector_from_idx_map.cc index c0cbd9facd..8b0a0be799 100644 --- a/lib/utils/src/utils/containers/vector_from_idx_map.cc +++ b/lib/utils/src/utils/containers/vector_from_idx_map.cc @@ -8,4 +8,7 @@ using T = value_type<0>; template std::optional> vector_from_idx_map(std::unordered_map const &); +template std::optional> + vector_from_idx_map(std::map const &); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/vector_of.cc b/lib/utils/src/utils/containers/vector_of.cc index 451539248c..0c23655b0f 100644 --- a/lib/utils/src/utils/containers/vector_of.cc +++ b/lib/utils/src/utils/containers/vector_of.cc @@ -2,16 +2,18 @@ #include "utils/archetypes/value_type.h" #include #include +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { +using O_T = ordered_value_type<0>; using T = value_type<0>; template std::vector vector_of(std::vector const &); -template std::vector vector_of(std::unordered_set const &); +template std::vector vector_of(std::set const &); -template std::vector vector_of(std::set const &); +template std::vector vector_of(std::unordered_set const &); template std::vector vector_of(std::optional const &); diff --git a/lib/utils/src/utils/containers/without_nullopts.cc b/lib/utils/src/utils/containers/without_nullopts.cc index 48830d6cf6..25a3c85526 100644 --- a/lib/utils/src/utils/containers/without_nullopts.cc +++ b/lib/utils/src/utils/containers/without_nullopts.cc @@ -2,8 +2,8 @@ namespace FlexFlow { -template std::unordered_set - without_nullopts(std::unordered_set> const &); +template std::set + without_nullopts(std::set> const &); template std::vector without_nullopts(std::vector> const &); diff --git a/lib/utils/src/utils/containers/zip_values_strict_with.cc b/lib/utils/src/utils/containers/zip_values_strict_with.cc index c3fc73a5f9..d7614ab177 100644 --- a/lib/utils/src/utils/containers/zip_values_strict_with.cc +++ b/lib/utils/src/utils/containers/zip_values_strict_with.cc @@ -1,15 +1,16 @@ #include "utils/containers/zip_values_strict_with.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using K = value_type<0>; +using K = ordered_value_type<0>; using V1 = value_type<1>; using V2 = value_type<2>; using Out = value_type<3>; using F = std::function; -template std::unordered_map zip_values_strict_with( - std::unordered_map const &, std::unordered_map const &, F &&); +template std::map zip_values_strict_with( + std::map const &, std::map const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/disjoint_set.cc b/lib/utils/src/utils/disjoint_set.cc index 199e485319..feb8867fa8 100644 --- a/lib/utils/src/utils/disjoint_set.cc +++ b/lib/utils/src/utils/disjoint_set.cc @@ -1,9 +1,9 @@ #include "utils/disjoint_set.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T = value_type<0>; +using T = ordered_value_type<0>; template class m_disjoint_set; diff --git a/lib/utils/src/utils/fmt/unordered_map.cc b/lib/utils/src/utils/fmt/unordered_map.cc index f8746e85a0..21db320044 100644 --- a/lib/utils/src/utils/fmt/unordered_map.cc +++ b/lib/utils/src/utils/fmt/unordered_map.cc @@ -1 +1 @@ -#include "utils/fmt/unordered_map.h" +#include "utils/fmt/map.h" diff --git a/lib/utils/src/utils/fmt/unordered_multiset.cc b/lib/utils/src/utils/fmt/unordered_multiset.cc index cf463296cc..9f20c0d9d1 100644 --- a/lib/utils/src/utils/fmt/unordered_multiset.cc +++ b/lib/utils/src/utils/fmt/unordered_multiset.cc @@ -1 +1 @@ -#include "utils/fmt/unordered_multiset.h" +#include "utils/fmt/multiset.h" diff --git a/lib/utils/src/utils/fmt/unordered_set.cc b/lib/utils/src/utils/fmt/unordered_set.cc index 354eb2f9e7..857367af48 100644 --- a/lib/utils/src/utils/fmt/unordered_set.cc +++ b/lib/utils/src/utils/fmt/unordered_set.cc @@ -1 +1 @@ -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" diff --git a/lib/utils/src/utils/full_binary_tree/find_paths_to_leaf.cc b/lib/utils/src/utils/full_binary_tree/find_paths_to_leaf.cc index 47845720ed..dd9ad385de 100644 --- a/lib/utils/src/utils/full_binary_tree/find_paths_to_leaf.cc +++ b/lib/utils/src/utils/full_binary_tree/find_paths_to_leaf.cc @@ -7,7 +7,7 @@ using Tree = value_type<0>; using Parent = value_type<1>; using Leaf = value_type<2>; -template std::unordered_set +template std::set find_paths_to_leaf(Tree const &, FullBinaryTreeImplementation const &, Leaf const &); diff --git a/lib/utils/src/utils/full_binary_tree/get_all_leaf_paths.cc b/lib/utils/src/utils/full_binary_tree/get_all_leaf_paths.cc index b4d8aa1011..e5c4ffa7da 100644 --- a/lib/utils/src/utils/full_binary_tree/get_all_leaf_paths.cc +++ b/lib/utils/src/utils/full_binary_tree/get_all_leaf_paths.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -template std::unordered_set +template std::set get_all_leaf_paths(value_type<0> const &, FullBinaryTreeImplementation, value_type<1>, diff --git a/lib/utils/src/utils/full_binary_tree/get_leaves.cc b/lib/utils/src/utils/full_binary_tree/get_leaves.cc index 0d7e9106f6..9e16012a97 100644 --- a/lib/utils/src/utils/full_binary_tree/get_leaves.cc +++ b/lib/utils/src/utils/full_binary_tree/get_leaves.cc @@ -1,13 +1,14 @@ #include "utils/full_binary_tree/get_leaves.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { using Tree = value_type<0>; using Parent = value_type<1>; -using Leaf = value_type<2>; +using Leaf = ordered_value_type<2>; -template std::unordered_multiset +template std::multiset get_leaves(Tree const &, FullBinaryTreeImplementation const &); diff --git a/lib/utils/src/utils/full_binary_tree/get_path_to_leaf_map.cc b/lib/utils/src/utils/full_binary_tree/get_path_to_leaf_map.cc index bb39fc739d..3aba26cdfa 100644 --- a/lib/utils/src/utils/full_binary_tree/get_path_to_leaf_map.cc +++ b/lib/utils/src/utils/full_binary_tree/get_path_to_leaf_map.cc @@ -7,7 +7,7 @@ using Tree = value_type<0>; using Parent = value_type<1>; using Leaf = value_type<2>; -template std::unordered_map get_path_to_leaf_map( +template std::map get_path_to_leaf_map( Tree const &, FullBinaryTreeImplementation const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/algorithms.cc b/lib/utils/src/utils/graph/algorithms.cc index 1206c56375..80e073d6f3 100644 --- a/lib/utils/src/utils/graph/algorithms.cc +++ b/lib/utils/src/utils/graph/algorithms.cc @@ -5,7 +5,7 @@ #include "utils/containers/set_difference.h" #include "utils/containers/set_of.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/exception.h" #include "utils/graph/digraph/algorithms/get_incoming_edges.h" @@ -50,13 +50,13 @@ std::vector add_nodes(DiGraph &g, int num_nodes) { struct GetNodesFunctor { template - std::unordered_set operator()(T const &t) { + std::set operator()(T const &t) { return get_nodes(t); } }; -std::unordered_set query_nodes(GraphView const &g, - std::unordered_set const &nodes) { +std::set query_nodes(GraphView const &g, + std::set const &nodes) { NodeQuery query = NodeQuery{ query_set::match_values_in(set_of(nodes)), }; @@ -121,14 +121,14 @@ void add_edges(UndirectedGraph &g, add_edges(g, std::vector{edges}); } -void add_edges(DiGraph &g, std::unordered_set const &edges) { +void add_edges(DiGraph &g, std::set const &edges) { for (DirectedEdge const &e : edges) { g.add_edge(e); } } void add_edges(UndirectedGraph &g, - std::unordered_set const &edges) { + std::set const &edges) { for (UndirectedEdge const &e : edges) { g.add_edge(e); } @@ -149,7 +149,7 @@ bool contains_edge(UndirectedGraphView const &g, UndirectedEdge const &e) { return contains(g.query_edges(q), e); } -void remove_edges(DiGraph &g, std::unordered_set const &edges) { +void remove_edges(DiGraph &g, std::set const &edges) { for (DirectedEdge const &e : edges) { ASSERT(contains_edge(g, e), "remove_edges expected edge to exist in DiGraph"); @@ -158,7 +158,7 @@ void remove_edges(DiGraph &g, std::unordered_set const &edges) { } void remove_edges(UndirectedGraph &g, - std::unordered_set const &edges) { + std::set const &edges) { for (UndirectedEdge const &e : edges) { ASSERT(contains_edge(g, e), "remove_edges expected edge to exist in UndirectedGraph"); @@ -166,7 +166,7 @@ void remove_edges(UndirectedGraph &g, } } -std::unordered_set get_node_edges(UndirectedGraphView const &g, +std::set get_node_edges(UndirectedGraphView const &g, Node const &n) { UndirectedEdgeQuery query = UndirectedEdgeQuery{ query_set::match_single_value(n), @@ -176,31 +176,31 @@ std::unordered_set get_node_edges(UndirectedGraphView const &g, } std::vector get_unchecked_dfs_ordering( - DiGraphView const &g, std::unordered_set const &starting_points) { + DiGraphView const &g, std::set const &starting_points) { UncheckedDFSView dfs_view = unchecked_dfs(g, starting_points); return {dfs_view.begin(), dfs_view.end()}; } std::vector get_dfs_ordering(DiGraphView const &g, - std::unordered_set const &starting_points) { + std::set const &starting_points) { CheckedDFSView dfs_view = dfs(g, starting_points); return {dfs_view.begin(), dfs_view.end()}; } std::vector get_bfs_ordering(DiGraphView const &g, - std::unordered_set const &starting_points) { + std::set const &starting_points) { BFSView bfs_view = bfs(g, starting_points); return {bfs_view.begin(), bfs_view.end()}; } -std::unordered_set get_neighbors(DiGraphView const &g, Node const &n) { +std::set get_neighbors(DiGraphView const &g, Node const &n) { UndirectedGraphView undirected = as_undirected(g); return get_neighbors(undirected, n); } -std::unordered_set get_neighbors(UndirectedGraphView const &g, +std::set get_neighbors(UndirectedGraphView const &g, Node const &n) { return flatmap(get_node_edges(g, n), [&](UndirectedEdge const &edge) { return set_difference(get_endpoints(edge), {n}); @@ -208,12 +208,12 @@ std::unordered_set get_neighbors(UndirectedGraphView const &g, } UndirectedGraphView get_subgraph(UndirectedGraphView const &g, - std::unordered_set const &nodes) { + std::set const &nodes) { return UndirectedGraphView::create(g, nodes); } DiGraphView get_subgraph(DiGraphView const &g, - std::unordered_set const &nodes) { + std::set const &nodes) { return DiGraphView::create(g, nodes); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc index 072104ae35..35f67ac808 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc @@ -7,7 +7,7 @@ namespace FlexFlow { -std::unordered_set get_edges(DataflowGraphView const &g) { +std::set get_edges(DataflowGraphView const &g) { return g.query_edges(dataflow_edge_query_all()); } @@ -34,7 +34,7 @@ std::vector get_outputs(DataflowGraphView const &g, }); } -std::unordered_set +std::set get_all_dataflow_outputs(DataflowGraphView const &g) { return g.query_outputs(dataflow_output_query_all()); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphism.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphism.cc index 0e4d0c6759..5c93e96b3a 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphism.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphism.cc @@ -7,7 +7,7 @@ namespace FlexFlow { std::optional find_isomorphism(DataflowGraphView const &src, DataflowGraphView const &dst) { - std::unordered_set all_isomorphisms = + std::set all_isomorphisms = find_isomorphisms(src, dst); if (all_isomorphisms.empty()) { diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc index 20eb69210d..241ad7efc7 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc @@ -5,10 +5,10 @@ namespace FlexFlow { -std::unordered_set +std::set find_isomorphisms(DataflowGraphView const &src, DataflowGraphView const &dst) { - std::unordered_set open_isomorphisms = + std::set open_isomorphisms = find_isomorphisms(view_as_open_dataflow_graph(src), view_as_open_dataflow_graph(dst)); diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc index ceca6982a6..9e06c21500 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc @@ -2,7 +2,7 @@ namespace FlexFlow { -std::unordered_set get_dataflow_edges_from_node_to_node( +std::set get_dataflow_edges_from_node_to_node( DataflowGraphView const &g, Node const &src, Node const &dst) { return g.query_edges(DataflowEdgeQuery{ /*src_nodes=*/query_set::match_single_value(src), diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc index 42bba1892f..6c6f2fc3db 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc @@ -17,9 +17,9 @@ std::vector get_incoming_edges(DataflowGraphView const &g, }); } -std::unordered_set +std::set get_incoming_edges(DataflowGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { DataflowEdgeQuery query = DataflowEdgeQuery{ query_set::matchall(), query_set::matchall(), diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc index f958b8e085..d64ce3c17c 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_set get_outgoing_edges(DataflowGraphView const &g, +std::set get_outgoing_edges(DataflowGraphView const &g, Node const &n) { return g.query_edges(DataflowEdgeQuery{ query_set::match_single_value(n), @@ -14,9 +14,9 @@ std::unordered_set get_outgoing_edges(DataflowGraphView const &g, }); } -std::unordered_set +std::set get_outgoing_edges(DataflowGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { DataflowEdgeQuery query = DataflowEdgeQuery{ query_set::match_values_in(set_of(ns)), query_set::matchall(), diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc index e89189776e..8d322c3c1d 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc @@ -5,11 +5,11 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_incoming_edges(DataflowGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { - std::unordered_set all_nodes = get_nodes(g); + std::set all_nodes = get_nodes(g); query_set src_query = query_set::match_values_in(set_of(set_minus(all_nodes, ns))); diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc index c958a2e248..6dfed0a9af 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc @@ -5,11 +5,11 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_outgoing_edges(DataflowGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { - std::unordered_set all_nodes = get_nodes(g); + std::set all_nodes = get_nodes(g); query_set dst_query = query_set::match_values_in(set_of(set_minus(all_nodes, ns))); diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_boundary_nodes_for_split.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_boundary_nodes_for_split.cc index 70a66c9a21..97f498c291 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_boundary_nodes_for_split.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_boundary_nodes_for_split.cc @@ -6,13 +6,13 @@ namespace FlexFlow { SplitBoundaryNodes get_transitive_reduced_boundary_nodes_for_split( TransitiveReducedDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { - std::unordered_set edges = + std::set edges = get_transitive_reduced_edges_across_split(tr_g, split); - std::unordered_set src_boundary_nodes = + std::set src_boundary_nodes = transform(edges, [](DataflowEdge const &e) { return e.src.node; }); - std::unordered_set dst_boundary_nodes = + std::set dst_boundary_nodes = transform(edges, [](DataflowEdge const &e) { return e.dst.node; }); return SplitBoundaryNodes{ diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc index 8a4adf0b3a..4c02f9c54e 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc @@ -6,15 +6,15 @@ namespace FlexFlow { -std::unordered_set get_transitive_reduced_edges_across_split( +std::set get_transitive_reduced_edges_across_split( TransitiveReducedDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { - std::unordered_set src_subgraph = - unordered_set_of(get_leaves(split.get_left_child())); - std::unordered_set dst_subgraph = - unordered_set_of(get_leaves(split.get_right_child())); + std::set src_subgraph = + set_of(get_leaves(split.get_left_child())); + std::set dst_subgraph = + set_of(get_leaves(split.get_right_child())); - std::unordered_set raw_edges = + std::set raw_edges = get_edges_from_subgraph_to_subgraph( tr_g.transitive_reduction, src_subgraph, dst_subgraph); diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc index 0bb94c87f4..d2a8943819 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_set get_transitive_reduced_outputs_across_split( +std::set get_transitive_reduced_outputs_across_split( TransitiveReducedDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { return transform(get_transitive_reduced_edges_across_split(tr_g, split), diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc index 703db4bf91..3c8fabb51e 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc @@ -7,28 +7,28 @@ ViewDataflowGraphAsOpenDataflowGraph::ViewDataflowGraphAsOpenDataflowGraph( DataflowGraphView const &g) : g(g) {} -std::unordered_set ViewDataflowGraphAsOpenDataflowGraph::query_nodes( +std::set ViewDataflowGraphAsOpenDataflowGraph::query_nodes( NodeQuery const &q) const { return this->g.query_nodes(q); } -std::unordered_set +std::set ViewDataflowGraphAsOpenDataflowGraph::query_edges( OpenDataflowEdgeQuery const &q) const { - std::unordered_set closed_edges = + std::set closed_edges = this->g.query_edges(q.standard_edge_query); return transform(closed_edges, [](DataflowEdge const &e) { return OpenDataflowEdge{e}; }); } -std::unordered_set +std::set ViewDataflowGraphAsOpenDataflowGraph::query_outputs( DataflowOutputQuery const &q) const { return this->g.query_outputs(q); } -std::unordered_set +std::set ViewDataflowGraphAsOpenDataflowGraph::get_inputs() const { return {}; } diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.cc index 90f2f5134c..967c5a81bc 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_from_dataflow_graph_data.cc @@ -10,21 +10,21 @@ ViewFromDataflowGraphData::ViewFromDataflowGraphData( DataflowGraphData const &data) : data(data) {} -std::unordered_set +std::set ViewFromDataflowGraphData::query_nodes(NodeQuery const &query) const { - return apply_node_query(query, this->data.nodes); + return apply_node_query(query, set_of(this->data.nodes)); } -std::unordered_set ViewFromDataflowGraphData::query_edges( +std::set ViewFromDataflowGraphData::query_edges( DataflowEdgeQuery const &query) const { - return filter(this->data.edges, [&](DataflowEdge const &e) { + return filter(set_of(this->data.edges), [&](DataflowEdge const &e) { return dataflow_edge_query_includes_dataflow_edge(query, e); }); } -std::unordered_set ViewFromDataflowGraphData::query_outputs( +std::set ViewFromDataflowGraphData::query_outputs( DataflowOutputQuery const &query) const { - return filter(this->data.outputs, [&](DataflowOutput const &o) { + return filter(set_of(this->data.outputs), [&](DataflowOutput const &o) { return dataflow_output_query_includes_dataflow_output(query, o); }); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph.cc b/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph.cc index 8ed36135e1..36ecf002c2 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph.cc @@ -15,16 +15,16 @@ void DataflowGraph::add_node_unsafe( return this->get_interface().add_node_unsafe(node, inputs, outputs); } -std::unordered_set DataflowGraph::query_nodes(NodeQuery const &q) const { +std::set DataflowGraph::query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } -std::unordered_set +std::set DataflowGraph::query_edges(DataflowEdgeQuery const &q) const { return this->get_interface().query_edges(q); } -std::unordered_set +std::set DataflowGraph::query_outputs(DataflowOutputQuery const &q) const { return this->get_interface().query_outputs(q); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc b/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc index 460f55b64a..6fa21a5a54 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc @@ -2,17 +2,17 @@ namespace FlexFlow { -std::unordered_set +std::set DataflowGraphView::query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } -std::unordered_set +std::set DataflowGraphView::query_edges(DataflowEdgeQuery const &q) const { return this->get_interface().query_edges(q); } -std::unordered_set +std::set DataflowGraphView::query_outputs(DataflowOutputQuery const &q) const { return this->get_interface().query_outputs(q); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/dataflow_output_query.cc b/lib/utils/src/utils/graph/dataflow_graph/dataflow_output_query.cc index eb1dfacd8f..849f909f22 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/dataflow_output_query.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/dataflow_output_query.cc @@ -28,9 +28,9 @@ DataflowOutputQuery dataflow_output_query_for_output(DataflowOutput const &o) { }; } -std::unordered_set +std::set apply_dataflow_output_query(DataflowOutputQuery const &q, - std::unordered_set const &os) { + std::set const &os) { return filter(os, [&](DataflowOutput const &o) { return dataflow_output_query_includes_dataflow_output(q, o); }); diff --git a/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc b/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc index ef9412b939..628679ce30 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set +std::set IDataflowGraphView::query_edges(DirectedEdgeQuery const &q) const { DataflowEdgeQuery dataflow_query = DataflowEdgeQuery{ q.srcs, @@ -11,7 +11,7 @@ std::unordered_set q.dsts, matchall(), }; - std::unordered_set dataflow_edges = + std::set dataflow_edges = this->query_edges(dataflow_query); return transform(dataflow_edges, [](DataflowEdge const &e) { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/apply_contraction.cc b/lib/utils/src/utils/graph/digraph/algorithms/apply_contraction.cc index 8124ae3505..4efbe28044 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/apply_contraction.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/apply_contraction.cc @@ -7,7 +7,7 @@ namespace FlexFlow { DiGraphView apply_contraction(DiGraphView const &g, - std::unordered_map const &nodes) { + std::map const &nodes) { auto get_dst = [&](Node const &src) { Node result = src; while (contains_key(nodes, result)) { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/calculate_topo_rank.cc b/lib/utils/src/utils/graph/digraph/algorithms/calculate_topo_rank.cc index 7e0606d338..9d1599d2a1 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/calculate_topo_rank.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/calculate_topo_rank.cc @@ -3,9 +3,9 @@ namespace FlexFlow { -std::unordered_map calculate_topo_rank(DiGraphView const &g) { +std::map calculate_topo_rank(DiGraphView const &g) { std::vector topo_ordering = get_topological_ordering(g); - std::unordered_map topo_rank; + std::map topo_rank; for (int i = 0; i < topo_ordering.size(); i++) { topo_rank[topo_ordering[i]] = i; } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc index f3ccd24536..b71e92b629 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc @@ -3,14 +3,14 @@ #include "utils/containers/filter.h" #include "utils/containers/maybe_get_only.h" #include "utils/containers/transform.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" #include namespace FlexFlow { std::optional get_component_containing_node_in_head( CompleteBipartiteCompositeDecomposition const &cbc, Node const &n) { - std::unordered_set found = + std::set found = filter(cbc.subgraphs, [&](BipartiteComponent const &bc) { return contains(bc.head_nodes, n); }); @@ -20,7 +20,7 @@ std::optional get_component_containing_node_in_head( std::optional get_component_containing_node_in_tail( CompleteBipartiteCompositeDecomposition const &cbc, Node const &n) { - std::unordered_set found = + std::set found = filter(cbc.subgraphs, [&](BipartiteComponent const &bc) { return contains(bc.tail_nodes, n); }); @@ -28,13 +28,13 @@ std::optional get_component_containing_node_in_tail( return maybe_get_only(found); } -std::unordered_set> +std::set> get_head_subcomponents(CompleteBipartiteCompositeDecomposition const &cbc) { return transform(cbc.subgraphs, [](BipartiteComponent const &bc) { return bc.head_nodes; }); } -std::unordered_set> +std::set> get_tail_subcomponents(CompleteBipartiteCompositeDecomposition const &cbc) { return transform(cbc.subgraphs, [](BipartiteComponent const &bc) { return bc.tail_nodes; }); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc index 4900cdaa10..a42feea337 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc @@ -31,10 +31,10 @@ std::optional edges_to_process.push(e); } - std::unordered_set already_in_a_head = {}; - std::unordered_set already_in_a_tail = {}; + std::set already_in_a_head = {}; + std::set already_in_a_tail = {}; - std::unordered_set already_processed = {}; + std::set already_processed = {}; CompleteBipartiteCompositeDecomposition result = CompleteBipartiteCompositeDecomposition{{}}; @@ -46,14 +46,14 @@ std::optional continue; } - std::unordered_set head = get_predecessors(g, e.dst); - std::unordered_set tail = get_successors(g, e.src); + std::set head = get_predecessors(g, e.dst); + std::set tail = get_successors(g, e.src); if (!are_disjoint(head, tail)) { return std::nullopt; } - std::unordered_set from_head_to_tail = + std::set from_head_to_tail = g.query_edges(DirectedEdgeQuery{ query_set::match_values_in(set_of(head)), query_set::match_values_in(set_of(tail)), diff --git a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc index 647754f522..ae9ad21d5f 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc @@ -11,12 +11,12 @@ bool is_complete_bipartite_digraph(DiGraphView const &g) { } bool is_complete_bipartite_digraph(DiGraphView const &g, - std::unordered_set const &srcs) { - std::unordered_set sinks = set_minus(get_nodes(g), srcs); + std::set const &srcs) { + std::set sinks = set_minus(get_nodes(g), srcs); - std::unordered_set edges = get_edges(g); + std::set edges = get_edges(g); - std::unordered_set expected_edges; + std::set expected_edges; for (Node const &src : srcs) { for (Node const &sink : sinks) { expected_edges.insert(DirectedEdge{src, sink}); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc b/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc index 5b0a45f6d8..d192dc8cdf 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set +std::set ContractNodeView::query_edges(DirectedEdgeQuery const &q) const { return transform(g.query_edges(q), [&](DirectedEdge const &e) { DirectedEdge result = e; @@ -17,7 +17,7 @@ std::unordered_set }); } -std::unordered_set +std::set ContractNodeView::query_nodes(NodeQuery const &q) const { return transform(g.query_nodes(q), [&](Node const &n) { if (n == this->from) { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc b/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc index 0858a6d1dc..61f735db24 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc @@ -5,15 +5,15 @@ namespace FlexFlow { FlippedView::FlippedView(DiGraphView const &g) : g(g) {} -std::unordered_set +std::set FlippedView::query_edges(DirectedEdgeQuery const &query) const { - std::unordered_set result = + std::set result = this->g.query_edges(DirectedEdgeQuery{query.dsts, query.srcs}); return transform( result, [](DirectedEdge const &e) { return flipped_directed_edge(e); }); } -std::unordered_set +std::set FlippedView::query_nodes(NodeQuery const &query) const { return this->g.query_nodes(query); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc index 96a34e6f0b..4af4620070 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc @@ -4,7 +4,7 @@ #include "utils/graph/digraph/algorithms/is_acyclic.h" namespace FlexFlow { -std::unordered_set get_ancestors(DiGraphView const &g, +std::set get_ancestors(DiGraphView const &g, Node const &starting_node) { assert(is_acyclic(g)); return get_descendants(flipped(g), starting_node); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc index 284794a49c..6004b1a4d1 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc @@ -1,6 +1,6 @@ #include "utils/graph/digraph/algorithms/get_descendants.h" #include "utils/containers/contains.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_successors.h" #include "utils/graph/digraph/algorithms/is_acyclic.h" @@ -8,12 +8,12 @@ #include "utils/graph/node/algorithms.h" namespace FlexFlow { -std::unordered_set get_descendants(DiGraphView const &g, +std::set get_descendants(DiGraphView const &g, Node const &starting_node) { assert(is_acyclic(g)); assert(contains(get_nodes(g), starting_node)); - return unordered_set_of( + return set_of( get_bfs_ordering(g, get_successors(g, starting_node))); }; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc index 3837c93af2..aaf0a850ef 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc @@ -3,21 +3,21 @@ #include "utils/containers/set_intersection.h" #include "utils/containers/values.h" #include "utils/graph/digraph/algorithms/get_dominators_map.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" #include "utils/optional.h" #include namespace FlexFlow { -std::unordered_set get_dominators(DiGraphView const &g, Node const &n) { +std::set get_dominators(DiGraphView const &g, Node const &n) { return get_dominators_map(g).at(n); } -std::unordered_set get_dominators(DiGraphView const &g, - std::unordered_set const &n) { +std::set get_dominators(DiGraphView const &g, + std::set const &n) { ASSERT(n.size() > 0, "Cannot find dominators of no nodes"); - std::optional> result = + std::optional> result = set_intersection(values(restrict_keys(get_dominators_map(g), n))); return assert_unwrap(result); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc index fc13a9d0d1..aaaf2f47a6 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/get_dominators_map.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_intersection.h" #include "utils/containers/transform.h" @@ -9,15 +9,15 @@ #include "utils/graph/digraph/algorithms/get_successors.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" #include "utils/graph/node/algorithms.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" #include #include namespace FlexFlow { -std::unordered_map> +std::map> get_dominators_map(DiGraphView const &g) { - std::unordered_set initial_nodes = get_initial_nodes(g); + std::set initial_nodes = get_initial_nodes(g); std::queue queue; @@ -25,18 +25,18 @@ std::unordered_map> queue.push(src); } - std::unordered_map> result = - generate_unordered_map(get_nodes(g), [&](Node const &) { return get_nodes(g); }); + std::map> result = + generate_map(get_nodes(g), [&](Node const &) { return get_nodes(g); }); while (!queue.empty()) { Node n = queue.front(); queue.pop(); - std::unordered_set old_result_entry = result.at(n); + std::set old_result_entry = result.at(n); result.at(n) = set_intersection(transform(get_predecessors(g, n), [&](Node const &n) { return result.at(n); - })).value_or(std::unordered_set{}); + })).value_or(std::set{}); result.at(n).insert(n); if (result.at(n) != old_result_entry) { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_edges.cc index 2f95c78d2a..a8446f3ad1 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_edges.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set get_edges(DiGraphView const &g) { +std::set get_edges(DiGraphView const &g) { return g.query_edges(directed_edge_query_all()); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc index e8d450d3a2..f7e07905aa 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc @@ -4,10 +4,10 @@ namespace FlexFlow { -std::unordered_set get_edges_from_subgraph_to_subgraph( +std::set get_edges_from_subgraph_to_subgraph( DiGraphView const &g, - std::unordered_set const &src_subgraph, - std::unordered_set const &dst_subgraph) { + std::set const &src_subgraph, + std::set const &dst_subgraph) { if (!are_disjoint(src_subgraph, dst_subgraph)) { throw mk_runtime_error( fmt::format("get_edges_from_subgraph_to_subgraph(DiGraphView, ...) " diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc index 42a9830837..21d598eca6 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc @@ -1,7 +1,7 @@ #include "utils/graph/digraph/algorithms/get_imm_dominators_map.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/filter_values.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_element_counts.h" #include "utils/containers/get_only.h" #include "utils/containers/keys.h" @@ -12,29 +12,29 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_imm_dominators_map(DiGraphView const &g) { - std::unordered_map> node_to_its_dominators = + std::map> node_to_its_dominators = get_dominators_map(g); auto get_imm_dominator = [&](Node const &n) { - std::unordered_set n_dominators = node_to_its_dominators.at(n); + std::set n_dominators = node_to_its_dominators.at(n); n_dominators.erase(n); std::vector recursive_dominator_list = concat_vectors( transform(vector_of(n_dominators), [&](Node const &dominator) { return vector_of(node_to_its_dominators.at(dominator)); })); - std::unordered_map dominator_counts = + std::map dominator_counts = get_element_counts(recursive_dominator_list); - std::unordered_set imm_dominators = unordered_keys( + std::set imm_dominators = keys( filter_values(dominator_counts, [](int count) { return count <= 1; })); ASSERT(imm_dominators.size() <= 1); return maybe_get_only(imm_dominators); }; - return generate_unordered_map(get_nodes(g), get_imm_dominator); + return generate_map(get_nodes(g), get_imm_dominator); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc index e043a185f9..ea95b6829f 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/get_imm_post_dominator.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_one_of.h" #include "utils/containers/get_only.h" #include "utils/containers/restrict_keys.h" @@ -20,7 +20,7 @@ std::optional get_imm_post_dominator(DiGraphView const &g, std::optional get_imm_post_dominator(DiGraphView const &g, - std::unordered_set const &nodes) { + std::set const &nodes) { if (nodes.empty()) { throw mk_runtime_error("Cannot get imm_post_dominator of no nodes"); @@ -31,8 +31,8 @@ std::optional } Node contracted_node = get_one_of(nodes); - std::unordered_map contraction = - generate_unordered_map(nodes, [&](Node const &) { return contracted_node; }); + std::map contraction = + generate_map(nodes, [&](Node const &) { return contracted_node; }); return get_imm_post_dominator(apply_contraction(g, contraction), contracted_node); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc index 0f7bf3a299..ebd84eed45 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_imm_post_dominators_map(DiGraphView const &g) { return get_imm_dominators_map(flipped(g)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc index e6c9d5e557..ad0f8cffaa 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc @@ -2,12 +2,11 @@ #include "utils/containers/group_by.h" #include "utils/containers/map_values.h" #include "utils/containers/set_of.h" -#include "utils/nonempty_unordered_set/nonempty_unordered_set.h" -#include "utils/containers/unordered_map_from_map.h" +#include "utils/nonempty_set/nonempty_set.h" namespace FlexFlow { -std::unordered_set get_incoming_edges(DiGraphView const &g, +std::set get_incoming_edges(DiGraphView const &g, Node const &n) { return g.query_edges(DirectedEdgeQuery{ query_set::matchall(), @@ -15,11 +14,11 @@ std::unordered_set get_incoming_edges(DiGraphView const &g, }); } -std::unordered_map> +std::map> get_incoming_edges(DiGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { - std::map> by_dst = + std::map> by_dst = group_by(g.query_edges(DirectedEdgeQuery{ query_set::matchall(), query_set::match_values_in(set_of(ns)), @@ -27,18 +26,18 @@ std::unordered_map> [](DirectedEdge const &e) { return e.dst; }) .l_to_r(); - std::map> result = + std::map> result = map_values(by_dst, [](nonempty_set const &s) - -> std::unordered_set { - return s.unwrap_as_unordered_set(); + -> std::set { + return s.unwrap_as_set(); }); for (Node const &n : ns) { result[n]; } - return unordered_map_from_map(result); + return result; } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_initial_nodes.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_initial_nodes.cc index 6be4f8a592..3dad5abd1e 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_initial_nodes.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_initial_nodes.cc @@ -6,9 +6,9 @@ namespace FlexFlow { -std::unordered_set get_initial_nodes(DiGraphView const &g) { - std::unordered_set all_nodes = get_nodes(g); - std::unordered_set with_incoming_edge = +std::set get_initial_nodes(DiGraphView const &g) { + std::set all_nodes = get_nodes(g); + std::set with_incoming_edge = transform(get_edges(g), [](DirectedEdge const &e) { return e.dst; }); return set_minus(all_nodes, with_incoming_edge); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc index 036d320f05..a6a736a83b 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc @@ -6,21 +6,21 @@ #include "utils/graph/digraph/algorithms/get_predecessors.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" #include "utils/graph/digraph/algorithms/is_acyclic.h" -#include +#include namespace FlexFlow { -std::unordered_map get_weighted_longest_path_lengths_from_root( - DiGraphView const &g, std::unordered_map const &node_costs) { +std::map get_weighted_longest_path_lengths_from_root( + DiGraphView const &g, std::map const &node_costs) { assert(is_acyclic(g)); assert(all_of(values(node_costs), [&](float cost) { return cost >= 0; })); std::vector topo_order = get_topological_ordering(g); - std::unordered_map longest_path_lengths; + std::map longest_path_lengths; for (Node const &n : topo_order) { - std::unordered_set predecessor_path_lengths = + std::set predecessor_path_lengths = transform(get_predecessors(g, n), [&](Node const &pred) { return longest_path_lengths.at(pred); }); @@ -32,16 +32,16 @@ std::unordered_map get_weighted_longest_path_lengths_from_root( return longest_path_lengths; } -std::unordered_map +std::map get_longest_path_lengths_from_root(DiGraphView const &g) { assert(is_acyclic(g)); std::vector topo_order = get_topological_ordering(g); - std::unordered_map longest_path_lengths; + std::map longest_path_lengths; for (Node const &n : topo_order) { - std::unordered_set predecessor_path_lengths = + std::set predecessor_path_lengths = transform(get_predecessors(g, n), [&](Node const &pred) { return longest_path_lengths.at(pred); }); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc index dbcd07b0ff..a87170397c 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc @@ -8,32 +8,32 @@ #include "utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h" #include "utils/graph/digraph/algorithms/is_acyclic.h" #include "utils/graph/node/algorithms.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" #include "utils/nonnegative_int/nonnegative_int.h" #include namespace FlexFlow { -std::optional> +std::optional> get_lowest_common_ancestors(DiGraphView const &g, - std::unordered_set const &nodes) { + std::set const &nodes) { ASSERT(is_acyclic(g)); ASSERT(is_subseteq_of(nodes, get_nodes(g))); if (num_nodes(g) == 0 || nodes.size() == 0) { return std::nullopt; } - std::unordered_set> ancestors = + std::set> ancestors = transform(nodes, [&](Node const &n) { return set_union(get_ancestors(g, n), {n}); }); - std::unordered_set common_ancestors = + std::set common_ancestors = set_intersection(ancestors).value(); if (common_ancestors.empty()) { - return std::unordered_set{}; + return std::set{}; } - std::unordered_map depth_levels = + std::map depth_levels = get_longest_path_lengths_from_root(g); nonnegative_int largest_depth_for_common_ancestors = maximum(transform( diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.cc index da494f19d2..96dff18333 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_node_with_greatest_topo_rank.cc @@ -3,9 +3,9 @@ namespace FlexFlow { -Node get_node_with_greatest_topo_rank(std::unordered_set const &nodes, +Node get_node_with_greatest_topo_rank(std::set const &nodes, DiGraphView const &g) { - std::unordered_map topo_rank = calculate_topo_rank(g); + std::map topo_rank = calculate_topo_rank(g); return *std::max_element(nodes.cbegin(), nodes.cend(), [&topo_rank](Node const &lhs, Node const &rhs) { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc index 883d8e3725..ba42fca1b5 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc @@ -2,14 +2,13 @@ #include "utils/containers/group_by.h" #include "utils/containers/map_values.h" #include "utils/containers/set_of.h" -#include "utils/nonempty_unordered_set/nonempty_unordered_set.h" -#include "utils/containers/unordered_map_from_map.h" +#include "utils/nonempty_set/nonempty_set.h" namespace FlexFlow { -std::unordered_map> +std::map> get_outgoing_edges(DiGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { std::map> by_src = group_by(g.query_edges(DirectedEdgeQuery{ @@ -19,20 +18,20 @@ std::unordered_map> [](DirectedEdge const &e) { return e.src; }) .l_to_r(); - std::map> result = + std::map> result = map_values(by_src, - [](nonempty_set const &s) -> std::unordered_set { - return s.unwrap_as_unordered_set(); + [](nonempty_set const &s) -> std::set { + return s.unwrap_as_set(); }); for (Node const &n : ns) { result[n]; } - return unordered_map_from_map(result); + return result; } -std::unordered_set get_outgoing_edges(DiGraphView const &g, +std::set get_outgoing_edges(DiGraphView const &g, Node const &n) { return g.query_edges(DirectedEdgeQuery{ query_set::match_single_value(n), diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc index 04bd1e6e02..2b0b43c9a8 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_set get_post_dominators(DiGraphView const &g, +std::set get_post_dominators(DiGraphView const &g, Node const &n) { return get_post_dominators_map(g).at(n); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc index 6ec0cc6ac5..c48e69cdbd 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_post_dominators_map(DiGraphView const &g) { return get_dominators_map(flipped(g)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc index 1f9cb4e006..d7f99477d2 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc @@ -6,19 +6,19 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_predecessors(DiGraphView const &g) { return get_predecessors(g, get_nodes(g)); } -std::unordered_set get_predecessors(DiGraphView const &g, Node const &n) { - return get_predecessors(g, std::unordered_set{n}).at(n); +std::set get_predecessors(DiGraphView const &g, Node const &n) { + return get_predecessors(g, std::set{n}).at(n); } -std::unordered_map> - get_predecessors(DiGraphView const &g, std::unordered_set const &ns) { +std::map> + get_predecessors(DiGraphView const &g, std::set const &ns) { return map_values(get_incoming_edges(g, ns), - [](std::unordered_set const &es) { + [](std::set const &es) { return transform( es, [](DirectedEdge const &e) { return e.src; }); }); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc index 6314a339c9..9dce4f8472 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc @@ -4,9 +4,9 @@ namespace FlexFlow { -std::unordered_set get_strict_dominators(DiGraphView const &g, +std::set get_strict_dominators(DiGraphView const &g, Node const &n) { - std::unordered_set result = get_dominators(g, {n}); + std::set result = get_dominators(g, {n}); result.erase(n); return result; } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc index d4fd60c898..07acbd274e 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc @@ -4,11 +4,11 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_strict_dominators_map(DiGraphView const &g) { return transform(get_dominators_map(g), - [](Node const &n, std::unordered_set const &doms) { - std::unordered_set result = doms; + [](Node const &n, std::set const &doms) { + std::set result = doms; result.erase(n); return std::make_pair(n, result); }); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc index 3067e25d2a..6889fc10f6 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc @@ -5,9 +5,9 @@ namespace FlexFlow { -std::unordered_set get_subgraph_outgoing_edges( - DiGraphView const &g, std::unordered_set const &subgraph_nodes) { - std::unordered_set external_nodes = +std::set get_subgraph_outgoing_edges( + DiGraphView const &g, std::set const &subgraph_nodes) { + std::set external_nodes = set_minus(get_nodes(g), subgraph_nodes); DirectedEdgeQuery query = DirectedEdgeQuery{ query_set::match_values_in(set_of(subgraph_nodes)), diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc index e860fb11b1..ac525353c0 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc @@ -3,10 +3,10 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_successors(DiGraphView const &g, - std::unordered_set const &subgraph_nodes) { - std::unordered_set successors = + std::set const &subgraph_nodes) { + std::set successors = transform(get_subgraph_outgoing_edges(g, subgraph_nodes), [](DirectedEdge const &e) { return e.dst; }); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc index 9c7a027b6e..8864e27417 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc @@ -4,17 +4,17 @@ namespace FlexFlow { -std::unordered_map> +std::map> get_successors(DiGraphView const &g) { return get_predecessors(flipped(g)); } -std::unordered_set get_successors(DiGraphView const &g, Node const &n) { +std::set get_successors(DiGraphView const &g, Node const &n) { return get_predecessors(flipped(g), n); } -std::unordered_map> - get_successors(DiGraphView const &g, std::unordered_set const &ns) { +std::map> + get_successors(DiGraphView const &g, std::set const &ns) { return get_predecessors(flipped(g), ns); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc index 688c369e36..e83e13ae75 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_set get_terminal_nodes(DiGraphView const &g) { +std::set get_terminal_nodes(DiGraphView const &g) { return get_initial_nodes(flipped(g)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering.cc index 58ae8de110..42aa38aa27 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering.cc @@ -14,8 +14,8 @@ static std::vector get_unchecked_topological_ordering(DiGraphView const &g) { auto dfs_view = unchecked_dfs(g, get_initial_nodes(g)); std::vector order; - std::unordered_set seen; - std::unordered_map> predecessors = + std::set seen; + std::map> predecessors = get_predecessors(g, get_nodes(g)); auto all_predecessors_seen = [&](Node const &n) -> bool { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering_from_starting_node.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering_from_starting_node.cc index 590548af66..a0ddce8480 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering_from_starting_node.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_topological_ordering_from_starting_node.cc @@ -12,7 +12,7 @@ namespace FlexFlow { static std::vector get_unchecked_topological_ordering_from_starting_node( DiGraphView const &g, Node const &starting_node) { - std::unordered_set descendants = get_descendants(g, starting_node); + std::set descendants = get_descendants(g, starting_node); descendants.insert(starting_node); return get_topological_ordering(get_subgraph(g, descendants)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc index 7771cc59b5..48c1e87a45 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_set> +std::set> get_weakly_connected_components(DiGraphView const &g) { return get_connected_components(as_undirected(g)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc b/lib/utils/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc index 7bc612aade..fff9bfb7a7 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc @@ -43,8 +43,8 @@ std::optional return get_component_containing_node_in_tail(cbc_decomposition, n).value(); }; - std::unordered_set initial_nodes = get_initial_nodes(view); - std::unordered_set terminal_nodes = get_terminal_nodes(view); + std::set initial_nodes = get_initial_nodes(view); + std::set terminal_nodes = get_terminal_nodes(view); auto src_for_node = [&](Node const &v) -> Node { if (contains(initial_nodes, v)) { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc b/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc index 66c04ec59c..5c0f14eb76 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/is_acyclic.cc @@ -1,8 +1,8 @@ #include "utils/graph/digraph/algorithms/is_acyclic.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/digraph/algorithms/get_successors.h" #include "utils/graph/node/algorithms.h" -#include +#include namespace FlexFlow { @@ -10,8 +10,8 @@ enum class ExplorationStatus { NOT_EXPLORED, BEING_EXPLORED, FULLY_EXPLORED }; bool is_acyclic(DiGraphView const &g) { - std::unordered_map status = - generate_unordered_map(get_nodes(g), [](Node const &n) { + std::map status = + generate_map(get_nodes(g), [](Node const &n) { return ExplorationStatus::NOT_EXPLORED; }); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc b/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc index 45e937aab2..2081e50cf0 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc @@ -20,7 +20,7 @@ DiGraphView transitive_closure(DiGraphView const &g) { bidict nodes = transform_keys(bidict_from_enumerating(get_nodes(g)), [](nonnegative_int x) { return x.unwrap_nonnegative(); }); - std::unordered_set edges = get_edges(g); + std::set edges = get_edges(g); int num_nodes = nodes.size(); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc b/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc index b7c304dcd1..f32de6c469 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc @@ -14,15 +14,15 @@ namespace FlexFlow { DirectedEdgeMaskView::DirectedEdgeMaskView( - DiGraphView const &g, std::unordered_set const &edge_mask) + DiGraphView const &g, std::set const &edge_mask) : g(g), edge_mask(edge_mask) {} -std::unordered_set +std::set DirectedEdgeMaskView::query_edges(DirectedEdgeQuery const &q) const { return set_intersection(g.query_edges(q), this->edge_mask); } -std::unordered_set +std::set DirectedEdgeMaskView::query_nodes(NodeQuery const &q) const { return g.query_nodes(q); } @@ -73,7 +73,7 @@ DiGraph transitive_reduction(DiGraphView const &g) { DiGraph result = materialize_digraph_view(g); // compute transitive reduction // see https://stackoverflow.com/a/6702198 - std::unordered_set edge_mask = get_edges(g); + std::set edge_mask = get_edges(g); for (int j = 0; j < num_nodes; j++) { for (int i = 0; i < num_nodes; i++) { if (has_edge(i, j)) { diff --git a/lib/utils/src/utils/graph/digraph/digraph.cc b/lib/utils/src/utils/graph/digraph/digraph.cc index 24015dc1f3..41070d73b5 100644 --- a/lib/utils/src/utils/graph/digraph/digraph.cc +++ b/lib/utils/src/utils/graph/digraph/digraph.cc @@ -22,11 +22,11 @@ void DiGraph::remove_edge(DirectedEdge const &e) { return this->get_ptr().remove_edge(e); } -std::unordered_set DiGraph::query_nodes(NodeQuery const &q) const { +std::set DiGraph::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } -std::unordered_set +std::set DiGraph::query_edges(DirectedEdgeQuery const &q) const { return this->get_ptr().query_edges(q); } diff --git a/lib/utils/src/utils/graph/digraph/digraph_view.cc b/lib/utils/src/utils/graph/digraph/digraph_view.cc index fb6de481d6..53cf868514 100644 --- a/lib/utils/src/utils/graph/digraph/digraph_view.cc +++ b/lib/utils/src/utils/graph/digraph/digraph_view.cc @@ -2,11 +2,11 @@ namespace FlexFlow { -std::unordered_set DiGraphView::query_nodes(NodeQuery const &q) const { +std::set DiGraphView::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } -std::unordered_set +std::set DiGraphView::query_edges(EdgeQuery const &query) const { return get_ptr().query_edges(query); } diff --git a/lib/utils/src/utils/graph/digraph/directed_edge_query.cc b/lib/utils/src/utils/graph/digraph/directed_edge_query.cc index b7aabb14be..33a2e04eea 100644 --- a/lib/utils/src/utils/graph/digraph/directed_edge_query.cc +++ b/lib/utils/src/utils/graph/digraph/directed_edge_query.cc @@ -13,7 +13,7 @@ bool matches_edge(DirectedEdgeQuery const &q, DirectedEdge const &e) { DirectedEdgeQuery query_intersection(DirectedEdgeQuery const &lhs, DirectedEdgeQuery const &rhs) { - std::unordered_set result_srcs; + std::set result_srcs; if (is_matchall(lhs.srcs) && !is_matchall(rhs.srcs)) { result_srcs = allowed_values(rhs.srcs); } else if (!is_matchall(lhs.srcs) && is_matchall(rhs.srcs)) { @@ -22,7 +22,7 @@ DirectedEdgeQuery query_intersection(DirectedEdgeQuery const &lhs, result_srcs = allowed_values(query_intersection(lhs.srcs, rhs.srcs)); } - std::unordered_set result_dsts; + std::set result_dsts; if (is_matchall(lhs.dsts) && !is_matchall(rhs.dsts)) { result_dsts = allowed_values(rhs.dsts); } else if (!is_matchall(lhs.dsts) && is_matchall(rhs.dsts)) { diff --git a/lib/utils/src/utils/graph/instances/adjacency_digraph.cc b/lib/utils/src/utils/graph/instances/adjacency_digraph.cc index 16590ec8c8..dd4e6672d1 100644 --- a/lib/utils/src/utils/graph/instances/adjacency_digraph.cc +++ b/lib/utils/src/utils/graph/instances/adjacency_digraph.cc @@ -8,7 +8,7 @@ AdjacencyDiGraph::AdjacencyDiGraph() {} AdjacencyDiGraph::AdjacencyDiGraph( NodeSource const &node_source, - std::unordered_map> const &adjacency) + std::map> const &adjacency) : node_source(node_source), adjacency(adjacency) {} AdjacencyDiGraph *AdjacencyDiGraph::clone() const { @@ -41,9 +41,9 @@ void AdjacencyDiGraph::remove_edge(DirectedEdge const &e) { this->adjacency.at(e.src).erase(e.dst); } -std::unordered_set +std::set AdjacencyDiGraph::query_edges(DirectedEdgeQuery const &query) const { - std::unordered_set result; + std::set result; for (auto const &src_kv : query_keys(query.srcs, this->adjacency)) { for (auto const &dst : apply_query(query.dsts, src_kv.second)) { result.insert(DirectedEdge{src_kv.first, dst}); @@ -52,7 +52,7 @@ std::unordered_set return result; } -std::unordered_set +std::set AdjacencyDiGraph::query_nodes(NodeQuery const &query) const { return apply_query(query.nodes, keys(this->adjacency)); } diff --git a/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc b/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc index 903ba0c589..faf4fea868 100644 --- a/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc +++ b/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc @@ -1,12 +1,12 @@ #include "utils/graph/instances/adjacency_multidigraph.h" #include "utils/containers/contains_key.h" #include "utils/containers/extend.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/values.h" #include "utils/graph/multidigraph/algorithms/get_edges.h" #include "utils/graph/node/algorithms.h" -#include "utils/hash/unordered_set.h" -#include "utils/containers/unordered_keys.h" +#include "utils/hash/set.h" +#include "utils/containers/keys.h" namespace FlexFlow { @@ -15,20 +15,20 @@ AdjacencyMultiDiGraph::AdjacencyMultiDiGraph() {} AdjacencyMultiDiGraph::AdjacencyMultiDiGraph( NodeSource const &node_source, MultiDiEdgeSource const &edge_source, - std::unordered_map< + std::map< Node, - std::unordered_map>> const + std::map>> const &adjacency, - std::unordered_map> const &edge_nodes) + std::map> const &edge_nodes) : node_source(node_source), edge_source(edge_source), adjacency(adjacency), edge_nodes(edge_nodes) {} Node AdjacencyMultiDiGraph::add_node() { Node new_node = this->node_source.new_node(); - std::unordered_set all_nodes = - set_union(unordered_keys(this->adjacency), {new_node}); - this->adjacency[new_node] = generate_unordered_map(all_nodes, [](Node const &) { - return std::unordered_set{}; + std::set all_nodes = + set_union(keys(this->adjacency), {new_node}); + this->adjacency[new_node] = generate_map(all_nodes, [](Node const &) { + return std::set{}; }); for (Node const &n : all_nodes) { @@ -49,9 +49,9 @@ MultiDiEdge AdjacencyMultiDiGraph::add_edge(Node const &src, Node const &dst) { void AdjacencyMultiDiGraph::remove_node(Node const &n) { assert(contains_key(this->adjacency, n)); - std::unordered_set outgoing = + std::set outgoing = set_union(values(this->adjacency.at(n))); - std::unordered_set incoming; + std::set incoming; for (auto const &[k, v] : this->adjacency) { if (k != n) { extend(incoming, v.at(n)); @@ -75,17 +75,17 @@ void AdjacencyMultiDiGraph::remove_edge(MultiDiEdge const &e) { this->adjacency.at(src).at(dst).erase(e); } -std::unordered_set +std::set AdjacencyMultiDiGraph::query_nodes(NodeQuery const &q) const { - return apply_query(q.nodes, unordered_keys(this->adjacency)); + return apply_query(q.nodes, keys(this->adjacency)); } -std::unordered_set +std::set AdjacencyMultiDiGraph::query_edges(MultiDiEdgeQuery const &q) const { - std::unordered_set result; + std::set result; - std::unordered_set srcs = apply_query(q.srcs, unordered_keys(this->adjacency)); - std::unordered_set dsts = apply_query(q.dsts, unordered_keys(this->adjacency)); + std::set srcs = apply_query(q.srcs, keys(this->adjacency)); + std::set dsts = apply_query(q.dsts, keys(this->adjacency)); for (Node const &src : srcs) { for (Node const &dst : dsts) { extend(result, this->adjacency.at(src).at(dst)); @@ -105,12 +105,12 @@ Node AdjacencyMultiDiGraph::get_multidiedge_dst(MultiDiEdge const &e) const { void AdjacencyMultiDiGraph::inplace_materialize_from( MultiDiGraphView const &g) { - std::unordered_set nodes = get_nodes(g); - std::unordered_set edges = get_edges(g); + std::set nodes = get_nodes(g); + std::set edges = get_edges(g); - this->adjacency = generate_unordered_map(nodes, [&](Node const &) { - return generate_unordered_map( - nodes, [&](Node const &) { return std::unordered_set{}; }); + this->adjacency = generate_map(nodes, [&](Node const &) { + return generate_map( + nodes, [&](Node const &) { return std::set{}; }); }); this->edge_nodes.clear(); diff --git a/lib/utils/src/utils/graph/instances/hashmap_undirected_graph.cc b/lib/utils/src/utils/graph/instances/hashmap_undirected_graph.cc index 6713fafe41..9407603b5b 100644 --- a/lib/utils/src/utils/graph/instances/hashmap_undirected_graph.cc +++ b/lib/utils/src/utils/graph/instances/hashmap_undirected_graph.cc @@ -43,15 +43,15 @@ void HashmapUndirectedGraph::add_edge(UndirectedEdge const &e) { } void HashmapUndirectedGraph::remove_edge(UndirectedEdge const &e) { - std::unordered_set &max_map = this->adjacency.at(e.endpoints.max()); + std::set &max_map = this->adjacency.at(e.endpoints.max()); max_map.erase(e.endpoints.min()); - std::unordered_set &min_map = this->adjacency.at(e.endpoints.min()); + std::set &min_map = this->adjacency.at(e.endpoints.min()); min_map.erase(e.endpoints.max()); } -std::unordered_set HashmapUndirectedGraph::query_edges( +std::set HashmapUndirectedGraph::query_edges( UndirectedEdgeQuery const &query) const { - std::unordered_set result; + std::set result; for (auto const &src_kv : query_keys(query.nodes, this->adjacency)) { for (auto const &dst : apply_query(query.nodes, src_kv.second)) { result.insert(make_undirected_edge(src_kv.first, dst)); @@ -60,7 +60,7 @@ std::unordered_set HashmapUndirectedGraph::query_edges( return result; } -std::unordered_set +std::set HashmapUndirectedGraph::query_nodes(NodeQuery const &query) const { return apply_query(query.nodes, keys(this->adjacency)); } diff --git a/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc b/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc index a5a1fb82bf..9da333ced0 100644 --- a/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc @@ -3,7 +3,7 @@ #include "utils/containers/enumerate_vector.h" #include "utils/containers/extend.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/graph/dataflow_graph/algorithms.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge.h" @@ -16,10 +16,10 @@ UnorderedSetDataflowGraph::UnorderedSetDataflowGraph() {} UnorderedSetDataflowGraph::UnorderedSetDataflowGraph( NodeSource const &node_source, DataflowGraphInputSource const &graph_input_source, - std::unordered_set const &nodes, - std::unordered_set const &edges, - std::unordered_set const &outputs, - std::unordered_set const &graph_inputs) + std::set const &nodes, + std::set const &edges, + std::set const &outputs, + std::set const &graph_inputs) : node_source(node_source), graph_input_source(graph_input_source), nodes(nodes), edges(edges), outputs(outputs), graph_inputs(graph_inputs) { } @@ -54,26 +54,26 @@ DataflowGraphInput UnorderedSetDataflowGraph::add_input() { return new_graph_input; } -std::unordered_set +std::set UnorderedSetDataflowGraph::query_nodes(NodeQuery const &q) const { return apply_query(q.nodes, this->nodes); } -std::unordered_set UnorderedSetDataflowGraph::query_edges( +std::set UnorderedSetDataflowGraph::query_edges( OpenDataflowEdgeQuery const &q) const { return filter(this->edges, [&](OpenDataflowEdge const &e) { return open_dataflow_edge_query_includes(q, e); }); } -std::unordered_set UnorderedSetDataflowGraph::query_outputs( +std::set UnorderedSetDataflowGraph::query_outputs( DataflowOutputQuery const &q) const { return filter(this->outputs, [&](DataflowOutput const &o) { return includes(q.nodes, o.node) && includes(q.output_idxs, o.idx); }); } -std::unordered_set +std::set UnorderedSetDataflowGraph::get_inputs() const { return this->graph_inputs; } @@ -92,7 +92,7 @@ void UnorderedSetDataflowGraph::add_node_unsafe( std::vector const &inputs, std::vector const &outputs) { assert(!contains(this->nodes, node)); - assert(are_disjoint(this->outputs, unordered_set_of(outputs))); + assert(are_disjoint(this->outputs, set_of(outputs))); this->nodes.insert(node); @@ -106,9 +106,9 @@ void UnorderedSetDataflowGraph::add_node_unsafe( void UnorderedSetDataflowGraph::inplace_materialize_from( DataflowGraphView const &view) { - std::unordered_set nodes = get_nodes(view); - std::unordered_set edges = get_edges(view); - std::unordered_set outputs = get_all_dataflow_outputs(view); + std::set nodes = get_nodes(view); + std::set edges = get_edges(view); + std::set outputs = get_all_dataflow_outputs(view); this->nodes = nodes; this->edges = transform( diff --git a/lib/utils/src/utils/graph/instances/unordered_set_undirected_graph.cc b/lib/utils/src/utils/graph/instances/unordered_set_undirected_graph.cc index cb44f4636d..3be8122cad 100644 --- a/lib/utils/src/utils/graph/instances/unordered_set_undirected_graph.cc +++ b/lib/utils/src/utils/graph/instances/unordered_set_undirected_graph.cc @@ -8,8 +8,8 @@ UnorderedSetUndirectedGraph::UnorderedSetUndirectedGraph() {} UnorderedSetUndirectedGraph::UnorderedSetUndirectedGraph( NodeSource const &node_source, - std::unordered_set const &nodes, - std::unordered_set const &edges) + std::set const &nodes, + std::set const &edges) : node_source(node_source), nodes(nodes), edges(edges) {} Node UnorderedSetUndirectedGraph::add_node() { @@ -36,12 +36,12 @@ void UnorderedSetUndirectedGraph::remove_edge(UndirectedEdge const &e) { this->edges.erase(e); } -std::unordered_set +std::set UnorderedSetUndirectedGraph::query_nodes(NodeQuery const &q) const { return apply_node_query(q, this->nodes); } -std::unordered_set UnorderedSetUndirectedGraph::query_edges( +std::set UnorderedSetUndirectedGraph::query_edges( UndirectedEdgeQuery const &q) const { return filter(this->edges, [&](UndirectedEdge const &e) { return matches_edge(q, e); }); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc index 1e55bf407e..62268cc868 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc @@ -8,6 +8,6 @@ using SlotName = ordered_value_type<0>; template DataflowGraphData dataflow_graph_data_from_kwarg_dataflow_graph_data( KwargDataflowGraphData const &, std::function< - std::vector(std::unordered_set const &)> const &); + std::vector(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc index bdbb6c1d6a..2d10dc66fe 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc @@ -8,6 +8,6 @@ using SlotName = ordered_value_type<0>; template DataflowGraphView dataflow_graph_from_kwarg_dataflow_graph( KwargDataflowGraphView const &, std::function< - std::vector(std::unordered_set const &)> const &); + std::vector(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.cc index f0aaf62581..f318875f5b 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_all_kwarg_dataflow_edges(KwargDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.cc index 343bc1c228..a520299045 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_inputs.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_all_kwarg_dataflow_inputs(KwargDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.cc index 523f2c2357..116946882c 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_all_kwarg_dataflow_outputs(KwargDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.cc index b01efa80ab..d746649849 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_edges_for_node.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_map> +template std::map> get_incoming_kwarg_dataflow_edges_for_node( KwargDataflowGraphView const &, Node const &); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.cc index cbb3589fde..a4d30a673a 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_kwarg_dataflow_outputs_for_node.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_map> +template std::map> get_incoming_kwarg_dataflow_outputs_for_node( KwargDataflowGraphView const &, Node const &); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.cc index 816e529be4..4922f0ca73 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set +template std::set get_incoming_slots_for_node(KwargDataflowGraphView const &, Node); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.cc index e42913922c..8885d8dbf4 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_edges_from_node_to_node.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_kwarg_dataflow_edges_from_node_to_node( KwargDataflowGraphView const &, Node const &, Node const &); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc index a17f751e1e..cd9177d1f0 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc @@ -7,6 +7,6 @@ using SlotName = ordered_value_type<0>; template KwargDataflowGraphView get_kwarg_dataflow_graph_subgraph(KwargDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc index e5dfa38d7f..d9a278099d 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc @@ -5,9 +5,9 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_kwarg_dataflow_subgraph_incoming_edges( KwargDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc index 4653e7339e..3333408b2a 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc @@ -5,9 +5,9 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_kwarg_dataflow_subgraph_outgoing_edges( KwargDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc index b1d2988223..58ef205766 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_value_uses.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_kwarg_dataflow_value_uses(KwargDataflowGraphView const &, KwargDataflowOutput const &); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.cc index 3b1454912e..cc83512d3e 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_map> +template std::map> get_outgoing_kwarg_dataflow_outputs_for_node( KwargDataflowGraphView const &, Node const &); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.cc index 98c2e3895e..1cc443bfe9 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set +template std::set get_outgoing_slots_for_node(KwargDataflowGraphView const &, Node); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc index b9585b562a..06c2a96bf4 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc @@ -12,6 +12,6 @@ template std::string kwarg_dataflow_graph_as_dot( &, std::function const &, std::function< - std::vector(std::unordered_set const &)> const &); + std::vector(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.cc index 1ebf9f9861..de1aec18dd 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_transitive_reduced_kwarg_dataflow_edges_across_split( TransitiveReducedKwargDataflowGraphView const &, BinarySeriesSplit const &); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.cc index ae287be1b6..ae431e67c8 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_outputs_across_split.cc @@ -5,7 +5,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; -template std::unordered_set> +template std::set> get_transitive_reduced_kwarg_dataflow_outputs_across_split( TransitiveReducedKwargDataflowGraphView const &, BinarySeriesSplit const &); diff --git a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc index dfeeab0aa0..f25239d487 100644 --- a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc +++ b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc @@ -8,7 +8,7 @@ using NodeLabel = value_type<0>; using OutputLabel = value_type<1>; using SlotName = ordered_value_type<2>; -template std::unordered_map +template std::map get_labelled_kwarg_dataflow_graph_node_label_map( LabelledKwargDataflowGraphView const &); diff --git a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc index bd287b5342..306ae3fe02 100644 --- a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc +++ b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc @@ -8,7 +8,7 @@ using NodeLabel = value_type<0>; using OutputLabel = value_type<1>; using SlotName = ordered_value_type<2>; -template std::unordered_map, OutputLabel> +template std::map, OutputLabel> get_labelled_kwarg_dataflow_graph_output_label_map( LabelledKwargDataflowGraphView const &); diff --git a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc index 03d19cd52e..45eadb84de 100644 --- a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc @@ -12,6 +12,6 @@ template LabelledKwargDataflowGraphView get_labelled_kwarg_dataflow_graph_subgraph( LabelledKwargDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc index 7901cdef2f..e9b3c6aab6 100644 --- a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc +++ b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc @@ -11,7 +11,7 @@ using SlotName = ordered_value_type<2>; template LabelledKwargDataflowGraphView kwarg_dataflow_graph_view_with_labelling( KwargDataflowGraphView const &, - std::unordered_map const &, - std::unordered_map, OutputLabel> const &); + std::map const &, + std::map, OutputLabel> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc index f1b9c13e17..1ef9fb2782 100644 --- a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc +++ b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc @@ -14,6 +14,6 @@ template std::string labelled_kwarg_dataflow_graph_view_as_dot( std::function const &, std::function const &, std::function< - std::vector(std::unordered_set const &)> const &); + std::vector(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc index d4a580eaab..74187baa0b 100644 --- a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc +++ b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc @@ -18,6 +18,6 @@ template std::string labelled_open_kwarg_dataflow_graph_view_as_dot( std::function const &, std::function const &, std::function< - std::vector(std::unordered_set const &)> const &); + std::vector(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc index 91ca81e877..1ac5c9f239 100644 --- a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc +++ b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc @@ -25,8 +25,8 @@ template LabelledOpenKwargDataflowGraphView open_kwarg_dataflow_graph_view_with_labelling( OpenKwargDataflowGraphView const &, - std::unordered_map const &, - std::unordered_map, + std::map const &, + std::map, ValueLabel> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_edge_counts.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_edge_counts.cc index 53497a715d..35f42bcf27 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_edge_counts.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_edge_counts.cc @@ -7,7 +7,7 @@ namespace FlexFlow { -std::unordered_map +std::map get_edge_counts(MultiDiGraphView const &g) { return get_element_counts( transform(vector_of(get_edges(g)), diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_edges.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_edges.cc index 4ad45cef9a..068a2e6e16 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_edges.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_edges.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set get_edges(MultiDiGraphView const &g) { +std::set get_edges(MultiDiGraphView const &g) { return g.query_edges(multidiedge_query_all()); } diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc index ddc3f4f7ae..8d98e1c50d 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc @@ -7,11 +7,10 @@ #include "utils/graph/multidigraph/multidiedge_query.dtg.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/query_set.h" -#include "utils/containers/unordered_map_from_map.h" namespace FlexFlow { -std::unordered_set get_incoming_edges(MultiDiGraphView const &g, +std::set get_incoming_edges(MultiDiGraphView const &g, Node const &n) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::matchall(), @@ -21,28 +20,28 @@ std::unordered_set get_incoming_edges(MultiDiGraphView const &g, return g.query_edges(query); } -std::unordered_map> +std::map> get_incoming_edges(MultiDiGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::matchall(), query_set::match_values_in(set_of(ns)), }; - std::map> result = map_values( + std::map> result = map_values( group_by(g.query_edges(query), [&](MultiDiEdge const &e) { return g.get_multidiedge_dst(e); }) .l_to_r(), [](nonempty_set const &s) - -> std::unordered_set { - return s.unwrap_as_unordered_set(); + -> std::set { + return s.unwrap_as_set(); }); for (Node const &n : ns) { result[n]; } - return unordered_map_from_map(result); + return result; } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc index 466bb903d9..d291b1e458 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.cc @@ -1,13 +1,13 @@ #include "utils/graph/multidigraph/algorithms/get_multidiedge_to_diedge_map.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/multidigraph/algorithms/get_directed_edge.h" #include "utils/graph/multidigraph/algorithms/get_edges.h" namespace FlexFlow { -std::unordered_map +std::map get_multidiedge_to_diedge_map(MultiDiGraphView const &g) { - return generate_unordered_map(get_edges(g), [&](MultiDiEdge const &e) { + return generate_map(get_edges(g), [&](MultiDiEdge const &e) { return get_directed_edge(g, e); }); } diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc index 143e59b3db..c66eed3dca 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc @@ -4,12 +4,11 @@ #include "utils/containers/set_of.h" #include "utils/graph/multidigraph/algorithms/get_edges.h" #include "utils/graph/node/algorithms.h" -#include -#include "utils/containers/unordered_map_from_map.h" +#include namespace FlexFlow { -std::unordered_set get_outgoing_edges(MultiDiGraphView const &g, +std::set get_outgoing_edges(MultiDiGraphView const &g, Node const &n) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::match_single_value(n), @@ -19,28 +18,28 @@ std::unordered_set get_outgoing_edges(MultiDiGraphView const &g, return g.query_edges(query); } -std::unordered_map> +std::map> get_outgoing_edges(MultiDiGraphView const &g, - std::unordered_set const &ns) { + std::set const &ns) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::match_values_in(set_of(ns)), query_set::matchall(), }; - std::map> result = map_values( + std::map> result = map_values( group_by(g.query_edges(query), [&](MultiDiEdge const &e) { return g.get_multidiedge_src(e); }) .l_to_r(), [](nonempty_set const &s) - -> std::unordered_set { - return s.unwrap_as_unordered_set(); + -> std::set { + return s.unwrap_as_set(); }); for (Node const &n : ns) { result[n]; } - return unordered_map_from_map(result); + return result; } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/multidigraph/i_multidigraph_view.cc b/lib/utils/src/utils/graph/multidigraph/i_multidigraph_view.cc index 62096f153d..7983d7c190 100644 --- a/lib/utils/src/utils/graph/multidigraph/i_multidigraph_view.cc +++ b/lib/utils/src/utils/graph/multidigraph/i_multidigraph_view.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set +std::set IMultiDiGraphView::query_edges(DirectedEdgeQuery const &q) const { return transform(this->query_edges(MultiDiEdgeQuery{q.srcs, q.dsts}), [&](MultiDiEdge const &e) { diff --git a/lib/utils/src/utils/graph/multidigraph/multidigraph.cc b/lib/utils/src/utils/graph/multidigraph/multidigraph.cc index 1c2d92982d..240f50405c 100644 --- a/lib/utils/src/utils/graph/multidigraph/multidigraph.cc +++ b/lib/utils/src/utils/graph/multidigraph/multidigraph.cc @@ -18,11 +18,11 @@ void MultiDiGraph::remove_edge(MultiDiEdge const &e) { this->get_interface().remove_edge(e); } -std::unordered_set MultiDiGraph::query_nodes(NodeQuery const &q) const { +std::set MultiDiGraph::query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } -std::unordered_set +std::set MultiDiGraph::query_edges(MultiDiEdgeQuery const &q) const { return this->get_interface().query_edges(q); } diff --git a/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc b/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc index 911a154405..b308293a49 100644 --- a/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc +++ b/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc @@ -2,12 +2,12 @@ namespace FlexFlow { -std::unordered_set +std::set MultiDiGraphView::query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } -std::unordered_set +std::set MultiDiGraphView::query_edges(MultiDiEdgeQuery const &q) const { return this->get_interface().query_edges(q); } diff --git a/lib/utils/src/utils/graph/node/algorithms.cc b/lib/utils/src/utils/graph/node/algorithms.cc index 61a4d9d9af..a1460ba506 100644 --- a/lib/utils/src/utils/graph/node/algorithms.cc +++ b/lib/utils/src/utils/graph/node/algorithms.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set get_nodes(GraphView const &g) { +std::set get_nodes(GraphView const &g) { return g.query_nodes(node_query_all()); } diff --git a/lib/utils/src/utils/graph/node/graph.cc b/lib/utils/src/utils/graph/node/graph.cc index 6c01ae2caf..2a67ae3d84 100644 --- a/lib/utils/src/utils/graph/node/graph.cc +++ b/lib/utils/src/utils/graph/node/graph.cc @@ -14,7 +14,7 @@ void Graph::remove_node_unsafe(Node const &node) { get_ptr().remove_node_unsafe(node); } -std::unordered_set Graph::query_nodes(NodeQuery const &q) const { +std::set Graph::query_nodes(NodeQuery const &q) const { return get_ptr().query_nodes(q); } diff --git a/lib/utils/src/utils/graph/node/graph_view.cc b/lib/utils/src/utils/graph/node/graph_view.cc index 5404e29f23..5252b598ef 100644 --- a/lib/utils/src/utils/graph/node/graph_view.cc +++ b/lib/utils/src/utils/graph/node/graph_view.cc @@ -4,7 +4,7 @@ namespace FlexFlow { GraphView::GraphView(cow_ptr_t ptr) : ptr(ptr) {} -std::unordered_set GraphView::query_nodes(NodeQuery const &g) const { +std::set GraphView::query_nodes(NodeQuery const &g) const { return this->ptr->query_nodes(g); } diff --git a/lib/utils/src/utils/graph/node/node_query.cc b/lib/utils/src/utils/graph/node/node_query.cc index aa24da42ae..79b53cb686 100644 --- a/lib/utils/src/utils/graph/node/node_query.cc +++ b/lib/utils/src/utils/graph/node/node_query.cc @@ -9,7 +9,7 @@ NodeQuery node_query_all() { NodeQuery query_intersection(NodeQuery const &lhs, NodeQuery const &rhs) { - std::unordered_set nodes; + std::set nodes; if (is_matchall(lhs.nodes) && !is_matchall(rhs.nodes)) { nodes = allowed_values(rhs.nodes); @@ -28,8 +28,8 @@ NodeQuery query_union(NodeQuery const &lhs, NodeQuery const &rhs) { NOT_IMPLEMENTED(); } -std::unordered_set apply_node_query(NodeQuery const &query, - std::unordered_set const &ns) { +std::set apply_node_query(NodeQuery const &query, + std::set const &ns) { return apply_query(query.nodes, ns); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphism.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphism.cc index d75a447127..0b34dfcb8c 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphism.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphism.cc @@ -7,7 +7,7 @@ namespace FlexFlow { std::optional find_isomorphism(OpenDataflowGraphView const &src, OpenDataflowGraphView const &dst) { - std::unordered_set all_isomorphisms = + std::set all_isomorphisms = find_isomorphisms(src, dst); if (all_isomorphisms.empty()) { diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc index 5747b834d1..78a160c390 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc @@ -32,33 +32,33 @@ static std::optional bidict const &unused_graph_inputs_mapping) { { - std::unordered_set already_mapped_src_nodes = + std::set already_mapped_src_nodes = left_entries(sink_node_mapping); - std::unordered_set src_g_sink_nodes = get_terminal_nodes(src_g); - assert(already_mapped_src_nodes == src_g_sink_nodes); + std::set src_g_sink_nodes = set_of(get_terminal_nodes(src_g)); + ASSERT(already_mapped_src_nodes == src_g_sink_nodes); } { - std::unordered_set already_mapped_dst_nodes = + std::set already_mapped_dst_nodes = right_entries(sink_node_mapping); - std::unordered_set dst_g_sink_nodes = get_terminal_nodes(dst_g); - assert(already_mapped_dst_nodes == dst_g_sink_nodes); + std::set dst_g_sink_nodes = set_of(get_terminal_nodes(dst_g)); + ASSERT(already_mapped_dst_nodes == dst_g_sink_nodes); } { - std::unordered_set already_mapped_src_inputs = + std::set already_mapped_src_inputs = right_entries(unused_graph_inputs_mapping); - std::unordered_set src_g_unused_inputs = - get_unused_open_dataflow_graph_inputs(src_g); - assert(already_mapped_src_inputs == src_g_unused_inputs); + std::set src_g_unused_inputs = + set_of(get_unused_open_dataflow_graph_inputs(src_g)); + ASSERT(already_mapped_src_inputs == src_g_unused_inputs); } { - std::unordered_set already_mapped_dst_inputs = + std::set already_mapped_dst_inputs = right_entries(unused_graph_inputs_mapping); - std::unordered_set dst_g_unused_inputs = - get_unused_open_dataflow_graph_inputs(dst_g); - assert(already_mapped_dst_inputs == dst_g_unused_inputs); + std::set dst_g_unused_inputs = + set_of(get_unused_open_dataflow_graph_inputs(dst_g)); + ASSERT(already_mapped_dst_inputs == dst_g_unused_inputs); } std::optional result = @@ -141,9 +141,9 @@ static std::optional return; } - assert(get_open_dataflow_edge_dst(src_edge).idx == + ASSERT(get_open_dataflow_edge_dst(src_edge).idx == get_open_dataflow_edge_dst(dst_edge).idx); - assert( + ASSERT( get_open_dataflow_edge_dst(src_edge).node == result->node_mapping.at_r(get_open_dataflow_edge_dst(dst_edge).node)); @@ -196,13 +196,13 @@ static std::optional return result; } -std::unordered_set +std::set find_isomorphisms(OpenDataflowGraphView const &src, OpenDataflowGraphView const &dst) { - std::unordered_set result; + std::set result; std::vector src_sink_nodes = vector_of(get_terminal_nodes(src)); - std::unordered_set dst_sink_nodes = get_terminal_nodes(dst); + std::set dst_sink_nodes = get_terminal_nodes(dst); if (src_sink_nodes.size() != dst_sink_nodes.size()) { return {}; @@ -210,7 +210,7 @@ std::unordered_set std::vector src_unused_graph_inputs = vector_of(get_unused_open_dataflow_graph_inputs(src)); - std::unordered_set dst_unused_graph_inputs = + std::set dst_unused_graph_inputs = get_unused_open_dataflow_graph_inputs(dst); if (src_unused_graph_inputs.size() != dst_unused_graph_inputs.size()) { @@ -235,7 +235,7 @@ std::unordered_set src, dst, sink_node_mapping, unused_graph_inputs_mapping); if (found.has_value()) { - assert(is_isomorphic_under(src, dst, found.value())); + ASSERT(is_isomorphic_under(src, dst, found.value())); result.insert(found.value()); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc index c4b5befcbc..c36bfee6ab 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc @@ -9,22 +9,22 @@ FromOpenDataflowGraphDataView::FromOpenDataflowGraphDataView( OpenDataflowGraphData const &data) : data(data) {} -std::unordered_set +std::set FromOpenDataflowGraphDataView::query_nodes(NodeQuery const &q) const { return apply_node_query(q, this->data.nodes); } -std::unordered_set FromOpenDataflowGraphDataView::query_edges( +std::set FromOpenDataflowGraphDataView::query_edges( OpenDataflowEdgeQuery const &q) const { return apply_open_dataflow_edge_query(q, this->data.edges); } -std::unordered_set FromOpenDataflowGraphDataView::query_outputs( +std::set FromOpenDataflowGraphDataView::query_outputs( DataflowOutputQuery const &q) const { return apply_dataflow_output_query(q, this->data.outputs); } -std::unordered_set +std::set FromOpenDataflowGraphDataView::get_inputs() const { return this->data.inputs; } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_edges.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_edges.cc index 610239feff..28c667ee19 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_edges.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_edges.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set get_edges(OpenDataflowGraphView const &g) { +std::set get_edges(OpenDataflowGraphView const &g) { return g.query_edges(open_dataflow_edge_query_all()); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edge.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edge.cc index ac1aae1168..275d25cf47 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edge.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edge.cc @@ -7,7 +7,7 @@ namespace FlexFlow { OpenDataflowEdge get_incoming_edge(OpenDataflowGraphView const &g, DataflowInput const &i) { OpenDataflowEdgeQuery query = open_dataflow_edge_query_all_incoming_to(i); - std::unordered_set query_result = g.query_edges(query); + std::set query_result = g.query_edges(query); return get_only(query_result); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc index b68aa4b1d8..72de3b1b4d 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc @@ -1,5 +1,5 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/sorted_by.h" #include "utils/containers/transform.h" #include "utils/graph/dataflow_graph/dataflow_edge_query.h" @@ -8,9 +8,9 @@ namespace FlexFlow { -std::unordered_set +std::set get_incoming_edges(OpenDataflowGraphView const &g) { - std::unordered_set raw_edges = + std::set raw_edges = g.query_edges(OpenDataflowEdgeQuery{ dataflow_input_edge_query_all(), dataflow_edge_query_none(), @@ -42,10 +42,10 @@ std::vector get_incoming_edges(OpenDataflowGraphView const &g, }); } -std::unordered_map> +std::map> get_incoming_edges(OpenDataflowGraphView const &g, - std::unordered_set const &ns) { - return generate_unordered_map(ns, + std::set const &ns) { + return generate_map(ns, [&](Node const &n) { return get_incoming_edges(g, n); }); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc index 78c7677de9..00de5dd2f9 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc @@ -2,7 +2,7 @@ namespace FlexFlow { -std::unordered_set +std::set get_open_dataflow_graph_inputs(OpenDataflowGraphView const &g) { return g.get_inputs(); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc index 12795b8f7e..60c8d1c6a5 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc @@ -5,10 +5,10 @@ namespace FlexFlow { -std::unordered_set +std::set get_open_dataflow_value_uses(OpenDataflowGraphView const &view, OpenDataflowValue const &value) { - std::unordered_set edges = + std::set edges = view.query_edges(open_dataflow_edge_query_all_outgoing_from(value)); return transform(edges, get_open_dataflow_edge_dst); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.cc index 0aa1bdb054..77e90fc06b 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.cc @@ -4,11 +4,11 @@ namespace FlexFlow { -std::unordered_set +std::set get_open_dataflow_values(OpenDataflowGraphView const &g) { return set_union( transform( - unordered_set_of(g.get_inputs()), + set_of(g.get_inputs()), [](DataflowGraphInput const &gi) { return OpenDataflowValue{gi}; }), transform(get_all_dataflow_outputs(g), [](DataflowOutput const &o) { return OpenDataflowValue{o}; })); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.cc index 14099e1c64..e51c5459bf 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_source_nodes.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_set get_source_nodes(OpenDataflowGraphView const &g) { +std::set get_source_nodes(OpenDataflowGraphView const &g) { auto is_source_node = [&](Node const &n) { std::vector incoming_edges = get_incoming_edges(g, n); return incoming_edges.empty(); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc index e8989c9ee1..a88dcdb4cb 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc @@ -2,7 +2,7 @@ #include "utils/bidict/generate_bidict.h" #include "utils/containers/enumerate_vector.h" #include "utils/containers/is_subseteq_of.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/graph/dataflow_graph/dataflow_output_query.h" #include "utils/graph/node/algorithms.h" @@ -19,7 +19,7 @@ namespace FlexFlow { OpenDataflowSubgraphResult get_subgraph(OpenDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { bidict full_graph_values_to_subgraph_inputs = get_full_graph_values_to_subgraph_inputs(g, subgraph_nodes); @@ -35,7 +35,7 @@ OpenDataflowSubgraphResult bidict get_full_graph_values_to_subgraph_inputs( OpenDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { DataflowGraphInputSource input_source; return generate_bidict(get_subgraph_inputs(g, subgraph_nodes), [&](OpenDataflowValue const &v) -> DataflowGraphInput { @@ -50,10 +50,10 @@ bidict OpenDataflowGraphData get_subgraph_data(OpenDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes, + std::set const &subgraph_nodes, bidict const &full_graph_values_to_subgraph_inputs) { - std::unordered_set subgraph_input_edges = + std::set subgraph_input_edges = transform(get_subgraph_incoming_edges(g, subgraph_nodes), [&](OpenDataflowEdge const &edge) { return edge.visit( @@ -84,12 +84,12 @@ OpenDataflowGraphData query_set::matchall(), }, }; - std::unordered_set subgraph_interior_edges = + std::set subgraph_interior_edges = g.query_edges(subgraph_interior_edges_query); - std::unordered_set subgraph_inputs = - unordered_set_of(values(full_graph_values_to_subgraph_inputs)); - std::unordered_set subgraph_outputs = + std::set subgraph_inputs = + set_of(values(full_graph_values_to_subgraph_inputs)); + std::set subgraph_outputs = filter(g.query_outputs(dataflow_output_query_all()), [&](DataflowOutput const &o) { return contains(subgraph_nodes, o.node); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.cc index fe2b46cd64..9d0fa06afd 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_incoming_edges.cc @@ -5,10 +5,10 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_incoming_edges(OpenDataflowGraphView const &g, - std::unordered_set const &ns) { - std::unordered_set nodes_not_in_ns = set_minus(get_nodes(g), ns); + std::set const &ns) { + std::set nodes_not_in_ns = set_minus(get_nodes(g), ns); OpenDataflowEdgeQuery query = OpenDataflowEdgeQuery{ DataflowInputEdgeQuery{ diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.cc index 08dda09698..ab7422fa94 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.cc @@ -10,10 +10,10 @@ namespace FlexFlow { -std::unordered_set +std::set get_subgraph_inputs(OpenDataflowGraphView const &g, - std::unordered_set const &subgraph_nodes) { - std::unordered_set relevant_edges; + std::set const &subgraph_nodes) { + std::set relevant_edges; for (std::vector const &incoming : values(get_incoming_edges(g, subgraph_nodes))) { auto comes_from_outside_subgraph = [&](OpenDataflowEdge const &e) -> bool { diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc index 8fbe7ae5bc..284dffe590 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc @@ -4,7 +4,7 @@ namespace FlexFlow { -std::unordered_set +std::set get_unused_open_dataflow_graph_inputs(OpenDataflowGraphView const &g) { return filter( get_open_dataflow_graph_inputs(g), [&](DataflowGraphInput const &i) { diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc b/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc index 59720d843f..5fb5d4ce4e 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc @@ -4,14 +4,14 @@ namespace FlexFlow { -std::unordered_set +std::set IOpenDataflowGraphView::query_edges(DataflowEdgeQuery const &q) const { OpenDataflowEdgeQuery open_query = OpenDataflowEdgeQuery{ dataflow_input_edge_query_none(), q, }; - std::unordered_set open_edges = + std::set open_edges = this->query_edges(open_query); return transform(open_edges, [](OpenDataflowEdge const &e) { diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc index 4882c3e143..7899ad1c36 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc @@ -58,9 +58,9 @@ OpenDataflowEdgeQuery }; } -std::unordered_set apply_open_dataflow_edge_query( +std::set apply_open_dataflow_edge_query( OpenDataflowEdgeQuery const &q, - std::unordered_set const &es) { + std::set const &es) { return filter(es, [&](OpenDataflowEdge const &e) { return open_dataflow_edge_query_includes(q, e); }); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc index 96ad920557..339087c1bd 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc @@ -2,12 +2,12 @@ namespace FlexFlow { -std::unordered_set +std::set OpenDataflowGraphView::get_inputs() const { return this->get_interface().get_inputs(); } -std::unordered_set +std::set OpenDataflowGraphView::query_edges(OpenDataflowEdgeQuery const &q) const { return this->get_interface().query_edges(q); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc b/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc index 171b321c66..be1478d38b 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc @@ -8,11 +8,11 @@ UnorderedSetOpenDataflowGraph::UnorderedSetOpenDataflowGraph() {} UnorderedSetOpenDataflowGraph::UnorderedSetOpenDataflowGraph( NodeSource const &node_source, DataflowGraphInputSource const &input_source, - std::unordered_set const &nodes, - std::unordered_set const &standard_edges, - std::unordered_set const &input_edges, - std::unordered_set const &outputs, - std::unordered_set const &graph_inputs) + std::set const &nodes, + std::set const &standard_edges, + std::set const &input_edges, + std::set const &outputs, + std::set const &graph_inputs) : node_source(node_source), input_source(input_source), nodes(nodes), standard_edges(standard_edges), input_edges(input_edges), outputs(outputs), graph_inputs(graph_inputs) {} @@ -22,21 +22,21 @@ NodeAddedResult UnorderedSetOpenDataflowGraph::add_node( NOT_IMPLEMENTED(); } -std::unordered_set +std::set UnorderedSetOpenDataflowGraph::query_nodes(NodeQuery const &q) const { return apply_query(q.nodes, this->nodes); } -std::unordered_set UnorderedSetOpenDataflowGraph::query_edges( +std::set UnorderedSetOpenDataflowGraph::query_edges( OpenDataflowEdgeQuery const &q) const { - std::unordered_set standard_edges = + std::set standard_edges = filter(this->standard_edges, [&](DataflowEdge const &e) { return includes(q.standard_edge_query.src_nodes, e.src.node) && includes(q.standard_edge_query.dst_nodes, e.dst.node) && includes(q.standard_edge_query.src_idxs, e.src.idx) && includes(q.standard_edge_query.dst_idxs, e.dst.idx); }); - std::unordered_set input_edges = + std::set input_edges = filter(this->input_edges, [&](DataflowInputEdge const &e) { return includes(q.input_edge_query.srcs, e.src) && includes(q.input_edge_query.dst_nodes, e.dst.node) && @@ -50,14 +50,14 @@ std::unordered_set UnorderedSetOpenDataflowGraph::query_edges( })); } -std::unordered_set UnorderedSetOpenDataflowGraph::query_outputs( +std::set UnorderedSetOpenDataflowGraph::query_outputs( DataflowOutputQuery const &q) const { return filter(this->outputs, [&](DataflowOutput const &o) { return includes(q.nodes, o.node) && includes(q.output_idxs, o.idx); }); } -std::unordered_set +std::set UnorderedSetOpenDataflowGraph::get_inputs() const { return this->graph_inputs; } diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc index 3765fd7a91..3e9a67f4e6 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template std::unordered_set> +template std::set> find_isomorphisms_between_open_kwarg_dataflow_graphs( OpenKwargDataflowGraphView const &, OpenKwargDataflowGraphView const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.cc index f5004a7a12..32276a37ba 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -std::unordered_set> +std::set> get_all_kwarg_dataflow_graph_inputs( OpenKwargDataflowGraphView const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc index ddc6fad77b..9527e53104 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc @@ -3,7 +3,7 @@ namespace FlexFlow { template -std::unordered_set> +std::set> get_all_open_kwarg_dataflow_edges( OpenKwargDataflowGraphView const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.cc index 2805d4a3d4..91eca4e843 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template std::unordered_set> +template std::set> get_all_open_kwarg_dataflow_values( OpenKwargDataflowGraphView const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc index 41eea02bd6..3bffb4f9f5 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template std::unordered_map> get_incoming_open_kwarg_dataflow_edges_for_node( OpenKwargDataflowGraphView const &, diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc index ead2334cc5..be7affcded 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; using GraphInputName = ordered_value_type<1>; -template std::unordered_map> get_incoming_open_kwarg_dataflow_values_for_node( OpenKwargDataflowGraphView const &, diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc index bdeea26c47..064eb36c46 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc @@ -9,20 +9,20 @@ using SlotName = ordered_value_type<1>; template OpenKwargDataflowSubgraphResult get_open_kwarg_dataflow_graph_subgraph( OpenKwargDataflowGraphView const &, - std::unordered_set const &, + std::set const &, std::function const &); template bidict, KwargDataflowGraphInput> get_full_kwarg_dataflow_graph_values_to_subgraph_inputs( OpenKwargDataflowGraphView const &, - std::unordered_set const &, + std::set const &, std::function const &); template OpenKwargDataflowGraphData get_open_kwarg_dataflow_subgraph_data( OpenKwargDataflowGraphView const &, - std::unordered_set const &, + std::set const &, bidict, KwargDataflowGraphInput> const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.cc index 1381b83c27..1841ad89c7 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_incoming_edges.cc @@ -6,9 +6,9 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template std::unordered_set> +template std::set> get_open_kwarg_dataflow_subgraph_incoming_edges( OpenKwargDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.cc index 6ca0911aa0..a06d81ff69 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_subgraph_inputs.cc @@ -6,9 +6,9 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -std::unordered_set> +std::set> get_open_kwarg_dataflow_subgraph_inputs( OpenKwargDataflowGraphView const &, - std::unordered_set const &); + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.cc index c265288d60..1e91450d88 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template std::unordered_set> +template std::set> get_open_kwarg_dataflow_value_uses( OpenKwargDataflowGraphView const &, OpenKwargDataflowValue const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.cc index ab26bd0356..12a8c4a259 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_unused_open_kwarg_dataflow_graph_inputs.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template std::unordered_set> +template std::set> get_unused_open_kwarg_dataflow_graph_inputs( OpenKwargDataflowGraphView const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc index 27113566d5..cc2fe65cc2 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc @@ -14,6 +14,6 @@ template std::string open_kwarg_dataflow_graph_as_dot( OpenKwargDataflowValue const &)> const &, std::function const &, std::function< - std::vector(std::unordered_set const &)> const &); + std::vector(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/render_dot.cc b/lib/utils/src/utils/graph/render_dot.cc index d5a1562ac1..6f163f73d8 100644 --- a/lib/utils/src/utils/graph/render_dot.cc +++ b/lib/utils/src/utils/graph/render_dot.cc @@ -22,7 +22,7 @@ std::string escape_dot_string(std::string const &s) { } std::string render_dot_node_attrs( - std::unordered_map const &node_attrs) { + std::map const &node_attrs) { std::ostringstream oss; for (auto const &[k, v] : node_attrs) { oss << fmt::format( @@ -32,7 +32,7 @@ std::string render_dot_node_attrs( } std::string render_node_label( - LabelledDataflowGraphView, + LabelledDataflowGraphView, std::string> const &g, Node const &n) { std::vector n_inputs = get_dataflow_inputs(g, n); @@ -60,13 +60,13 @@ std::string render_node_label( } std::string render_dot( - LabelledDataflowGraphView, + LabelledDataflowGraphView, std::string> const &g) { std::vector lines; lines.push_back("digraph {"); for (Node const &n : get_nodes(g)) { - std::unordered_map node_attrs = g.at(n); + std::map node_attrs = g.at(n); node_attrs.at("label") = render_node_label(g, n); node_attrs["shape"] = "record"; diff --git a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc index 398abb3faf..9a76b940d5 100644 --- a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc +++ b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc @@ -2,7 +2,7 @@ #include "utils/containers/get_only.h" #include "utils/containers/slice.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/vector_of.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_parallel_split.dtg.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.dtg.h" diff --git a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc index c11968e8b9..cb0505a6a9 100644 --- a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc +++ b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc @@ -66,7 +66,7 @@ bool is_binary_sp_tree_right_associative( generic_impl_for_binary_sp_tree()); } -std::unordered_multiset +std::multiset get_leaves(BinarySPDecompositionTree const &tree) { return get_leaves(tree, generic_impl_for_binary_sp_tree()); } diff --git a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.cc b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.cc index 07e2c3e3e3..9b0e98f23b 100644 --- a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.cc +++ b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/find_paths_to_leaf.cc @@ -8,7 +8,7 @@ using Series = value_type<1>; using Parallel = value_type<2>; using Leaf = value_type<3>; -template std::unordered_set find_paths_to_leaf( +template std::set find_paths_to_leaf( Tree const &, GenericBinarySPDecompositionTreeImplementation; using Parallel = value_type<2>; using Leaf = value_type<3>; -template std::unordered_set get_all_leaf_paths( +template std::set get_all_leaf_paths( Tree const &tree, GenericBinarySPDecompositionTreeImplementation; -using Series = value_type<1>; -using Parallel = value_type<2>; -using Leaf = value_type<3>; +using Series = ordered_value_type<1>; +using Parallel = ordered_value_type<2>; +using Leaf = ordered_value_type<3>; -template std::unordered_multiset +template std::multiset get_leaves(Tree const &, GenericBinarySPDecompositionTreeImplementation; using Parallel = value_type<2>; using Leaf = value_type<3>; -template std::unordered_map get_path_to_leaf_map( +template std::map get_path_to_leaf_map( Tree const &, GenericBinarySPDecompositionTreeImplementation parallel_extend(DiGraph &g, +std::map parallel_extend(DiGraph &g, DiGraphView const &ext) { - std::unordered_map node_map; + std::map node_map; for (Node const &node : get_nodes(ext)) { node_map.emplace(node, g.add_node()); } @@ -26,10 +26,10 @@ std::unordered_map parallel_extend(DiGraph &g, return node_map; } -std::unordered_map serial_extend(DiGraph &g, +std::map serial_extend(DiGraph &g, DiGraphView const &ext) { - std::unordered_set original_sinks = get_terminal_nodes(g); - std::unordered_map node_map = parallel_extend(g, ext); + std::set original_sinks = get_terminal_nodes(g); + std::map node_map = parallel_extend(g, ext); for (Node const &node1 : original_sinks) { for (Node const &node2 : get_initial_nodes(ext)) { g.add_edge(DirectedEdge{node1, node_map.at(node2)}); @@ -58,7 +58,7 @@ DiGraph series_composition(std::vector const &graphs) { return g; } -// TODO(@pietro): should be std::unordered_set, but DiGraphs are +// TODO(@pietro): should be std::set, but DiGraphs are // currently non-hashable DiGraph parallel_composition(std::vector const &graphs) { DiGraph g = DiGraph::create(); diff --git a/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc b/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc index f0df7b6391..baffa96c53 100644 --- a/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc +++ b/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc @@ -4,35 +4,35 @@ #include "utils/containers/get_only.h" #include "utils/containers/set_union.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/variant.h" #include namespace FlexFlow { -std::unordered_set get_ancestors(SeriesParallelDecomposition const &sp, +std::set get_ancestors(SeriesParallelDecomposition const &sp, Node const &node); -static std::unordered_set get_ancestors(Node const &, Node const &node) { +static std::set get_ancestors(Node const &, Node const &node) { return {}; } -static std::unordered_set get_ancestors(SeriesSplit const &serial, +static std::set get_ancestors(SeriesSplit const &serial, Node const &node) { - std::unordered_set ancestors{}; + std::set ancestors{}; for (std::variant const &child : serial.children) { SeriesParallelDecomposition child_sp = widen(child); if (contains(get_nodes(child_sp), node)) { return set_union(ancestors, get_ancestors(child_sp, node)); } - ancestors = set_union(ancestors, unordered_set_of(get_nodes(child_sp))); + ancestors = set_union(ancestors, set_of(get_nodes(child_sp))); } PANIC("Node not found in SeriesSplit"); } -static std::unordered_set get_ancestors(ParallelSplit const ¶llel, +static std::set get_ancestors(ParallelSplit const ¶llel, Node const &node) { SeriesParallelDecomposition branch = get_only(filter(transform(parallel.get_children(), @@ -45,10 +45,10 @@ static std::unordered_set get_ancestors(ParallelSplit const ¶llel, return get_ancestors(branch, node); } -std::unordered_set get_ancestors(SeriesParallelDecomposition const &sp, +std::set get_ancestors(SeriesParallelDecomposition const &sp, Node const &node) { assert(contains(get_nodes(sp), node)); - return sp.visit>( + return sp.visit>( [&](auto const &t) { return get_ancestors(t, node); }); } diff --git a/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc index 33bdd74787..5391fafd56 100644 --- a/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc @@ -2,7 +2,7 @@ #include "utils/containers/get_only.h" #include "utils/containers/map_values.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" #include "utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.h" #include "utils/graph/digraph/algorithms/transitive_reduction.h" #include "utils/graph/instances/adjacency_multidigraph.h" @@ -37,10 +37,10 @@ std::optional MultiDiGraph ttsp = MultiDiGraph::materialize_copy_of( inverse_line_graph_result.graph); - std::unordered_map + std::map ttsp_edge_to_sp_tree = map_values( inverse_line_graph_result.inverse_edge_to_line_node_bidict - .as_unordered_map(), + .as_map(), [](Node const &n) { return SeriesParallelDecomposition{n}; }); auto perform_extended_parallel_reduction = @@ -49,7 +49,7 @@ std::optional apply_extended_parallel_reduction(ttsp, parallel_reduction); SeriesParallelDecomposition new_tree = parallel_composition(transform( - unordered_multiset_of(parallel_reduction.edges), + multiset_of(parallel_reduction.edges), [&](MultiDiEdge const &e) { return ttsp_edge_to_sp_tree.at(e); })); for (MultiDiEdge const &e : parallel_reduction.edges) { @@ -81,7 +81,7 @@ std::optional while (true) { bool reduction_has_happened = false; - std::unordered_set parallel_reductions = + std::set parallel_reductions = find_all_extended_parallel_reductions(ttsp); if (!parallel_reductions.empty()) { @@ -91,7 +91,7 @@ std::optional reduction_has_happened = true; } - std::unordered_set series_reductions = + std::set series_reductions = find_all_extended_series_reductions(ttsp); if (!series_reductions.empty()) { for (ExtendedSeriesReduction series_reduction : series_reductions) { @@ -132,10 +132,10 @@ std::optional MultiDiGraph ttsp = MultiDiGraph::materialize_copy_of( inverse_line_graph_result.graph); - std::unordered_map + std::map ttsp_edge_to_sp_tree = map_values( inverse_line_graph_result.inverse_edge_to_line_node_bidict - .as_unordered_map(), + .as_map(), [](Node const &n) { return BinarySPDecompositionTree{n}; }); while (true) { diff --git a/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc index b5a07a5a4d..624f06d760 100644 --- a/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc @@ -11,7 +11,7 @@ #include "utils/graph/series_parallel/series_split.dtg.h" #include "utils/overload.h" #include "utils/variant.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -36,16 +36,16 @@ NonNormalSPDecomposition non_normal_series_composition( } NonNormalSPDecomposition non_normal_parallel_composition( - std::unordered_multiset const &sp_compositions) { + std::multiset const &sp_compositions) { - std::unordered_multiset< + std::multiset< std::variant<::FlexFlow::NonNormalSeriesSplit, ::FlexFlow::Node>> composition{}; for (NonNormalSPDecomposition const &sp_comp : sp_compositions) { if (sp_comp.has()) { composition = multiset_union( - composition, unordered_multiset_of(sp_comp.get().get_children())); + composition, multiset_of(sp_comp.get().get_children())); } else if (sp_comp.has()) { composition.insert(sp_comp.get()); } else { @@ -72,7 +72,7 @@ static NonNormalSeriesSplit as_non_normal(SeriesSplit const &s) { static NonNormalParallelSplit as_non_normal(ParallelSplit const &p) { return non_normal_parallel_composition( - unordered_multiset_of(transform(p.get_children(), + multiset_of(transform(p.get_children(), [](std::variant const &child) { return as_non_normal( widen(child)); diff --git a/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc index 5eda579f81..25a8ca1cfa 100644 --- a/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc @@ -6,7 +6,7 @@ #include "utils/graph/series_parallel/non_normal_sp_decomposition.h" #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/variant.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -55,7 +55,7 @@ static SeriesParallelDecomposition if (normalized_children.size() == 1) { return get_only(normalized_children); } - return parallel_composition(unordered_multiset_of(normalized_children)); + return parallel_composition(multiset_of(normalized_children)); } SeriesParallelDecomposition diff --git a/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc b/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc index cf03db0e8a..69acec155b 100644 --- a/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc +++ b/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc @@ -4,7 +4,7 @@ #include "utils/containers/get_one_of.h" #include "utils/containers/group_by.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/graph/digraph/directed_edge.dtg.h" #include "utils/graph/multidigraph/algorithms/get_directed_edge.h" @@ -13,9 +13,9 @@ #include "utils/graph/multidigraph/multidigraph.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/series_parallel/extended_parallel_reduction.dtg.h" -#include "utils/hash/unordered_set.h" -#include -#include +#include "utils/hash/set.h" +#include +#include namespace FlexFlow { @@ -27,7 +27,7 @@ ParallelReduction make_parallel_reduction(MultiDiEdge const &e1, std::optional find_parallel_reduction(MultiDiGraphView const &g) { - std::unordered_map seen; + std::map seen; for (MultiDiEdge const &edge : get_edges(g)) { DirectedEdge diedge = get_directed_edge(g, edge); if (contains_key(seen, diedge)) { @@ -38,20 +38,20 @@ std::optional return std::nullopt; } -std::unordered_set +std::set find_all_extended_parallel_reductions(MultiDiGraphView const &g) { - std::unordered_map> + std::map> reduction_groups; for (MultiDiEdge const &edge : get_edges(g)) { reduction_groups[get_directed_edge(g, edge)].insert(edge); } - std::unordered_set> reductions = filter( - unordered_set_of(values(reduction_groups)), - [](std::unordered_set const &s) { return s.size() > 1; }); + std::set> reductions = filter( + set_of(values(reduction_groups)), + [](std::set const &s) { return s.size() > 1; }); return transform(reductions, - [&](std::unordered_set const &edges) { + [&](std::set const &edges) { return ExtendedParallelReduction{edges}; }); } diff --git a/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc index 8c9f655fe5..918c13d57d 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc @@ -6,16 +6,16 @@ #include "utils/containers/set_union.h" #include "utils/containers/sum.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/values.h" #include "utils/containers/vector_of.h" #include "utils/exception.h" #include "utils/graph/series_parallel/intermediate_sp_decomposition_tree.h" #include "utils/graph/series_parallel/series_parallel_metrics.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" #include "utils/nonnegative_int/nonnegative_int.h" #include "utils/variant.h" -#include +#include #include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -58,21 +58,21 @@ SeriesParallelDecomposition to_final_ast( internal_to_final_ast(ast)); } -std::unordered_multiset get_nodes(SeriesParallelDecomposition const &sp) { - return sp.visit>( +std::multiset get_nodes(SeriesParallelDecomposition const &sp) { + return sp.visit>( [](auto &&t) { return get_nodes(t); }); } -std::unordered_multiset get_nodes(SeriesSplit const &serial) { +std::multiset get_nodes(SeriesSplit const &serial) { return multiset_union(transform( serial.children, [](std::variant const &child) - -> std::unordered_multiset { + -> std::multiset { return std::visit([](auto &&t) { return get_nodes(t); }, child); })); } -std::unordered_multiset get_nodes(ParallelSplit const ¶llel) { +std::multiset get_nodes(ParallelSplit const ¶llel) { return multiset_union(transform( vector_of(parallel.get_children()), [](std::variant const &child) { @@ -80,7 +80,7 @@ std::unordered_multiset get_nodes(ParallelSplit const ¶llel) { })); } -std::unordered_multiset get_nodes(Node const &node) { +std::multiset get_nodes(Node const &node) { return {node}; } @@ -122,7 +122,7 @@ SeriesParallelDecomposition series_composition( } SeriesParallelDecomposition parallel_composition( - std::unordered_multiset const + std::multiset const &sp_compositions) { ASSERT(sp_compositions.size() > 0, @@ -132,13 +132,13 @@ SeriesParallelDecomposition parallel_composition( return get_only(sp_compositions); } - std::unordered_multiset< + std::multiset< std::variant<::FlexFlow::SeriesSplit, ::FlexFlow::Node>> composition{}; for (SeriesParallelDecomposition const &sp_comp : sp_compositions) { if (sp_comp.has()) { composition = multiset_union(composition, - unordered_multiset_of(sp_comp.get().get_children())); + multiset_of(sp_comp.get().get_children())); } else if (sp_comp.has()) { composition.insert(sp_comp.get()); } else { diff --git a/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc b/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc index 590912e93f..12f5eb2582 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc @@ -4,7 +4,7 @@ #include "utils/containers/transform.h" #include "utils/containers/values.h" #include "utils/containers/vector_of.h" -#include "utils/fmt/unordered_multiset.h" +#include "utils/fmt/multiset.h" #include "utils/fmt/multiset.h" #include "utils/graph/digraph/algorithms/get_edges.h" #include "utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h" @@ -15,61 +15,61 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/nonnegative_int/nonnegative_int.h" #include "utils/variant.h" -#include +#include namespace FlexFlow { -static std::unordered_map +static std::map get_num_occurrences_of_nodes(Node const &node) { return {{node, 1_n}}; } template -static std::unordered_map +static std::map get_num_occurrences_of_nodes_impl(T const &t) { - std::unordered_map counter; + std::map counter; for (Node const &node : get_nodes(t)) { counter.emplace(node, 0_n).first->second += 1_n; } return counter; } -static std::unordered_map +static std::map get_num_occurrences_of_nodes(ParallelSplit const ¶llel) { return get_num_occurrences_of_nodes_impl(parallel); } -static std::unordered_map +static std::map get_num_occurrences_of_nodes(SeriesSplit const &serial) { return get_num_occurrences_of_nodes_impl(serial); } -std::unordered_map +std::map get_num_occurrences_of_nodes(SeriesParallelDecomposition const &sp) { return get_num_occurrences_of_nodes_impl(sp); } float work_cost(SeriesParallelDecomposition const &sp, - std::unordered_map cost_map) { + std::map cost_map) { return sum(transform(get_nodes(sp), [&](Node const &node) { return cost_map.at(node); })); } float work_cost(DiGraphView const &g, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return sum(transform(vector_of(get_nodes(g)), [&](Node const &node) { return cost_map.at(node); })); } static float critical_path_cost(Node const &node, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return cost_map.at(node); } static float critical_path_cost(SeriesSplit const &serial, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return sum(transform( serial.children, [&](std::variant const &child) { return critical_path_cost(widen(child), @@ -79,7 +79,7 @@ static float static float critical_path_cost(ParallelSplit const ¶llel, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return maximum(transform(parallel.get_children(), [&](std::variant const &child) { return critical_path_cost( @@ -89,13 +89,13 @@ static float } float critical_path_cost(SeriesParallelDecomposition const &sp, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return sp.visit( [&](auto const &t) { return critical_path_cost(t, cost_map); }); } float critical_path_cost(DiGraphView const &g, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return maximum( values(get_weighted_longest_path_lengths_from_root(g, cost_map))); } @@ -110,14 +110,14 @@ nonnegative_int num_dependencies(DiGraphView const &g) { float relative_work_increase(DiGraphView const &g, SeriesParallelDecomposition const &sp, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return work_cost(sp, cost_map) / work_cost(g, cost_map); } float relative_critical_path_cost_increase( DiGraphView const &g, SeriesParallelDecomposition const &sp, - std::unordered_map const &cost_map) { + std::map const &cost_map) { return critical_path_cost(sp, cost_map) / critical_path_cost(g, cost_map); } diff --git a/lib/utils/src/utils/graph/series_parallel/series_reduction.cc b/lib/utils/src/utils/graph/series_parallel/series_reduction.cc index 459e61be71..c8aa6db3c8 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_reduction.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_reduction.cc @@ -4,7 +4,7 @@ #include "utils/containers/get_only.h" #include "utils/containers/require_same.h" #include "utils/containers/slice.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/graph/digraph/algorithms/get_predecessors.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" @@ -16,8 +16,8 @@ #include "utils/graph/multidigraph/multidigraph_view.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/series_parallel/extended_series_reduction.dtg.h" -#include "utils/hash/unordered_set.h" -#include +#include "utils/hash/set.h" +#include namespace FlexFlow { @@ -51,14 +51,14 @@ std::optional return std::nullopt; } -std::unordered_set +std::set find_all_extended_series_reductions(MultiDiGraphView const &g) { auto incoming_edges_map = get_incoming_edges(g, get_nodes(g)); auto outgoing_edges_map = get_outgoing_edges(g, get_nodes(g)); - std::unordered_map> strands; - std::unordered_map node_to_head_of_strand; + std::map> strands; + std::map node_to_head_of_strand; for (Node const &n : get_topological_ordering(g)) { if ((incoming_edges_map.at(n).size() == 1) && @@ -81,7 +81,7 @@ std::unordered_set } } - return transform(unordered_set_of(values(strands)), [&](auto const &edges) { + return transform(set_of(values(strands)), [&](auto const &edges) { return ExtendedSeriesReduction{edges}; }); } diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.cc index b7d2952fdc..f20e2db3c3 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/dependencies_are_maintained.cc @@ -1,6 +1,6 @@ #include "utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h" #include "utils/containers/is_subseteq_of.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/graph/digraph/algorithms/get_ancestors.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/series_parallel/get_ancestors.h" @@ -12,13 +12,13 @@ namespace FlexFlow { bool dependencies_are_maintained(DiGraphView const &g, SeriesParallelDecomposition const &sp) { ASSERT(has_no_duplicate_nodes(sp)); - if (unordered_set_of(get_nodes(sp)) != get_nodes(g)) { + if (set_of(get_nodes(sp)) != get_nodes(g)) { return false; } for (Node const &n : get_nodes(g)) { - std::unordered_set ancestors_in_g = get_ancestors(g, n); - std::unordered_set ancestors_in_sp = get_ancestors(sp, n); + std::set ancestors_in_g = get_ancestors(g, n); + std::set ancestors_in_sp = get_ancestors(sp, n); if (!is_subseteq_of(ancestors_in_g, ancestors_in_sp)) { return false; } diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc index 56076fb4ea..1721d16083 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc @@ -9,7 +9,7 @@ #include "utils/containers/set_union.h" #include "utils/containers/transform.h" #include "utils/containers/values.h" -#include "utils/fmt/unordered_multiset.h" +#include "utils/fmt/multiset.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_descendants.h" #include "utils/graph/digraph/algorithms/get_edges.h" @@ -33,28 +33,28 @@ #include "utils/positive_int/positive_int.h" #include -#include -#include +#include +#include namespace FlexFlow { -static std::unordered_set filter_out_sync_nodes( - std::unordered_set const &nodes, - std::unordered_map const &node_roles) { +static std::set filter_out_sync_nodes( + std::set const &nodes, + std::map const &node_roles) { return filter( nodes, [&](Node const &n) { return node_roles.at(n) != NodeRole::SYNC; }); } static nonnegative_int get_max_depth(DiGraph const &sp, - std::unordered_map const &depth_map) { + std::map const &depth_map) { return maximum(values(filter_keys( depth_map, [&](Node const &n) { return contains(get_nodes(sp), n); }))); } DiGraph add_dummy_nodes(DiGraph g, - std::unordered_map &node_roles) { - std::unordered_map depth_map = + std::map &node_roles) { + std::map depth_map = get_longest_path_lengths_from_root(g); for (DirectedEdge const &e : get_edges(g)) { @@ -87,11 +87,11 @@ DiGraph add_dummy_nodes(DiGraph g, return g; } -std::unordered_set +std::set get_component(DiGraph const &g, Node const &node, - std::unordered_map const &depth_map, - std::unordered_map const &node_roles) { + std::map const &depth_map, + std::map const &node_roles) { nonnegative_int max_depth = get_max_depth(g, depth_map); auto is_in_last_2_strata = [&](Node const &n) { @@ -109,38 +109,38 @@ std::unordered_set } }; - std::unordered_set last_two_layers_nodes = + std::set last_two_layers_nodes = filter(get_nodes(g), is_in_last_2_strata); DiGraphView subgraph = get_subgraph(g, last_two_layers_nodes); - std::unordered_set component = + std::set component = get_only(filter(get_weakly_connected_components(subgraph), - [&](std::unordered_set const &component) { + [&](std::set const &component) { return contains(component, node); })); - std::unordered_set component_without_sync_nodes = + std::set component_without_sync_nodes = filter_out_sync_nodes(component, node_roles); return component_without_sync_nodes; } -static std::unordered_set +static std::set get_forest_escribano(DiGraph const &g, Node const &handle, - std::unordered_set const &component, - std::unordered_map const &node_roles) { - std::unordered_set> subtrees = + std::set const &component, + std::map const &node_roles) { + std::set> subtrees = transform(get_successors(g, handle), [&](Node const &n) { return set_union(get_descendants(g, n), {n}); }); auto subtrees_overlapping_with_component = - filter(subtrees, [&](std::unordered_set subtree) { + filter(subtrees, [&](std::set subtree) { return set_intersection(subtree, component).size() > 0; }); - std::unordered_set forest = + std::set forest = set_union(subtrees_overlapping_with_component); forest.insert(handle); @@ -150,8 +150,8 @@ static std::unordered_set static std::pair, nonempty_set> get_up_and_down_sets( DiGraph const &g, - std::unordered_set const &forest, - std::unordered_map const &depth_map) { + std::set const &forest, + std::map const &depth_map) { nonnegative_int max_depth = get_max_depth(g, depth_map); @@ -163,11 +163,11 @@ static std::pair, nonempty_set> grouped_by_depth.at_l(max_depth)); } -static std::unordered_set +static std::set edges_to_remove(DiGraph const &g, - std::unordered_set const &up, - std::unordered_set const &down) { - std::unordered_set to_remove; + std::set const &up, + std::set const &down) { + std::set to_remove; for (Node const &u : up) { to_remove = set_union(to_remove, get_outgoing_edges(g, u)); @@ -179,9 +179,9 @@ static std::unordered_set return to_remove; } -static std::unordered_set - edges_to_add_escribano(std::unordered_set const &up, - std::unordered_set const &down, +static std::set + edges_to_add_escribano(std::set const &up, + std::set const &down, Node const &sync_node) { return set_union(transform(up, [&](Node const &u) { @@ -193,7 +193,7 @@ static std::unordered_set } static Node add_sync_node(DiGraph &sp, - std::unordered_map &node_roles) { + std::map &node_roles) { Node sync_node = sp.add_node(); node_roles[sync_node] = NodeRole::SYNC; return sync_node; @@ -203,10 +203,10 @@ SeriesParallelDecomposition escribano_sp_ization(DiGraph g) { ASSERT(is_2_terminal_dag(g)); ASSERT(is_acyclic(g)); - std::unordered_map node_roles = get_initial_node_role_map(g); + std::map node_roles = get_initial_node_role_map(g); g = add_dummy_nodes(g, node_roles); - std::unordered_map depth_map = + std::map depth_map = get_longest_path_lengths_from_root(g); DiGraph sp = DiGraph::create(); @@ -223,18 +223,18 @@ SeriesParallelDecomposition escribano_sp_ization(DiGraph g) { sp.add_node_unsafe(node); add_edges(sp, get_incoming_edges(g, node)); - std::unordered_set component = + std::set component = get_component(sp, node, depth_map, node_roles); Node handle = get_only(get_lowest_common_ancestors(sp, component).value()); - std::unordered_set forest = + std::set forest = get_forest_escribano(sp, handle, component, node_roles); std::pair, nonempty_set> up_down_sets = get_up_and_down_sets(sp, forest, depth_map); - std::unordered_set up = up_down_sets.first.unwrap_as_unordered_set(); - std::unordered_set down = - up_down_sets.second.unwrap_as_unordered_set(); + std::set up = up_down_sets.first.unwrap_as_set(); + std::set down = + up_down_sets.second.unwrap_as_set(); remove_edges(sp, edges_to_remove(sp, up, down)); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc index c091155a85..9da2d523a6 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc @@ -3,7 +3,7 @@ #include "utils/containers/argmin.h" #include "utils/containers/contains.h" #include "utils/containers/filter.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_only.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/keys.h" @@ -38,38 +38,38 @@ #include "utils/graph/series_parallel/sp_ization/up_down_partition.h" #include -#include -#include +#include +#include namespace FlexFlow { -static std::unordered_set - get_component(DiGraph const &sp, std::unordered_set const &nodes) { - std::unordered_set parents = set_union( +static std::set + get_component(DiGraph const &sp, std::set const &nodes) { + std::set parents = set_union( transform(nodes, [&](Node const &n) { return get_predecessors(sp, n); })); - std::unordered_set children = set_union(transform( + std::set children = set_union(transform( parents, [&](Node const &p) { return get_descendants(sp, p); })); - std::unordered_set other_parents = set_union(transform( + std::set other_parents = set_union(transform( children, [&](Node const &c) { return get_predecessors(sp, c); })); return set_union(set_union(parents, children), other_parents); } -static std::unordered_set +static std::set get_forest_flexible(DiGraph const &sp, Node const &handle, - std::unordered_set const &component, - std::unordered_map const &node_roles) { - std::unordered_set> subtrees = + std::set const &component, + std::map const &node_roles) { + std::set> subtrees = transform(get_successors(sp, handle), [&](Node const &n) { return set_union(get_descendants(sp, n), {n}); }); - std::unordered_set> overlapping_subtrees = - filter(subtrees, [&](std::unordered_set const &subtree) { + std::set> overlapping_subtrees = + filter(subtrees, [&](std::set const &subtree) { return !set_intersection(subtree, component).empty(); }); - std::unordered_set forest = set_union(overlapping_subtrees); + std::set forest = set_union(overlapping_subtrees); forest.insert(handle); return filter(forest, [&](Node const &n) { @@ -79,34 +79,34 @@ static std::unordered_set static UpDownPartition get_up_and_down_sets(DiGraph const &sp, - std::unordered_set const &nodes, - std::unordered_set const &forest, - std::unordered_map const &cost_map, - std::unordered_map const &node_roles) { + std::set const &nodes, + std::set const &forest, + std::map const &cost_map, + std::map const &node_roles) { DiGraph sp_pure = contract_out_nodes_of_given_role( materialize_digraph_view(sp), NodeRole::SYNC, node_roles); - std::unordered_set base_down = nodes; - std::unordered_set base_up = set_intersection( + std::set base_down = nodes; + std::set base_up = set_intersection( set_union(transform( nodes, [&](Node const &n) { return get_ancestors(sp_pure, n); })), forest); - std::unordered_set assignable_nodes = + std::set assignable_nodes = set_difference(forest, set_union(base_up, base_down)); DiGraphView forest_subgraph = get_subgraph(sp_pure, forest); - std::unordered_map critical_path_cost_map = + std::map critical_path_cost_map = get_weighted_longest_path_lengths_from_root(forest_subgraph, cost_map); auto get_partition_with_max_up_cost = [&](float reference_cost) -> UpDownPartition { - std::unordered_set up = + std::set up = set_union(base_up, filter(assignable_nodes, [&](Node const &n) { return critical_path_cost_map.at(n) <= reference_cost; })); - std::unordered_set down = + std::set down = set_difference(set_union(base_down, assignable_nodes), up); return UpDownPartition{up, down}; }; @@ -134,14 +134,14 @@ static UpDownPartition return true; }; - std::unordered_set partitions = + std::set partitions = transform(assignable_nodes, [&](Node const &n) { return get_partition_with_max_up_cost(critical_path_cost_map.at(n)); }); partitions.insert( UpDownPartition{base_up, set_union(base_down, assignable_nodes)}); - std::unordered_set valid_partitions = + std::set valid_partitions = filter(partitions, is_valid); ASSERT(!valid_partitions.empty()); @@ -155,12 +155,12 @@ static UpDownPartition return argmin(valid_partitions, partition_cost); } -static std::unordered_set edges_to_remove_flexible( +static std::set edges_to_remove_flexible( DiGraph const &sp, - std::unordered_set const &up, - std::unordered_set const &down, - std::unordered_map const &node_roles) { - std::unordered_set to_remove; + std::set const &up, + std::set const &down, + std::map const &node_roles) { + std::set to_remove; // from up to down for (Node const &u : up) { @@ -173,8 +173,8 @@ static std::unordered_set edges_to_remove_flexible( for (Node const &node : get_nodes(sp)) { if (node_roles.at(node) == NodeRole::SYNC) { - std::unordered_set preds = get_predecessors(sp, node); - std::unordered_set succs = get_successors(sp, node); + std::set preds = get_predecessors(sp, node); + std::set succs = get_successors(sp, node); if (is_subseteq_of(preds, up) && is_subseteq_of(succs, down)) { to_remove = set_union(to_remove, get_incoming_edges(sp, node)); to_remove = set_union(to_remove, get_outgoing_edges(sp, node)); @@ -185,12 +185,12 @@ static std::unordered_set edges_to_remove_flexible( return to_remove; } -static std::unordered_set +static std::set edges_to_add_flexible(DiGraph const &sp, UpDownPartition const &partition, Node const &sync_node) { - std::unordered_set up_frontier = get_up_frontier(sp, partition); - std::unordered_set down_frontier = get_down_frontier(sp, partition); + std::set up_frontier = get_up_frontier(sp, partition); + std::set down_frontier = get_down_frontier(sp, partition); return set_union(transform(up_frontier, [&](Node const &u) { @@ -202,39 +202,39 @@ static std::unordered_set } static Node add_sync_node(DiGraph &sp, - std::unordered_map &node_roles, - std::unordered_map &cost_map) { + std::map &node_roles, + std::map &cost_map) { Node sync_node = sp.add_node(); node_roles[sync_node] = NodeRole::SYNC; cost_map[sync_node] = 0.0f; return sync_node; } -static std::unordered_set +static std::set get_next_nodes(DiGraph const &sp, DiGraph const &g, - std::unordered_map const &cost_map) { - std::unordered_map sp_longest_paths = + std::map const &cost_map) { + std::map sp_longest_paths = get_weighted_longest_path_lengths_from_root(sp, cost_map); - std::unordered_set sp_nodes = get_nodes(sp); - std::unordered_set g_nodes = get_nodes(g); + std::set sp_nodes = get_nodes(sp); + std::set g_nodes = get_nodes(g); // candidate nodes: not in sp but all predecessors in sp - std::unordered_set candidate_nodes = + std::set candidate_nodes = filter(g_nodes, [&](Node const &node) { if (contains(sp_nodes, node)) { return false; } - std::unordered_set preds = get_predecessors(g, node); + std::set preds = get_predecessors(g, node); return is_subseteq_of(preds, sp_nodes); }); ASSERT(!candidate_nodes.empty()); - std::unordered_map critical_path_costs = - generate_unordered_map(candidate_nodes, [&](Node const &node) { - std::unordered_set preds = get_predecessors(g, node); + std::map critical_path_costs = + generate_map(candidate_nodes, [&](Node const &node) { + std::set preds = get_predecessors(g, node); float max_parent_cost = maximum(transform(preds, [&](Node const &pred) { return sp_longest_paths.at(pred); })); @@ -245,15 +245,15 @@ static std::unordered_set return std::make_pair(critical_path_costs.at(n), n.raw_uid); }); - std::unordered_set ref_preds = get_predecessors(g, ref_node); + std::set ref_preds = get_predecessors(g, ref_node); return filter(candidate_nodes, [&](Node const &node) { return get_predecessors(g, node) == ref_preds; }); } static bool cost_map_is_valid(DiGraphView const &g, - std::unordered_map const &cost_map) { - bool has_correct_nodes = (get_nodes(g) == unordered_keys(cost_map)); + std::map const &cost_map) { + bool has_correct_nodes = (get_nodes(g) == keys(cost_map)); bool has_nonnegative_costs = all_of(values(cost_map), [&](float const &cost) { return cost >= 0.0f; }); return has_correct_nodes && has_nonnegative_costs; @@ -261,11 +261,11 @@ static bool cost_map_is_valid(DiGraphView const &g, SeriesParallelDecomposition flexible_sync_unchecked(DiGraphView const &g, - std::unordered_map cost_map) { + std::map cost_map) { DiGraph g_reduced = materialize_digraph_view(transitive_reduction(g)); - std::unordered_map node_roles = + std::map node_roles = get_initial_node_role_map(g_reduced); DiGraph sp = DiGraph::create(); @@ -273,7 +273,7 @@ SeriesParallelDecomposition sp.add_node_unsafe(root); while (!is_subseteq_of(get_nodes(g_reduced), get_nodes(sp))) { - std::unordered_set nodes = get_next_nodes(sp, g_reduced, cost_map); + std::set nodes = get_next_nodes(sp, g_reduced, cost_map); for (Node const &node : nodes) { // here we add node unsafe so that we don't have to keep around a mapping @@ -288,9 +288,9 @@ SeriesParallelDecomposition // added edges sp = transitive_reduction(sp); - std::unordered_set component = get_component(sp, nodes); + std::set component = get_component(sp, nodes); Node handle = get_only(get_lowest_common_ancestors(sp, component).value()); - std::unordered_set forest = + std::set forest = get_forest_flexible(sp, handle, component, node_roles); UpDownPartition partition = @@ -316,7 +316,7 @@ SeriesParallelDecomposition SeriesParallelDecomposition flexible_sp_ization(DiGraphView const &g, - std::unordered_map const &cost_map) { + std::map const &cost_map) { ASSERT(is_2_terminal_dag(g)); ASSERT(is_acyclic(g)); ASSERT(cost_map_is_valid(g, cost_map)); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc index 4ebaf89756..cdc209f9d1 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc @@ -3,8 +3,8 @@ #include "utils/containers/maximum.h" #include "utils/containers/range.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_multiset_of.h" -#include "utils/fmt/unordered_multiset.h" +#include "utils/containers/multiset_of.h" +#include "utils/fmt/multiset.h" #include "utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h" #include "utils/graph/digraph/algorithms/is_acyclic.h" #include "utils/graph/series_parallel/non_normal_sp_decomposition.h" @@ -12,16 +12,16 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h" #include -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" namespace FlexFlow { -std::vector> +std::vector> stratum_split_assuming_unit_cost(DiGraphView const &g) { - std::unordered_map node_to_stratum = + std::map node_to_stratum = get_longest_path_lengths_from_root(g); - std::unordered_set nodes = unordered_keys(node_to_stratum); + std::set nodes = keys(node_to_stratum); OneToMany strata_to_nodes = group_by(nodes, [&](Node const &n) { return node_to_stratum.at(n); }); @@ -29,16 +29,16 @@ std::vector> return transform(range(1, num_strata.unwrap_nonnegative() + 1), [&](int depth) { - return unordered_multiset_of( + return multiset_of( strata_to_nodes.at_l(nonnegative_int{depth})); }); } static SeriesParallelDecomposition naive_stratum_merge( - std::vector> stratum_split) { + std::vector> stratum_split) { auto merge_one_stratum = - [&](std::unordered_multiset const &stratum_nodes) { + [&](std::multiset const &stratum_nodes) { auto as_singleton_sp = [](Node const &node) { return NonNormalSPDecomposition{node}; }; @@ -57,7 +57,7 @@ static SeriesParallelDecomposition naive_stratum_merge( SeriesParallelDecomposition naive_stratum_sync_sp_ization_unchecked(DiGraphView const &g) { - std::vector> stratum_split = + std::vector> stratum_split = stratum_split_assuming_unit_cost(g); return naive_stratum_merge(stratum_split); } diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc index a6d4183a23..d70abfac58 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/sp_ization/node_role.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_predecessors.h" #include "utils/graph/digraph/algorithms/get_successors.h" @@ -8,16 +8,16 @@ namespace FlexFlow { -std::unordered_map +std::map get_initial_node_role_map(DiGraphView const &g) { - return generate_unordered_map(get_nodes(g), + return generate_map(get_nodes(g), [](Node const &) { return NodeRole::PURE; }); } DiGraph contract_out_nodes_of_given_role( DiGraph g, NodeRole const &role, - std::unordered_map const &node_roles) { + std::map const &node_roles) { for (Node const &n : get_nodes(g)) { if (node_roles.at(n) == role) { for (Node const &pred : get_predecessors(g, n)) { diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc index 87a7eaa9b4..d4c62007dd 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc @@ -6,7 +6,7 @@ namespace FlexFlow { -std::unordered_set get_up_frontier(DiGraph const &sp, +std::set get_up_frontier(DiGraph const &sp, UpDownPartition const &partition) { DiGraphView up_subgraph = get_subgraph(sp, partition.up); return filter(partition.up, [&](Node const &node) { @@ -14,7 +14,7 @@ std::unordered_set get_up_frontier(DiGraph const &sp, }); } -std::unordered_set get_down_frontier(DiGraph const &sp, +std::set get_down_frontier(DiGraph const &sp, UpDownPartition const &partition) { DiGraphView down_subgraph = get_subgraph(sp, partition.down); return filter(partition.down, [&](Node const &node) { diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc index 7423437b1c..db28ee9970 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc @@ -4,7 +4,7 @@ #include "utils/containers/group_by.h" #include "utils/containers/slice.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" #include "utils/fmt/variant.h" #include "utils/graph/digraph/algorithms/get_initial_nodes.h" #include "utils/graph/digraph/algorithms/get_predecessors.h" @@ -18,7 +18,7 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/variant.h" #include -#include +#include namespace FlexFlow { @@ -34,7 +34,7 @@ static NonNormalSeriesSplit cut_off_head(NonNormalSeriesSplit const &s) { * with coalescing: S(1, P( S(2,5), S(3,4) )) */ static NonNormalSPDecomposition parallel_composition_with_coalescing( - std::unordered_set const &strands) { + std::set const &strands) { if (strands.size() == 1) { return NonNormalSPDecomposition{get_only(strands)}; } @@ -48,18 +48,18 @@ static NonNormalSPDecomposition parallel_composition_with_coalescing( return strand.children.at(0); }; - std::unordered_set non_empty_strands = + std::set non_empty_strands = filter(strands, is_non_empty_strand); OneToMany, NonNormalSeriesSplit> strands_grouped_by_head = group_by(non_empty_strands, strand_head); // recursively coalesce the strands - std::unordered_multiset coalesced_strands; + std::multiset coalesced_strands; for (auto const &[head, strands_with_head] : strands_grouped_by_head.l_to_r()) { - std::unordered_set tails = - transform(strands_with_head.unwrap_as_unordered_set(), cut_off_head); + std::set tails = + transform(strands_with_head.unwrap_as_set(), cut_off_head); NonNormalSPDecomposition parallel_comp = parallel_composition_with_coalescing(tails); @@ -78,7 +78,7 @@ static NonNormalSPDecomposition parallel_composition_with_coalescing( static SeriesParallelDecomposition work_duplicating_sp_ization_unchecked_with_coalescing( DiGraphView const &g) { - std::unordered_map node_to_sp; + std::map node_to_sp; Node source = get_only(get_initial_nodes(g)); node_to_sp.emplace(source, NonNormalSeriesSplit{{source}}); @@ -87,7 +87,7 @@ static SeriesParallelDecomposition if (node == source) { continue; } - std::unordered_set predecessors_as_sp = + std::set predecessors_as_sp = transform(get_predecessors(g, node), [&](Node const &p) { return node_to_sp.at(p); }); @@ -108,12 +108,12 @@ static SeriesParallelDecomposition static SeriesParallelDecomposition work_duplicating_sp_ization_unchecked(DiGraphView const &g) { - std::unordered_map node_to_sp; + std::map node_to_sp; for (Node const &node : get_topological_ordering(g)) { - std::unordered_multiset predecessors_as_sp = - unordered_multiset_of( + std::multiset predecessors_as_sp = + multiset_of( transform(get_predecessors(g, node), [&](Node const &p) { return node_to_sp.at(p); })); diff --git a/lib/utils/src/utils/graph/traversal.cc b/lib/utils/src/utils/graph/traversal.cc index a4df327b2a..beffcf90d2 100644 --- a/lib/utils/src/utils/graph/traversal.cc +++ b/lib/utils/src/utils/graph/traversal.cc @@ -11,7 +11,7 @@ unchecked_dfs_iterator::unchecked_dfs_iterator(DiGraphView const &g, : stack(stack), graph(g) {} unchecked_dfs_iterator::unchecked_dfs_iterator( - DiGraphView const &g, std::unordered_set const &starting_points) + DiGraphView const &g, std::set const &starting_points) : graph(g) { for (Node const &n : starting_points) { this->stack.push_back(n); @@ -30,7 +30,7 @@ unchecked_dfs_iterator &unchecked_dfs_iterator::operator++() { Node const last = this->operator*(); this->stack.pop_back(); - std::unordered_set outgoing = get_outgoing_edges(graph, last); + std::set outgoing = get_outgoing_edges(graph, last); for (DirectedEdge const &e : outgoing) { auto it = std::find(stack.begin(), stack.end(), e.dst); if (it == stack.end()) { @@ -65,11 +65,11 @@ bool unchecked_dfs_iterator::operator!=( checked_dfs_iterator::checked_dfs_iterator(DiGraphView const &g, std::vector const &stack, - std::unordered_set const &seen) + std::set const &seen) : iter(g, stack), seen(seen) {} checked_dfs_iterator::checked_dfs_iterator( - DiGraphView const &g, std::unordered_set const &starting_points) + DiGraphView const &g, std::set const &starting_points) : iter(g, starting_points), seen{} {} checked_dfs_iterator::reference checked_dfs_iterator::operator*() const { @@ -104,12 +104,12 @@ bool checked_dfs_iterator::operator!=(checked_dfs_iterator const &other) const { bfs_iterator::bfs_iterator(DiGraphView const &g, std::queue const &q, - std::optional> const &seen) + std::optional> const &seen) : graph(g), q(q), seen(seen) {} bfs_iterator::bfs_iterator(DiGraphView const &g, - std::unordered_set const &starting_points) - : graph(g), seen(std::unordered_set{}) { + std::set const &starting_points) + : graph(g), seen(std::set{}) { for (Node const &n : starting_points) { this->q.push(n); } @@ -129,7 +129,7 @@ bfs_iterator &bfs_iterator::operator++() { this->seen.value().insert(current); this->q.pop(); - std::unordered_set outgoing = + std::set outgoing = get_outgoing_edges(graph, {current}); for (DirectedEdge const &e : outgoing) { if (!contains(this->seen.value(), e.dst)) { @@ -165,7 +165,7 @@ bool bfs_iterator::operator!=(bfs_iterator const &other) const { } CheckedDFSView::CheckedDFSView(DiGraphView const &g, - std::unordered_set const &starting_points) + std::set const &starting_points) : graph(g), starting_points(starting_points) {} checked_dfs_iterator CheckedDFSView::cbegin() const { @@ -185,12 +185,12 @@ checked_dfs_iterator CheckedDFSView::end() const { } CheckedDFSView dfs(DiGraphView const &g, - std::unordered_set const &starting_points) { + std::set const &starting_points) { return CheckedDFSView(g, starting_points); } UncheckedDFSView::UncheckedDFSView( - DiGraphView const &g, std::unordered_set const &starting_points) + DiGraphView const &g, std::set const &starting_points) : graph(g), starting_points(starting_points) {} unchecked_dfs_iterator UncheckedDFSView::cbegin() const { @@ -211,12 +211,12 @@ unchecked_dfs_iterator UncheckedDFSView::end() const { UncheckedDFSView unchecked_dfs(DiGraphView const &g, - std::unordered_set const &starting_points) { + std::set const &starting_points) { return UncheckedDFSView(g, starting_points); } BFSView::BFSView(DiGraphView const &g, - std::unordered_set const &starting_points) + std::set const &starting_points) : graph(g), starting_points(starting_points) {} bfs_iterator BFSView::cbegin() const { @@ -236,7 +236,7 @@ bfs_iterator BFSView::end() const { } BFSView bfs(DiGraphView const &g, - std::unordered_set const &starting_points) { + std::set const &starting_points) { return BFSView(g, starting_points); } diff --git a/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc b/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc index 361b82f746..562fc1f906 100644 --- a/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc +++ b/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc @@ -1,18 +1,18 @@ #include "utils/graph/undirected/algorithms/get_connected_components.h" #include "utils/graph/algorithms.h" #include "utils/graph/node/algorithms.h" -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" namespace FlexFlow { -std::unordered_set> +std::set> get_connected_components(UndirectedGraphView const &g) { - std::unordered_set> components; - std::unordered_set visited; + std::set> components; + std::set visited; for (Node const &node : get_nodes(g)) { - std::unordered_set component = - unordered_set_of(get_bfs_ordering(as_digraph(g), {node})); + std::set component = + set_of(get_bfs_ordering(as_digraph(g), {node})); components.insert(component); visited = set_union(visited, component); } diff --git a/lib/utils/src/utils/graph/undirected/algorithms/get_edges.cc b/lib/utils/src/utils/graph/undirected/algorithms/get_edges.cc index 8ae825c1ab..19d1e078a7 100644 --- a/lib/utils/src/utils/graph/undirected/algorithms/get_edges.cc +++ b/lib/utils/src/utils/graph/undirected/algorithms/get_edges.cc @@ -3,7 +3,7 @@ namespace FlexFlow { -std::unordered_set get_edges(UndirectedGraphView const &g) { +std::set get_edges(UndirectedGraphView const &g) { return g.query_edges(undirected_edge_query_all()); } diff --git a/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc b/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc index d28818c5b4..8bf76965a3 100644 --- a/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc +++ b/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc @@ -3,14 +3,14 @@ namespace FlexFlow { -std::unordered_set get_neighboring_nodes(UndirectedGraphView const &g, +std::set get_neighboring_nodes(UndirectedGraphView const &g, Node const &n) { - std::unordered_set edges = g.query_edges( + std::set edges = g.query_edges( UndirectedEdgeQuery{query_set::match_single_value(n)}); - std::unordered_set result = + std::set result = set_union(transform(vector_of(edges), [](UndirectedEdge const &e) { - return std::unordered_set{e.endpoints.max(), e.endpoints.max()}; + return std::set{e.endpoints.max(), e.endpoints.max()}; })); result.erase(n); return result; diff --git a/lib/utils/src/utils/graph/undirected/undirected_edge.cc b/lib/utils/src/utils/graph/undirected/undirected_edge.cc index cc3d18ce97..a6203151a2 100644 --- a/lib/utils/src/utils/graph/undirected/undirected_edge.cc +++ b/lib/utils/src/utils/graph/undirected/undirected_edge.cc @@ -8,7 +8,7 @@ bool is_connected_to(UndirectedEdge const &e, Node const &n) { return e.endpoints.min() == n || e.endpoints.max() == n; } -std::unordered_set get_endpoints(UndirectedEdge const &e) { +std::set get_endpoints(UndirectedEdge const &e) { return {e.endpoints.min(), e.endpoints.max()}; } diff --git a/lib/utils/src/utils/graph/undirected/undirected_graph.cc b/lib/utils/src/utils/graph/undirected/undirected_graph.cc index 32c9468ec3..12662fd6f0 100644 --- a/lib/utils/src/utils/graph/undirected/undirected_graph.cc +++ b/lib/utils/src/utils/graph/undirected/undirected_graph.cc @@ -32,12 +32,12 @@ IUndirectedGraph &UndirectedGraph::get_ptr() { GraphView::ptr.get_mutable()); } -std::unordered_set +std::set UndirectedGraph::query_edges(UndirectedEdgeQuery const &q) const { return this->get_ptr().query_edges(q); } -std::unordered_set +std::set UndirectedGraph::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } diff --git a/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc b/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc index b74ccfa322..2daa8dcbfd 100644 --- a/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc +++ b/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc @@ -2,12 +2,12 @@ namespace FlexFlow { -std::unordered_set +std::set UndirectedGraphView::query_edges(UndirectedEdgeQuery const &q) const { return this->get_ptr().query_edges(q); } -std::unordered_set +std::set UndirectedGraphView::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } diff --git a/lib/utils/src/utils/graph/views/views.cc b/lib/utils/src/utils/graph/views/views.cc index 6efbb45a17..5977d800af 100644 --- a/lib/utils/src/utils/graph/views/views.cc +++ b/lib/utils/src/utils/graph/views/views.cc @@ -12,14 +12,14 @@ namespace FlexFlow { UndirectedSubgraphView::UndirectedSubgraphView( UndirectedGraphView const &g, - std::unordered_set const &subgraph_nodes) + std::set const &subgraph_nodes) : g(g), subgraph_nodes(subgraph_nodes) {} UndirectedSubgraphView *UndirectedSubgraphView::clone() const { return new UndirectedSubgraphView(g, subgraph_nodes); } -std::unordered_set UndirectedSubgraphView::query_edges( +std::set UndirectedSubgraphView::query_edges( UndirectedEdgeQuery const &query) const { UndirectedEdgeQuery subgraph_query = UndirectedEdgeQuery{ query_set::match_values_in(set_of(this->subgraph_nodes)), @@ -27,7 +27,7 @@ std::unordered_set UndirectedSubgraphView::query_edges( return this->g.query_edges(query_intersection(query, subgraph_query)); } -std::unordered_set +std::set UndirectedSubgraphView::query_nodes(NodeQuery const &query) const { NodeQuery subgraph_query = NodeQuery{ query_set::match_values_in(set_of(this->subgraph_nodes)), @@ -37,10 +37,10 @@ std::unordered_set } DiSubgraphView::DiSubgraphView(DiGraphView const &g, - std::unordered_set const &subgraph_nodes) + std::set const &subgraph_nodes) : g(g), subgraph_nodes(subgraph_nodes) {} -std::unordered_set +std::set DiSubgraphView::query_edges(DirectedEdgeQuery const &query) const { DirectedEdgeQuery subgraph_query = DirectedEdgeQuery{ query_set::match_values_in(set_of(this->subgraph_nodes)), @@ -49,7 +49,7 @@ std::unordered_set return this->g.query_edges(query_intersection(query, subgraph_query)); } -std::unordered_set +std::set DiSubgraphView::query_nodes(NodeQuery const &query) const { NodeQuery subgraph_query = NodeQuery{ query_set::match_values_in(set_of(this->subgraph_nodes)), @@ -64,12 +64,12 @@ DiSubgraphView *DiSubgraphView::clone() const { UndirectedGraphView view_subgraph(UndirectedGraphView const &g, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { return UndirectedGraphView::create(g, subgraph_nodes); } DiGraphView view_subgraph(DiGraphView const &g, - std::unordered_set const &subgraph_nodes) { + std::set const &subgraph_nodes) { return DiGraphView::create(g, subgraph_nodes); } @@ -77,20 +77,20 @@ UndirectedEdge to_undirected_edge(DirectedEdge const &e) { return make_undirected_edge(e.src, e.dst); } -std::unordered_set to_undirected_edges( - std::unordered_set const &directed_edges) { +std::set to_undirected_edges( + std::set const &directed_edges) { return transform(directed_edges, [](DirectedEdge const &e) { return to_undirected_edge(e); }); } -std::unordered_set to_directed_edges(UndirectedEdge const &e) { - return std::unordered_set{ +std::set to_directed_edges(UndirectedEdge const &e) { + return std::set{ DirectedEdge{e.endpoints.min(), e.endpoints.max()}, DirectedEdge{e.endpoints.max(), e.endpoints.min()}}; } -std::unordered_set to_directed_edges( - std::unordered_set const &undirected_edges) { +std::set to_directed_edges( + std::set const &undirected_edges) { return flatmap(undirected_edges, [](UndirectedEdge const &e) { return to_directed_edges(e); }); } @@ -98,7 +98,7 @@ std::unordered_set to_directed_edges( ViewDiGraphAsUndirectedGraph::ViewDiGraphAsUndirectedGraph(DiGraphView const &g) : g(g) {} -std::unordered_set ViewDiGraphAsUndirectedGraph::query_edges( +std::set ViewDiGraphAsUndirectedGraph::query_edges( UndirectedEdgeQuery const &undirected_query) const { DirectedEdgeQuery q1{undirected_query.nodes, query_set::matchall()}; DirectedEdgeQuery q2{query_set::matchall(), undirected_query.nodes}; @@ -106,7 +106,7 @@ std::unordered_set ViewDiGraphAsUndirectedGraph::query_edges( set_union(this->g.query_edges(q1), this->g.query_edges(q2))); } -std::unordered_set ViewDiGraphAsUndirectedGraph::query_nodes( +std::set ViewDiGraphAsUndirectedGraph::query_nodes( NodeQuery const &node_query) const { return this->g.query_nodes(node_query); } @@ -123,18 +123,18 @@ ViewUndirectedGraphAsDiGraph *ViewUndirectedGraphAsDiGraph::clone() const { return new ViewUndirectedGraphAsDiGraph(g); } -std::unordered_set ViewUndirectedGraphAsDiGraph::query_edges( +std::set ViewUndirectedGraphAsDiGraph::query_edges( DirectedEdgeQuery const &q) const { - std::unordered_set undirected_edges = + std::set undirected_edges = g.query_edges(UndirectedEdgeQuery{query_union(q.srcs, q.dsts)}); - std::unordered_set directed_edges = + std::set directed_edges = flatmap(undirected_edges, [](UndirectedEdge const &e) { return to_directed_edges(e); }); return filter(directed_edges, [&](DirectedEdge const &e) { return matches_edge(q, e); }); } -std::unordered_set +std::set ViewUndirectedGraphAsDiGraph::query_nodes(NodeQuery const &q) const { return g.query_nodes(q); } diff --git a/lib/utils/src/utils/hash/unordered_map.cc b/lib/utils/src/utils/hash/unordered_map.cc index 52c4140641..1efd05e9b1 100644 --- a/lib/utils/src/utils/hash/unordered_map.cc +++ b/lib/utils/src/utils/hash/unordered_map.cc @@ -1 +1 @@ -#include "utils/hash/unordered_map.h" +#include "utils/hash/map.h" diff --git a/lib/utils/src/utils/hash/unordered_multiset.cc b/lib/utils/src/utils/hash/unordered_multiset.cc index 7f6f73f428..d84ca7d614 100644 --- a/lib/utils/src/utils/hash/unordered_multiset.cc +++ b/lib/utils/src/utils/hash/unordered_multiset.cc @@ -1 +1 @@ -#include "utils/hash/unordered_multiset.h" +#include "utils/hash/multiset.h" diff --git a/lib/utils/src/utils/hash/unordered_set.cc b/lib/utils/src/utils/hash/unordered_set.cc index 907555d908..ac109492a0 100644 --- a/lib/utils/src/utils/hash/unordered_set.cc +++ b/lib/utils/src/utils/hash/unordered_set.cc @@ -1 +1 @@ -#include "utils/hash/unordered_set.h" +#include "utils/hash/set.h" diff --git a/lib/utils/src/utils/many_to_one/many_to_one.cc b/lib/utils/src/utils/many_to_one/many_to_one.cc index 52ae52153b..933d32a176 100644 --- a/lib/utils/src/utils/many_to_one/many_to_one.cc +++ b/lib/utils/src/utils/many_to_one/many_to_one.cc @@ -1,27 +1,27 @@ #include "utils/many_to_one/many_to_one.h" #include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/rapidcheckable_value_type.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" using namespace ::FlexFlow; namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template struct ManyToOne; -template std::unordered_map, R> +template std::map, R> format_as(ManyToOne const &); template std::ostream &operator<<(std::ostream &, ManyToOne const &); -template std::unordered_set> +template std::set> unstructured_relation_from_many_to_one(ManyToOne const &); template ManyToOne many_to_one_from_unstructured_relation( - std::unordered_set> const &); + std::set> const &); } // namespace FlexFlow @@ -46,8 +46,8 @@ template struct Arbitrary<::FlexFlow::ManyToOne>; namespace std { -using L = ::FlexFlow::value_type<0>; -using R = ::FlexFlow::value_type<1>; +using L = ::FlexFlow::ordered_value_type<0>; +using R = ::FlexFlow::ordered_value_type<1>; template struct hash>; diff --git a/lib/utils/src/utils/many_to_one/many_to_one_from_bidict.cc b/lib/utils/src/utils/many_to_one/many_to_one_from_bidict.cc index 030ab95de5..b740a7dfaa 100644 --- a/lib/utils/src/utils/many_to_one/many_to_one_from_bidict.cc +++ b/lib/utils/src/utils/many_to_one/many_to_one_from_bidict.cc @@ -1,10 +1,10 @@ #include "utils/many_to_one/many_to_one_from_bidict.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using L = value_type<0>; -using R = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; template ManyToOne many_to_one_from_bidict(bidict const &); diff --git a/lib/utils/src/utils/many_to_one/many_to_one_from_map.cc b/lib/utils/src/utils/many_to_one/many_to_one_from_map.cc index a898bb9e00..1f984b12b6 100644 --- a/lib/utils/src/utils/many_to_one/many_to_one_from_map.cc +++ b/lib/utils/src/utils/many_to_one/many_to_one_from_map.cc @@ -1,18 +1,13 @@ #include "utils/many_to_one/many_to_one_from_map.h" #include "utils/archetypes/ordered_value_type.h" -#include "utils/archetypes/value_type.h" namespace FlexFlow { -using L1 = value_type<0>; -using R1 = value_type<1>; +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; -template ManyToOne - many_to_one_from_map(std::unordered_map const &); +template ManyToOne many_to_one_from_map(std::map const &); -using L2 = ordered_value_type<0>; -using R2 = ordered_value_type<1>; - -template ManyToOne many_to_one_from_map(std::map const &); +template ManyToOne many_to_one_from_map(std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/one_to_many/one_to_many.cc b/lib/utils/src/utils/one_to_many/one_to_many.cc index ce6220c509..2094a65ece 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many.cc @@ -18,7 +18,7 @@ template std::map> template std::ostream &operator<<(std::ostream &, OneToMany const &); -template std::unordered_set> +template std::set> unstructured_relation_from_one_to_many(OneToMany const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc b/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc index 124adb20c3..eddc23fe5f 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc @@ -7,6 +7,6 @@ using L = ordered_value_type<0>; using R = ordered_value_type<1>; template OneToMany one_to_many_from_l_to_r_mapping( - std::unordered_map> const &); + std::map> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/orthotope/dim_coord.cc b/lib/utils/src/utils/orthotope/dim_coord.cc index 75bacacb38..cf3c12712e 100644 --- a/lib/utils/src/utils/orthotope/dim_coord.cc +++ b/lib/utils/src/utils/orthotope/dim_coord.cc @@ -1,30 +1,30 @@ #include "utils/orthotope/dim_coord.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T = value_type<0>; +using T = ordered_value_type<0>; -template std::unordered_set get_coord_dims(DimCoord const &); +template std::set get_coord_dims(DimCoord const &); template DimCoord restrict_coord_to_dims(DimCoord const &, - std::unordered_set const &); + std::set const &); template OrthotopeCoord orthotope_coord_from_dim_coord(DimCoord const &, DimOrdering const &); template DimCoord dim_coord_from_orthotope_coord(OrthotopeCoord const &, - std::unordered_set const &, + std::set const &, DimOrdering const &); template DimCoord lift_dim_coord(DimCoord const &, - std::unordered_set const &); + std::set const &); -template std::unordered_set> +template std::set> get_coords_in_dim_domain(DimDomain const &); -template std::unordered_set> +template std::set> get_coords_in_minimal_dim_domain(MinimalDimDomain const &); template DimCoord get_maximum_coord_in_domain(DimDomain const &); diff --git a/lib/utils/src/utils/orthotope/dim_domain.cc b/lib/utils/src/utils/orthotope/dim_domain.cc index 3a410d31cb..a79bb9ee55 100644 --- a/lib/utils/src/utils/orthotope/dim_domain.cc +++ b/lib/utils/src/utils/orthotope/dim_domain.cc @@ -9,20 +9,20 @@ template DimDomain empty_dim_domain(); template nonnegative_int dim_domain_num_dims(DimDomain const &); -template std::unordered_set get_domain_dims(DimDomain const &); +template std::set get_domain_dims(DimDomain const &); -template std::unordered_set get_trivial_domain_dims(DimDomain const &); +template std::set get_trivial_domain_dims(DimDomain const &); -template std::unordered_set get_nontrivial_domain_dims(DimDomain const &); +template std::set get_nontrivial_domain_dims(DimDomain const &); template DimDomain restrict_domain_to_dims(DimDomain const &, - std::unordered_set const &); + std::set const &); template Orthotope orthotope_from_dim_domain(DimDomain const &, DimOrdering const &); template DimDomain dim_domain_from_orthotope(Orthotope const &, - std::unordered_set const &, + std::set const &, DimOrdering const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/orthotope/dim_projection.cc b/lib/utils/src/utils/orthotope/dim_projection.cc index 9fa250fa47..9c4c5e0c78 100644 --- a/lib/utils/src/utils/orthotope/dim_projection.cc +++ b/lib/utils/src/utils/orthotope/dim_projection.cc @@ -13,10 +13,10 @@ template DimProjection DimOrdering const &, DimOrdering const &); -template std::unordered_set +template std::set input_dims_of_projection(DimProjection const &); -template std::unordered_set +template std::set output_dims_of_projection(DimProjection const &); template DimProjection invert_dim_projection(DimProjection const &); diff --git a/lib/utils/src/utils/orthotope/down_projection.cc b/lib/utils/src/utils/orthotope/down_projection.cc index 684521a5de..271640605b 100644 --- a/lib/utils/src/utils/orthotope/down_projection.cc +++ b/lib/utils/src/utils/orthotope/down_projection.cc @@ -8,10 +8,10 @@ using R = ordered_value_type<1>; template DownProjection make_empty_down_projection(); -template std::unordered_set +template std::set input_dims_of_down_projection(DownProjection const &); -template std::unordered_set +template std::set output_dims_of_down_projection(DownProjection const &); template DimCoord compute_down_projection(DownProjection const &, @@ -20,7 +20,7 @@ template DimCoord compute_down_projection(DownProjection const &, DimOrdering const &); template void project_dims(DownProjection &, - std::unordered_set const &, + std::set const &, R const &); template UpProjection diff --git a/lib/utils/src/utils/orthotope/eq_projection.cc b/lib/utils/src/utils/orthotope/eq_projection.cc index a6965dc36e..7deffc7711 100644 --- a/lib/utils/src/utils/orthotope/eq_projection.cc +++ b/lib/utils/src/utils/orthotope/eq_projection.cc @@ -8,10 +8,10 @@ using R = ordered_value_type<1>; template EqProjection make_empty_eq_projection(); -template std::unordered_set +template std::set input_dims_of_eq_projection(EqProjection const &); -template std::unordered_set +template std::set output_dims_of_eq_projection(EqProjection const &); template void project_dims(EqProjection &, L const &, R const &); diff --git a/lib/utils/src/utils/orthotope/minimal_dim_domain.cc b/lib/utils/src/utils/orthotope/minimal_dim_domain.cc index de60fad397..5bff87fc27 100644 --- a/lib/utils/src/utils/orthotope/minimal_dim_domain.cc +++ b/lib/utils/src/utils/orthotope/minimal_dim_domain.cc @@ -1,9 +1,9 @@ #include "utils/orthotope/minimal_dim_domain.h" -#include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -using T = value_type<0>; +using T = ordered_value_type<0>; template MinimalDimDomain empty_minimal_dim_domain(); @@ -20,14 +20,14 @@ template MinimalDimDomain template DimDomain dim_domain_from_minimal_dim_domain(MinimalDimDomain const &, - std::unordered_set const &); + std::set const &); -template std::unordered_set +template std::set get_minimal_domain_dims(MinimalDimDomain const &); template MinimalDimDomain restrict_minimal_domain_to_dims(MinimalDimDomain const &, - std::unordered_set const &); + std::set const &); template MinimalOrthotope minimal_orthotope_from_minimal_dim_domain(MinimalDimDomain const &, @@ -35,7 +35,7 @@ template MinimalOrthotope template MinimalDimDomain minimal_dim_domain_from_minimal_orthotope(MinimalOrthotope const &, - std::unordered_set const &, + std::set const &, DimOrdering const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc b/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc index 5d4c47b491..c281d31d66 100644 --- a/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc +++ b/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc @@ -19,8 +19,8 @@ template MinimalDimDomainMapping template DimDomainMapping dim_domain_mapping_from_minimal_dim_domain( MinimalDimDomainMapping const &, - std::unordered_set const &, - std::unordered_set const &); + std::set const &, + std::set const &); template MinimalDimDomainMapping minimal_dim_domain_mapping_identity_map(MinimalDimDomain const &, diff --git a/lib/utils/src/utils/orthotope/orthotope.cc b/lib/utils/src/utils/orthotope/orthotope.cc index bb1f0d9ae8..0a53a113d6 100644 --- a/lib/utils/src/utils/orthotope/orthotope.cc +++ b/lib/utils/src/utils/orthotope/orthotope.cc @@ -9,7 +9,7 @@ #include "utils/containers/slice.h" #include "utils/containers/sum.h" #include "utils/containers/transform.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/containers/zip3_with_strict.h" #include "utils/containers/zip_strict.h" #include "utils/containers/zip_with_strict.h" @@ -28,14 +28,14 @@ positive_int orthotope_get_volume(Orthotope const &orthotope) { return product(orthotope.dims); } -std::unordered_set +std::set get_all_coords_in_orthotope(Orthotope const &orthotope) { - std::unordered_multiset> raw_coords = + std::multiset> raw_coords = cartesian_product(transform(orthotope.dims, [](positive_int dim_size) { return nonnegative_range(dim_size); })); - return unordered_set_of( + return set_of( transform(raw_coords, [](std::vector const &raw_coord) { return OrthotopeCoord{raw_coord}; })); diff --git a/lib/utils/src/utils/orthotope/up_projection.cc b/lib/utils/src/utils/orthotope/up_projection.cc index 0c8909dffd..587ccdeb19 100644 --- a/lib/utils/src/utils/orthotope/up_projection.cc +++ b/lib/utils/src/utils/orthotope/up_projection.cc @@ -14,10 +14,10 @@ template UpProjection using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template std::unordered_set +template std::set input_dims_of_up_projection(UpProjection const &); -template std::unordered_set +template std::set output_dims_of_up_projection(UpProjection const &); template DimCoord compute_up_projection(UpProjection const &, @@ -29,7 +29,7 @@ template UpProjection make_empty_up_projection(); template void project_dims(UpProjection &, L const &, - std::unordered_set const &); + std::set const &); template DownProjection invert_up_projection(UpProjection const &); diff --git a/lib/utils/src/utils/record_formatter.cc b/lib/utils/src/utils/record_formatter.cc index b44ed87338..c5d7a3c708 100644 --- a/lib/utils/src/utils/record_formatter.cc +++ b/lib/utils/src/utils/record_formatter.cc @@ -90,6 +90,6 @@ template RecordFormatter mk_kv_record(std::string const &, using K = ordered_value_type<0>; using V = value_type<0>; -template RecordFormatter mk_record_for_map(std::unordered_map const &); +template RecordFormatter mk_record_for_map(std::map const &); } // namespace FlexFlow diff --git a/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h b/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h index b44fc5e5af..659c23814c 100644 --- a/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h +++ b/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h @@ -3,8 +3,8 @@ #include #include #include -#include -#include +#include +#include #include using namespace FlexFlow; diff --git a/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_map.h b/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_map.h index 4fd5d15009..94cbbe341c 100644 --- a/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_map.h +++ b/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_map.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_DOCTEST_FMT_UNORDERED_MAP_H #define _FLEXFLOW_LIB_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_DOCTEST_FMT_UNORDERED_MAP_H -#include "utils/fmt/unordered_map.h" +#include "utils/fmt/map.h" #include namespace doctest { template -struct StringMaker> { - static String convert(std::unordered_map const &m) { +struct StringMaker> { + static String convert(std::map const &m) { return toString(fmt::to_string(m)); } }; diff --git a/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_multiset.h b/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_multiset.h index 94dae42239..acb81e4916 100644 --- a/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_multiset.h +++ b/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_multiset.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_DOCTEST_FMT_UNORDERED_MULTISET_H #define _FLEXFLOW_LIB_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_DOCTEST_FMT_UNORDERED_MULTISET_H -#include "utils/fmt/unordered_multiset.h" +#include "utils/fmt/multiset.h" #include namespace doctest { template -struct StringMaker> { - static String convert(std::unordered_multiset const &m) { +struct StringMaker> { + static String convert(std::multiset const &m) { return toString(fmt::to_string(m)); } }; diff --git a/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_set.h b/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_set.h index 441590365d..978dc457df 100644 --- a/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_set.h +++ b/lib/utils/test/common/include/test/utils/doctest/fmt/unordered_set.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_DOCTEST_FMT_UNORDERED_SET_H #define _FLEXFLOW_LIB_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_DOCTEST_FMT_UNORDERED_SET_H -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include namespace doctest { template -struct StringMaker> { - static String convert(std::unordered_set const &m) { +struct StringMaker> { + static String convert(std::set const &m) { return toString(fmt::to_string(m)); } }; diff --git a/lib/utils/test/common/include/test/utils/rapidcheck/gen.h b/lib/utils/test/common/include/test/utils/rapidcheck/gen.h index ad9879d6e7..e6796fff0d 100644 --- a/lib/utils/test/common/include/test/utils/rapidcheck/gen.h +++ b/lib/utils/test/common/include/test/utils/rapidcheck/gen.h @@ -2,14 +2,14 @@ #define _FLEXFLOW_UTILS_LIB_TEST_COMMON_INCLUDE_UTILS_TEST_RAPIDCHECK_GEN_H #include -#include +#include namespace rc { template -Gen> subset_of(C const &sets) { +Gen> subset_of(C const &sets) { return gen::exec([&] { - std::unordered_set result; + std::set result; for (auto const &elem : sets) { if (*gen::arbitrary()) { result.insert(elem); diff --git a/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_map.cc b/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_map.cc index b893e632ed..976e65cfca 100644 --- a/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_map.cc +++ b/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_map.cc @@ -1 +1 @@ -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" diff --git a/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_multiset.cc b/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_multiset.cc index 55d2e69056..9c5b2f4d1e 100644 --- a/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_multiset.cc +++ b/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_multiset.cc @@ -1 +1 @@ -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" diff --git a/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_set.cc b/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_set.cc index 13ad811e63..9ec70698bc 100644 --- a/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_set.cc +++ b/lib/utils/test/common/src/test/utils/doctest/fmt/unordered_set.cc @@ -1 +1 @@ -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc index 98bce0a904..4d7f0ad495 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc @@ -1,5 +1,5 @@ #include "utils/bidict/algorithms/bidict_from_enumerating.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/bidict/algorithms/left_entries.h" #include "utils/bidict/algorithms/right_entries.h" #include @@ -30,20 +30,20 @@ TEST_SUITE(FF_TEST_SUITE) { } } - TEST_CASE("bidict_from_enumerating(std::unordered_set)") { - std::unordered_set input = {"zero", "one", "two"}; + TEST_CASE("bidict_from_enumerating(std::set)") { + std::set input = {"zero", "one", "two"}; bidict result = bidict_from_enumerating(input); - std::unordered_set result_left_entries = + std::set result_left_entries = left_entries(result); - std::unordered_set correct_left_entries = {0_n, 1_n, 2_n}; + std::set correct_left_entries = {0_n, 1_n, 2_n}; CHECK(result_left_entries == correct_left_entries); - std::unordered_set result_right_entries = + std::set result_right_entries = right_entries(result); - std::unordered_set correct_right_entries = input; + std::set correct_right_entries = input; CHECK(result_right_entries == correct_right_entries); } diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_map.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_map.cc index 029b7d0461..767046636d 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_map.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_map.cc @@ -4,9 +4,9 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("bidict_from_map(std::unordered_map)") { + TEST_CASE("bidict_from_map(std::map)") { SUBCASE("map values do not contain duplicates") { - std::unordered_map input = { + std::map input = { {1, "one"}, {2, "two"}, {3, "three"}, @@ -23,7 +23,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("map values contain duplicates") { - std::unordered_map input = { + std::map input = { {1, "odd"}, {2, "even"}, {3, "odd"}, diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc index 98c4e27faf..26f35693ff 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc @@ -6,7 +6,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("bidict_from_unstructured_relation") { SUBCASE("relation is one-to-one") { - std::unordered_set> input = { + std::set> input = { {1, "one"}, {2, "two"}, }; @@ -22,7 +22,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("relation is one-to-many") { - std::unordered_set> input = { + std::set> input = { {1, "one"}, {1, "ONE"}, {2, "two"}, @@ -32,7 +32,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("relation is many-to-one") { - std::unordered_set> input = { + std::set> input = { {1, "odd"}, {2, "even"}, {3, "odd"}, @@ -42,7 +42,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("relation is none of the above") { - std::unordered_set> input = { + std::set> input = { {1, "odd"}, {1, "ODD"}, {2, "even"}, diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_unordered_set_of.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_unordered_set_of.cc index d44a2fe62b..61fa613bae 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_unordered_set_of.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_unordered_set_of.cc @@ -4,7 +4,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("unordered_set_of(bidict)") { - CHECK_MESSAGE(false, "TODO: unordered_set_of(bidict)"); + TEST_CASE("set_of(bidict)") { + CHECK_MESSAGE(false, "TODO: set_of(bidict)"); } } diff --git a/lib/utils/test/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc b/lib/utils/test/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc index a4aaf553a4..6ca2e9cc44 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/unstructured_relation_from_bidict.cc @@ -1,6 +1,6 @@ #include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" #include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; @@ -12,9 +12,9 @@ TEST_SUITE(FF_TEST_SUITE) { {2, "two"}, }; - std::unordered_set> result = + std::set> result = unstructured_relation_from_bidict(input); - std::unordered_set> correct = { + std::set> correct = { {1, "one"}, {2, "two"}, }; diff --git a/lib/utils/test/src/utils/bidict/bidict.cc b/lib/utils/test/src/utils/bidict/bidict.cc index ead45fe86a..6db4bd1fbc 100644 --- a/lib/utils/test/src/utils/bidict/bidict.cc +++ b/lib/utils/test/src/utils/bidict/bidict.cc @@ -1,6 +1,6 @@ #include "utils/bidict/bidict.h" #include "test/utils/doctest/check_without_stringify.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/vector.h" #include "test/utils/rapidcheck.h" #include @@ -84,16 +84,16 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(dict.size() == 2); } - SUBCASE("implicitly convert to std::unordered_map") { - std::unordered_map res = dict; - std::unordered_map expected = {{1, "one"}, {2, "two"}}; + SUBCASE("implicitly convert to std::map") { + std::map res = dict; + std::map expected = {{1, "one"}, {2, "two"}}; CHECK(res == expected); } SUBCASE("bidict::begin") { auto it = dict.begin(); - CHECK(it->first == 2); - CHECK(it->second == "two"); + CHECK(it->first == 1); + CHECK(it->second == "one"); } SUBCASE("bidict::end") { @@ -104,7 +104,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("fmt::to_string(bidict)") { std::string result = fmt::to_string(dict); - std::string correct = fmt::to_string(dict.as_unordered_map()); + std::string correct = fmt::to_string(dict.as_map()); CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/commutative_pair.cc b/lib/utils/test/src/utils/commutative_pair.cc index af015e2b8c..1dde0399b4 100644 --- a/lib/utils/test/src/utils/commutative_pair.cc +++ b/lib/utils/test/src/utils/commutative_pair.cc @@ -150,7 +150,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("fmt::to_string") { std::string result = fmt::to_string(x); - std::unordered_set correct_options = {"{2, 1}", "{1, 2}"}; + std::set correct_options = {"{2, 1}", "{1, 2}"}; CHECK(contains(correct_options, result)); } @@ -158,7 +158,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::ostringstream oss; oss << x; std::string result = oss.str(); - std::unordered_set correct_options = {"{2, 1}", "{1, 2}"}; + std::set correct_options = {"{2, 1}", "{1, 2}"}; CHECK(contains(correct_options, result)); } } diff --git a/lib/utils/test/src/utils/containers/are_disjoint.cc b/lib/utils/test/src/utils/containers/are_disjoint.cc index 17516dbf13..d9cf4513d4 100644 --- a/lib/utils/test/src/utils/containers/are_disjoint.cc +++ b/lib/utils/test/src/utils/containers/are_disjoint.cc @@ -1,30 +1,30 @@ #include "utils/containers/are_disjoint.h" #include -#include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("are_disjoint") { SUBCASE("disjoint") { - std::unordered_set l = {1, 2, 3}; - std::unordered_set r = {4, 5, 6}; + std::set l = {1, 2, 3}; + std::set r = {4, 5, 6}; CHECK(are_disjoint(l, r)); } SUBCASE("not disjoint") { - std::unordered_set l = {1, 2, 3, 4}; - std::unordered_set r = {3, 4, 5, 6}; + std::set l = {1, 2, 3, 4}; + std::set r = {3, 4, 5, 6}; CHECK_FALSE(are_disjoint(l, r)); } SUBCASE("one empty set") { - std::unordered_set l = {1, 2}; - std::unordered_set r = {}; + std::set l = {1, 2}; + std::set r = {}; CHECK(are_disjoint(l, r)); } SUBCASE("both empty sets") { - std::unordered_set l = {}; - std::unordered_set r = {}; + std::set l = {}; + std::set r = {}; CHECK(are_disjoint(l, r)); } } diff --git a/lib/utils/test/src/utils/containers/argmax.cc b/lib/utils/test/src/utils/containers/argmax.cc index fcfda0c200..b5837e329b 100644 --- a/lib/utils/test/src/utils/containers/argmax.cc +++ b/lib/utils/test/src/utils/containers/argmax.cc @@ -1,7 +1,7 @@ #include "utils/containers/argmax.h" #include "utils/containers/contains.h" #include -#include +#include #include using namespace FlexFlow; @@ -20,9 +20,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("ties") { - std::unordered_set input = {-1, 1, 2}; + std::set input = {-1, 1, 2}; int result = argmax(input, [](int x) { return -(x * x); }); - CHECK(contains(std::unordered_set{-1, 1}, result)); + CHECK(contains(std::set{-1, 1}, result)); } } } diff --git a/lib/utils/test/src/utils/containers/argmin.cc b/lib/utils/test/src/utils/containers/argmin.cc index 1a5bf4014e..d62a06ddbc 100644 --- a/lib/utils/test/src/utils/containers/argmin.cc +++ b/lib/utils/test/src/utils/containers/argmin.cc @@ -1,7 +1,7 @@ #include "utils/containers/argmin.h" #include "utils/containers/contains.h" #include -#include +#include #include using namespace FlexFlow; @@ -20,9 +20,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("ties") { - std::unordered_set input = {-1, 1, 2}; + std::set input = {-1, 1, 2}; int result = argmin(input, [](int x) { return x * x; }); - CHECK(contains(std::unordered_set{-1, 1}, result)); + CHECK(contains(std::set{-1, 1}, result)); } } } diff --git a/lib/utils/test/src/utils/containers/binary_cartesian_product.cc b/lib/utils/test/src/utils/containers/binary_cartesian_product.cc index bbf79252af..2b7f5b6a31 100644 --- a/lib/utils/test/src/utils/containers/binary_cartesian_product.cc +++ b/lib/utils/test/src/utils/containers/binary_cartesian_product.cc @@ -1,6 +1,6 @@ #include "utils/containers/binary_cartesian_product.h" #include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/hash/pair.h" #include #include @@ -10,13 +10,13 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("binary_cartesian_product") { SUBCASE("both lhs and rhs are nonempty") { - std::unordered_set lhs = {1, 3}; - std::unordered_set rhs = {"a", "b", "c"}; + std::set lhs = {1, 3}; + std::set rhs = {"a", "b", "c"}; - std::unordered_set> result = + std::set> result = binary_cartesian_product(lhs, rhs); - std::unordered_set> correct = { + std::set> correct = { {1, "a"}, {3, "b"}, {1, "c"}, @@ -29,25 +29,25 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("lhs is empty") { - std::unordered_set lhs = {}; - std::unordered_set rhs = {"a", "b", "c"}; + std::set lhs = {}; + std::set rhs = {"a", "b", "c"}; - std::unordered_set> result = + std::set> result = binary_cartesian_product(lhs, rhs); - std::unordered_set> correct = {}; + std::set> correct = {}; CHECK(result == correct); } SUBCASE("rhs is empty") { - std::unordered_set lhs = {1, 3}; - std::unordered_set rhs = {}; + std::set lhs = {1, 3}; + std::set rhs = {}; - std::unordered_set> result = + std::set> result = binary_cartesian_product(lhs, rhs); - std::unordered_set> correct = {}; + std::set> correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc index 250d1c7f69..d4487343f2 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc @@ -1,34 +1,34 @@ -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/containers/binary_merge_disjoint_maps.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("binary_merge_disjoint_unordered_maps") { - std::unordered_map l_map = { + TEST_CASE("binary_merge_disjoint_maps") { + std::map l_map = { {1, "one"}, {2, "two"}, }; - std::unordered_map r_map = { + std::map r_map = { {3, "three"}, }; - std::unordered_map correct = { + std::map correct = { {1, "one"}, {2, "two"}, {3, "three"}, }; SUBCASE("maps are disjoint") { - std::unordered_map result = - binary_merge_disjoint_unordered_maps(l_map, r_map); + std::map result = + binary_merge_disjoint_maps(l_map, r_map); CHECK(result == correct); } SUBCASE("maps are not disjoint") { - CHECK_THROWS(binary_merge_disjoint_unordered_maps(l_map, l_map)); + CHECK_THROWS(binary_merge_disjoint_maps(l_map, l_map)); } } } diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc index 4e825b99e2..6a848565e1 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc @@ -1,57 +1,57 @@ -#include "utils/containers/binary_merge_unordered_maps_with.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/containers/binary_merge_maps_with.h" +#include "test/utils/doctest/fmt/map.h" #include #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("binary_merge_unordered_maps_with") { + TEST_CASE("binary_merge_maps_with") { auto fail_if_called = [](std::string const &, std::string const &) -> std::string { PANIC(); }; SUBCASE("lhs and rhs do not overlap") { - std::unordered_map lhs = { + std::map lhs = { {1, "lhs_one."}, {4, "lhs_four."}, }; - std::unordered_map rhs = { + std::map rhs = { {2, "rhs_two."}, {5, "rhs_five."}, }; - std::unordered_map correct = { + std::map correct = { {1, "lhs_one."}, {2, "rhs_two."}, {4, "lhs_four."}, {5, "rhs_five."}, }; - std::unordered_map result = - binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + std::map result = + binary_merge_maps_with(lhs, rhs, fail_if_called); CHECK(result == correct); } SUBCASE("lhs and rhs overlap") { - std::unordered_map lhs = { + std::map lhs = { {1, "lhs_one."}, {4, "lhs_four."}, }; - std::unordered_map rhs = { + std::map rhs = { {2, "rhs_two."}, {4, "rhs_four."}, {5, "rhs_five."}, }; - std::unordered_map result = binary_merge_unordered_maps_with( + std::map result = binary_merge_maps_with( lhs, rhs, [](std::string const &l, std::string const &r) { return l + r; }); - std::unordered_map correct = { + std::map correct = { {1, "lhs_one."}, {2, "rhs_two."}, {4, "lhs_four.rhs_four."}, @@ -62,47 +62,47 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("lhs is empty") { - std::unordered_map lhs = {}; + std::map lhs = {}; - std::unordered_map rhs = { + std::map rhs = { {2, "rhs_two."}, {4, "rhs_four."}, {5, "rhs_five."}, }; - std::unordered_map result = - binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + std::map result = + binary_merge_maps_with(lhs, rhs, fail_if_called); - std::unordered_map correct = rhs; + std::map correct = rhs; CHECK(result == correct); } SUBCASE("rhs is empty") { - std::unordered_map lhs = { + std::map lhs = { {1, "lhs_one."}, {4, "lhs_four."}, }; - std::unordered_map rhs = {}; + std::map rhs = {}; - std::unordered_map result = - binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + std::map result = + binary_merge_maps_with(lhs, rhs, fail_if_called); - std::unordered_map correct = lhs; + std::map correct = lhs; CHECK(result == correct); } SUBCASE("both lhs and rhs are empty") { - std::unordered_map lhs = {}; + std::map lhs = {}; - std::unordered_map rhs = {}; + std::map rhs = {}; - std::unordered_map result = - binary_merge_unordered_maps_with(lhs, rhs, fail_if_called); + std::map result = + binary_merge_maps_with(lhs, rhs, fail_if_called); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc index d857cf2d91..fe152dd832 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc @@ -1,30 +1,30 @@ -#include "utils/containers/binary_merge_unordered_maps_with_left_dominating.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/containers/binary_merge_maps_with_left_dominating.h" +#include "test/utils/doctest/fmt/map.h" #include #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("binary_merge_unordered_maps_with_left_dominating") { - std::unordered_map l_map = { + TEST_CASE("binary_merge_maps_with_left_dominating") { + std::map l_map = { {1, "one"}, {2, "left_two"}, }; - std::unordered_map r_map = { + std::map r_map = { {2, "right_two"}, {3, "three"}, }; - std::unordered_map correct = { + std::map correct = { {1, "one"}, {2, "left_two"}, {3, "three"}, }; - std::unordered_map result = - binary_merge_unordered_maps_with_left_dominating(l_map, r_map); + std::map result = + binary_merge_maps_with_left_dominating(l_map, r_map); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc index 71f50c4dac..c107f2b7ff 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc @@ -1,30 +1,30 @@ -#include "utils/containers/binary_merge_unordered_maps_with_right_dominating.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/containers/binary_merge_maps_with_right_dominating.h" +#include "test/utils/doctest/fmt/map.h" #include #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("binary_merge_unordered_maps_with_right_dominating") { - std::unordered_map l_map = { + TEST_CASE("binary_merge_maps_with_right_dominating") { + std::map l_map = { {1, "one"}, {2, "left_two"}, }; - std::unordered_map r_map = { + std::map r_map = { {2, "right_two"}, {3, "three"}, }; - std::unordered_map correct = { + std::map correct = { {1, "one"}, {2, "right_two"}, {3, "three"}, }; - std::unordered_map result = - binary_merge_unordered_maps_with_right_dominating(l_map, r_map); + std::map result = + binary_merge_maps_with_right_dominating(l_map, r_map); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/cartesian_product.cc b/lib/utils/test/src/utils/containers/cartesian_product.cc index 773d94c8d0..e25ffd67f7 100644 --- a/lib/utils/test/src/utils/containers/cartesian_product.cc +++ b/lib/utils/test/src/utils/containers/cartesian_product.cc @@ -1,8 +1,8 @@ #include "utils/containers/cartesian_product.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/vector.h" #include -#include +#include #include using namespace FlexFlow; @@ -12,59 +12,59 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("empty") { std::vector> containers = {}; - std::unordered_multiset> result = + std::multiset> result = cartesian_product(containers); - std::unordered_multiset> correct = {{}}; + std::multiset> correct = {{}}; CHECK(result == correct); } SUBCASE("single container, one element") { std::vector> containers = {{1}}; - std::unordered_multiset> result = + std::multiset> result = cartesian_product(containers); - std::unordered_multiset> correct = {{1}}; + std::multiset> correct = {{1}}; CHECK(result == correct); } SUBCASE("single container, multiple elements") { std::vector> containers = {{1, 2, 3}}; - std::unordered_multiset> result = + std::multiset> result = cartesian_product(containers); - std::unordered_multiset> correct = {{1}, {2}, {3}}; + std::multiset> correct = {{1}, {2}, {3}}; CHECK(result == correct); } SUBCASE("multiple containers, one element each") { std::vector> containers = {{1}, {2}, {3}}; - std::unordered_multiset> result = + std::multiset> result = cartesian_product(containers); - std::unordered_multiset> correct = {{1, 2, 3}}; + std::multiset> correct = {{1, 2, 3}}; CHECK(result == correct); } SUBCASE("multiple containers, multiple elements") { std::vector> containers = {{1, 2}, {3, 4}}; - std::unordered_multiset> result = + std::multiset> result = cartesian_product(containers); - std::unordered_multiset> correct = { + std::multiset> correct = { {1, 3}, {1, 4}, {2, 3}, {2, 4}}; CHECK(result == correct); } SUBCASE("multiple containers, duplicate elements") { std::vector> containers = {{1, 1}, {2, 3}}; - std::unordered_multiset> result = + std::multiset> result = cartesian_product(containers); - std::unordered_multiset> correct = { + std::multiset> correct = { {1, 2}, {1, 3}, {1, 3}, {1, 2}}; CHECK(result == correct); } SUBCASE("1 empty container, 1 non-empty container") { std::vector> containers = {{}, {2, 3}}; - std::unordered_multiset> result = + std::multiset> result = cartesian_product(containers); - std::unordered_multiset> correct = {}; + std::multiset> correct = {}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/contains.cc b/lib/utils/test/src/utils/containers/contains.cc index 9d686ab814..f2da08b430 100644 --- a/lib/utils/test/src/utils/containers/contains.cc +++ b/lib/utils/test/src/utils/containers/contains.cc @@ -1,6 +1,6 @@ #include "utils/containers/contains.h" #include -#include +#include #include using namespace FlexFlow; @@ -13,8 +13,8 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK_FALSE(contains(v, 6)); } - SUBCASE("std::unordered_set") { - std::unordered_set s = {1, 2, 3, 4, 5}; + SUBCASE("std::set") { + std::set s = {1, 2, 3, 4, 5}; CHECK(contains(s, 3)); CHECK_FALSE(contains(s, 6)); } diff --git a/lib/utils/test/src/utils/containers/contains_duplicates.cc b/lib/utils/test/src/utils/containers/contains_duplicates.cc index 3a0ffd17f2..27764c5aaa 100644 --- a/lib/utils/test/src/utils/containers/contains_duplicates.cc +++ b/lib/utils/test/src/utils/containers/contains_duplicates.cc @@ -7,7 +7,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE_TEMPLATE("contains_duplicates(T)", T, std::vector, - std::unordered_multiset, + std::multiset, std::multiset) { SUBCASE("container has duplicates") { T input = {2, 7, 3, 4, 2, 1}; diff --git a/lib/utils/test/src/utils/containers/contains_key.cc b/lib/utils/test/src/utils/containers/contains_key.cc index da099113a6..48933e3fab 100644 --- a/lib/utils/test/src/utils/containers/contains_key.cc +++ b/lib/utils/test/src/utils/containers/contains_key.cc @@ -2,13 +2,13 @@ #include #include #include -#include +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("contains_key(std::unordered_map, K)") { - std::unordered_map m = { + TEST_CASE("contains_key(std::map, K)") { + std::map m = { {1, "one"}, }; CHECK(contains_key(m, 1)); diff --git a/lib/utils/test/src/utils/containers/contains_value.cc b/lib/utils/test/src/utils/containers/contains_value.cc index 136ef3b304..9a678e894e 100644 --- a/lib/utils/test/src/utils/containers/contains_value.cc +++ b/lib/utils/test/src/utils/containers/contains_value.cc @@ -5,8 +5,8 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("contains_value(std::unordered_map, V)") { - std::unordered_map m = { + TEST_CASE("contains_value(std::map, V)") { + std::map m = { {1, "one"}, {3, "three"}, {4, "three"}, diff --git a/lib/utils/test/src/utils/containers/enumerate.cc b/lib/utils/test/src/utils/containers/enumerate.cc index 22bae1f613..17fcdbc046 100644 --- a/lib/utils/test/src/utils/containers/enumerate.cc +++ b/lib/utils/test/src/utils/containers/enumerate.cc @@ -1,11 +1,11 @@ #include "utils/containers/enumerate.h" #include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/multiset.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" -#include "utils/containers/unordered_keys.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/keys.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/values.h" #include "utils/containers/vector_of.h" #include @@ -43,14 +43,14 @@ TEST_SUITE(FF_TEST_SUITE) { } } - TEST_CASE("enumerate(std::unordered_set)") { - std::unordered_set input = {"A", "B", "C", "D"}; + TEST_CASE("enumerate(std::set)") { + std::set input = {"A", "B", "C", "D"}; - std::unordered_set correct_keys = {0_n, 1_n, 2_n, 3_n}; - std::unordered_multiset correct_values = {"A", "B", "C", "D"}; + std::set correct_keys = {0_n, 1_n, 2_n, 3_n}; + std::multiset correct_values = {"A", "B", "C", "D"}; std::map result = enumerate(input); - CHECK(unordered_keys(result) == correct_keys); - CHECK(unordered_multiset_of(values(result)) == correct_values); + CHECK(keys(result) == correct_keys); + CHECK(multiset_of(values(result)) == correct_values); } } diff --git a/lib/utils/test/src/utils/containers/extend.cc b/lib/utils/test/src/utils/containers/extend.cc index ad0a276a5d..a75171374d 100644 --- a/lib/utils/test/src/utils/containers/extend.cc +++ b/lib/utils/test/src/utils/containers/extend.cc @@ -1,5 +1,5 @@ #include "utils/containers/extend.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include @@ -16,12 +16,12 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } - TEST_CASE("extend(std::unordered_set &, C)") { - std::unordered_set result = {1, 2, 3}; + TEST_CASE("extend(std::set &, C)") { + std::set result = {1, 2, 3}; std::vector rhs = {3, 3, 4, 5}; extend(result, rhs); - std::unordered_set correct = {1, 2, 3, 4, 5}; + std::set correct = {1, 2, 3, 4, 5}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/filter.cc b/lib/utils/test/src/utils/containers/filter.cc index 9462d30024..cffd3676b7 100644 --- a/lib/utils/test/src/utils/containers/filter.cc +++ b/lib/utils/test/src/utils/containers/filter.cc @@ -1,9 +1,9 @@ #include "utils/containers/filter.h" #include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_map.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/multiset.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "test/utils/rapidcheck.h" @@ -13,9 +13,9 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE_TEMPLATE("filter(T, F)", T, std::vector, - std::unordered_set, std::set, - std::unordered_map, + std::set, + std::map, std::map) { RC_SUBCASE("filter returns empty for predicate always_false", [](T const &t) { @@ -41,12 +41,12 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } - TEST_CASE("filter(std::unordered_set, F)") { - std::unordered_set input = {1, 2, 3, 4, 5, 6, 7, 8}; + TEST_CASE("filter(std::set, F)") { + std::set input = {1, 2, 3, 4, 5, 6, 7, 8}; auto predicate = [](int x) { return x % 2 == 0; }; - std::unordered_set result = filter(input, predicate); - std::unordered_set correct = {2, 4, 6, 8}; + std::set result = filter(input, predicate); + std::set correct = {2, 4, 6, 8}; CHECK(result == correct); } @@ -59,8 +59,8 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } - TEST_CASE("filter(std::unordered_map, F)") { - std::unordered_map input = { + TEST_CASE("filter(std::map, F)") { + std::map input = { {3, "4"}, {1, "1"}, {2, "9"}, @@ -70,8 +70,8 @@ TEST_SUITE(FF_TEST_SUITE) { return std::to_string(x.first) == x.second; }; - std::unordered_map result = filter(input, predicate); - std::unordered_map correct = { + std::map result = filter(input, predicate); + std::map correct = { {1, "1"}, {4, "4"}, }; @@ -97,12 +97,12 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } - TEST_CASE("filter(std::unordered_multiset, F)") { - std::unordered_multiset input = {1, 1, 2, 2, 2, 3, 4, 5, 6, 7, 8, 8}; + TEST_CASE("filter(std::multiset, F)") { + std::multiset input = {1, 1, 2, 2, 2, 3, 4, 5, 6, 7, 8, 8}; auto predicate = [](int x) { return x % 2 == 0; }; - std::unordered_multiset result = filter(input, predicate); - std::unordered_multiset correct = {2, 2, 2, 4, 6, 8, 8}; + std::multiset result = filter(input, predicate); + std::multiset correct = {2, 2, 2, 4, 6, 8, 8}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/filter_keys.cc b/lib/utils/test/src/utils/containers/filter_keys.cc index 00e327a6f1..6de906c491 100644 --- a/lib/utils/test/src/utils/containers/filter_keys.cc +++ b/lib/utils/test/src/utils/containers/filter_keys.cc @@ -1,18 +1,18 @@ #include "utils/containers/filter_keys.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include -#include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("filter_keys") { - std::unordered_map m = { + std::map m = { {1, "one"}, {2, "two"}, {3, "three"}}; auto f = [](int x) { return x % 2 == 1; }; - std::unordered_map result = filter_keys(m, f); - std::unordered_map correct = {{1, "one"}, {3, "three"}}; + std::map result = filter_keys(m, f); + std::map correct = {{1, "one"}, {3, "three"}}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/filtermap_keys.cc b/lib/utils/test/src/utils/containers/filtermap_keys.cc index 582e94392b..b71bb8a052 100644 --- a/lib/utils/test/src/utils/containers/filtermap_keys.cc +++ b/lib/utils/test/src/utils/containers/filtermap_keys.cc @@ -1,17 +1,17 @@ #include "utils/containers/filtermap_keys.h" #include "test/utils/doctest/fmt/map.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("filtermap_keys(std::unordered_map, F)") { - std::unordered_map input = { + TEST_CASE("filtermap_keys(std::map, F)") { + std::map input = { {1, "one"}, {2, "two"}, }; - std::unordered_map result = + std::map result = filtermap_keys(input, [](int k) -> std::optional { if (k == 1) { return std::nullopt; @@ -21,7 +21,7 @@ TEST_SUITE(FF_TEST_SUITE) { return oss.str(); } }); - std::unordered_map correct = { + std::map correct = { {"3", "two"}, }; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/containers/filtermap_values.cc b/lib/utils/test/src/utils/containers/filtermap_values.cc index 8db6d6a964..f6d4335405 100644 --- a/lib/utils/test/src/utils/containers/filtermap_values.cc +++ b/lib/utils/test/src/utils/containers/filtermap_values.cc @@ -1,17 +1,17 @@ #include "utils/containers/filtermap_values.h" #include "test/utils/doctest/fmt/map.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("filtermap_values(std::unordered_map, F)") { - std::unordered_map input = { + TEST_CASE("filtermap_values(std::map, F)") { + std::map input = { {1, "one"}, {2, "two"}, }; - std::unordered_map result = + std::map result = filtermap_values(input, [](std::string const &v) -> std::optional { if (v == "two") { return std::nullopt; @@ -19,7 +19,7 @@ TEST_SUITE(FF_TEST_SUITE) { return v.size() + 1; } }); - std::unordered_map correct = { + std::map correct = { {1, 4}, }; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/containers/filtrans.cc b/lib/utils/test/src/utils/containers/filtrans.cc index cd1c2f896c..4aab65d528 100644 --- a/lib/utils/test/src/utils/containers/filtrans.cc +++ b/lib/utils/test/src/utils/containers/filtrans.cc @@ -1,6 +1,6 @@ #include "utils/containers/filtrans.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include @@ -23,9 +23,9 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } - TEST_CASE("filtrans(std::unordered_set, F)") { - std::unordered_set input = {1, 2, 3, 4}; - std::unordered_set result = + TEST_CASE("filtrans(std::set, F)") { + std::set input = {1, 2, 3, 4}; + std::set result = filtrans(input, [](int x) -> std::optional { if ((x % 2) == 0) { return std::to_string(x); @@ -34,7 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { } }); - std::unordered_set correct = {"2", "4"}; + std::set correct = {"2", "4"}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/find.cc b/lib/utils/test/src/utils/containers/find.cc index 36d4b771d8..b3fc17eb82 100644 --- a/lib/utils/test/src/utils/containers/find.cc +++ b/lib/utils/test/src/utils/containers/find.cc @@ -3,7 +3,7 @@ #include #include #include -#include +#include #include using namespace FlexFlow; @@ -27,8 +27,8 @@ TEST_SUITE(FF_TEST_SUITE) { } } - SUBCASE("unordered_set") { - std::unordered_set s = {1, 2, 3, 4, 5}; + SUBCASE("set") { + std::set s = {1, 2, 3, 4, 5}; SUBCASE("element in container") { CHECK_WITHOUT_STRINGIFY(find(s, 3) == std::find(s.begin(), s.end(), 3)); diff --git a/lib/utils/test/src/utils/containers/flatmap.cc b/lib/utils/test/src/utils/containers/flatmap.cc index 6a6d3c86a8..18e0cf88a3 100644 --- a/lib/utils/test/src/utils/containers/flatmap.cc +++ b/lib/utils/test/src/utils/containers/flatmap.cc @@ -1,7 +1,7 @@ #include "utils/containers/flatmap.h" #include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_map.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "utils/containers/map_keys.h" #include "utils/hash/pair.h" @@ -44,9 +44,9 @@ TEST_SUITE(FF_TEST_SUITE) { } } - TEST_CASE("flatmap(std::unordered_set, F)") { + TEST_CASE("flatmap(std::set, F)") { auto get_chars = [](std::string const &s) { - std::unordered_set result; + std::set result; for (char c : s) { result.insert(c); } @@ -54,20 +54,20 @@ TEST_SUITE(FF_TEST_SUITE) { }; SUBCASE("type changing") { - std::unordered_set input = {"hello", " ", "", "world", "!"}; + std::set input = {"hello", " ", "", "world", "!"}; - std::unordered_set result = flatmap(input, get_chars); - std::unordered_set correct = { + std::set result = flatmap(input, get_chars); + std::set correct = { 'h', 'e', 'l', 'o', ' ', 'w', 'r', 'd', '!'}; CHECK(result == correct); } SUBCASE("input is empty") { - std::unordered_set input = {}; + std::set input = {}; - std::unordered_set result = flatmap(input, get_chars); - std::unordered_set correct = {}; + std::set result = flatmap(input, get_chars); + std::set correct = {}; CHECK(result == correct); } @@ -105,24 +105,24 @@ TEST_SUITE(FF_TEST_SUITE) { } } - TEST_CASE("flatmap(std::unordered_map, F)") { + TEST_CASE("flatmap(std::map, F)") { auto de_nest_keys = [](int k1, - std::unordered_map const &v) { + std::map const &v) { return map_keys(v, [&](int k2) { return std::pair{k1, k2}; }); }; SUBCASE("input is empty") { - std::unordered_map> input = {}; + std::map> input = {}; - std::unordered_map, std::string> result = + std::map, std::string> result = flatmap(input, de_nest_keys); - std::unordered_map, std::string> correct = {}; + std::map, std::string> correct = {}; CHECK(result == correct); } SUBCASE("input is not empty") { - std::unordered_map> input = { + std::map> input = { { 1, { @@ -142,9 +142,9 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_map, std::string> result = + std::map, std::string> result = flatmap(input, de_nest_keys); - std::unordered_map, std::string> correct = { + std::map, std::string> correct = { {{1, 2}, "a"}, {{1, 3}, "b"}, {{3, 3}, "a"}, @@ -155,12 +155,12 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("duplicate result keys") { auto always_return_same_map = [](int, std::string const &) { - return std::unordered_map{ + return std::map{ {"mykey", 10000}, }; }; - std::unordered_map input = { + std::map input = { {1, "a"}, {2, "b"}, }; diff --git a/lib/utils/test/src/utils/containers/get_all_assignments.cc b/lib/utils/test/src/utils/containers/get_all_assignments.cc index d5f989318f..bcf69cdc87 100644 --- a/lib/utils/test/src/utils/containers/get_all_assignments.cc +++ b/lib/utils/test/src/utils/containers/get_all_assignments.cc @@ -1,6 +1,6 @@ #include "utils/containers/get_all_assignments.h" -#include "test/utils/doctest/fmt/unordered_map.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; @@ -8,24 +8,24 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_all_assignments") { SUBCASE("empty input") { - std::unordered_map> input = {}; + std::map> input = {}; - std::unordered_set> result = + std::set> result = get_all_assignments(input); - std::unordered_set> correct = {{}}; + std::set> correct = {{}}; CHECK(result == correct); } SUBCASE("non-empty input") { - std::unordered_map> input = { + std::map> input = { {"a", {1, 2, 3}}, {"b", {2, 3}}, }; - std::unordered_set> result = + std::set> result = get_all_assignments(input); - std::unordered_set> correct = { + std::set> correct = { {{"a", 1}, {"b", 2}}, {{"a", 1}, {"b", 3}}, {{"a", 2}, {"b", 2}}, @@ -38,14 +38,14 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("one possible-values set is empty") { - std::unordered_map> input = { + std::map> input = { {"a", {}}, {"b", {2, 3}}, }; - std::unordered_set> result = + std::set> result = get_all_assignments(input); - std::unordered_set> correct = {}; + std::set> correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/get_all_permutations.cc b/lib/utils/test/src/utils/containers/get_all_permutations.cc index cc5edb4075..2c21f1fe58 100644 --- a/lib/utils/test/src/utils/containers/get_all_permutations.cc +++ b/lib/utils/test/src/utils/containers/get_all_permutations.cc @@ -1,7 +1,7 @@ #include "utils/containers/get_all_permutations.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/vector.h" -#include "utils/containers/unordered_multiset_of.h" +#include "utils/containers/multiset_of.h" #include "utils/hash/vector.h" #include @@ -12,9 +12,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("input size 1") { std::vector input = {1}; - std::unordered_multiset> result = - unordered_multiset_of(get_all_permutations(input)); - std::unordered_multiset> correct = {{1}}; + std::multiset> result = + multiset_of(get_all_permutations(input)); + std::multiset> correct = {{1}}; CHECK(result == correct); } @@ -22,9 +22,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("input size 3") { std::vector input = {2, 1, 3}; - std::unordered_multiset> result = - unordered_multiset_of(get_all_permutations(input)); - std::unordered_multiset> correct = { + std::multiset> result = + multiset_of(get_all_permutations(input)); + std::multiset> correct = { {1, 2, 3}, {1, 3, 2}, {2, 1, 3}, @@ -39,9 +39,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("elements repeated") { std::vector input = {1, 2, 2}; - std::unordered_multiset> result = - unordered_multiset_of(get_all_permutations(input)); - std::unordered_multiset> correct = { + std::multiset> result = + multiset_of(get_all_permutations(input)); + std::multiset> correct = { {1, 2, 2}, {2, 1, 2}, {2, 2, 1}, diff --git a/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc b/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc index 3ec51ec2f6..5e75a7a25d 100644 --- a/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc +++ b/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc @@ -1,5 +1,5 @@ #include "utils/containers/get_all_permutations_with_repetition.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/vector.h" #include "utils/hash/vector.h" #include @@ -12,9 +12,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("output vector has only one element") { std::vector input = {1, 2, 3}; - std::unordered_multiset> result = + std::multiset> result = get_all_permutations_with_repetition(input, 1_n); - std::unordered_multiset> correct = { + std::multiset> correct = { {1}, {2}, {3}, @@ -26,9 +26,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("input vector has only one element") { std::vector input = {1}; - std::unordered_multiset> result = + std::multiset> result = get_all_permutations_with_repetition(input, 2_n); - std::unordered_multiset> correct = { + std::multiset> correct = { {1, 1}, }; @@ -38,9 +38,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("input, output vectors have more than 1 element") { std::vector input = {1, 2}; - std::unordered_multiset> result = + std::multiset> result = get_all_permutations_with_repetition(input, 3_n); - std::unordered_multiset> correct = { + std::multiset> correct = { {1, 1, 1}, {1, 1, 2}, {1, 2, 1}, @@ -57,9 +57,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("duplicate elements") { std::vector input = {1, 2, 2}; - std::unordered_multiset> result = + std::multiset> result = get_all_permutations_with_repetition(input, 2_n); - std::unordered_multiset> correct = {{1, 1}, + std::multiset> correct = {{1, 1}, {1, 2}, {1, 2}, {2, 1}, @@ -75,9 +75,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("n == 0") { std::vector input = {1, 2, 3}; - std::unordered_multiset> result = + std::multiset> result = get_all_permutations_with_repetition(input, 0_n); - std::unordered_multiset> correct = {{}}; + std::multiset> correct = {{}}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/get_element_counts.cc b/lib/utils/test/src/utils/containers/get_element_counts.cc index 8fc87dba90..e9bd4c55dd 100644 --- a/lib/utils/test/src/utils/containers/get_element_counts.cc +++ b/lib/utils/test/src/utils/containers/get_element_counts.cc @@ -1,5 +1,5 @@ #include "utils/containers/get_element_counts.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; @@ -7,8 +7,8 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_element_counts") { std::vector input = {1, 2, 3, 2, 3, 3, 2, 3}; - std::unordered_map result = get_element_counts(input); - std::unordered_map correct = {{1, 1}, {2, 3}, {3, 4}}; + std::map result = get_element_counts(input); + std::map correct = {{1, 1}, {2, 3}, {3, 4}}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/get_one_of.cc b/lib/utils/test/src/utils/containers/get_one_of.cc index 326a292560..8ca3ee7807 100644 --- a/lib/utils/test/src/utils/containers/get_one_of.cc +++ b/lib/utils/test/src/utils/containers/get_one_of.cc @@ -1,18 +1,18 @@ #include "utils/containers/get_one_of.h" #include "utils/containers/contains.h" #include -#include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_one_of") { SUBCASE("non-empty set") { - std::unordered_set s = {1, 2, 3}; + std::set s = {1, 2, 3}; CHECK(contains(s, get_one_of(s))); } SUBCASE("empty set") { - std::unordered_set s = {}; + std::set s = {}; CHECK_THROWS(get_one_of(s)); } } diff --git a/lib/utils/test/src/utils/containers/get_only.cc b/lib/utils/test/src/utils/containers/get_only.cc index 675236e539..88c1f65437 100644 --- a/lib/utils/test/src/utils/containers/get_only.cc +++ b/lib/utils/test/src/utils/containers/get_only.cc @@ -12,7 +12,7 @@ TEST_SUITE(FF_TEST_SUITE) { } TEST_CASE("get_only") { - std::unordered_set input = {5}; + std::set input = {5}; int result = get_only(input); int correct = 5; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/containers/group_by.cc b/lib/utils/test/src/utils/containers/group_by.cc index 77eb6ea78a..1f9848260a 100644 --- a/lib/utils/test/src/utils/containers/group_by.cc +++ b/lib/utils/test/src/utils/containers/group_by.cc @@ -1,15 +1,15 @@ #include "utils/containers/group_by.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_map.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("group_by(std::unordered_set, F)") { - std::unordered_set input = {0, 3, 2, 9, 8}; + TEST_CASE("group_by(std::set, F)") { + std::set input = {0, 3, 2, 9, 8}; OneToMany result = group_by(input, [](int x) { return x % 3; }); OneToMany correct = { @@ -23,9 +23,9 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("group_by(std::vector, F)") { std::vector input = {0, 3, 0, 2, 2, 9, 8, 9}; - std::unordered_map> result = + std::map> result = group_by(input, [](int x) { return x % 3; }); - std::unordered_map> correct = { + std::map> correct = { {0, {0, 3, 0, 9, 9}}, {2, {2, 2, 8}}, }; diff --git a/lib/utils/test/src/utils/containers/inplace_filter.cc b/lib/utils/test/src/utils/containers/inplace_filter.cc index ac430279b0..43c3c01444 100644 --- a/lib/utils/test/src/utils/containers/inplace_filter.cc +++ b/lib/utils/test/src/utils/containers/inplace_filter.cc @@ -1,8 +1,8 @@ #include "utils/containers/inplace_filter.h" #include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_map.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "test/utils/rapidcheck.h" #include @@ -13,9 +13,9 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE_TEMPLATE("inplace_filter(T, F)", T, std::vector, - std::unordered_set, std::set, - std::unordered_map, + std::set, + std::map, std::map) { RC_SUBCASE("inplace_filter returns empty for predicate always_false", [](T t) { @@ -42,12 +42,12 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(input == correct); } - TEST_CASE("inplace_filter(std::unordered_set &, F)") { - std::unordered_set input = {1, 2, 3, 4, 5, 6, 7, 8}; + TEST_CASE("inplace_filter(std::set &, F)") { + std::set input = {1, 2, 3, 4, 5, 6, 7, 8}; auto predicate = [](int x) { return x % 2 == 0; }; inplace_filter(input, predicate); - std::unordered_set correct = {2, 4, 6, 8}; + std::set correct = {2, 4, 6, 8}; CHECK(input == correct); } @@ -60,8 +60,8 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(input == correct); } - TEST_CASE("inplace_filter(std::unordered_map &, F)") { - std::unordered_map input = { + TEST_CASE("inplace_filter(std::map &, F)") { + std::map input = { {3, "4"}, {1, "1"}, {2, "9"}, @@ -72,7 +72,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; inplace_filter(input, predicate); - std::unordered_map correct = { + std::map correct = { {1, "1"}, {4, "4"}, }; diff --git a/lib/utils/test/src/utils/containers/is_submapeq_of.cc b/lib/utils/test/src/utils/containers/is_submapeq_of.cc index df89444235..f84672c38b 100644 --- a/lib/utils/test/src/utils/containers/is_submapeq_of.cc +++ b/lib/utils/test/src/utils/containers/is_submapeq_of.cc @@ -1,38 +1,38 @@ #include "utils/containers/is_submapeq_of.h" #include #include -#include +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("is_submapeq_of") { - std::unordered_map super = { + std::map super = { {1, "one"}, {2, "two"}, {3, "three"}}; SUBCASE("keys and values match") { - std::unordered_map sub = {{1, "one"}, {2, "two"}}; + std::map sub = {{1, "one"}, {2, "two"}}; CHECK(is_submapeq_of(sub, super)); } SUBCASE("keys and values don't match") { - std::unordered_map sub = {{1, "one"}, {4, "four"}}; + std::map sub = {{1, "one"}, {4, "four"}}; CHECK_FALSE(is_submapeq_of(sub, super)); } SUBCASE("keys match but values don't") { - std::unordered_map sub = {{1, "wrong_value"}, + std::map sub = {{1, "wrong_value"}, {2, "two"}}; CHECK_FALSE(is_submapeq_of(sub, super)); } SUBCASE("values match but keys don't") { - std::unordered_map sub = {{5, "one"}, {6, "two"}}; + std::map sub = {{5, "one"}, {6, "two"}}; CHECK_FALSE(is_submapeq_of(sub, super)); } SUBCASE("sub is a superset of super") { - std::unordered_map sub = { + std::map sub = { {1, "one"}, {2, "two"}, {3, "three"}, {4, "four"}}; CHECK_FALSE(is_submapeq_of(sub, super)); } diff --git a/lib/utils/test/src/utils/containers/is_subseteq_of.cc b/lib/utils/test/src/utils/containers/is_subseteq_of.cc index d762f171b6..ecd28ff113 100644 --- a/lib/utils/test/src/utils/containers/is_subseteq_of.cc +++ b/lib/utils/test/src/utils/containers/is_subseteq_of.cc @@ -1,13 +1,13 @@ #include "utils/containers/is_subseteq_of.h" #include -#include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("is_subseteq_of") { - std::unordered_set s1 = {1, 2}; - std::unordered_set s2 = {1, 2, 3}; + std::set s1 = {1, 2}; + std::set s2 = {1, 2, 3}; CHECK(is_subseteq_of(s1, s2) == true); CHECK(is_subseteq_of(s2, s1) == false); CHECK(is_subseteq_of(s1, s1) == true); diff --git a/lib/utils/test/src/utils/containers/is_superseteq_of.cc b/lib/utils/test/src/utils/containers/is_superseteq_of.cc index e3b429fa64..b82b4f5572 100644 --- a/lib/utils/test/src/utils/containers/is_superseteq_of.cc +++ b/lib/utils/test/src/utils/containers/is_superseteq_of.cc @@ -1,20 +1,20 @@ #include "utils/containers/is_superseteq_of.h" #include -#include +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("is_superseteq_of") { - std::unordered_set super = {1, 2, 3, 4}; + std::set super = {1, 2, 3, 4}; SUBCASE("true containment") { - std::unordered_set sub = {1, 2, 3}; + std::set sub = {1, 2, 3}; CHECK(is_superseteq_of(super, sub)); } SUBCASE("false containment") { - std::unordered_set sub = {1, 2, 5}; + std::set sub = {1, 2, 5}; CHECK_FALSE(is_superseteq_of(super, sub)); } diff --git a/lib/utils/test/src/utils/containers/keys.cc b/lib/utils/test/src/utils/containers/keys.cc index d2ac3dbfba..ce455a608c 100644 --- a/lib/utils/test/src/utils/containers/keys.cc +++ b/lib/utils/test/src/utils/containers/keys.cc @@ -2,8 +2,8 @@ #include "test/utils/doctest/fmt/set.h" #include #include -#include -#include +#include +#include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/containers/lift_optional_through_map.cc b/lib/utils/test/src/utils/containers/lift_optional_through_map.cc index 7d1c4820c7..e9a84ca392 100644 --- a/lib/utils/test/src/utils/containers/lift_optional_through_map.cc +++ b/lib/utils/test/src/utils/containers/lift_optional_through_map.cc @@ -1,6 +1,6 @@ #include "utils/containers/lift_optional_through_map.h" #include "test/utils/doctest/fmt/optional.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; @@ -8,7 +8,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("lift_optional_through_map") { SUBCASE("throws if only some of the values are nullopt") { - std::unordered_map> input = { + std::map> input = { {1, std::nullopt}, {2, "two"}, }; @@ -17,31 +17,31 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("returns nullopt if all of the values are nullopt") { - std::unordered_map> input = { + std::map> input = { {1, std::nullopt}, {2, std::nullopt}, }; - std::optional> result = + std::optional> result = lift_optional_through_map(input); - std::optional> correct = + std::optional> correct = std::nullopt; CHECK(result == correct); } SUBCASE("returns the map if all of the values are not nullopt") { - std::unordered_map> input = { + std::map> input = { {1, "one"}, {2, "two"}, }; - std::optional> result = + std::optional> result = lift_optional_through_map(input); - std::optional> correct = - std::unordered_map{ + std::optional> correct = + std::map{ {1, "one"}, {2, "two"}, }; @@ -50,7 +50,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("throws if the input is an empty map") { - std::unordered_map> input = {}; + std::map> input = {}; CHECK_THROWS(lift_optional_through_map(input)); } diff --git a/lib/utils/test/src/utils/containers/lookup_in_map.cc b/lib/utils/test/src/utils/containers/lookup_in_map.cc index 9ca356ee4b..9bd8b19da5 100644 --- a/lib/utils/test/src/utils/containers/lookup_in_map.cc +++ b/lib/utils/test/src/utils/containers/lookup_in_map.cc @@ -9,7 +9,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("lookup_in_map") { - std::unordered_map map = {{"a", 1}, {"b", 2}}; + std::map map = {{"a", 1}, {"b", 2}}; SUBCASE("existing keys") { std::function func = lookup_in_map(map); @@ -23,7 +23,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("empty map") { - std::unordered_map map = {}; + std::map map = {}; std::function func = lookup_in_map(map); CHECK_THROWS(func("a")); } diff --git a/lib/utils/test/src/utils/containers/map_keys.cc b/lib/utils/test/src/utils/containers/map_keys.cc index 5c0a81d5e6..47ba041b18 100644 --- a/lib/utils/test/src/utils/containers/map_keys.cc +++ b/lib/utils/test/src/utils/containers/map_keys.cc @@ -1,23 +1,23 @@ #include "utils/containers/map_keys.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include -#include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("map_keys") { SUBCASE("Distinct keys after transformation") { - std::unordered_map m = {{1, "one"}, {2, "two"}}; + std::map m = {{1, "one"}, {2, "two"}}; auto f = [](int x) { return x * x; }; - std::unordered_map result = map_keys(m, f); - std::unordered_map correct = {{1, "one"}, {4, "two"}}; + std::map result = map_keys(m, f); + std::map correct = {{1, "one"}, {4, "two"}}; CHECK(correct == result); } SUBCASE("Non-distinct keys after transformation") { - std::unordered_map m = { + std::map m = { {1, "one"}, {2, "two"}, {-1, "minus one"}}; auto f = [](int x) { return std::abs(x); }; CHECK_THROWS(map_keys(m, f)); diff --git a/lib/utils/test/src/utils/containers/map_keys2.cc b/lib/utils/test/src/utils/containers/map_keys2.cc index 6e7110543b..206370f133 100644 --- a/lib/utils/test/src/utils/containers/map_keys2.cc +++ b/lib/utils/test/src/utils/containers/map_keys2.cc @@ -1,5 +1,5 @@ #include "utils/containers/map_keys2.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; @@ -7,7 +7,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("map_keys2") { SUBCASE("output keys are unique") { - std::unordered_map m = { + std::map m = { {1, "aa"}, {2, "aaaaa"}, }; @@ -16,9 +16,9 @@ TEST_SUITE(FF_TEST_SUITE) { return std::to_string(k + v.size()); }; - std::unordered_map result = map_keys2(m, f); + std::map result = map_keys2(m, f); - std::unordered_map correct = { + std::map correct = { {"3", "aa"}, {"7", "aaaaa"}, }; @@ -27,7 +27,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("output keys are non-unique") { - std::unordered_map m = { + std::map m = { {1, "aa"}, {2, "aaaaa"}, }; diff --git a/lib/utils/test/src/utils/containers/map_keys_and_values.cc b/lib/utils/test/src/utils/containers/map_keys_and_values.cc index d50ed82bfc..998e973513 100644 --- a/lib/utils/test/src/utils/containers/map_keys_and_values.cc +++ b/lib/utils/test/src/utils/containers/map_keys_and_values.cc @@ -1,5 +1,5 @@ #include "utils/containers/map_keys_and_values.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace FlexFlow; @@ -7,16 +7,16 @@ using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("map_keys_and_values") { SUBCASE("Distinct keys after transformation") { - std::unordered_map m = {{1, "one"}, {2, "three"}}; + std::map m = {{1, "one"}, {2, "three"}}; auto fk = [](int x) { return x * x; }; auto fv = [](std::string const &s) { return s.size(); }; - std::unordered_map result = map_keys_and_values(m, fk, fv); - std::unordered_map correct = {{1, 3}, {4, 5}}; + std::map result = map_keys_and_values(m, fk, fv); + std::map correct = {{1, 3}, {4, 5}}; CHECK(correct == result); } SUBCASE("Non-distinct keys after transformation") { - std::unordered_map m = { + std::map m = { {1, "one"}, {2, "two"}, {-1, "minus one"}}; auto fk = [](int x) { return std::abs(x); }; auto fv = [](std::string const &s) { return s.size(); }; diff --git a/lib/utils/test/src/utils/containers/map_values.cc b/lib/utils/test/src/utils/containers/map_values.cc index a21645d0d5..6ed3e8cb81 100644 --- a/lib/utils/test/src/utils/containers/map_values.cc +++ b/lib/utils/test/src/utils/containers/map_values.cc @@ -1,17 +1,17 @@ #include "utils/containers/map_values.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include -#include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("map_values") { - std::unordered_map m = {{1, "one"}, {3, "three"}}; + std::map m = {{1, "one"}, {3, "three"}}; auto f = [](std::string const &s) { return s.size(); }; - std::unordered_map result = map_values(m, f); - std::unordered_map correct = {{1, 3}, {3, 5}}; + std::map result = map_values(m, f); + std::map correct = {{1, 3}, {3, 5}}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/map_values2.cc b/lib/utils/test/src/utils/containers/map_values2.cc index 5dd0fa9883..07783c1dcd 100644 --- a/lib/utils/test/src/utils/containers/map_values2.cc +++ b/lib/utils/test/src/utils/containers/map_values2.cc @@ -1,5 +1,5 @@ #include "utils/containers/map_values2.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include @@ -7,7 +7,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("map_values2") { - std::unordered_map m = { + std::map m = { {1, "aa"}, {2, "aaaaa"}, {4, "bbb"}, @@ -15,9 +15,9 @@ TEST_SUITE(FF_TEST_SUITE) { auto f = [](int k, std::string const &v) -> int { return k + v.size(); }; - std::unordered_map result = map_values2(m, f); + std::map result = map_values2(m, f); - std::unordered_map correct = { + std::map correct = { {1, 3}, {2, 7}, {4, 7}, diff --git a/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc index ceecea96c1..bf4b2202d7 100644 --- a/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc +++ b/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc @@ -1,37 +1,37 @@ -#include "utils/containers/merge_disjoint_unordered_maps.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/containers/merge_disjoint_maps.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("merge_disjoint_unordered_maps") { - std::unordered_map m1 = { + TEST_CASE("merge_disjoint_maps") { + std::map m1 = { {4, "four"}, {2, "two"}, }; - std::unordered_map m2 = { + std::map m2 = { {3, "four"}, }; - std::unordered_map m3 = { + std::map m3 = { {1, "one"}, }; - std::unordered_map m4 = {}; + std::map m4 = {}; SUBCASE("maps are disjoint") { - std::vector> input = { + std::vector> input = { m1, m2, m3, m4, }; - std::unordered_map result = merge_disjoint_unordered_maps(input); + std::map result = merge_disjoint_maps(input); - std::unordered_map correct = { + std::map correct = { {4, "four"}, {2, "two"}, {3, "four"}, @@ -42,12 +42,12 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("maps are not disjoint") { - std::unordered_map m5 = { + std::map m5 = { {4, "five"}, {6, "six"}, }; - std::vector> input = { + std::vector> input = { m1, m2, m3, @@ -55,16 +55,16 @@ TEST_SUITE(FF_TEST_SUITE) { m5, }; - CHECK_THROWS(merge_disjoint_unordered_maps(input)); + CHECK_THROWS(merge_disjoint_maps(input)); } SUBCASE("maps are not disjoint but have identical values") { - std::unordered_map m5 = { + std::map m5 = { {4, "four"}, {6, "six"}, }; - std::vector> input = { + std::vector> input = { m1, m2, m3, @@ -72,7 +72,7 @@ TEST_SUITE(FF_TEST_SUITE) { m5, }; - CHECK_THROWS(merge_disjoint_unordered_maps(input)); + CHECK_THROWS(merge_disjoint_maps(input)); } } } diff --git a/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc b/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc index 66827ca453..fd73a9345e 100644 --- a/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc +++ b/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc @@ -1,50 +1,50 @@ -#include "utils/containers/merge_unordered_maps_with.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/containers/merge_maps_with.h" +#include "test/utils/doctest/fmt/map.h" #include "test/utils/rapidcheck.h" -#include "utils/containers/binary_merge_unordered_maps_with.h" +#include "utils/containers/binary_merge_maps_with.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("merge_unordered_maps_with") { + TEST_CASE("merge_maps_with") { auto string_concat = [](std::string const &l, std::string const &r) { return l + r; }; RC_SUBCASE( - "with two inputs, matches binary_merge_unordered_maps_with", - [&](std::unordered_map const &lhs, - std::unordered_map const &rhs) { - std::unordered_map from_merge_unordered_maps_with = - merge_unordered_maps_with(std::vector{lhs, rhs}, string_concat); + "with two inputs, matches binary_merge_maps_with", + [&](std::map const &lhs, + std::map const &rhs) { + std::map from_merge_maps_with = + merge_maps_with(std::vector{lhs, rhs}, string_concat); - std::unordered_map from_binary_merge_unordered_maps_with = - binary_merge_unordered_maps_with(lhs, rhs, string_concat); + std::map from_binary_merge_maps_with = + binary_merge_maps_with(lhs, rhs, string_concat); - CHECK(from_merge_unordered_maps_with == from_binary_merge_unordered_maps_with); + CHECK(from_merge_maps_with == from_binary_merge_maps_with); }); SUBCASE("maps overlap") { - std::unordered_map map1 = { + std::map map1 = { {1, "map1_one."}, {4, "map1_four."}, }; - std::unordered_map map2 = { + std::map map2 = { {2, "map2_two."}, {4, "map2_four."}, {5, "map2_five."}, }; - std::unordered_map map3 = { + std::map map3 = { {1, "map3_one."}, }; - std::unordered_map result = - merge_unordered_maps_with(std::vector{map1, map2, map3}, string_concat); + std::map result = + merge_maps_with(std::vector{map1, map2, map3}, string_concat); - std::unordered_map correct = { + std::map correct = { {1, "map1_one.map3_one."}, {2, "map2_two."}, {4, "map1_four.map2_four."}, @@ -58,24 +58,24 @@ TEST_SUITE(FF_TEST_SUITE) { std::string const &) -> std::string { PANIC(); }; SUBCASE("maps do not overlap") { - std::unordered_map map1 = { + std::map map1 = { {8, "map1_eight."}, {4, "map1_four."}, }; - std::unordered_map map2 = { + std::map map2 = { {2, "map2_two."}, {5, "map2_five."}, }; - std::unordered_map map3 = { + std::map map3 = { {1, "map3_one."}, }; - std::unordered_map result = - merge_unordered_maps_with(std::vector{map1, map2, map3}, fail_if_called); + std::map result = + merge_maps_with(std::vector{map1, map2, map3}, fail_if_called); - std::unordered_map correct = { + std::map correct = { {1, "map3_one."}, {2, "map2_two."}, {4, "map1_four."}, @@ -87,12 +87,12 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("no maps are provided") { - std::vector> maps = {}; + std::vector> maps = {}; - std::unordered_map result = - merge_unordered_maps_with(maps, fail_if_called); + std::map result = + merge_maps_with(maps, fail_if_called); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/multiset_union.cc b/lib/utils/test/src/utils/containers/multiset_union.cc index 8c40bf55ab..7eeb24df99 100644 --- a/lib/utils/test/src/utils/containers/multiset_union.cc +++ b/lib/utils/test/src/utils/containers/multiset_union.cc @@ -1,18 +1,18 @@ #include "utils/containers/multiset_union.h" #include "test/utils/doctest/fmt/multiset.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("multiset_union(std::unordered_multiset, " - "std::unordered_multiset)") { - std::unordered_multiset input_lhs = {1, 2, 2, 3}; - std::unordered_multiset input_rhs = {1, 2, 5}; + TEST_CASE("multiset_union(std::multiset, " + "std::multiset)") { + std::multiset input_lhs = {1, 2, 2, 3}; + std::multiset input_rhs = {1, 2, 5}; - std::unordered_multiset result = multiset_union(input_lhs, input_rhs); - std::unordered_multiset correct = {1, 1, 2, 2, 2, 3, 5}; + std::multiset result = multiset_union(input_lhs, input_rhs); + std::multiset correct = {1, 1, 2, 2, 2, 3, 5}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/permute_with_key.cc b/lib/utils/test/src/utils/containers/permute_with_key.cc index c7dab61b4d..1b5df88175 100644 --- a/lib/utils/test/src/utils/containers/permute_with_key.cc +++ b/lib/utils/test/src/utils/containers/permute_with_key.cc @@ -1,10 +1,10 @@ #include "utils/containers/permute_with_key.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "test/utils/rapidcheck/doctest.h" #include "utils/containers/get_all_permutations.h" #include "utils/containers/range.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/hash/vector.h" #include @@ -21,12 +21,12 @@ TEST_SUITE(FF_TEST_SUITE) { }; int max_permutations = 4 * 3 * 2 * 1; - std::unordered_set> generated_permutations = - unordered_set_of(transform(range(max_permutations), [&](int key) { + std::set> generated_permutations = + set_of(transform(range(max_permutations), [&](int key) { return permute_with_key(key, input); })); - std::unordered_set> all_permutations = - unordered_set_of(get_all_permutations(input)); + std::set> all_permutations = + set_of(get_all_permutations(input)); CHECK(generated_permutations == all_permutations); } diff --git a/lib/utils/test/src/utils/containers/product.cc b/lib/utils/test/src/utils/containers/product.cc index 2278bfba17..ff6c515870 100644 --- a/lib/utils/test/src/utils/containers/product.cc +++ b/lib/utils/test/src/utils/containers/product.cc @@ -3,7 +3,7 @@ #include #include #include -#include +#include #include using namespace ::FlexFlow; @@ -15,7 +15,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector, std::vector, std::set, - std::unordered_set) { + std::set) { SUBCASE("non-empty container") { C input = {1, -2, 3, 5}; diff --git a/lib/utils/test/src/utils/containers/range.cc b/lib/utils/test/src/utils/containers/range.cc index a5b9481860..95439a0598 100644 --- a/lib/utils/test/src/utils/containers/range.cc +++ b/lib/utils/test/src/utils/containers/range.cc @@ -1,8 +1,8 @@ #include "utils/containers/range.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include -#include +#include #include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/containers/repeat_element.cc b/lib/utils/test/src/utils/containers/repeat_element.cc index 08bee8bec8..ba90df570d 100644 --- a/lib/utils/test/src/utils/containers/repeat_element.cc +++ b/lib/utils/test/src/utils/containers/repeat_element.cc @@ -1,8 +1,8 @@ #include "utils/containers/repeat_element.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include -#include +#include using namespace FlexFlow; @@ -14,11 +14,11 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector correct = {42, 42, 42, 42, 42}; CHECK(result == correct); } - SUBCASE("unordered_set") { - std::unordered_set x = {1.0, 1.5}; - std::vector> result = + SUBCASE("set") { + std::set x = {1.0, 1.5}; + std::vector> result = repeat_element(nonnegative_int{3}, x); - std::vector> correct = { + std::vector> correct = { {1.0, 1.5}, {1.0, 1.5}, {1.0, 1.5}}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/require_all_same1.cc b/lib/utils/test/src/utils/containers/require_all_same1.cc index 4c49124077..6e4c24a3cd 100644 --- a/lib/utils/test/src/utils/containers/require_all_same1.cc +++ b/lib/utils/test/src/utils/containers/require_all_same1.cc @@ -3,14 +3,14 @@ #include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/optional.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/multiset.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "utils/expected.h" #include #include #include -#include +#include using namespace ::FlexFlow; @@ -18,8 +18,8 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE_TEMPLATE("require_all_same1(T)", T, std::vector, - std::unordered_set, - std::unordered_multiset, + std::set, + std::multiset, std::set, std::multiset) { SUBCASE("input is empty") { diff --git a/lib/utils/test/src/utils/containers/require_no_duplicates.cc b/lib/utils/test/src/utils/containers/require_no_duplicates.cc index 67733d791a..b84c093e0c 100644 --- a/lib/utils/test/src/utils/containers/require_no_duplicates.cc +++ b/lib/utils/test/src/utils/containers/require_no_duplicates.cc @@ -1,34 +1,34 @@ #include "utils/containers/require_no_duplicates.h" #include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/multiset.h" +#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("require_no_duplicates(std::unordered_multiset)") { + TEST_CASE("require_no_duplicates(std::multiset)") { SUBCASE("empty") { - std::unordered_multiset input = {}; + std::multiset input = {}; - std::unordered_set result = require_no_duplicates(input); - std::unordered_set correct = {}; + std::set result = require_no_duplicates(input); + std::set correct = {}; CHECK(result == correct); } SUBCASE("input has duplicates") { - std::unordered_multiset input = {1, 2, 2}; + std::multiset input = {1, 2, 2}; CHECK_THROWS(require_no_duplicates(input)); } SUBCASE("input does not have duplicates") { - std::unordered_multiset input = {1, 2, 4}; + std::multiset input = {1, 2, 4}; - std::unordered_set result = require_no_duplicates(input); - std::unordered_set correct = {1, 2, 4}; + std::set result = require_no_duplicates(input); + std::set correct = {1, 2, 4}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/require_only_key.cc b/lib/utils/test/src/utils/containers/require_only_key.cc index 3b30edcdf2..3767c8ded7 100644 --- a/lib/utils/test/src/utils/containers/require_only_key.cc +++ b/lib/utils/test/src/utils/containers/require_only_key.cc @@ -6,7 +6,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("require_only_key") { SUBCASE("input has one key that matches") { - std::unordered_map m = {{2, "a"}}; + std::map m = {{2, "a"}}; std::string result = require_only_key(m, 2); std::string correct = "a"; @@ -15,19 +15,19 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("input has one key that does not match") { - std::unordered_map m = {{3, "a"}}; + std::map m = {{3, "a"}}; CHECK_THROWS(require_only_key(m, 2)); } SUBCASE("input is empty") { - std::unordered_map m = {}; + std::map m = {}; CHECK_THROWS(require_only_key(m, 2)); } SUBCASE("input has more than one key") { - std::unordered_map m = { + std::map m = { {2, "a"}, {3, "b"}, }; diff --git a/lib/utils/test/src/utils/containers/require_two_keys.cc b/lib/utils/test/src/utils/containers/require_two_keys.cc index ebcd799849..c2728c22fc 100644 --- a/lib/utils/test/src/utils/containers/require_two_keys.cc +++ b/lib/utils/test/src/utils/containers/require_two_keys.cc @@ -7,7 +7,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("require_two_keys") { SUBCASE("input is too small") { - std::unordered_map m = { + std::map m = { {2, "a"}, }; @@ -15,7 +15,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("input is too large") { - std::unordered_map m = { + std::map m = { {2, "a"}, {3, "b"}, {4, "c"}, @@ -25,7 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("input is correct size, but keys don't match") { - std::unordered_map m = { + std::map m = { {2, "a"}, {4, "c"}, }; @@ -34,7 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("input is correct size and both keys are the same") { - std::unordered_map m = { + std::map m = { {2, "a"}, {3, "b"}, }; @@ -43,7 +43,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("input is correct size and keys match") { - std::unordered_map m = { + std::map m = { {2, "a"}, {4, "c"}, }; diff --git a/lib/utils/test/src/utils/containers/restrict_keys.cc b/lib/utils/test/src/utils/containers/restrict_keys.cc index b4b376784e..2b78a59468 100644 --- a/lib/utils/test/src/utils/containers/restrict_keys.cc +++ b/lib/utils/test/src/utils/containers/restrict_keys.cc @@ -1,5 +1,5 @@ #include "utils/containers/restrict_keys.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include @@ -7,11 +7,11 @@ using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("restrict_keys") { - std::unordered_map m = { + std::map m = { {1, "one"}, {2, "two"}, {3, "three"}}; - std::unordered_set mask = {2, 3, 4}; - std::unordered_map result = restrict_keys(m, mask); - std::unordered_map correct = {{2, "two"}, {3, "three"}}; + std::set mask = {2, 3, 4}; + std::map result = restrict_keys(m, mask); + std::map correct = {{2, "two"}, {3, "three"}}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/set_intersection.cc b/lib/utils/test/src/utils/containers/set_intersection.cc index d206985792..69a3845ef8 100644 --- a/lib/utils/test/src/utils/containers/set_intersection.cc +++ b/lib/utils/test/src/utils/containers/set_intersection.cc @@ -1,14 +1,14 @@ #include "utils/containers/set_intersection.h" #include "test/utils/doctest/fmt/optional.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE_TEMPLATE( - "set_intersection(S, S)", S, std::unordered_set, std::set) { + "set_intersection(S, S)", S, std::set, std::set) { S input_l = {1, 2, 3}; S input_r = {2, 3, 5}; @@ -19,7 +19,7 @@ TEST_SUITE(FF_TEST_SUITE) { } TEST_CASE_TEMPLATE( - "set_intersection(C)", S, std::unordered_set, std::set) { + "set_intersection(C)", S, std::set, std::set) { SUBCASE("input is empty container") { std::vector input = {}; diff --git a/lib/utils/test/src/utils/containers/set_union.cc b/lib/utils/test/src/utils/containers/set_union.cc index d842e4df96..4afeb0fd22 100644 --- a/lib/utils/test/src/utils/containers/set_union.cc +++ b/lib/utils/test/src/utils/containers/set_union.cc @@ -1,16 +1,16 @@ #include "utils/containers/set_union.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include -#include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("set_union") { - std::unordered_set s1 = {1, 2, 3}; - std::unordered_set s2 = {2, 3, 4}; - std::unordered_set result = set_union(s1, s2); - std::unordered_set correct = {1, 2, 3, 4}; + std::set s1 = {1, 2, 3}; + std::set s2 = {2, 3, 4}; + std::set result = set_union(s1, s2); + std::set correct = {1, 2, 3, 4}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/sorted_by.cc b/lib/utils/test/src/utils/containers/sorted_by.cc index 0ae2e0da77..6cc63f760d 100644 --- a/lib/utils/test/src/utils/containers/sorted_by.cc +++ b/lib/utils/test/src/utils/containers/sorted_by.cc @@ -1,7 +1,7 @@ #include "utils/containers/sorted_by.h" #include "test/utils/doctest/fmt/vector.h" #include -#include +#include #include using namespace FlexFlow; @@ -9,7 +9,7 @@ using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("sorted_by") { SUBCASE("sort increasing") { - std::unordered_set s = {5, 2, 3, 4, 1}; + std::set s = {5, 2, 3, 4, 1}; std::vector result = sorted_by(s, [](int a, int b) { return a < b; }); std::vector correct = {1, 2, 3, 4, 5}; @@ -17,7 +17,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("sort decreasing") { - std::unordered_set input = {-5, -1, -3, -2, -4}; + std::set input = {-5, -1, -3, -2, -4}; std::vector result = sorted_by(input, [](int a, int b) { return a > b; }); std::vector correct = {-1, -2, -3, -4, -5}; diff --git a/lib/utils/test/src/utils/containers/sum_where.cc b/lib/utils/test/src/utils/containers/sum_where.cc index 7a909aea39..fbb4cf6510 100644 --- a/lib/utils/test/src/utils/containers/sum_where.cc +++ b/lib/utils/test/src/utils/containers/sum_where.cc @@ -1,6 +1,6 @@ #include "utils/containers/sum_where.h" #include -#include +#include #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/transform.cc b/lib/utils/test/src/utils/containers/transform.cc index 3122c67117..1aa83dd834 100644 --- a/lib/utils/test/src/utils/containers/transform.cc +++ b/lib/utils/test/src/utils/containers/transform.cc @@ -1,6 +1,6 @@ #include "utils/containers/transform.h" #include "test/utils/doctest/fmt/optional.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include @@ -15,11 +15,11 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(result == correct); } - TEST_CASE("transform(std::unordered_set, F)") { - std::unordered_set input = {1, 2, 3}; - std::unordered_set result = + TEST_CASE("transform(std::set, F)") { + std::set input = {1, 2, 3}; + std::set result = transform(input, [](int x) { return std::to_string(x); }); - std::unordered_set correct = {"1", "2", "3"}; + std::set correct = {"1", "2", "3"}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/try_at.cc b/lib/utils/test/src/utils/containers/try_at.cc index 548c9b0c79..f0ff3a7004 100644 --- a/lib/utils/test/src/utils/containers/try_at.cc +++ b/lib/utils/test/src/utils/containers/try_at.cc @@ -8,7 +8,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE_TEMPLATE("try_at(T, K)", T, - std::unordered_map, + std::map, std::map) { T m = {{1, "one"}, {2, "two"}}; diff --git a/lib/utils/test/src/utils/containers/try_get_one_of.cc b/lib/utils/test/src/utils/containers/try_get_one_of.cc index 6cf2513926..19e20da549 100644 --- a/lib/utils/test/src/utils/containers/try_get_one_of.cc +++ b/lib/utils/test/src/utils/containers/try_get_one_of.cc @@ -6,9 +6,9 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("try_get_one_of(std::unordered_set)") { + TEST_CASE("try_get_one_of(std::set)") { SUBCASE("input is empty") { - std::unordered_set input = {}; + std::set input = {}; std::optional result = try_get_one_of(input); std::optional correct = std::nullopt; @@ -17,7 +17,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("input is non-empty") { - std::unordered_set input = {1, 2, 3}; + std::set input = {1, 2, 3}; std::optional result = try_get_one_of(input); diff --git a/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc index b8a7a85f74..6804ea0243 100644 --- a/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc +++ b/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc @@ -1,26 +1,26 @@ -#include "utils/containers/try_merge_nondisjoint_unordered_maps.h" +#include "utils/containers/try_merge_nondisjoint_maps.h" #include "test/utils/doctest/fmt/optional.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("try_merge_nondisjoing_unordered_maps(std::unordered_map, " - "std::unordered_map)") { - std::unordered_map d1 = { + TEST_CASE("try_merge_nondisjoing_maps(std::map, " + "std::map)") { + std::map d1 = { {0, "zero"}, {1, "one"}, }; - std::unordered_map d2 = { + std::map d2 = { {0, "zero"}, {2, "two"}, }; SUBCASE("compatible neither superset") { - std::optional> result = - try_merge_nondisjoint_unordered_maps(d1, d2); - std::optional> correct = {{ + std::optional> result = + try_merge_nondisjoint_maps(d1, d2); + std::optional> correct = {{ {0, "zero"}, {1, "one"}, {2, "two"}, @@ -30,18 +30,18 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("mismatched key") { d1.insert({2, "three"}); - std::optional> result = - try_merge_nondisjoint_unordered_maps(d1, d2); - std::optional> correct = + std::optional> result = + try_merge_nondisjoint_maps(d1, d2); + std::optional> correct = std::nullopt; CHECK(result == correct); } SUBCASE("repeated value") { d1.insert({3, "one"}); - std::optional> result = - try_merge_nondisjoint_unordered_maps(d1, d2); - std::optional> correct = {{ + std::optional> result = + try_merge_nondisjoint_maps(d1, d2); + std::optional> correct = {{ {0, "zero"}, {1, "one"}, {2, "two"}, @@ -52,26 +52,26 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("left superset") { d1.insert({2, "two"}); - std::optional> result = - try_merge_nondisjoint_unordered_maps(d1, d2); - std::optional> correct = d1; + std::optional> result = + try_merge_nondisjoint_maps(d1, d2); + std::optional> correct = d1; CHECK(result == correct); } SUBCASE("right superset") { d2.insert({1, "one"}); - std::optional> result = - try_merge_nondisjoint_unordered_maps(d1, d2); - std::optional> correct = d2; + std::optional> result = + try_merge_nondisjoint_maps(d1, d2); + std::optional> correct = d2; CHECK(result == correct); } SUBCASE("equal") { d1.insert({2, "two"}); d2.insert({1, "one"}); - std::optional> result = - try_merge_nondisjoint_unordered_maps(d1, d2); - std::optional> correct = d1; + std::optional> result = + try_merge_nondisjoint_maps(d1, d2); + std::optional> correct = d1; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc b/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc index f0cdb19611..13d9dabeb9 100644 --- a/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc +++ b/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc @@ -1,5 +1,5 @@ -#include "utils/containers/unordered_map_from_pairs.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/containers/map_from_pairs.h" +#include "test/utils/doctest/fmt/map.h" #include "utils/containers/contains.h" #include #include @@ -8,16 +8,16 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("unordered_map_from_pairs") { + TEST_CASE("map_from_pairs") { SUBCASE("nonempty input") { std::vector> input = { {1, "hello"}, {3, "world"}, }; - std::unordered_map result = - unordered_map_from_pairs(input); - std::unordered_map correct = { + std::map result = + map_from_pairs(input); + std::map correct = { {1, "hello"}, {3, "world"}, }; @@ -28,9 +28,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("empty input") { std::vector> input = {}; - std::unordered_map result = - unordered_map_from_pairs(input); - std::unordered_map correct = {}; + std::map result = + map_from_pairs(input); + std::map correct = {}; CHECK(result == correct); } @@ -42,10 +42,10 @@ TEST_SUITE(FF_TEST_SUITE) { {1, "b"}, }; - std::unordered_map result = - unordered_map_from_pairs(input); + std::map result = + map_from_pairs(input); - std::vector> + std::vector> possible_correct_values = { {{1, "a"}, {2, "c"}}, {{1, "b"}, {2, "c"}}, diff --git a/lib/utils/test/src/utils/containers/unordered_multiset_of.cc b/lib/utils/test/src/utils/containers/unordered_multiset_of.cc index becb7fdce0..d44979f655 100644 --- a/lib/utils/test/src/utils/containers/unordered_multiset_of.cc +++ b/lib/utils/test/src/utils/containers/unordered_multiset_of.cc @@ -1,15 +1,15 @@ -#include "utils/containers/unordered_multiset_of.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "utils/containers/multiset_of.h" +#include "test/utils/doctest/fmt/multiset.h" #include #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("unordered_multiset_of") { + TEST_CASE("multiset_of") { std::vector input = {1, 2, 3, 3, 2, 3}; - std::unordered_multiset result = unordered_multiset_of(input); - std::unordered_multiset correct = {1, 2, 3, 3, 2, 3}; + std::multiset result = multiset_of(input); + std::multiset correct = {1, 2, 3, 3, 2, 3}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/unordered_set_of.cc b/lib/utils/test/src/utils/containers/unordered_set_of.cc index b8ca1d1797..69b762c6ab 100644 --- a/lib/utils/test/src/utils/containers/unordered_set_of.cc +++ b/lib/utils/test/src/utils/containers/unordered_set_of.cc @@ -1,15 +1,15 @@ -#include "utils/containers/unordered_set_of.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "utils/containers/set_of.h" +#include "test/utils/doctest/fmt/set.h" #include #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("unordered_set_of") { + TEST_CASE("set_of") { std::vector input = {1, 2, 3, 3, 2, 3}; - std::unordered_set result = unordered_set_of(input); - std::unordered_set correct = {1, 2, 3}; + std::set result = set_of(input); + std::set correct = {1, 2, 3}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/containers/unstructured_exhaustive_relational_join.cc b/lib/utils/test/src/utils/containers/unstructured_exhaustive_relational_join.cc index d756dababa..84e770bc78 100644 --- a/lib/utils/test/src/utils/containers/unstructured_exhaustive_relational_join.cc +++ b/lib/utils/test/src/utils/containers/unstructured_exhaustive_relational_join.cc @@ -1,6 +1,6 @@ #include "utils/containers/unstructured_exhaustive_relational_join.h" #include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/hash/pair.h" #include #include @@ -10,7 +10,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("unstructured_exhaustive_relational_join") { SUBCASE("join is exhaustive") { - std::unordered_set> lhs = { + std::set> lhs = { {1, "one"}, {1, "odd"}, {2, "two"}, @@ -18,16 +18,16 @@ TEST_SUITE(FF_TEST_SUITE) { {3, "odd"}, }; - std::unordered_set> rhs = { + std::set> rhs = { {"one", false}, {"odd", true}, {"two", true}, {"three", true}, }; - std::unordered_set> result = + std::set> result = unstructured_exhaustive_relational_join(lhs, rhs); - std::unordered_set> correct = { + std::set> correct = { {1, false}, {1, true}, {2, true}, @@ -38,14 +38,14 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("join is not exhaustive in lhs") { - std::unordered_set> lhs = { + std::set> lhs = { {1, "one"}, {1, "odd"}, {2, "two"}, {3, "odd"}, }; - std::unordered_set> rhs = { + std::set> rhs = { {"one", false}, {"odd", true}, {"two", true}, @@ -56,7 +56,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("join is not exhaustive in rhs") { - std::unordered_set> lhs = { + std::set> lhs = { {1, "one"}, {1, "odd"}, {2, "two"}, @@ -64,7 +64,7 @@ TEST_SUITE(FF_TEST_SUITE) { {3, "odd"}, }; - std::unordered_set> rhs = { + std::set> rhs = { {"one", false}, {"odd", true}, {"two", true}, diff --git a/lib/utils/test/src/utils/containers/values.cc b/lib/utils/test/src/utils/containers/values.cc index 5fe69ac5e9..e98882e4e9 100644 --- a/lib/utils/test/src/utils/containers/values.cc +++ b/lib/utils/test/src/utils/containers/values.cc @@ -1,18 +1,18 @@ #include "utils/containers/values.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include #include -#include +#include #include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("values") { - std::unordered_map m = { + std::map m = { {1, "one"}, {2, "two"}, {3, "three"}, {33, "three"}}; - std::unordered_multiset result = values(m); - std::unordered_multiset correct = { + std::multiset result = values(m); + std::multiset correct = { "one", "two", "three", "three"}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/zip_values_strict.cc b/lib/utils/test/src/utils/containers/zip_values_strict.cc index 2ba6890b9e..c778138dee 100644 --- a/lib/utils/test/src/utils/containers/zip_values_strict.cc +++ b/lib/utils/test/src/utils/containers/zip_values_strict.cc @@ -1,5 +1,5 @@ #include "utils/containers/zip_values_strict.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "test/utils/doctest/fmt/map.h" #include #include @@ -8,18 +8,18 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("zip_values_strict") { SUBCASE("key sets are the same") { - std::unordered_map m1 = { + std::map m1 = { {2, "two"}, {3, "three"}, }; - std::unordered_map m2 = { + std::map m2 = { {2, "TWO"}, {3, "THREE"}, }; - std::unordered_map> result = + std::map> result = zip_values_strict(m1, m2); - std::unordered_map> correct = { + std::map> correct = { {2, {"two", "TWO"}}, {3, {"three", "THREE"}}, }; @@ -28,11 +28,11 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("key sets are different but same size") { - std::unordered_map m1 = { + std::map m1 = { {2, "two"}, {3, "three"}, }; - std::unordered_map m2 = { + std::map m2 = { {2, "TWO"}, {4, "FOUR"}, }; @@ -41,11 +41,11 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("key sets are subset") { - std::unordered_map m1 = { + std::map m1 = { {2, "two"}, {3, "three"}, }; - std::unordered_map m2 = { + std::map m2 = { {2, "TWO"}, }; diff --git a/lib/utils/test/src/utils/dot/dot_file.cc b/lib/utils/test/src/utils/dot/dot_file.cc index 05720593dd..040982e858 100644 --- a/lib/utils/test/src/utils/dot/dot_file.cc +++ b/lib/utils/test/src/utils/dot/dot_file.cc @@ -67,8 +67,8 @@ TEST_SUITE(FF_TEST_SUITE) { std::string expectedOutput = R"EXPECTED_OUTPUT(digraph taskgraph { subgraph cluster_0 { -node1; node0; +node1; subgraph cluster_1 { node1; } diff --git a/lib/utils/test/src/utils/fmt/unordered_map.cc b/lib/utils/test/src/utils/fmt/unordered_map.cc index c980bc1e52..b82a13b6d2 100644 --- a/lib/utils/test/src/utils/fmt/unordered_map.cc +++ b/lib/utils/test/src/utils/fmt/unordered_map.cc @@ -1,18 +1,18 @@ -#include "utils/fmt/unordered_map.h" -#include "test/utils/doctest/fmt/unordered_map.h" +#include "utils/fmt/map.h" +#include "test/utils/doctest/fmt/map.h" #include "utils/containers/get_element_counts.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("fmt::to_string(std::unordered_map)") { - std::unordered_map input = {{0, 10}, {1, 1}, {3, 5}, {2, 8}}; + TEST_CASE("fmt::to_string(std::map)") { + std::map input = {{0, 10}, {1, 1}, {3, 5}, {2, 8}}; std::string result = fmt::to_string(input); std::string correct = "{{0, 10}, {1, 1}, {2, 8}, {3, 5}}"; - std::unordered_map result_char_counts = + std::map result_char_counts = get_element_counts(result); - std::unordered_map correct_char_counts = + std::map correct_char_counts = get_element_counts(correct); CHECK(result_char_counts == correct_char_counts); } diff --git a/lib/utils/test/src/utils/fmt/unordered_set.cc b/lib/utils/test/src/utils/fmt/unordered_set.cc index f4cf985d0f..3322e56ea4 100644 --- a/lib/utils/test/src/utils/fmt/unordered_set.cc +++ b/lib/utils/test/src/utils/fmt/unordered_set.cc @@ -1,6 +1,6 @@ #include "utils/fmt/unordered_set.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" -#include "utils/containers/unordered_multiset_of.h" +#include "test/utils/doctest/fmt/multiset.h" +#include "utils/containers/multiset_of.h" #include "utils/hash-utils.h" #include @@ -49,6 +49,6 @@ TEST_SUITE(FF_TEST_SUITE) { std::string result = fmt::to_string(input); std::string correct = "{0, 1, 2, 3}"; CHECK(result != correct); - CHECK(unordered_multiset_of(result) == unordered_multiset_of(correct)); + CHECK(multiset_of(result) == multiset_of(correct)); } } diff --git a/lib/utils/test/src/utils/graph/algorithms.cc b/lib/utils/test/src/utils/graph/algorithms.cc index 5917ab9b47..78ee4e704d 100644 --- a/lib/utils/test/src/utils/graph/algorithms.cc +++ b/lib/utils/test/src/utils/graph/algorithms.cc @@ -32,7 +32,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(3)}, DirectedEdge{n.at(2), n.at(3)}}); - std::unordered_set> corrects = { + std::set> corrects = { {n.at(0), n.at(1), n.at(3), n.at(2), n.at(3)}, {n.at(0), n.at(2), n.at(3), n.at(1), n.at(3)}}; std::vector result = get_unchecked_dfs_ordering(g, {n.at(0)}); @@ -52,7 +52,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(5)}}); SUBCASE("branching path") { - std::unordered_set> corrects = { + std::set> corrects = { {n.at(0), n.at(1), n.at(2), n.at(3), n.at(4), n.at(5)}, {n.at(0), n.at(2), n.at(1), n.at(3), n.at(4), n.at(5)}}; std::vector result = get_bfs_ordering(g, {n.at(0)}); @@ -75,7 +75,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(2), n.at(0)}, DirectedEdge{n.at(2), n.at(1)}}); - std::unordered_set> corrects = { + std::set> corrects = { {n.at(0), n.at(1), n.at(2)}, {n.at(0), n.at(2), n.at(1)}}; std::vector result = get_bfs_ordering(g, {n.at(0)}); CHECK(contains(corrects, result)); @@ -111,7 +111,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("branching") { g.add_edge(DirectedEdge{n.at(1), n.at(3)}); - std::unordered_set> corrects = { + std::set> corrects = { {n.at(0), n.at(1), n.at(2), n.at(3)}, {n.at(0), n.at(1), n.at(3), n.at(2)}}; std::vector result = get_dfs_ordering(g, {n.at(0)}); diff --git a/lib/utils/test/src/utils/graph/cow_ptr_t.cc b/lib/utils/test/src/utils/graph/cow_ptr_t.cc index e6a6f9661e..6feba34dab 100644 --- a/lib/utils/test/src/utils/graph/cow_ptr_t.cc +++ b/lib/utils/test/src/utils/graph/cow_ptr_t.cc @@ -1,7 +1,7 @@ #include "utils/graph/cow_ptr_t.h" #include #include -#include +#include #include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms.cc index 74fa0e0ac2..325002c18c 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms.cc @@ -1,7 +1,7 @@ #include "utils/graph/dataflow_graph/algorithms.h" #include "test/utils/doctest/fmt/vector.h" #include "utils/containers/get_only.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" #include "utils/graph/instances/unordered_set_dataflow_graph.h" #include diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc index 37d7dcd712..46fe51aeff 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.cc @@ -65,9 +65,9 @@ TEST_SUITE(FF_TEST_SUITE) { n4_2 n4 n4_0>,shape=plaintext]; + node0:o0 -> node3:i0; node1:o0 -> node3:i1; node2:o0 -> node3:i2; - node0:o0 -> node3:i0; })EXPECTED_OUTPUT"; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc index e619cc3b1c..b3f7c69036 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_dataflow_edges_from_node_to_node.cc @@ -19,9 +19,9 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult n2_added = g.add_node({n1_o0, n1_o0, n1_o1}, 0_n); Node n2 = n2_added.node; - std::unordered_set result = + std::set result = get_dataflow_edges_from_node_to_node(g, n1, n2); - std::unordered_set correct = { + std::set correct = { DataflowEdge{ n1_o0, DataflowInput{n2, 0_n}, @@ -52,9 +52,9 @@ TEST_SUITE(FF_TEST_SUITE) { Node n3 = n3_added.node; DataflowOutput o3 = get_only(n3_added.outputs); - std::unordered_set result = + std::set result = get_dataflow_edges_from_node_to_node(g, n1, n3); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } @@ -68,9 +68,9 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult n2_added = g.add_node({o1}, 0_n); Node n2 = n2_added.node; - std::unordered_set result = + std::set result = get_dataflow_edges_from_node_to_node(g, n2, n1); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } @@ -82,9 +82,9 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult n2_added = g.add_node({}, 1_n); Node n2 = n2_added.node; - std::unordered_set result = + std::set result = get_dataflow_edges_from_node_to_node(g, n1, n2); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } @@ -94,9 +94,9 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult n1_added = g.add_node({}, 1_n); Node n1 = n1_added.node; - std::unordered_set result = + std::set result = get_dataflow_edges_from_node_to_node(g, n1, n1); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc index c37dcf5be7..71eeda7601 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc @@ -27,16 +27,16 @@ TEST_SUITE(FF_TEST_SUITE) { DataflowOutput o4 = get_only(n4_added.outputs); SUBCASE("n2 - single outgoing edge") { - std::unordered_set result = get_outgoing_edges(g, n2); - std::unordered_set correct = { + std::set result = get_outgoing_edges(g, n2); + std::set correct = { DataflowEdge{o2, DataflowInput{n4, 0_n}}, }; CHECK(result == correct); } SUBCASE("n1 - multiple outgoing edges") { - std::unordered_set result = get_outgoing_edges(g, n1); - std::unordered_set correct = { + std::set result = get_outgoing_edges(g, n1); + std::set correct = { DataflowEdge{o1, DataflowInput{n2, 0_n}}, DataflowEdge{o1, DataflowInput{n3, 0_n}}, }; @@ -44,13 +44,13 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("n4 - no outgoing edges") { - std::unordered_set result = get_outgoing_edges(g, n4); - std::unordered_set correct = {}; + std::set result = get_outgoing_edges(g, n4); + std::set correct = {}; CHECK(result == correct); } } - TEST_CASE("get_outgoing_edges(DataflowGraphView, std::unordered_set)") { + TEST_CASE("get_outgoing_edges(DataflowGraphView, std::set)") { DataflowGraph g = DataflowGraph::create(); NodeAddedResult n1_added = g.add_node({}, 1_n); @@ -70,9 +70,9 @@ TEST_SUITE(FF_TEST_SUITE) { DataflowOutput o4 = get_only(n4_added.outputs); SUBCASE("multiple nodes - combined outgoing edges") { - std::unordered_set nodes = {n1, n2}; - std::unordered_set result = get_outgoing_edges(g, nodes); - std::unordered_set correct = { + std::set nodes = {n1, n2}; + std::set result = get_outgoing_edges(g, nodes); + std::set correct = { DataflowEdge{o1, DataflowInput{n2, 0_n}}, DataflowEdge{o1, DataflowInput{n3, 0_n}}, DataflowEdge{o2, DataflowInput{n4, 0_n}}, @@ -81,9 +81,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("multiple nodes - no outgoing edges") { - std::unordered_set nodes = {n3, n4}; - std::unordered_set result = get_outgoing_edges(g, nodes); - std::unordered_set correct = {}; + std::set nodes = {n3, n4}; + std::set result = get_outgoing_edges(g, nodes); + std::set correct = {}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc index 6c770a9d29..001de8c60c 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc @@ -8,7 +8,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_subgraph_incoming_edges(DataflowGraphView, " - "std::unordered_set") { + "std::set") { DataflowGraph g = DataflowGraph::create(); NodeAddedResult n1_added = g.add_node({}, 1_n); @@ -27,12 +27,12 @@ TEST_SUITE(FF_TEST_SUITE) { Node n4 = n4_added.node; DataflowOutput o4 = get_only(n4_added.outputs); - std::unordered_set input_node_set = {n2, n3}; + std::set input_node_set = {n2, n3}; - std::unordered_set result = + std::set result = get_subgraph_incoming_edges(g, input_node_set); - std::unordered_set correct = { + std::set correct = { DataflowEdge{o1, DataflowInput{n2, 0_n}}, DataflowEdge{o1, DataflowInput{n3, 0_n}}, DataflowEdge{o1, DataflowInput{n3, 2_n}}, diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc index bb7f3c4c30..ca65433a27 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc @@ -8,7 +8,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_subgraph_outgoing_edges(DataflowGraphView, " - "std::unordered_set") { + "std::set") { DataflowGraph g = DataflowGraph::create(); NodeAddedResult n1_added = g.add_node({}, 1_n); @@ -27,12 +27,12 @@ TEST_SUITE(FF_TEST_SUITE) { Node n4 = n4_added.node; DataflowOutput o4 = get_only(n4_added.outputs); - std::unordered_set input_node_set = {n2, n3}; + std::set input_node_set = {n2, n3}; - std::unordered_set result = + std::set result = get_subgraph_outgoing_edges(g, input_node_set); - std::unordered_set correct = { + std::set correct = { DataflowEdge{o2, DataflowInput{n4, 1_n}}, DataflowEdge{o3, DataflowInput{n4, 2_n}}, }; diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc index dbd5278d72..9f47cc8648 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc @@ -49,9 +49,9 @@ TEST_SUITE(FF_TEST_SUITE) { make_parallel_split(make_leaf(n3), make_leaf(n4)), }; - std::unordered_set result = + std::set result = get_transitive_reduced_edges_across_split(tr_g, split); - std::unordered_set correct = { + std::set correct = { DataflowEdge{ o1, DataflowInput{n3, 1_n}, @@ -86,9 +86,9 @@ TEST_SUITE(FF_TEST_SUITE) { make_leaf(n2), }; - std::unordered_set result = + std::set result = get_transitive_reduced_edges_across_split(tr_g, split); - std::unordered_set correct = { + std::set correct = { DataflowEdge{ n1_o1, DataflowInput{n2, 0_n}, @@ -131,9 +131,9 @@ TEST_SUITE(FF_TEST_SUITE) { make_series_split(make_leaf(n3), make_leaf(n4)), }; - std::unordered_set result = + std::set result = get_transitive_reduced_edges_across_split(tr_g, split); - std::unordered_set correct = { + std::set correct = { DataflowEdge{ o2, DataflowInput{n3, 1_n}, diff --git a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc index 14ee5dd18c..8997d13261 100644 --- a/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc +++ b/lib/utils/test/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_outputs_across_split.cc @@ -43,9 +43,9 @@ TEST_SUITE(FF_TEST_SUITE) { make_series_split(make_leaf(n3), make_leaf(n4)), }; - std::unordered_set result = + std::set result = get_transitive_reduced_outputs_across_split(tr_g, split); - std::unordered_set correct = {o2}; + std::set correct = {o2}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/apply_contraction.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/apply_contraction.cc index 49a28f3c19..d8315ea6dc 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/apply_contraction.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/apply_contraction.cc @@ -35,14 +35,14 @@ TEST_SUITE(FF_TEST_SUITE) { }); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = {n.at(2), n.at(4), n.at(5)}; + std::set result_nodes = get_nodes(result); + std::set correct_nodes = {n.at(2), n.at(4), n.at(5)}; CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(2), n.at(4)}, DirectedEdge{n.at(2), n.at(2)}, DirectedEdge{n.at(4), n.at(2)}, diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc index edaffb0a93..78be249a07 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc @@ -1,7 +1,7 @@ #include "utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.h" #include "utils/fmt/optional.h" -#include "utils/fmt/unordered_set.h" -#include "utils/hash/unordered_set.h" +#include "utils/fmt/set.h" +#include "utils/hash/set.h" #include using namespace ::FlexFlow; @@ -48,17 +48,17 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_head_subcomponents") { - std::unordered_set> result = + std::set> result = get_head_subcomponents(cbc); - std::unordered_set> correct = {bc1.head_nodes, + std::set> correct = {bc1.head_nodes, bc2.head_nodes}; CHECK(result == correct); } SUBCASE("get_tail_subcomponents") { - std::unordered_set> result = + std::set> result = get_tail_subcomponents(cbc); - std::unordered_set> correct = {bc1.tail_nodes, + std::set> correct = {bc1.tail_nodes, bc2.tail_nodes}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc index 1eb6ed4602..0235e071ac 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc @@ -19,7 +19,7 @@ TEST_SUITE(FF_TEST_SUITE) { // source of bugs in the past auto check_cbc_decomposition_is_edge_order_invariant = [](DiGraphView const &g) { - std::unordered_set edges = get_edges(g); + std::set edges = get_edges(g); std::vector edge_order1 = vector_of(edges); std::vector edge_order2 = reversed(edge_order1); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc index 740ad52def..7cea620170 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/is_complete_bipartite_digraph.cc @@ -8,7 +8,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("is_complete_bipartite_digraph(UndirectedGraphView, " - "std::unordered_set)") { + "std::set)") { DiGraph g = DiGraph::create(); SUBCASE("simple bipartite graph") { @@ -25,7 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { }); SUBCASE("source group") { - std::unordered_set group1 = {n.at(0), n.at(1), n.at(2)}; + std::set group1 = {n.at(0), n.at(1), n.at(2)}; bool result = is_complete_bipartite_digraph(g, group1); bool correct = true; @@ -34,7 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("sink group") { - std::unordered_set group1 = {n.at(3), n.at(4)}; + std::set group1 = {n.at(3), n.at(4)}; bool result = is_complete_bipartite_digraph(g, group1); bool correct = false; @@ -52,7 +52,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(0), n.at(3)}, DirectedEdge{n.at(1), n.at(3)}, }); - std::unordered_set group1 = {n.at(0), n.at(1)}; + std::set group1 = {n.at(0), n.at(1)}; bool result = is_complete_bipartite_digraph(g, group1); bool correct = false; @@ -71,7 +71,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(3)}, DirectedEdge{n.at(2), n.at(3)}, }); - std::unordered_set group1 = {n.at(0), n.at(1)}; + std::set group1 = {n.at(0), n.at(1)}; bool result = is_complete_bipartite_digraph(g, group1); bool correct = false; @@ -89,7 +89,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(2), n.at(1)}, DirectedEdge{n.at(1), n.at(3)}, }); - std::unordered_set group1 = {n.at(0), n.at(1)}; + std::set group1 = {n.at(0), n.at(1)}; bool result = is_complete_bipartite_digraph(g, group1); bool correct = false; @@ -107,7 +107,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(1), n.at(3)}, }); - std::unordered_set group1 = {n.at(0)}; + std::set group1 = {n.at(0)}; bool result = is_complete_bipartite_digraph(g, group1); bool correct = false; diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc index 769d33da9f..d4c1bdd0a0 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc @@ -28,15 +28,15 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView result = contract_node(g, n.at(0), n.at(3)); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = { + std::set result_nodes = get_nodes(result); + std::set correct_nodes = { n.at(1), n.at(2), n.at(3), n.at(4)}; CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(3), n.at(1)}, DirectedEdge{n.at(3), n.at(2)}, DirectedEdge{n.at(3), n.at(4)}, diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/flipped.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/flipped.cc index c74b9b2aae..4a8c2abd55 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/flipped.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/flipped.cc @@ -34,13 +34,13 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView result = flipped(g); SUBCASE("nodes") { - std::unordered_set correct_nodes = unordered_set_of(n); - std::unordered_set result_nodes = get_nodes(result); + std::set correct_nodes = set_of(n); + std::set result_nodes = get_nodes(result); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set correct_edges = { + std::set correct_edges = { DirectedEdge{n.at(1), n.at(0)}, DirectedEdge{n.at(2), n.at(1)}, DirectedEdge{n.at(3), n.at(1)}, @@ -49,7 +49,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(3)}, DirectedEdge{n.at(4), n.at(3)}, }; - std::unordered_set result_edges = get_edges(result); + std::set result_edges = get_edges(result); CHECK(result_edges == correct_edges); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_descendants.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_descendants.cc index a115569139..5af8683c2f 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_descendants.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_descendants.cc @@ -12,8 +12,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("single node") { std::vector n = add_nodes(g, 1); - std::unordered_set correct = {}; - std::unordered_set result = get_descendants(g, n.at(0)); + std::set correct = {}; + std::set result = get_descendants(g, n.at(0)); CHECK(correct == result); } @@ -25,26 +25,26 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(2), n.at(3)}}); SUBCASE("n.at(0)") { - std::unordered_set correct = {n.at(1), n.at(2), n.at(3)}; - std::unordered_set result = get_descendants(g, n.at(0)); + std::set correct = {n.at(1), n.at(2), n.at(3)}; + std::set result = get_descendants(g, n.at(0)); CHECK(correct == result); } SUBCASE("n.at(1)") { - std::unordered_set correct = {n.at(2), n.at(3)}; - std::unordered_set result = get_descendants(g, n.at(1)); + std::set correct = {n.at(2), n.at(3)}; + std::set result = get_descendants(g, n.at(1)); CHECK(correct == result); } SUBCASE("n.at(2)") { - std::unordered_set correct = {n.at(3)}; - std::unordered_set result = get_descendants(g, n.at(2)); + std::set correct = {n.at(3)}; + std::set result = get_descendants(g, n.at(2)); CHECK(correct == result); } SUBCASE("n.at(3)") { - std::unordered_set correct = {}; - std::unordered_set result = get_descendants(g, n.at(3)); + std::set correct = {}; + std::set result = get_descendants(g, n.at(3)); CHECK(correct == result); } } @@ -60,26 +60,26 @@ TEST_SUITE(FF_TEST_SUITE) { }); SUBCASE("n.at(0)") { - std::unordered_set correct = {n.at(1), n.at(2), n.at(3)}; - std::unordered_set result = get_descendants(g, n.at(0)); + std::set correct = {n.at(1), n.at(2), n.at(3)}; + std::set result = get_descendants(g, n.at(0)); CHECK(correct == result); } SUBCASE("n.at(1)") { - std::unordered_set correct = {n.at(3)}; - std::unordered_set result = get_descendants(g, n.at(1)); + std::set correct = {n.at(3)}; + std::set result = get_descendants(g, n.at(1)); CHECK(correct == result); } SUBCASE("n.at(2)") { - std::unordered_set correct = {n.at(3)}; - std::unordered_set result = get_descendants(g, n.at(2)); + std::set correct = {n.at(3)}; + std::set result = get_descendants(g, n.at(2)); CHECK(correct == result); } SUBCASE("n.at(3)") { - std::unordered_set correct = {}; - std::unordered_set result = get_descendants(g, n.at(3)); + std::set correct = {}; + std::set result = get_descendants(g, n.at(3)); CHECK(correct == result); } } @@ -94,32 +94,32 @@ TEST_SUITE(FF_TEST_SUITE) { }); SUBCASE("n.at(0)") { - std::unordered_set correct = {n.at(1), n.at(2)}; - std::unordered_set result = get_descendants(g, n.at(0)); + std::set correct = {n.at(1), n.at(2)}; + std::set result = get_descendants(g, n.at(0)); CHECK(correct == result); } SUBCASE("n.at(1)") { - std::unordered_set correct = {n.at(2)}; - std::unordered_set result = get_descendants(g, n.at(1)); + std::set correct = {n.at(2)}; + std::set result = get_descendants(g, n.at(1)); CHECK(correct == result); } SUBCASE("n.at(2)") { - std::unordered_set correct = {}; - std::unordered_set result = get_descendants(g, n.at(2)); + std::set correct = {}; + std::set result = get_descendants(g, n.at(2)); CHECK(correct == result); } SUBCASE("n.at(3)") { - std::unordered_set correct = {n.at(4)}; - std::unordered_set result = get_descendants(g, n.at(3)); + std::set correct = {n.at(4)}; + std::set result = get_descendants(g, n.at(3)); CHECK(correct == result); } SUBCASE("n.at(4)") { - std::unordered_set correct = {}; - std::unordered_set result = get_descendants(g, n.at(4)); + std::set correct = {}; + std::set result = get_descendants(g, n.at(4)); CHECK(correct == result); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators.cc index 17bea2210f..ac38ba1acc 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators.cc @@ -21,15 +21,15 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_dominators(DiGraph, Node)") { Node node = n.at(2); - std::unordered_set correct = {n.at(0), n.at(2)}; - std::unordered_set result = get_dominators(g, node); + std::set correct = {n.at(0), n.at(2)}; + std::set result = get_dominators(g, node); CHECK(correct == result); } - SUBCASE("get_dominators(DiGraph, std::unordered_set)") { - std::unordered_set nodes = {n.at(1), n.at(3)}; - std::unordered_set result = get_dominators(g, nodes); - std::unordered_set correct = {n.at(0)}; + SUBCASE("get_dominators(DiGraph, std::set)") { + std::set nodes = {n.at(1), n.at(3)}; + std::set result = get_dominators(g, nodes); + std::set correct = {n.at(0)}; CHECK(correct == result); } } @@ -54,14 +54,14 @@ TEST_SUITE(FF_TEST_SUITE) { }); SUBCASE("node 1") { - std::unordered_set result = get_dominators(g, n.at(1)); - std::unordered_set correct = {n.at(0), n.at(1)}; + std::set result = get_dominators(g, n.at(1)); + std::set correct = {n.at(0), n.at(1)}; CHECK(result == correct); } SUBCASE("node 3") { - std::unordered_set result = get_dominators(g, n.at(3)); - std::unordered_set correct = {n.at(0), n.at(1), n.at(3)}; + std::set result = get_dominators(g, n.at(3)); + std::set correct = {n.at(0), n.at(1), n.at(3)}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc index 6b814c1920..c2eeef7a1a 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc @@ -25,7 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {n.at(0)}}, {n.at(1), {n.at(0), n.at(1)}}, {n.at(2), {n.at(0), n.at(1), n.at(2)}}, @@ -34,7 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {n.at(0), n.at(1), n.at(5)}}, }; - std::unordered_map> result = + std::map> result = get_dominators_map(g); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges.cc index 182e3295c3..bcc9ee6971 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/get_edges.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/digraph.h" #include "utils/graph/instances/adjacency_digraph.h" @@ -21,32 +21,32 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); SUBCASE("Base") { - std::unordered_set correct = unordered_set_of(e); - std::unordered_set result = get_edges(g); + std::set correct = set_of(e); + std::set result = get_edges(g); CHECK(result == correct); } SUBCASE("Adding an edge") { g.add_edge(DirectedEdge{n.at(3), n.at(1)}); - std::unordered_set correct = { + std::set correct = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(0), n.at(2)}, DirectedEdge{n.at(0), n.at(3)}, DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(3), n.at(1)}, }; - std::unordered_set result = get_edges(g); + std::set result = get_edges(g); CHECK(result == correct); } SUBCASE("Removing an edge") { g.remove_edge(DirectedEdge{n.at(0), n.at(3)}); - std::unordered_set correct = { + std::set correct = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(0), n.at(2)}, DirectedEdge{n.at(1), n.at(2)}, }; - std::unordered_set result = get_edges(g); + std::set result = get_edges(g); CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc index 5a1ea99671..1eb08a54cc 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc @@ -11,8 +11,8 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector n = add_nodes(g, 5); SUBCASE("basic tests") { - std::unordered_set src_subgraph = {n.at(0), n.at(1), n.at(4)}; - std::unordered_set dst_subgraph = {n.at(2), n.at(3)}; + std::set src_subgraph = {n.at(0), n.at(1), n.at(4)}; + std::set dst_subgraph = {n.at(2), n.at(3)}; SUBCASE("returns all edges between subgraphs") { std::vector e = { @@ -24,9 +24,9 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); - std::unordered_set result = + std::set result = get_edges_from_subgraph_to_subgraph(g, src_subgraph, dst_subgraph); - std::unordered_set correct = unordered_set_of(e); + std::set correct = set_of(e); CHECK(result == correct); } @@ -39,9 +39,9 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); - std::unordered_set result = + std::set result = get_edges_from_subgraph_to_subgraph(g, src_subgraph, dst_subgraph); - std::unordered_set correct = {e.at(0)}; + std::set correct = {e.at(0)}; CHECK(result == correct); } @@ -54,9 +54,9 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); - std::unordered_set result = + std::set result = get_edges_from_subgraph_to_subgraph(g, src_subgraph, dst_subgraph); - std::unordered_set correct = {e.at(1)}; + std::set correct = {e.at(1)}; CHECK(result == correct); } @@ -70,9 +70,9 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); - std::unordered_set result = + std::set result = get_edges_from_subgraph_to_subgraph(g, src_subgraph, dst_subgraph); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } @@ -88,25 +88,25 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); SUBCASE("returns no edges if no nodes in src_subgraph") { - std::unordered_set result = - get_edges_from_subgraph_to_subgraph(g, {}, unordered_set_of(n)); - std::unordered_set correct = {}; + std::set result = + get_edges_from_subgraph_to_subgraph(g, {}, set_of(n)); + std::set correct = {}; CHECK(result == correct); } SUBCASE("returns no edges if no nodes in dst_subgraph") { - std::unordered_set result = - get_edges_from_subgraph_to_subgraph(g, unordered_set_of(n), {}); - std::unordered_set correct = {}; + std::set result = + get_edges_from_subgraph_to_subgraph(g, set_of(n), {}); + std::set correct = {}; CHECK(result == correct); } SUBCASE("returns no edges if both subgraphs are empty") { - std::unordered_set result = + std::set result = get_edges_from_subgraph_to_subgraph(g, {}, {}); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } @@ -122,19 +122,19 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); - std::unordered_set src_subgraph = {n.at(0)}; - std::unordered_set dst_subgraph = {n.at(3)}; + std::set src_subgraph = {n.at(0)}; + std::set dst_subgraph = {n.at(3)}; - std::unordered_set result = + std::set result = get_edges_from_subgraph_to_subgraph(g, src_subgraph, dst_subgraph); - std::unordered_set correct = {e.at(1)}; + std::set correct = {e.at(1)}; CHECK(result == correct); } SUBCASE("throws an error if subgraphs are not disjoint") { - std::unordered_set src_subgraph = {n.at(0), n.at(1), n.at(2)}; - std::unordered_set dst_subgraph = {n.at(1), n.at(3)}; + std::set src_subgraph = {n.at(0), n.at(1), n.at(2)}; + std::set dst_subgraph = {n.at(1), n.at(3)}; CHECK_THROWS( get_edges_from_subgraph_to_subgraph(g, src_subgraph, dst_subgraph)); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc index d830af61f7..3d71bc836d 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc @@ -25,7 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), std::nullopt}, {n.at(1), n.at(0)}, {n.at(2), n.at(1)}, @@ -34,7 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), n.at(1)}, }; - std::unordered_map> result = + std::map> result = get_imm_dominators_map(g); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc index e92a37169b..a22a737df3 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_post_dominators_map.cc @@ -12,10 +12,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("single-node graph") { std::vector n = add_nodes(g, 1); - std::unordered_map> result = + std::map> result = get_imm_post_dominators_map(g); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), std::nullopt}, }; @@ -27,10 +27,10 @@ TEST_SUITE(FF_TEST_SUITE) { g.add_edge(DirectedEdge{n.at(0), n.at(1)}); - std::unordered_map> result = + std::map> result = get_imm_post_dominators_map(g); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {n.at(1)}}, {n.at(1), std::nullopt}, }; @@ -57,9 +57,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(8), n.at(9)}, }); - std::unordered_map> result = + std::map> result = get_imm_post_dominators_map(g); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), n.at(9)}, {n.at(1), n.at(7)}, {n.at(2), n.at(8)}, @@ -92,7 +92,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), n.at(1)}, {n.at(1), n.at(5)}, {n.at(2), n.at(4)}, @@ -101,7 +101,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), std::nullopt}, }; - std::unordered_map> result = + std::map> result = get_imm_post_dominators_map(g); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_incoming_edges.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_incoming_edges.cc index abf0df8607..1de1c558b8 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_incoming_edges.cc @@ -23,7 +23,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {}}, {n.at(1), {DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(4), n.at(1)}}}, @@ -34,7 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {DirectedEdge{n.at(1), n.at(5)}}}, }; - std::unordered_map> result = + std::map> result = get_incoming_edges(g, get_nodes(g)); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_initial_nodes.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_initial_nodes.cc index 787b09c21d..fe24c85635 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_initial_nodes.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_initial_nodes.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/get_initial_nodes.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/digraph.h" #include "utils/graph/instances/adjacency_digraph.h" @@ -21,29 +21,29 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); SUBCASE("Base") { - std::unordered_set correct = {n.at(0)}; - std::unordered_set result = get_initial_nodes(g); + std::set correct = {n.at(0)}; + std::set result = get_initial_nodes(g); CHECK(result == correct); } SUBCASE("Adding an edge to remove a source") { g.add_edge(DirectedEdge{n.at(2), n.at(0)}); - std::unordered_set correct = {}; - std::unordered_set result = get_initial_nodes(g); + std::set correct = {}; + std::set result = get_initial_nodes(g); CHECK(result == correct); } SUBCASE("Removing an edge to create a new source") { g.remove_edge(DirectedEdge{n.at(0), n.at(1)}); - std::unordered_set correct = {n.at(0), n.at(1)}; - std::unordered_set result = get_initial_nodes(g); + std::set correct = {n.at(0), n.at(1)}; + std::set result = get_initial_nodes(g); CHECK(result == correct); } SUBCASE("Creating a cycle") { g.add_edge(DirectedEdge{n.at(2), n.at(0)}); - std::unordered_set result = get_initial_nodes(g); - std::unordered_set correct = {}; + std::set result = get_initial_nodes(g); + std::set correct = {}; CHECK(result.empty()); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc index 9dd44fe4ec..40fcb074fc 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.cc @@ -20,7 +20,7 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, edges); - std::unordered_map expected_lengths = { + std::map expected_lengths = { {n.at(0), 1_n}, {n.at(1), 2_n}, {n.at(2), 3_n}, @@ -46,7 +46,7 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, edges); - std::unordered_map expected_lengths = { + std::map expected_lengths = { {n.at(0), 1_n}, {n.at(1), 2_n}, {n.at(2), 3_n}, diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc index 3fbc3b0cc2..69ccfadc5c 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc @@ -11,8 +11,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("returns nullopt for empty input") { SUBCASE("empty graph") { - std::optional> correct = std::nullopt; - std::optional> result = + std::optional> correct = std::nullopt; + std::optional> result = get_lowest_common_ancestors(g, {}); CHECK(correct == result); } @@ -22,8 +22,8 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges( g, {DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(0), n.at(2)}}); - std::optional> correct = std::nullopt; - std::optional> result = + std::optional> correct = std::nullopt; + std::optional> result = get_lowest_common_ancestors(g, {}); CHECK(correct == result); } @@ -32,9 +32,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("trees") { SUBCASE("single node") { std::vector n = add_nodes(g, 1); - std::optional> correct = - std::unordered_set{n.at(0)}; - std::optional> result = + std::optional> correct = + std::set{n.at(0)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(0)}); CHECK(correct == result); } @@ -46,25 +46,25 @@ TEST_SUITE(FF_TEST_SUITE) { {DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(0), n.at(2)}}); SUBCASE("LCA of siblings is parent") { - std::optional> correct = - std::unordered_set{n.at(0)}; - std::optional> result = + std::optional> correct = + std::set{n.at(0)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(1), n.at(2)}); CHECK(correct == result); } SUBCASE("LCA of a single node is itself") { - std::optional> correct = - std::unordered_set{n.at(1)}; - std::optional> result = + std::optional> correct = + std::set{n.at(1)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(1)}); CHECK(correct == result); } SUBCASE("LCA of another single node is itself") { - std::optional> correct = - std::unordered_set{n.at(2)}; - std::optional> result = + std::optional> correct = + std::set{n.at(2)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(2)}); CHECK(correct == result); } @@ -80,33 +80,33 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(3), n.at(5)}}); SUBCASE("LCA of nodes at different depths (root is LCA)") { - std::optional> correct = - std::unordered_set{n.at(0)}; - std::optional> result = + std::optional> correct = + std::set{n.at(0)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(5), n.at(2)}); CHECK(correct == result); } SUBCASE("LCA of node and its ancestor is the ancestor") { - std::optional> correct = - std::unordered_set{n.at(3)}; - std::optional> result = + std::optional> correct = + std::set{n.at(3)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(5), n.at(3)}); CHECK(correct == result); } SUBCASE("LCA of siblings at depth 2") { - std::optional> correct = - std::unordered_set{n.at(1)}; - std::optional> result = + std::optional> correct = + std::set{n.at(1)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(3), n.at(4)}); CHECK(correct == result); } SUBCASE("LCA of multiple nodes across different branches") { - std::optional> correct = - std::unordered_set{n.at(0)}; - std::optional> result = + std::optional> correct = + std::set{n.at(0)}; + std::optional> result = get_lowest_common_ancestors( g, {n.at(1), n.at(2), n.at(3), n.at(4), n.at(5)}); CHECK(correct == result); @@ -121,25 +121,25 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(2), n.at(3)}}); SUBCASE("LCA of adjacent nodes in a path") { - std::optional> correct = - std::unordered_set{n.at(2)}; - std::optional> result = + std::optional> correct = + std::set{n.at(2)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(2), n.at(3)}); CHECK(correct == result); } SUBCASE("LCA of non-adjacent nodes in a path") { - std::optional> correct = - std::unordered_set{n.at(1)}; - std::optional> result = + std::optional> correct = + std::set{n.at(1)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(1), n.at(3)}); CHECK(correct == result); } SUBCASE("LCA of multiple nodes in a path") { - std::optional> correct = - std::unordered_set{n.at(1)}; - std::optional> result = + std::optional> correct = + std::set{n.at(1)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(1), n.at(2), n.at(3)}); CHECK(correct == result); } @@ -154,9 +154,9 @@ TEST_SUITE(FF_TEST_SUITE) { g, {DirectedEdge{n.at(0), n.at(2)}, DirectedEdge{n.at(1), n.at(2)}}); - std::optional> correct = - std::unordered_set{}; - std::optional> result = + std::optional> correct = + std::set{}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(0), n.at(1)}); CHECK(correct == result); } @@ -169,9 +169,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(0), n.at(3)}, DirectedEdge{n.at(1), n.at(3)}}); - std::optional> correct = - std::unordered_set{n.at(0), n.at(1)}; - std::optional> result = + std::optional> correct = + std::set{n.at(0), n.at(1)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(2), n.at(3)}); CHECK(correct == result); } @@ -187,9 +187,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(3), n.at(5)}, DirectedEdge{n.at(1), n.at(5)}}); - std::optional> correct = - std::unordered_set{n.at(3)}; - std::optional> result = + std::optional> correct = + std::set{n.at(3)}; + std::optional> result = get_lowest_common_ancestors(g, {n.at(4), n.at(5)}); CHECK(correct == result); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc index 54ac3ac0f9..6bee8367bd 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc @@ -23,7 +23,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {DirectedEdge{n.at(0), n.at(1)}}}, {n.at(1), {DirectedEdge{n.at(1), n.at(2)}, @@ -35,7 +35,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {}}, }; - std::unordered_map> result = + std::map> result = get_outgoing_edges(g, get_nodes(g)); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc index 9e6e2a8c41..0b21b7dfa9 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc @@ -14,9 +14,9 @@ TEST_SUITE(FF_TEST_SUITE) { g.add_edge(DirectedEdge{n.at(0), n.at(1)}); - std::unordered_map> result = + std::map> result = get_post_dominators_map(g); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {n.at(0), n.at(1)}}, {n.at(1), {n.at(1)}}, }; @@ -41,9 +41,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(8), n.at(9)}, }); - std::unordered_map> result = + std::map> result = get_post_dominators_map(g); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {n.at(0), n.at(9)}}, {n.at(1), {n.at(1), n.at(7), n.at(9)}}, {n.at(2), {n.at(2), n.at(8), n.at(9)}}, @@ -76,7 +76,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {n.at(0), n.at(1), n.at(5)}}, {n.at(1), {n.at(1), n.at(5)}}, {n.at(2), {n.at(1), n.at(2), n.at(4), n.at(5)}}, @@ -85,7 +85,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {n.at(5)}}, }; - std::unordered_map> result = + std::map> result = get_post_dominators_map(g); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc index 3ba06101c0..8241b8891f 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc @@ -22,7 +22,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {}}, {n.at(1), {n.at(0), n.at(4)}}, {n.at(2), {n.at(1)}}, @@ -31,7 +31,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {n.at(1)}}, }; - std::unordered_map> result = + std::map> result = get_predecessors(g); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc index 12db835108..5688494233 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc @@ -22,7 +22,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(1)}, }); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {n.at(1)}}, {n.at(1), {n.at(2), n.at(3), n.at(5)}}, {n.at(2), {n.at(4)}}, @@ -31,7 +31,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {}}, }; - std::unordered_map> result = + std::map> result = get_successors(g); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc index e6ee9d08b4..0d1e3d1a66 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_terminal_nodes.cc @@ -1,5 +1,5 @@ #include "utils/graph/digraph/algorithms/get_terminal_nodes.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/containers/set_of.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/digraph.h" #include "utils/graph/instances/adjacency_digraph.h" @@ -21,22 +21,22 @@ TEST_SUITE(FF_TEST_SUITE) { add_edges(g, e); SUBCASE("Base") { - std::unordered_set correct = {n.at(2), n.at(3)}; - std::unordered_set result = get_terminal_nodes(g); + std::set correct = {n.at(2), n.at(3)}; + std::set result = get_terminal_nodes(g); CHECK(result == correct); } SUBCASE("Adding an edge to remove a terminal node") { g.add_edge(DirectedEdge{n.at(3), n.at(2)}); - std::unordered_set correct = {n.at(2)}; - std::unordered_set result = get_terminal_nodes(g); + std::set correct = {n.at(2)}; + std::set result = get_terminal_nodes(g); CHECK(result == correct); } SUBCASE("Creating a cycle") { g.add_edge(DirectedEdge{n.at(2), n.at(0)}); - std::unordered_set result = get_terminal_nodes(g); - std::unordered_set correct = {n.at(3)}; + std::set result = get_terminal_nodes(g); + std::set correct = {n.at(3)}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc index 58ae971d18..ee407b2179 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc @@ -13,9 +13,9 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("single node") { std::vector n = add_nodes(g, 1); - std::unordered_set> result = + std::set> result = get_weakly_connected_components(g); - std::unordered_set> correct = {{n.at(0)}}; + std::set> correct = {{n.at(0)}}; CHECK(result == correct); } @@ -26,9 +26,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(0), n.at(0)}, }); - std::unordered_set> result = + std::set> result = get_weakly_connected_components(g); - std::unordered_set> correct = {{n.at(0)}}; + std::set> correct = {{n.at(0)}}; CHECK(result == correct); } @@ -40,9 +40,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(1)}, }); - std::unordered_set> result = + std::set> result = get_weakly_connected_components(g); - std::unordered_set> correct = {{n.at(0)}, + std::set> correct = {{n.at(0)}, {n.at(1)}}; CHECK(result == correct); } @@ -54,9 +54,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(0), n.at(1)}, }); - std::unordered_set> result = + std::set> result = get_weakly_connected_components(g); - std::unordered_set> correct = { + std::set> correct = { {n.at(0), n.at(1)}}; CHECK(result == correct); } @@ -69,9 +69,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(0)}, }); - std::unordered_set> result = + std::set> result = get_weakly_connected_components(g); - std::unordered_set> correct = { + std::set> correct = { {n.at(0), n.at(1)}}; CHECK(result == correct); } @@ -90,9 +90,9 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(3)}, }); - std::unordered_set> result = + std::set> result = get_weakly_connected_components(g); - std::unordered_set> correct = { + std::set> correct = { {n.at(0), n.at(1), n.at(2)}, {n.at(3), n.at(4)}, }; diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc index 54b934eb2b..89b24f6e95 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc @@ -57,15 +57,15 @@ TEST_SUITE(FF_TEST_SUITE) { REQUIRE(maybe_result.has_value()); InverseLineGraphResult result = maybe_result.value(); - std::unordered_set result_nodes = get_nodes(result.graph); + std::set result_nodes = get_nodes(result.graph); REQUIRE(result_nodes.size() == 5); std::vector inv = get_topological_ordering(result.graph); SUBCASE("edges") { - std::unordered_map result_edges = + std::map result_edges = get_edge_counts(result.graph); - std::unordered_map correct_edges = { + std::map correct_edges = { {DirectedEdge{inv.at(0), inv.at(1)}, 1}, {DirectedEdge{inv.at(1), inv.at(2)}, 1}, {DirectedEdge{inv.at(1), inv.at(3)}, 1}, @@ -76,13 +76,13 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("inverse_edge_to_line_node_bidict") { - std::unordered_map result_bidict = + std::map result_bidict = map_values(result.inverse_edge_to_line_node_bidict.reversed() - .as_unordered_map(), + .as_map(), [&](MultiDiEdge const &e) { return get_directed_edge(result.graph, e); }); - std::unordered_map correct_bidict = { + std::map correct_bidict = { {n.at(0), DirectedEdge{inv.at(0), inv.at(1)}}, {n.at(1), DirectedEdge{inv.at(1), inv.at(2)}}, {n.at(2), DirectedEdge{inv.at(1), inv.at(3)}}, @@ -119,28 +119,28 @@ TEST_SUITE(FF_TEST_SUITE) { REQUIRE(maybe_result.has_value()); InverseLineGraphResult result = maybe_result.value(); - std::unordered_set result_nodes = get_nodes(result.graph); + std::set result_nodes = get_nodes(result.graph); REQUIRE(result_nodes.size() == 2); std::vector inv = get_topological_ordering(result.graph); SUBCASE("edges") { - std::unordered_map result_edges = + std::map result_edges = get_edge_counts(result.graph); - std::unordered_map correct_edges = { + std::map correct_edges = { {DirectedEdge{inv.at(0), inv.at(1)}, 2}, }; CHECK(result_edges == correct_edges); } SUBCASE("inverse_edge_to_line_node_bidict") { - std::unordered_map result_bidict = + std::map result_bidict = map_values(result.inverse_edge_to_line_node_bidict.reversed() - .as_unordered_map(), + .as_map(), [&](MultiDiEdge const &e) { return get_directed_edge(result.graph, e); }); - std::unordered_map correct_bidict = { + std::map correct_bidict = { {n.at(0), DirectedEdge{inv.at(0), inv.at(1)}}, {n.at(1), DirectedEdge{inv.at(0), inv.at(1)}}, }; diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_closure.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_closure.cc index 36db627ad8..b181dbaffb 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_closure.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_closure.cc @@ -25,14 +25,14 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView result = transitive_closure(g); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = unordered_set_of(n); + std::set result_nodes = get_nodes(result); + std::set correct_nodes = set_of(n); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(0), n.at(2)}, DirectedEdge{n.at(0), n.at(3)}, diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_reduction.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_reduction.cc index 7e47bc470c..d44ea5ce2d 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_reduction.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/transitive_reduction.cc @@ -25,14 +25,14 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView result = transitive_reduction(g); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = unordered_set_of(n); + std::set result_nodes = get_nodes(result); + std::set correct_nodes = set_of(n); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(1), n.at(2)}, }; @@ -51,8 +51,8 @@ TEST_SUITE(FF_TEST_SUITE) { }); DiGraphView result = transitive_reduction(g); - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(2), n.at(3)}, @@ -74,8 +74,8 @@ TEST_SUITE(FF_TEST_SUITE) { }); DiGraphView result = transitive_reduction(g); - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(2), n.at(3)}, @@ -105,14 +105,14 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView result = transitive_reduction(g); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = unordered_set_of(n); + std::set result_nodes = get_nodes(result); + std::set correct_nodes = set_of(n); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(0), n.at(2)}, DirectedEdge{n.at(1), n.at(3)}, @@ -138,14 +138,14 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView result = transitive_reduction(g); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = unordered_set_of(n); + std::set result_nodes = get_nodes(result); + std::set correct_nodes = set_of(n); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(2), n.at(3)}, @@ -168,14 +168,14 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView result = transitive_reduction(g); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(result); - std::unordered_set correct_nodes = unordered_set_of(n); + std::set result_nodes = get_nodes(result); + std::set correct_nodes = set_of(n); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(result); - std::unordered_set correct_edges = { + std::set result_edges = get_edges(result); + std::set correct_edges = { DirectedEdge{n.at(0), n.at(2)}, DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(1), n.at(3)}, diff --git a/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc b/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc index d07e5d5703..9ab7ee8d2d 100644 --- a/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc +++ b/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc @@ -32,8 +32,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("query_nodes") { SUBCASE("query_all") { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {n[0], n[1], n[2], n[3], n[4]}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {n[0], n[1], n[2], n[3], n[4]}; CHECK(result == correct); } @@ -43,8 +43,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n[0], n[2]}), }; - std::unordered_set result = g.query_nodes(query); - std::unordered_set correct = {n[0], n[2]}; + std::set result = g.query_nodes(query); + std::set correct = {n[0], n[2]}; CHECK(result == correct); } @@ -53,10 +53,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("query_edges") { SUBCASE("query_all") { - std::unordered_set result = + std::set result = g.query_edges(directed_edge_query_all()); - std::unordered_set correct = { + std::set correct = { e.at(0), e.at(1), e.at(2), @@ -74,9 +74,9 @@ TEST_SUITE(FF_TEST_SUITE) { }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = - std::unordered_set{e[0]}; + std::set result = g.query_edges(query); + std::set correct = + std::set{e[0]}; CHECK(result == correct); } } @@ -85,33 +85,33 @@ TEST_SUITE(FF_TEST_SUITE) { g.remove_node_unsafe(n[0]); CHECK(g.query_nodes(node_query_all()) == - std::unordered_set{n[1], n[2], n[3], n[4]}); + std::set{n[1], n[2], n[3], n[4]}); // removing a node also removes its adjacent edges CHECK(g.query_edges(directed_edge_query_all()) == - std::unordered_set{e[2], e[3], e[4]}); + std::set{e[2], e[3], e[4]}); g.remove_node_unsafe(n[1]); CHECK(g.query_nodes(node_query_all()) == - std::unordered_set{n[2], n[3], n[4]}); + std::set{n[2], n[3], n[4]}); CHECK(g.query_edges(directed_edge_query_all()) == - std::unordered_set{e[3]}); + std::set{e[3]}); } SUBCASE("remove_edge") { g.remove_edge(e[0]); CHECK(g.query_edges(directed_edge_query_all()) == - std::unordered_set{e[1], e[2], e[3], e[4]}); + std::set{e[1], e[2], e[3], e[4]}); CHECK(g.query_nodes(node_query_all()) == - std::unordered_set{n[0], n[1], n[2], n[3], n[4]}); + std::set{n[0], n[1], n[2], n[3], n[4]}); g.remove_edge(e[1]); g.remove_edge(e[3]); CHECK(g.query_edges(directed_edge_query_all()) == - std::unordered_set{e[2], e[4]}); + std::set{e[2], e[4]}); } } } diff --git a/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc b/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc index d69e1ee71e..39230fb090 100644 --- a/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc +++ b/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc @@ -12,18 +12,18 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiGraph g = MultiDiGraph::create(); auto check_state = - [&](std::unordered_set const &correct_nodes, - std::unordered_set const &correct_edges) { + [&](std::set const &correct_nodes, + std::set const &correct_edges) { { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = correct_nodes; + std::set result = g.query_nodes(node_query_all()); + std::set correct = correct_nodes; REQUIRE(result == correct); } { - std::unordered_set result = + std::set result = g.query_edges(multidiedge_query_all()); - std::unordered_set correct = correct_edges; + std::set correct = correct_edges; REQUIRE(result == correct); } }; @@ -79,8 +79,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::matchall(), }; - std::unordered_set result = g.query_edges(input); - std::unordered_set correct = {e1, e2, e3}; + std::set result = g.query_edges(input); + std::set correct = {e1, e2, e3}; CHECK(result == correct); } @@ -90,8 +90,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n1), }; - std::unordered_set result = g.query_edges(input); - std::unordered_set correct = {e1, e2, e4}; + std::set result = g.query_edges(input); + std::set correct = {e1, e2, e4}; CHECK(result == correct); } @@ -100,8 +100,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n1), query_set::match_single_value(n2), }; - std::unordered_set result = g.query_edges(input); - std::unordered_set correct = {e3}; + std::set result = g.query_edges(input); + std::set correct = {e3}; CHECK(result == correct); } @@ -111,8 +111,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n1), }; - std::unordered_set result = g.query_edges(input); - std::unordered_set correct = {e1, e2}; + std::set result = g.query_edges(input); + std::set correct = {e1, e2}; CHECK(result == correct); } @@ -132,34 +132,34 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiGraphView g2 = g; SUBCASE("nodes") { g.add_node(); - std::unordered_set result = g2.query_nodes(node_query_all()); - std::unordered_set correct = {n1, n2}; + std::set result = g2.query_nodes(node_query_all()); + std::set correct = {n1, n2}; CHECK(result == correct); } SUBCASE("edges") { g.add_edge(n1, n2); - std::unordered_set result = + std::set result = g2.query_edges(multidiedge_query_all()); - std::unordered_set correct = {e1, e2, e3, e4}; + std::set correct = {e1, e2, e3, e4}; CHECK(result == correct); } } SUBCASE("materialize_copy_of") { - std::unordered_set correct_nodes = get_nodes(g); - std::unordered_map correct_edges = + std::set correct_nodes = get_nodes(g); + std::map correct_edges = get_multidiedge_to_diedge_map(g); MultiDiGraph g2 = MultiDiGraph::materialize_copy_of(g); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(g2); + std::set result_nodes = get_nodes(g2); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_map result_edges = + std::map result_edges = get_multidiedge_to_diedge_map(g); CHECK(result_edges == correct_edges); } diff --git a/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc b/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc index dd7402948b..cc1142bdcd 100644 --- a/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc @@ -12,60 +12,60 @@ TEST_SUITE(FF_TEST_SUITE) { DataflowGraph g = DataflowGraph::create(); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {}; REQUIRE(result == correct); } { - std::unordered_set result = + std::set result = g.query_edges(dataflow_edge_query_all()); - std::unordered_set correct = {}; + std::set correct = {}; REQUIRE(result == correct); } { - std::unordered_set result = + std::set result = g.query_outputs(dataflow_output_query_all()); - std::unordered_set correct = {}; + std::set correct = {}; REQUIRE(result == correct); } NodeAddedResult added = g.add_node({}, 2_n); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node}; REQUIRE(result == correct); } { - std::unordered_set result = + std::set result = g.query_edges(dataflow_edge_query_all()); - std::unordered_set correct = {}; + std::set correct = {}; REQUIRE(result == correct); } { - std::unordered_set result = + std::set result = g.query_outputs(dataflow_output_query_all()); - std::unordered_set correct = - unordered_set_of(added.outputs); + std::set correct = + set_of(added.outputs); REQUIRE(result == correct); } NodeAddedResult added2 = g.add_node(added.outputs, 3_n); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node, added2.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node, added2.node}; REQUIRE(result == correct); } { - std::unordered_set result = + std::set result = g.query_edges(dataflow_edge_query_all()); - std::unordered_set correct = { + std::set correct = { DataflowEdge{added.outputs.at(0), DataflowInput{added2.node, 0_n}}, DataflowEdge{added.outputs.at(1), DataflowInput{added2.node, 1_n}}, }; @@ -73,10 +73,10 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::unordered_set result = + std::set result = g.query_outputs(dataflow_output_query_all()); - std::unordered_set correct = set_union( - unordered_set_of(added.outputs), unordered_set_of(added2.outputs)); + std::set correct = set_union( + set_of(added.outputs), set_of(added2.outputs)); REQUIRE(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/instances/unordered_set_kwarg_dataflow_graph.cc b/lib/utils/test/src/utils/graph/instances/unordered_set_kwarg_dataflow_graph.cc index 1709940262..8b0eb918ae 100644 --- a/lib/utils/test/src/utils/graph/instances/unordered_set_kwarg_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/instances/unordered_set_kwarg_dataflow_graph.cc @@ -13,22 +13,22 @@ TEST_SUITE(FF_TEST_SUITE) { UnorderedSetKwargDataflowGraph>(); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_edges(kwarg_dataflow_edge_query_all()); - std::unordered_set> correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); - std::unordered_set> correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } @@ -50,23 +50,23 @@ TEST_SUITE(FF_TEST_SUITE) { added.outputs.at("output_3"); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_edges(kwarg_dataflow_edge_query_all()); - std::unordered_set> correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); - std::unordered_set> correct = - unordered_set_of(values(added.outputs)); + std::set> correct = + set_of(values(added.outputs)); REQUIRE(result == correct); } @@ -92,13 +92,13 @@ TEST_SUITE(FF_TEST_SUITE) { }; { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node, added2.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node, added2.node}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_edges(kwarg_dataflow_edge_query_all()); auto mk_edge = @@ -115,7 +115,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::unordered_set> correct = { + std::set> correct = { mk_edge(added_output_1, added2.node, "input_1"), mk_edge(added_output_3, added2.node, "input_2"), }; @@ -124,14 +124,14 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); auto get_output_set = [](KwargNodeAddedResult const &r) { - return unordered_set_of(values(r.outputs)); + return set_of(values(r.outputs)); }; - std::unordered_set> correct = + std::set> correct = set_union(get_output_set(added), get_output_set(added2)); REQUIRE(result == correct); diff --git a/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc b/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc index ddebe4e612..5448a81760 100644 --- a/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc @@ -19,24 +19,24 @@ TEST_SUITE(FF_TEST_SUITE) { std::string>>(); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {}; REQUIRE(result == correct); } { - std::unordered_set> + std::set> result = g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::unordered_set> + std::set> correct = {}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); - std::unordered_set> correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } @@ -69,25 +69,25 @@ TEST_SUITE(FF_TEST_SUITE) { added.outputs.at("output_3"); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node}; REQUIRE(result == correct); } { - std::unordered_set> + std::set> result = g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::unordered_set> + std::set> correct = {}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); - std::unordered_set> correct = - unordered_set_of(values(added.outputs)); + std::set> correct = + set_of(values(added.outputs)); REQUIRE(result == correct); } @@ -124,13 +124,13 @@ TEST_SUITE(FF_TEST_SUITE) { }; { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node, added2.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node, added2.node}; REQUIRE(result == correct); } { - std::unordered_set> + std::set> result = g.query_edges( open_kwarg_dataflow_edge_query_all()); @@ -164,7 +164,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::unordered_set> + std::set> correct = { internal_edge(added_output_1, added2.node, "input_1"), internal_edge(added_output_3, added2.node, "input_2"), @@ -174,14 +174,14 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); auto get_output_set = [](KwargNodeAddedResult const &r) { - return unordered_set_of(values(r.outputs)); + return set_of(values(r.outputs)); }; - std::unordered_set> correct = + std::set> correct = set_union(get_output_set(added), get_output_set(added2)); REQUIRE(result == correct); diff --git a/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc b/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc index 906cf841f8..354ea4ac85 100644 --- a/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc @@ -13,24 +13,24 @@ TEST_SUITE(FF_TEST_SUITE) { UnorderedSetOpenKwargDataflowGraph>(); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {}; REQUIRE(result == correct); } { - std::unordered_set> + std::set> result = g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::unordered_set> + std::set> correct = {}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); - std::unordered_set> correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } @@ -52,25 +52,25 @@ TEST_SUITE(FF_TEST_SUITE) { added.outputs.at("output_3"); { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node}; REQUIRE(result == correct); } { - std::unordered_set> + std::set> result = g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::unordered_set> + std::set> correct = {}; REQUIRE(result == correct); } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); - std::unordered_set> correct = - unordered_set_of(values(added.outputs)); + std::set> correct = + set_of(values(added.outputs)); REQUIRE(result == correct); } @@ -103,13 +103,13 @@ TEST_SUITE(FF_TEST_SUITE) { }; { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = {added.node, added2.node}; + std::set result = g.query_nodes(node_query_all()); + std::set correct = {added.node, added2.node}; REQUIRE(result == correct); } { - std::unordered_set> + std::set> result = g.query_edges( open_kwarg_dataflow_edge_query_all()); @@ -143,7 +143,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::unordered_set> + std::set> correct = { internal_edge(added_output_1, added2.node, "input_1"), internal_edge(added_output_3, added2.node, "input_2"), @@ -153,14 +153,14 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::unordered_set> result = + std::set> result = g.query_outputs(kwarg_dataflow_output_query_all()); auto get_output_set = [](KwargNodeAddedResult const &r) { - return unordered_set_of(values(r.outputs)); + return set_of(values(r.outputs)); }; - std::unordered_set> correct = + std::set> correct = set_union(get_output_set(added), get_output_set(added2)); REQUIRE(result == correct); diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc index c2fb348075..5d8983ff4e 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc @@ -67,8 +67,8 @@ TEST_SUITE(FF_TEST_SUITE) { }; std::function( - std::unordered_set const &)> - slot_ordering = [](std::unordered_set const &slots) + std::set const &)> + slot_ordering = [](std::set const &slots) -> std::vector { return reversed(sorted(slots)); }; DataflowGraphData result = diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc index 9ac43c2ee2..f46e62c588 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc @@ -22,9 +22,9 @@ TEST_SUITE(FF_TEST_SUITE) { UnorderedSetKwargDataflowGraph>(); KwargNodeAddedResult n0_added = g.add_node( - /*inputs=*/std::unordered_map>{}, - /*outputs=*/std::unordered_set{ + /*outputs=*/std::set{ "a", }); @@ -32,9 +32,9 @@ TEST_SUITE(FF_TEST_SUITE) { require_only_key(n0_added.outputs, std::string{"a"}); KwargNodeAddedResult n1_added = g.add_node( - /*inputs=*/std::unordered_map>{}, - /*outputs=*/std::unordered_set{ + /*outputs=*/std::set{ "b", "c", }); @@ -44,12 +44,12 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n2_added = g.add_node( /*inputs=*/ - std::unordered_map>{ + std::map>{ {"z", o1}, {"y", o2}, {"x", o0}, }, - /*outputs=*/std::unordered_set{ + /*outputs=*/std::set{ "d", }); @@ -57,8 +57,8 @@ TEST_SUITE(FF_TEST_SUITE) { }(); std::function( - std::unordered_set const &)> - slot_ordering = [](std::unordered_set const &slots) + std::set const &)> + slot_ordering = [](std::set const &slots) -> std::vector { return reversed(sorted(slots)); }; DataflowGraphView result = diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc index 8a752e9887..42be45f161 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc @@ -53,7 +53,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set is contains all graph nodes") { KwargDataflowGraphView result = get_kwarg_dataflow_graph_subgraph( - g, std::unordered_set{n1, n2, n3, n4, n5}); + g, std::set{n1, n2, n3, n4, n5}); KwargDataflowGraphData result_data = get_kwarg_dataflow_graph_data(result); @@ -64,7 +64,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set is overlapping") { KwargDataflowGraphView result = - get_kwarg_dataflow_graph_subgraph(g, std::unordered_set{n2, n3, n5}); + get_kwarg_dataflow_graph_subgraph(g, std::set{n2, n3, n5}); KwargDataflowGraphData result_data = get_kwarg_dataflow_graph_data(result); @@ -87,14 +87,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set is non-overlapping") { KwargDataflowGraphView result = - get_kwarg_dataflow_graph_subgraph(g, std::unordered_set{}); + get_kwarg_dataflow_graph_subgraph(g, std::set{}); KwargDataflowGraphData result_data = get_kwarg_dataflow_graph_data(result); KwargDataflowGraphData correct_data = KwargDataflowGraphData{ - /*nodes=*/std::unordered_set{}, - /*edges=*/std::unordered_set>{}, - /*outputs=*/std::unordered_set>{}, + /*nodes=*/std::set{}, + /*edges=*/std::set>{}, + /*outputs=*/std::set>{}, }; CHECK(result_data == correct_data); diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc index c4667bf746..8d7342e120 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc @@ -3,6 +3,7 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" #include "utils/graph/node/algorithms.h" #include +#include "test/utils/doctest/fmt/set.h" using namespace ::FlexFlow; @@ -38,16 +39,16 @@ TEST_SUITE(FF_TEST_SUITE) { Node n1 = Node{1}; Node n2 = Node{2}; - std::unordered_set all_nodes = {n0, n1, n2}; + std::set all_nodes = {n0, n1, n2}; - std::unordered_set>> all_edges = { + std::set>> all_edges = { mk_edge(n0, 1, n1, 0), mk_edge(n0, 1, n1, std::nullopt), mk_edge(n1, 2, n2, 3), mk_edge(n0, std::nullopt, n2, 1), }; - std::unordered_set>> all_outputs = { + std::set>> all_outputs = { mk_output(n0, 1), mk_output(n0, std::nullopt), mk_output(n0, 4), @@ -66,23 +67,23 @@ TEST_SUITE(FF_TEST_SUITE) { view_from_kwarg_dataflow_graph_data(data); SUBCASE("get_nodes") { - std::unordered_set result = get_nodes(g); - std::unordered_set correct = all_nodes; + std::set result = get_nodes(g); + std::set correct = all_nodes; ASSERT(result == correct); } SUBCASE("get_all_kwarg_dataflow_edges") { - std::unordered_set>> result = + std::set>> result = get_all_kwarg_dataflow_edges(g); - std::unordered_set>> correct = + std::set>> correct = all_edges; ASSERT(result == correct); } SUBCASE("get_all_kwarg_dataflow_outputs") { - std::unordered_set>> result = + std::set>> result = get_all_kwarg_dataflow_outputs(g); - std::unordered_set>> correct = + std::set>> correct = all_outputs; ASSERT(result == correct); } diff --git a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_data.cc b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_data.cc index d9da80772e..9d0cea8e96 100644 --- a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_data.cc +++ b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_data.cc @@ -23,10 +23,10 @@ TEST_SUITE(FF_TEST_SUITE) { LabelledKwargDataflowGraphData correct = LabelledKwargDataflowGraphData{ - /*node_data=*/std::unordered_map{}, - /*edges=*/std::unordered_set>{}, + /*node_data=*/std::map{}, + /*edges=*/std::set>{}, /*output_data=*/ - std::unordered_map, float>{}, + std::map, float>{}, }; ASSERT(result == correct); @@ -46,7 +46,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_label=*/n1_label, /*inputs=*/{}, /*output_labels=*/ - std::unordered_map{ + std::map{ {2, n1_t1_label}, }); Node n1 = n1_added.node; @@ -55,11 +55,11 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n2_added = g.add_node( /*node_label=*/n2_label, /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {0, n2_t1_label}, {1, n2_t2_label}, }); @@ -70,13 +70,13 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n3_added = g.add_node( /*node_label=*/n3_label, /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, {1, n1_t1}, {2, n2_t2}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {4, n3_t1_label}, }); Node n3 = n3_added.node; @@ -106,20 +106,20 @@ TEST_SUITE(FF_TEST_SUITE) { LabelledKwargDataflowGraphData correct = LabelledKwargDataflowGraphData{ - /*node_data=*/std::unordered_map{ + /*node_data=*/std::map{ {n1, n1_label}, {n2, n2_label}, {n3, n3_label}, }, /*edges=*/ - std::unordered_set>{ + std::set>{ mk_edge(n1, 2, n2, 3), mk_edge(n1, 2, n3, 3), mk_edge(n1, 2, n3, 1), mk_edge(n2, 1, n3, 2), }, /*output_data=*/ - std::unordered_map, float>{ + std::map, float>{ {n1_t1, n1_t1_label}, {n2_t1, n2_t1_label}, {n2_t2, n2_t2_label}, diff --git a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc index 5df26765b0..e6d5c0f7e7 100644 --- a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc +++ b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.cc @@ -16,12 +16,12 @@ TEST_SUITE(FF_TEST_SUITE) { int>>(); SUBCASE("graph is empty") { - std::unordered_map result = + std::map result = get_labelled_kwarg_dataflow_graph_node_label_map( static_cast< LabelledKwargDataflowGraphView>(g)); - std::unordered_map correct = {}; + std::map correct = {}; CHECK(result == correct); } @@ -35,7 +35,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_label=*/n1_label, /*inputs=*/{}, /*output_labels=*/ - std::unordered_map{ + std::map{ {2, 5.3}, }); Node n1 = n1_added.node; @@ -44,11 +44,11 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n2_added = g.add_node( /*node_label=*/n2_label, /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {0, 12.1}, {1, 3.2}, }); @@ -59,24 +59,24 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n3_added = g.add_node( /*node_label=*/n3_label, /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, {1, n1_t1}, {2, n2_t2}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {4, 1.7}, }); Node n3 = n3_added.node; KwargDataflowOutput n3_t1 = require_only_key(n3_added.outputs, 4); - std::unordered_map result = + std::map result = get_labelled_kwarg_dataflow_graph_node_label_map( static_cast< LabelledKwargDataflowGraphView>(g)); - std::unordered_map correct = { + std::map correct = { {n1, n1_label}, {n2, n2_label}, {n3, n3_label}, diff --git a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc index f4a63c2056..5f740b8768 100644 --- a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc +++ b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.cc @@ -16,12 +16,12 @@ TEST_SUITE(FF_TEST_SUITE) { int>>(); SUBCASE("graph is empty") { - std::unordered_map, float> result = + std::map, float> result = get_labelled_kwarg_dataflow_graph_output_label_map( static_cast< LabelledKwargDataflowGraphView>(g)); - std::unordered_map, float> correct = {}; + std::map, float> correct = {}; CHECK(result == correct); } @@ -36,7 +36,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_label=*/"n1", /*inputs=*/{}, /*output_labels=*/ - std::unordered_map{ + std::map{ {2, n1_t1_label}, }); Node n1 = n1_added.node; @@ -45,11 +45,11 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n2_added = g.add_node( /*node_label=*/"n2", /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {0, n2_t1_label}, {1, n2_t2_label}, }); @@ -60,24 +60,24 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n3_added = g.add_node( /*node_label=*/"n1", /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, {1, n1_t1}, {2, n2_t2}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {4, n3_t1_label}, }); Node n3 = n3_added.node; KwargDataflowOutput n3_t1 = require_only_key(n3_added.outputs, 4); - std::unordered_map, float> result = + std::map, float> result = get_labelled_kwarg_dataflow_graph_output_label_map( static_cast< LabelledKwargDataflowGraphView>(g)); - std::unordered_map, float> correct = { + std::map, float> correct = { {n1_t1, n1_t1_label}, {n2_t1, n2_t1_label}, {n2_t2, n2_t2_label}, diff --git a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc index 0489d36312..d24aed9903 100644 --- a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc @@ -30,7 +30,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_label=*/n1_label, /*inputs=*/{}, /*output_labels=*/ - std::unordered_map{ + std::map{ {2, n1_t1_label}, }); Node n1 = n1_added.node; @@ -39,11 +39,11 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n2_added = g.add_node( /*node_label=*/n2_label, /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {0, n2_t1_label}, {1, n2_t2_label}, }); @@ -54,13 +54,13 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n3_added = g.add_node( /*node_label=*/n3_label, /*inputs=*/ - std::unordered_map>{ + std::map>{ {3, n1_t1}, {1, n1_t1}, {2, n2_t2}, }, /*output_labels=*/ - std::unordered_map{ + std::map{ {4, n3_t1_label}, }); Node n3 = n3_added.node; @@ -88,7 +88,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set includes all graph nodes") { LabelledKwargDataflowGraphView result = get_labelled_kwarg_dataflow_graph_subgraph( - input, std::unordered_set{n1, n2, n3}); + input, std::set{n1, n2, n3}); LabelledKwargDataflowGraphData result_data = get_labelled_kwarg_dataflow_graph_data(result); @@ -121,7 +121,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set includes only some graph nodes") { LabelledKwargDataflowGraphView result = get_labelled_kwarg_dataflow_graph_subgraph( - input, std::unordered_set{n2, n3}); + input, std::set{n2, n3}); LabelledKwargDataflowGraphData result_data = get_labelled_kwarg_dataflow_graph_data(result); @@ -149,16 +149,16 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set includes no graph nodes") { LabelledKwargDataflowGraphView result = get_labelled_kwarg_dataflow_graph_subgraph( - input, std::unordered_set{}); + input, std::set{}); LabelledKwargDataflowGraphData result_data = get_labelled_kwarg_dataflow_graph_data(result); LabelledKwargDataflowGraphData correct_data = LabelledKwargDataflowGraphData{ - /*node_data=*/std::unordered_map{}, - /*edges=*/std::unordered_set>{}, + /*node_data=*/std::map{}, + /*edges=*/std::set>{}, /*output_data=*/ - std::unordered_map, float>{}, + std::map, float>{}, }; CHECK(result_data == correct_data); @@ -167,14 +167,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set includes nodes not in graph") { LabelledKwargDataflowGraphView with_invalid_node = get_labelled_kwarg_dataflow_graph_subgraph( - input, std::unordered_set{n2, n3, Node{100}}); + input, std::set{n2, n3, Node{100}}); LabelledKwargDataflowGraphData with_invalid_node_data = get_labelled_kwarg_dataflow_graph_data(with_invalid_node); LabelledKwargDataflowGraphView without_invalid_node = get_labelled_kwarg_dataflow_graph_subgraph( - input, std::unordered_set{n2, n3}); + input, std::set{n2, n3}); LabelledKwargDataflowGraphData without_invalid_node_data = get_labelled_kwarg_dataflow_graph_data(without_invalid_node); diff --git a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc index 0cb75bb5da..98b61d8542 100644 --- a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc +++ b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc @@ -67,7 +67,7 @@ TEST_SUITE(FF_TEST_SUITE) { float n4_label = 7.8; float n5_label = 2.2; - std::unordered_map node_labelling = { + std::map node_labelling = { {n1, 3.5}, {n2, 1.2}, {n3, 1.2}, @@ -80,7 +80,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::string n3_1_label = "c"; std::string n5_0_label = "d"; - std::unordered_map, std::string> value_labelling = + std::map, std::string> value_labelling = { {n1_0, n1_0_label}, {n2_3, n2_3_label}, diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc index d9d91a03e9..6c16277914 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc @@ -26,9 +26,9 @@ TEST_SUITE(FF_TEST_SUITE) { auto dst = [&](MultiDiEdge const &e) { return g.get_multidiedge_dst(e); }; SUBCASE("adds only those edges") { - std::unordered_set added = + std::set added = g.query_edges(multidiedge_query_all()); - std::unordered_set returned = unordered_set_of(result); + std::set returned = set_of(result); CHECK(returned == added); } @@ -37,7 +37,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("returns unique edges") { - CHECK(unordered_set_of(result).size() == result.size()); + CHECK(set_of(result).size() == result.size()); } SUBCASE("edge 0") { diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_nodes.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_nodes.cc index e3d9ee6a29..d6b57fbe19 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_nodes.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_nodes.cc @@ -9,8 +9,8 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("add_nodes(MultiDiGraph &, int)") { MultiDiGraph g = MultiDiGraph::create(); - std::unordered_set result = unordered_set_of(add_nodes(g, 3_n)); - std::unordered_set correct = g.query_nodes(node_query_all()); + std::set result = set_of(add_nodes(g, 3_n)); + std::set correct = g.query_nodes(node_query_all()); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_edges.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_edges.cc index 0dfcc8a851..fc827deb69 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_edges.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_edges.cc @@ -20,8 +20,8 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(0), n.at(0)}, }); - std::unordered_set result = get_edges(g); - std::unordered_set correct = unordered_set_of(e); + std::set result = get_edges(g); + std::set correct = set_of(e); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc index ef5cf3c502..4380c1c76e 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc @@ -5,7 +5,7 @@ #include "utils/graph/multidigraph/algorithms/add_nodes.h" #include "utils/graph/multidigraph/multidigraph.h" #include -#include +#include #include using namespace FlexFlow; @@ -25,25 +25,25 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_incoming_edges(MultiDiGraphView, Node)") { SUBCASE("node has incoming edges") { - std::unordered_set result = get_incoming_edges(g, n.at(1)); - std::unordered_set correct = {edges.at(1), edges.at(2)}; + std::set result = get_incoming_edges(g, n.at(1)); + std::set correct = {edges.at(1), edges.at(2)}; CHECK(result == correct); } SUBCASE("node has no incoming edges") { - std::unordered_set result = get_incoming_edges(g, n.at(2)); - std::unordered_set correct = {}; + std::set result = get_incoming_edges(g, n.at(2)); + std::set correct = {}; CHECK(result == correct); } } - SUBCASE("get_incoming_edges(MultiDiGraphView, std::unordered_set)") { + SUBCASE("get_incoming_edges(MultiDiGraphView, std::set)") { - std::unordered_set ns = {n.at(0), n.at(2)}; - std::unordered_map> result = + std::set ns = {n.at(0), n.at(2)}; + std::map> result = get_incoming_edges(g, ns); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {edges.at(0), edges.at(3), edges.at(4)}}, {n.at(2), {}}}; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc index 20011cb133..09ba24f997 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc @@ -5,7 +5,7 @@ #include "utils/graph/multidigraph/algorithms/add_nodes.h" #include "utils/graph/multidigraph/multidigraph.h" #include -#include +#include #include using namespace FlexFlow; @@ -27,26 +27,26 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_outgoing_edges(MultiDiGraphView, Node)") { SUBCASE("node has outgoing edges") { - std::unordered_set result = get_outgoing_edges(g, n.at(0)); - std::unordered_set correct = { + std::set result = get_outgoing_edges(g, n.at(0)); + std::set correct = { edges.at(0), edges.at(1), edges.at(2), edges.at(3)}; CHECK(result == correct); } SUBCASE("node has no outgoing edges") { - std::unordered_set result = get_outgoing_edges(g, n.at(2)); - std::unordered_set correct = {}; + std::set result = get_outgoing_edges(g, n.at(2)); + std::set correct = {}; CHECK(result == correct); } } - SUBCASE("get_outgoing_edges(MultiDiGraphView, std::unordered_set)") { + SUBCASE("get_outgoing_edges(MultiDiGraphView, std::set)") { - std::unordered_set ns = {n.at(0), n.at(1)}; - std::unordered_map> result = + std::set ns = {n.at(0), n.at(1)}; + std::map> result = get_outgoing_edges(g, ns); - std::unordered_map> correct = { + std::map> correct = { {n.at(0), {edges.at(0), edges.at(1), edges.at(2), edges.at(3)}}, {n.at(1), {edges.at(4)}}}; diff --git a/lib/utils/test/src/utils/graph/multidigraph/multidigraph.cc b/lib/utils/test/src/utils/graph/multidigraph/multidigraph.cc index f427a78b64..f5c39df6a7 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/multidigraph.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/multidigraph.cc @@ -4,7 +4,7 @@ #include "utils/graph/multidigraph/multidiedge_query.h" #include "utils/graph/query_set.h" #include -#include +#include #include using namespace FlexFlow; @@ -31,8 +31,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n3), }; - std::unordered_set result = g.query_nodes(query); - std::unordered_set correct = {n3}; + std::set result = g.query_nodes(query); + std::set correct = {n3}; CHECK(result == correct); } @@ -44,8 +44,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n1), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = {e7}; + std::set result = g.query_edges(query); + std::set correct = {e7}; CHECK(result == correct); } @@ -58,8 +58,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n1), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = {e7, e8}; + std::set result = g.query_edges(query); + std::set correct = {e7, e8}; CHECK(result == correct); } } @@ -70,45 +70,45 @@ TEST_SUITE(FF_TEST_SUITE) { NodeQuery node_query = NodeQuery{ query_set::match_single_value(n0), }; - std::unordered_set node_result = g.query_nodes(node_query); - std::unordered_set node_correct = {}; + std::set node_result = g.query_nodes(node_query); + std::set node_correct = {}; CHECK(node_result == node_correct); MultiDiEdgeQuery edge_query = MultiDiEdgeQuery{ query_set::match_single_value(n0), query_set::match_values_in(std::set{n1, n2}), }; - std::unordered_set edge_result = g.query_edges(edge_query); - std::unordered_set edge_correct = {}; + std::set edge_result = g.query_edges(edge_query); + std::set edge_correct = {}; CHECK(edge_result == edge_correct); } SUBCASE("remove_edge") { g.remove_edge(e3); - std::unordered_set result = g.query_edges(MultiDiEdgeQuery{ + std::set result = g.query_edges(MultiDiEdgeQuery{ query_set::match_single_value(n1), query_set::match_single_value(n2), }); - std::unordered_set correct = {e4}; + std::set correct = {e4}; CHECK(result == correct); SUBCASE("remove non-duplicate edge") { g.remove_edge(e0); - std::unordered_set result = g.query_edges(MultiDiEdgeQuery{ + std::set result = g.query_edges(MultiDiEdgeQuery{ query_set::match_single_value(n0), query_set::match_single_value(n2), }); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } SUBCASE("remove duplicate edge") { g.remove_edge(e1); - std::unordered_set result = g.query_edges(MultiDiEdgeQuery{ + std::set result = g.query_edges(MultiDiEdgeQuery{ query_set::match_single_value(n1), query_set::match_single_value(n0), }); - std::unordered_set correct = {e2}; + std::set correct = {e2}; CHECK(result == correct); } } @@ -119,8 +119,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n0, n1, n2}), }; - std::unordered_set result = g.query_nodes(query); - std::unordered_set correct = {n0, n1, n2}; + std::set result = g.query_nodes(query); + std::set correct = {n0, n1, n2}; CHECK(result == correct); } @@ -129,8 +129,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n0, n2}), }; - std::unordered_set result = g.query_nodes(query); - std::unordered_set correct = {n0, n2}; + std::set result = g.query_nodes(query); + std::set correct = {n0, n2}; CHECK(result == correct); } @@ -139,8 +139,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::matchall(), }; - std::unordered_set result = g.query_nodes(query); - std::unordered_set correct = {n0, n1, n2}; + std::set result = g.query_nodes(query); + std::set correct = {n0, n1, n2}; CHECK(result == correct); } @@ -152,8 +152,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n3, n4}), }; - std::unordered_set result = g.query_nodes(query); - std::unordered_set correct = {}; + std::set result = g.query_nodes(query); + std::set correct = {}; CHECK(result == correct); } } @@ -165,8 +165,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n0, n1, n2}), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = {e0, e1, e2, e3, e4, e5, e6}; + std::set result = g.query_edges(query); + std::set correct = {e0, e1, e2, e3, e4, e5, e6}; CHECK(result == correct); } @@ -176,8 +176,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n0, n1, n2}), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = {e1, e2, e3, e4}; + std::set result = g.query_edges(query); + std::set correct = {e1, e2, e3, e4}; CHECK(result == correct); } @@ -187,8 +187,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n2), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = {e0, e3, e4, e6}; + std::set result = g.query_edges(query); + std::set correct = {e0, e3, e4, e6}; CHECK(result == correct); } @@ -198,8 +198,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::matchall(), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = {e0, e1, e2, e3, e4, e5, e6}; + std::set result = g.query_edges(query); + std::set correct = {e0, e1, e2, e3, e4, e5, e6}; CHECK(result == correct); } @@ -212,8 +212,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n4), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = {}; + std::set result = g.query_edges(query); + std::set correct = {}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc index fd54b801ce..6a5c9fbc2a 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc @@ -15,9 +15,9 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult n0_added = g.add_node({}, 1_n); - std::unordered_set result = + std::set result = get_open_dataflow_graph_inputs(g); - std::unordered_set correct = {i0, i1}; + std::set correct = {i0, i1}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc index c7d294a588..b25f1f3349 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_value_uses.cc @@ -27,13 +27,13 @@ TEST_SUITE(FF_TEST_SUITE) { 1_n); Node n1 = n1_added.node; - std::unordered_set correct = { + std::set correct = { DataflowInput{n0, 0_n}, DataflowInput{n0, 2_n}, DataflowInput{n1, 2_n}, }; - std::unordered_set result = + std::set result = get_open_dataflow_value_uses(g, OpenDataflowValue{i0}); CHECK(result == correct); @@ -60,12 +60,12 @@ TEST_SUITE(FF_TEST_SUITE) { g.add_node({OpenDataflowValue{o0_1}, OpenDataflowValue{i0}}, 1_n); Node n2 = n2_added.node; - std::unordered_set correct = { + std::set correct = { DataflowInput{n1, 1_n}, DataflowInput{n2, 0_n}, }; - std::unordered_set result = + std::set result = get_open_dataflow_value_uses(g, OpenDataflowValue{o0_1}); CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc index c44e5f81b7..e29788c138 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc @@ -3,6 +3,7 @@ #include "utils/containers/contains.h" #include "utils/containers/get_only.h" #include "utils/graph/instances/unordered_set_dataflow_graph.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h" #include "utils/graph/open_dataflow_graph/open_dataflow_graph.h" @@ -12,7 +13,7 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_full_graph_values_to_subgraph_inputs(OpenDataflowGraphView, " - "std::unordered_set) ") { + "std::set) ") { OpenDataflowGraph graph = OpenDataflowGraph::create(); @@ -36,14 +37,14 @@ TEST_SUITE(FF_TEST_SUITE) { graph.add_node({OpenDataflowValue{i2}, v1, v2}, 1_n); Node n3 = n3_added.node; - std::unordered_set subgraph_nodes = {n1, n2, n3}; + std::set subgraph_nodes = {n1, n2, n3}; bidict full_graph_values_to_subgraph_inputs = get_full_graph_values_to_subgraph_inputs(graph, subgraph_nodes); SUBCASE("left entries are correct") { - std::unordered_set correct = { + std::set correct = { v0, OpenDataflowValue{i1}, OpenDataflowValue{i2}}; CHECK(left_entries(full_graph_values_to_subgraph_inputs) == correct); } @@ -53,13 +54,13 @@ TEST_SUITE(FF_TEST_SUITE) { i1); CHECK(full_graph_values_to_subgraph_inputs.at_l(OpenDataflowValue{i2}) == i2); - std::unordered_set inputs = {i1, i2}; + std::set inputs = {i1, i2}; CHECK(!contains(inputs, full_graph_values_to_subgraph_inputs.at_l(v0))); } } TEST_CASE( - "get_subgraph_data(OpenDataflowGraphView, std::unordered_set, " + "get_subgraph_data(OpenDataflowGraphView, std::set, " "bidict)") { SUBCASE("2-node graph without inputs") { OpenDataflowGraph graph = @@ -73,7 +74,7 @@ TEST_SUITE(FF_TEST_SUITE) { Node n1 = n1_added.node; SUBCASE("subgraph is full graph") { - std::unordered_set subgraph_nodes = {n0, n1}; + std::set subgraph_nodes = {n0, n1}; bidict full_graph_values_to_subgraph_inputs = @@ -100,7 +101,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("subgraph is n0") { - std::unordered_set subgraph_nodes = {n0}; + std::set subgraph_nodes = {n0}; bidict full_graph_values_to_subgraph_inputs = @@ -119,7 +120,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("subgraph is n1") { - std::unordered_set subgraph_nodes = {n1}; + std::set subgraph_nodes = {n1}; bidict full_graph_values_to_subgraph_inputs = @@ -144,7 +145,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("subgraph is empty") { - std::unordered_set subgraph_nodes = {}; + std::set subgraph_nodes = {}; bidict full_graph_values_to_subgraph_inputs = @@ -177,7 +178,7 @@ TEST_SUITE(FF_TEST_SUITE) { Node n2 = n2_added.node; SUBCASE("subgraph is full graph") { - std::unordered_set subgraph_nodes = {n0, n1, n2}; + std::set subgraph_nodes = {n0, n1, n2}; bidict full_graph_values_to_subgraph_inputs = @@ -215,7 +216,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("subgraph is (n0, n1) split") { - std::unordered_set subgraph_nodes = {n0, n1}; + std::set subgraph_nodes = {n0, n1}; bidict full_graph_values_to_subgraph_inputs = @@ -247,7 +248,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("subgraph is (n0, n1) split") { - std::unordered_set subgraph_nodes = {n0, n1}; + std::set subgraph_nodes = {n0, n1}; bidict full_graph_values_to_subgraph_inputs = @@ -279,7 +280,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("subgraph is (n0, n2) split") { - std::unordered_set subgraph_nodes = {n0, n2}; + std::set subgraph_nodes = {n0, n2}; bidict full_graph_values_to_subgraph_inputs = @@ -310,7 +311,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("subgraph is (n1, n2) split") { - std::unordered_set subgraph_nodes = {n1, n2}; + std::set subgraph_nodes = {n1, n2}; bidict full_graph_values_to_subgraph_inputs = diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc index e1a2062865..ec5ed15141 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_unused_open_dataflow_graph_inputs.cc @@ -15,10 +15,10 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult g_n1_added = g.add_node({OpenDataflowValue{g_i2}}, 1_n); - std::unordered_set result = + std::set result = get_unused_open_dataflow_graph_inputs(g); - std::unordered_set correct = {g_i1, g_i3}; + std::set correct = {g_i1, g_i3}; CHECK(result == correct); } @@ -30,10 +30,10 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult g_n1_added = g.add_node({OpenDataflowValue{g_i1}, OpenDataflowValue{g_i2}}, 1_n); - std::unordered_set result = + std::set result = get_unused_open_dataflow_graph_inputs(g); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc index 7466c23943..127f03c0d9 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc @@ -95,8 +95,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(n0), }; - std::unordered_set result_nodes = result.query_nodes(query); - std::unordered_set correct = {}; + std::set result_nodes = result.query_nodes(query); + std::set correct = {}; CHECK(result_nodes == correct); } @@ -105,8 +105,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_single_value(new_node0), }; - std::unordered_set result_nodes = result.query_nodes(query); - std::unordered_set correct = {new_node0}; + std::set result_nodes = result.query_nodes(query); + std::set correct = {new_node0}; CHECK(result_nodes == correct); } } @@ -119,9 +119,9 @@ TEST_SUITE(FF_TEST_SUITE) { dataflow_edge_query_for_edge( DataflowEdge{n0_output, DataflowInput{n1, 1_n}}), }; - std::unordered_set result_nodes = + std::set result_nodes = result.query_edges(query); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result_nodes == correct); } @@ -139,9 +139,9 @@ TEST_SUITE(FF_TEST_SUITE) { dataflow_edge_query_for_edge(new_standard_edge), }; - std::unordered_set result_nodes = + std::set result_nodes = result.query_edges(query); - std::unordered_set correct = { + std::set correct = { OpenDataflowEdge{new_standard_edge}, OpenDataflowEdge{new_input_edge}, }; @@ -156,10 +156,10 @@ TEST_SUITE(FF_TEST_SUITE) { DataflowOutputQuery query = dataflow_output_query_for_output(old_output); - std::unordered_set result_outputs = + std::set result_outputs = result.query_outputs(query); - std::unordered_set correct = {}; + std::set correct = {}; CHECK(result_outputs == correct); } @@ -169,10 +169,10 @@ TEST_SUITE(FF_TEST_SUITE) { DataflowOutputQuery query = dataflow_output_query_for_output(new_output); - std::unordered_set result_outputs = + std::set result_outputs = result.query_outputs(query); - std::unordered_set correct = {new_output}; + std::set correct = {new_output}; CHECK(result_outputs == correct); } diff --git a/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.cc b/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.cc index e7b4e176f1..0266ad4958 100644 --- a/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.cc +++ b/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.cc @@ -37,8 +37,8 @@ TEST_SUITE(FF_TEST_SUITE) { OpenKwargDataflowGraphView lhs = mk_graph(); OpenKwargDataflowGraphView rhs = mk_graph(); - std::unordered_set lhs_nodes = get_nodes(lhs); - std::unordered_set rhs_nodes = get_nodes(rhs); + std::set lhs_nodes = get_nodes(lhs); + std::set rhs_nodes = get_nodes(rhs); std::vector ordered_lhs_nodes = vector_of(lhs_nodes); diff --git a/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc b/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc index e96468ac7a..e17079b770 100644 --- a/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc +++ b/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc @@ -23,7 +23,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n1_added = g.add_node( /*inputs=*/ - std::unordered_map>{ + std::map>{ { 1, OpenKwargDataflowValue{input1}, @@ -37,7 +37,7 @@ TEST_SUITE(FF_TEST_SUITE) { OpenKwargDataflowValue{input1}, }, }, - /*outputs=*/std::unordered_set{ + /*outputs=*/std::set{ 5, }); @@ -46,7 +46,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult n2_added = g.add_node( /*inputs=*/ - std::unordered_map>{ + std::map>{ { 4, OpenKwargDataflowValue{input2}, @@ -56,7 +56,7 @@ TEST_SUITE(FF_TEST_SUITE) { OpenKwargDataflowValue{n1_output}, }, }, - /*outputs=*/std::unordered_set{ + /*outputs=*/std::set{ 5, }); @@ -80,7 +80,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult> input1_added = g.add_node( /*inputs=*/{}, - /*outputs=*/std::unordered_set>{ + /*outputs=*/std::set>{ std::nullopt, }); @@ -89,7 +89,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult> input2_added = g.add_node( /*inputs=*/{}, - /*outputs=*/std::unordered_set>{ + /*outputs=*/std::set>{ std::nullopt, }); @@ -98,7 +98,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult> n1_added = g.add_node( /*inputs=*/ - std::unordered_map, + std::map, KwargDataflowOutput>>{ { 1, @@ -113,7 +113,7 @@ TEST_SUITE(FF_TEST_SUITE) { input1, }, }, - /*outputs=*/std::unordered_set>{ + /*outputs=*/std::set>{ 5, }); @@ -122,7 +122,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult> n2_added = g.add_node( /*inputs=*/ - std::unordered_map, + std::map, KwargDataflowOutput>>{ { 4, @@ -133,7 +133,7 @@ TEST_SUITE(FF_TEST_SUITE) { n1_output, }, }, - /*outputs=*/std::unordered_set>{ + /*outputs=*/std::set>{ 5, }); diff --git a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc index c2643506eb..a1ae7254bf 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "test/utils/rapidcheck.h" #include "utils/containers/contains.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h" @@ -65,8 +65,8 @@ TEST_SUITE(FF_TEST_SUITE) { nonnegative_int expected_height = 2_n; CHECK(result_height == expected_height); - std::unordered_multiset result_nodes = get_leaves(result); - std::unordered_multiset expected_nodes = {n1, n2, n3, n4}; + std::multiset result_nodes = get_leaves(result); + std::multiset expected_nodes = {n1, n2, n3, n4}; CHECK(result_nodes == expected_nodes); } @@ -82,7 +82,7 @@ TEST_SUITE(FF_TEST_SUITE) { make_series_split(make_series_split(make_leaf(n1), make_leaf(n2)), make_series_split(make_leaf(n3), make_leaf(n4))); - std::unordered_set corrects = { + std::set corrects = { make_parallel_split(balanced_series, make_leaf(n5)), make_parallel_split(make_leaf(n5), balanced_series)}; diff --git a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc index 9ca869b2b0..1f8f37d967 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.dtg.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h" #include @@ -25,8 +25,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("leaf") { BinarySPDecompositionTree input = BinarySPDecompositionTree{n1}; - std::unordered_multiset result = generic_get_leaves(input); - std::unordered_multiset correct = {n1}; + std::multiset result = generic_get_leaves(input); + std::multiset correct = {n1}; CHECK(result == correct); } @@ -40,8 +40,8 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_multiset result = generic_get_leaves(input); - std::unordered_multiset correct = {n1, n2}; + std::multiset result = generic_get_leaves(input); + std::multiset correct = {n1, n2}; CHECK(result == correct); } @@ -54,8 +54,8 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_multiset result = generic_get_leaves(input); - std::unordered_multiset correct = {n1, n1}; + std::multiset result = generic_get_leaves(input); + std::multiset correct = {n1, n1}; CHECK(result == correct); } @@ -70,8 +70,8 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_multiset result = generic_get_leaves(input); - std::unordered_multiset correct = {n1, n2}; + std::multiset result = generic_get_leaves(input); + std::multiset correct = {n1, n2}; CHECK(result == correct); } @@ -84,8 +84,8 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_multiset result = generic_get_leaves(input); - std::unordered_multiset correct = {n1, n1}; + std::multiset result = generic_get_leaves(input); + std::multiset correct = {n1, n1}; CHECK(result == correct); } @@ -109,8 +109,8 @@ TEST_SUITE(FF_TEST_SUITE) { make_series_split(make_leaf(n2), make_leaf(n3))), make_parallel_split(make_leaf(n2), make_leaf(n1))); - std::unordered_multiset result = generic_get_leaves(input); - std::unordered_multiset correct = {n1, n1, n2, n2, n3}; + std::multiset result = generic_get_leaves(input); + std::multiset correct = {n1, n1, n2, n2, n3}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc index 08ffc44ff3..1589f682dc 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "test/utils/rapidcheck.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/nary_sp_tree_from_binary.h" @@ -67,8 +67,8 @@ TEST_SUITE(FF_TEST_SUITE) { // left-associative binary SP trees CHECK(is_binary_sp_tree_left_associative(result)); - std::unordered_multiset result_nodes = get_leaves(result); - std::unordered_multiset correct_nodes = {n1, n2, n3}; + std::multiset result_nodes = get_leaves(result); + std::multiset correct_nodes = {n1, n2, n3}; CHECK(result_nodes == correct_nodes); } @@ -96,8 +96,8 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(is_binary_sp_tree_left_associative(result)); - std::unordered_multiset result_nodes = get_leaves(result); - std::unordered_multiset correct_nodes = { + std::multiset result_nodes = get_leaves(result); + std::multiset correct_nodes = { n1, n2, n3, n3, n5, n6, n4, n5}; CHECK(result_nodes == correct_nodes); diff --git a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc index 7b43f52b8f..bb4aa5ec54 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.h" #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include @@ -65,8 +65,8 @@ TEST_SUITE(FF_TEST_SUITE) { // right-associative binary SP trees CHECK(is_binary_sp_tree_right_associative(result)); - std::unordered_multiset result_nodes = get_nodes(input); - std::unordered_multiset correct_nodes = {n1, n2, n3}; + std::multiset result_nodes = get_nodes(input); + std::multiset correct_nodes = {n1, n2, n3}; CHECK(result_nodes == correct_nodes); } @@ -94,8 +94,8 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(is_binary_sp_tree_right_associative(result)); - std::unordered_multiset result_nodes = get_nodes(input); - std::unordered_multiset correct_nodes = { + std::multiset result_nodes = get_nodes(input); + std::multiset correct_nodes = { n1, n2, n3, n3, n5, n6, n4, n5}; CHECK(result_nodes == correct_nodes); diff --git a/lib/utils/test/src/utils/graph/series_parallel/get_ancestors.cc b/lib/utils/test/src/utils/graph/series_parallel/get_ancestors.cc index 22e1fe51cf..ca0f821161 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/get_ancestors.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/get_ancestors.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/get_ancestors.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include #include @@ -12,24 +12,24 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Single Node") { SeriesParallelDecomposition sp = SeriesParallelDecomposition{n.at(0)}; - std::unordered_set correct = {}; - std::unordered_set result = get_ancestors(sp, n.at(0)); + std::set correct = {}; + std::set result = get_ancestors(sp, n.at(0)); CHECK(correct == result); } SUBCASE("Simple Series") { SeriesParallelDecomposition sp = SeriesParallelDecomposition{SeriesSplit{{n.at(0), n.at(1), n.at(2)}}}; - std::unordered_set correct = {n.at(0), n.at(1)}; - std::unordered_set result = get_ancestors(sp, n.at(2)); + std::set correct = {n.at(0), n.at(1)}; + std::set result = get_ancestors(sp, n.at(2)); CHECK(correct == result); } SUBCASE("Simple Parallel") { SeriesParallelDecomposition sp = SeriesParallelDecomposition{ ParallelSplit{{n.at(0), n.at(1), n.at(2)}}}; - std::unordered_set correct = {}; - std::unordered_set result = get_ancestors(sp, n.at(1)); + std::set correct = {}; + std::set result = get_ancestors(sp, n.at(1)); CHECK(correct == result); } @@ -37,16 +37,16 @@ TEST_SUITE(FF_TEST_SUITE) { SeriesParallelDecomposition sp = SeriesParallelDecomposition{SeriesSplit{ {n.at(0), ParallelSplit{{SeriesSplit{{n.at(1), n.at(2)}}, n.at(3)}}}}}; - std::unordered_set correct = {n.at(0), n.at(1)}; - std::unordered_set result = get_ancestors(sp, n.at(2)); + std::set correct = {n.at(0), n.at(1)}; + std::set result = get_ancestors(sp, n.at(2)); CHECK(correct == result); } SUBCASE("Rhombus") { SeriesParallelDecomposition sp = SeriesParallelDecomposition{ SeriesSplit{{n.at(0), ParallelSplit{{n.at(1), n.at(2)}}, n.at(3)}}}; - std::unordered_set correct = {n.at(0), n.at(1), n.at(2)}; - std::unordered_set result = get_ancestors(sp, n.at(3)); + std::set correct = {n.at(0), n.at(1), n.at(2)}; + std::set result = get_ancestors(sp, n.at(3)); CHECK(correct == result); } @@ -58,8 +58,8 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(1), ParallelSplit{{n.at(2), n.at(3)}}, n.at(4)}}, SeriesSplit{{n.at(5), n.at(6)}}}}, n.at(7)}}}; - std::unordered_set correct = {n.at(0), n.at(1), n.at(2), n.at(3)}; - std::unordered_set result = get_ancestors(sp, n.at(4)); + std::set correct = {n.at(0), n.at(1), n.at(2), n.at(3)}; + std::set result = get_ancestors(sp, n.at(4)); CHECK(correct == result); correct = {n.at(0), n.at(1)}; diff --git a/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc b/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc index a2f818b5e9..a33fbf3df3 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc @@ -76,7 +76,7 @@ TEST_SUITE(FF_TEST_SUITE) { }); std::optional result = find_parallel_reduction(g); - std::unordered_set correct_options = { + std::set correct_options = { make_parallel_reduction(e.at(0), e.at(1)), make_parallel_reduction(e.at(1), e.at(2)), make_parallel_reduction(e.at(0), e.at(2)), @@ -121,22 +121,22 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiEdge returned_edge = apply_parallel_reduction(g, input); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(g); - std::unordered_set correct_nodes = unordered_set_of(n); + std::set result_nodes = get_nodes(g); + std::set correct_nodes = set_of(n); CHECK(result_nodes == correct_nodes); } SUBCASE("edge shape") { - std::unordered_map result_edges = get_edge_counts(g); - std::unordered_map correct_edges = { + std::map result_edges = get_edge_counts(g); + std::map correct_edges = { {DirectedEdge{n.at(0), n.at(1)}, 1}, }; CHECK(result_edges == correct_edges); } SUBCASE("return value and edge ids") { - std::unordered_set result_edge_ids = get_edges(g); - std::unordered_set correct_edge_ids = {returned_edge}; + std::set result_edge_ids = get_edges(g); + std::set correct_edge_ids = {returned_edge}; CHECK(result_edge_ids == correct_edge_ids); } } @@ -154,7 +154,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(3), n.at(4)}, }); - std::unordered_map input_edge_counts = + std::map input_edge_counts = get_edge_counts(g); MultiDiEdge reduction_e1 = e.at(3); @@ -165,15 +165,15 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiEdge returned_edge = apply_parallel_reduction(g, input); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(g); - std::unordered_set correct_nodes = unordered_set_of(n); + std::set result_nodes = get_nodes(g); + std::set correct_nodes = set_of(n); CHECK(result_nodes == correct_nodes); } SUBCASE("edge shape") { - std::unordered_map result_edges = get_edge_counts(g); - std::unordered_map correct_edges = [&] { - std::unordered_map new_edge_counts = + std::map result_edges = get_edge_counts(g); + std::map correct_edges = [&] { + std::map new_edge_counts = input_edge_counts; new_edge_counts.at(get_directed_edge(g, reduction_e1))--; return new_edge_counts; @@ -182,9 +182,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("return value and edge ids") { - std::unordered_set result_edge_ids = get_edges(g); - std::unordered_set correct_edge_ids = [&] { - std::unordered_set new_edges = unordered_set_of(e); + std::set result_edge_ids = get_edges(g); + std::set correct_edge_ids = [&] { + std::set new_edges = set_of(e); new_edges.erase(reduction_e1); new_edges.erase(reduction_e2); new_edges.insert(returned_edge); diff --git a/lib/utils/test/src/utils/graph/series_parallel/series_parallel_decomposition.cc b/lib/utils/test/src/utils/graph/series_parallel/series_parallel_decomposition.cc index f5766c9fdd..a1e777abc2 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/series_parallel_decomposition.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/series_parallel_decomposition.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include using namespace ::FlexFlow; @@ -84,8 +84,8 @@ TEST_SUITE(FF_TEST_SUITE) { }}, }}}; - std::unordered_multiset result = get_nodes(input); - std::unordered_multiset correct = { + std::multiset result = get_nodes(input); + std::multiset correct = { Node{1}, Node{2}, Node{2}, @@ -112,8 +112,8 @@ TEST_SUITE(FF_TEST_SUITE) { Node{7}, }}; - std::unordered_multiset result = get_nodes(input); - std::unordered_multiset correct = { + std::multiset result = get_nodes(input); + std::multiset correct = { Node{1}, Node{2}, Node{3}, @@ -139,8 +139,8 @@ TEST_SUITE(FF_TEST_SUITE) { }}, }}; - std::unordered_multiset result = get_nodes(input); - std::unordered_multiset correct = { + std::multiset result = get_nodes(input); + std::multiset correct = { Node{1}, Node{2}, Node{4}, @@ -153,8 +153,8 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_nodes(Node)") { Node input = Node{5}; - std::unordered_multiset result = get_nodes(input); - std::unordered_multiset correct = {input}; + std::multiset result = get_nodes(input); + std::multiset correct = {input}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc b/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc index 78537a4342..d6b06d28d4 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc @@ -1,6 +1,6 @@ #include "utils/graph/series_parallel/series_reduction.h" #include "utils/containers/set_minus.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/fmt/vector.h" #include "utils/graph/instances/adjacency_multidigraph.h" #include "utils/graph/multidigraph/algorithms/add_edges.h" @@ -119,7 +119,7 @@ TEST_SUITE(FF_TEST_SUITE) { }); std::optional result = find_series_reduction(g); - std::unordered_set correct_options = { + std::set correct_options = { make_series_reduction(e.at(0), e.at(1)), make_series_reduction(e.at(1), e.at(2)), }; @@ -164,14 +164,14 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiEdge returned_edge = apply_series_reduction(g, reduction); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(g); - std::unordered_set correct_nodes = {n.at(0), n.at(2)}; + std::set result_nodes = get_nodes(g); + std::set correct_nodes = {n.at(0), n.at(2)}; CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(g); - std::unordered_set correct_edges = {returned_edge}; + std::set result_edges = get_edges(g); + std::set correct_edges = {returned_edge}; CHECK(result_edges == correct_edges); } @@ -212,16 +212,16 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiEdge returned_edge = apply_series_reduction(g, reduction); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(g); - std::unordered_set correct_nodes = - set_minus(unordered_set_of(n), {n.at(4)}); + std::set result_nodes = get_nodes(g); + std::set correct_nodes = + set_minus(set_of(n), {n.at(4)}); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(g); - std::unordered_set correct_edges = [&] { - std::unordered_set new_edges = unordered_set_of(e); + std::set result_edges = get_edges(g); + std::set correct_edges = [&] { + std::set new_edges = set_of(e); new_edges.erase(reduction_e1); new_edges.erase(reduction_e2); new_edges.insert(returned_edge); @@ -258,9 +258,9 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(2), n.at(3)}, }); - std::unordered_set result = + std::set result = find_all_extended_series_reductions(g); - std::unordered_set correct = { + std::set correct = { ExtendedSeriesReduction{{e.at(0), e.at(1), e.at(2)}}}; CHECK(result == correct); } @@ -273,9 +273,9 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(1), n.at(3)}, {n.at(2), n.at(3)}}); - std::unordered_set result = + std::set result = find_all_extended_series_reductions(g); - std::unordered_set correct = { + std::set correct = { ExtendedSeriesReduction{{e.at(0), e.at(2)}}, ExtendedSeriesReduction{{e.at(1), e.at(3)}}}; CHECK(result == correct); @@ -296,9 +296,9 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(6), n.at(8)}, {n.at(7), n.at(8)}}); - std::unordered_set result = + std::set result = find_all_extended_series_reductions(g); - std::unordered_set correct = { + std::set correct = { ExtendedSeriesReduction{{e.at(0), e.at(2), e.at(7)}}, ExtendedSeriesReduction{{e.at(3), e.at(6)}}, ExtendedSeriesReduction{{e.at(5), e.at(9)}}}; @@ -320,14 +320,14 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiEdge returned_edge = apply_extended_series_reduction(g, reduction); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(g); - std::unordered_set correct_nodes = {n.at(0), n.at(3)}; + std::set result_nodes = get_nodes(g); + std::set correct_nodes = {n.at(0), n.at(3)}; CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(g); - std::unordered_set correct_edges = {returned_edge}; + std::set result_edges = get_edges(g); + std::set correct_edges = {returned_edge}; CHECK(result_edges == correct_edges); } @@ -366,16 +366,16 @@ TEST_SUITE(FF_TEST_SUITE) { MultiDiEdge returned_edge = apply_extended_series_reduction(g, reduction); SUBCASE("nodes") { - std::unordered_set result_nodes = get_nodes(g); - std::unordered_set correct_nodes = - set_minus(unordered_set_of(n), {n.at(4), n.at(3)}); + std::set result_nodes = get_nodes(g); + std::set correct_nodes = + set_minus(set_of(n), {n.at(4), n.at(3)}); CHECK(result_nodes == correct_nodes); } SUBCASE("edges") { - std::unordered_set result_edges = get_edges(g); - std::unordered_set correct_edges = [&] { - std::unordered_set new_edges = unordered_set_of(e); + std::set result_edges = get_edges(g); + std::set correct_edges = [&] { + std::set new_edges = set_of(e); new_edges = set_minus(new_edges, {e.at(3), e.at(4), e.at(5)}); new_edges.insert(returned_edge); return new_edges; diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc index d7bd9ae31d..002e77828b 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc @@ -1,5 +1,5 @@ #include "utils/graph/series_parallel/sp_ization/escribano_algo.h" -#include "test/utils/doctest/fmt/unordered_multiset.h" +#include "test/utils/doctest/fmt/multiset.h" #include "utils/containers/values.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_edges.h" @@ -17,7 +17,7 @@ #include "utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h" #include "utils/graph/series_parallel/sp_ization/node_role.h" #include -#include +#include using namespace FlexFlow; @@ -33,7 +33,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(2), n.at(3)}, }); - std::unordered_map node_types = { + std::map node_types = { {n.at(0), NodeRole::PURE}, {n.at(1), NodeRole::PURE}, {n.at(2), NodeRole::PURE}, @@ -48,7 +48,7 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(node_types.size() == 6); CHECK(values(node_types) == - std::unordered_multiset{NodeRole::PURE, + std::multiset{NodeRole::PURE, NodeRole::PURE, NodeRole::PURE, NodeRole::PURE, @@ -69,19 +69,19 @@ TEST_SUITE(FF_TEST_SUITE) { {DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(1), n.at(3)}}); - std::unordered_map node_roles = { + std::map node_roles = { {n.at(0), NodeRole::PURE}, {n.at(1), NodeRole::SYNC}, {n.at(2), NodeRole::PURE}, {n.at(3), NodeRole::PURE}, }; - std::unordered_map depth_map = { + std::map depth_map = { {n.at(0), 0_n}, {n.at(2), 1_n}, {n.at(3), 1_n}, }; - std::unordered_set correct = {n.at(0), n.at(2), n.at(3)}; - std::unordered_set result = + std::set correct = {n.at(0), n.at(2), n.at(3)}; + std::set result = get_component(g, n.at(2), depth_map, node_roles); CHECK(correct == result); } @@ -94,7 +94,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(2), n.at(4)}, DirectedEdge{n.at(3), n.at(4)}, DirectedEdge{n.at(3), n.at(5)}}); - std::unordered_map node_roles = { + std::map node_roles = { {n.at(0), NodeRole::PURE}, {n.at(1), NodeRole::PURE}, {n.at(2), NodeRole::SYNC}, @@ -102,23 +102,23 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(4), NodeRole::PURE}, {n.at(5), NodeRole::PURE}, }; - std::unordered_map depth_map = { + std::map depth_map = { {n.at(0), 0_n}, {n.at(1), 0_n}, {n.at(4), 1_n}, {n.at(5), 1_n}, }; SUBCASE("n.at(4)'s component") { - std::unordered_set correct = { + std::set correct = { n.at(0), n.at(1), n.at(4), n.at(5)}; - std::unordered_set result = + std::set result = get_component(g, n.at(4), depth_map, node_roles); CHECK(correct == result); } SUBCASE("n.at(5)'s component") { - std::unordered_set correct = { + std::set correct = { n.at(0), n.at(1), n.at(4), n.at(5)}; - std::unordered_set result = + std::set result = get_component(g, n.at(5), depth_map, node_roles); CHECK(correct == result); } @@ -134,7 +134,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(3), n.at(4)}, DirectedEdge{n.at(4), n.at(5)}, DirectedEdge{n.at(4), n.at(6)}}); - std::unordered_map node_roles = { + std::map node_roles = { {n.at(0), NodeRole::PURE}, {n.at(1), NodeRole::SYNC}, {n.at(2), NodeRole::PURE}, @@ -143,23 +143,23 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), NodeRole::PURE}, {n.at(6), NodeRole::PURE}}; - std::unordered_map depth_map = {{n.at(0), 0_n}, + std::map depth_map = {{n.at(0), 0_n}, {n.at(2), 1_n}, {n.at(3), 1_n}, {n.at(5), 2_n}, {n.at(6), 2_n}}; SUBCASE("n.at(5)'s component") { - std::unordered_set correct = { + std::set correct = { n.at(2), n.at(3), n.at(5), n.at(6)}; - std::unordered_set result = + std::set result = get_component(g, n.at(5), depth_map, node_roles); CHECK(correct == result); } SUBCASE("n.at(6)'s component") { - std::unordered_set correct = { + std::set correct = { n.at(2), n.at(3), n.at(5), n.at(6)}; - std::unordered_set result = + std::set result = get_component(g, n.at(6), depth_map, node_roles); CHECK(correct == result); } @@ -180,7 +180,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(5), n.at(8)}, DirectedEdge{n.at(6), n.at(9)}, }); - std::unordered_map node_roles = { + std::map node_roles = { {n.at(0), NodeRole::PURE}, {n.at(1), NodeRole::SYNC}, {n.at(2), NodeRole::PURE}, @@ -193,7 +193,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(9), NodeRole::PURE}, }; - std::unordered_map depth_map = {{n.at(0), 0_n}, + std::map depth_map = {{n.at(0), 0_n}, {n.at(2), 1_n}, {n.at(3), 1_n}, {n.at(4), 1_n}, @@ -201,20 +201,20 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(8), 2_n}, {n.at(9), 2_n}}; SUBCASE("n.at(7)'s component") { - std::unordered_set correct = {n.at(2), n.at(7), n.at(8)}; - std::unordered_set result = + std::set correct = {n.at(2), n.at(7), n.at(8)}; + std::set result = get_component(g, n.at(7), depth_map, node_roles); CHECK(correct == result); } SUBCASE("n.at(8)'s component") { - std::unordered_set correct = {n.at(2), n.at(7), n.at(8)}; - std::unordered_set result = + std::set correct = {n.at(2), n.at(7), n.at(8)}; + std::set result = get_component(g, n.at(8), depth_map, node_roles); CHECK(correct == result); } SUBCASE("n.at(9)'s component") { - std::unordered_set correct = {n.at(3), n.at(4), n.at(9)}; - std::unordered_set result = + std::set correct = {n.at(3), n.at(4), n.at(9)}; + std::set result = get_component(g, n.at(9), depth_map, node_roles); CHECK(correct == result); } diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc index 110c03c7bb..452608e83c 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc @@ -16,8 +16,8 @@ #include "utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h" #include "utils/graph/series_parallel/sp_ization/node_role.dtg.h" #include -#include -#include +#include +#include using namespace FlexFlow; @@ -27,7 +27,7 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraph g = DiGraph::create(); Node n0 = g.add_node(); - std::unordered_map cost_map = {{n0, 1.0f}}; + std::map cost_map = {{n0, 1.0f}}; SeriesParallelDecomposition result = flexible_sp_ization(g, cost_map); SeriesParallelDecomposition correct = SeriesParallelDecomposition{n0}; @@ -45,7 +45,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[1], n[3]}, DirectedEdge{n[2], n[3]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 1.0f}, @@ -68,7 +68,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[1], n[2]}, DirectedEdge{n[2], n[3]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 1.0f}, @@ -95,7 +95,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[3], n[5]}, DirectedEdge{n[4], n[5]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 1.0f}, @@ -127,7 +127,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[3], n[5]}, DirectedEdge{n[4], n[5]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 10.0f}, @@ -159,7 +159,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[3], n[5]}, DirectedEdge{n[4], n[5]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 1000.0f}, @@ -192,7 +192,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[5], n[7]}, DirectedEdge{n[6], n[7]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 1.0f}, @@ -228,7 +228,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[0], n[5]}, DirectedEdge{n[5], n[6]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 100.0f}, @@ -268,7 +268,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[2], n[7]}, DirectedEdge{n[7], n[5]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 1.0f}, {n[2], 1.0f}, @@ -308,7 +308,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[2], n[7]}, DirectedEdge{n[7], n[5]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 100.0f}, {n[2], 1.0f}, @@ -350,7 +350,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{m[4], m[5]}, DirectedEdge{m[5], m[6]}}); - std::unordered_map cost_map2 = { + std::map cost_map2 = { {m[0], 1.0f}, {m[1], 1.0f}, {m[2], 1.0f}, @@ -389,7 +389,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[14], n[15]}, DirectedEdge{n[15], n[16]}, DirectedEdge{n[16], n[17]}}); - std::unordered_map cost_map; + std::map cost_map; for (int i = 0; i < 18; i++) { cost_map[n[i]] = 1.0f; } @@ -429,7 +429,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[13], n[15]}, DirectedEdge{n[14], n[15]}, DirectedEdge{n[15], n[16]}, DirectedEdge{n[16], n[17]}}); - std::unordered_map cost_map; + std::map cost_map; for (int i = 0; i < 18; i++) { cost_map[n[i]] = 1.0f; } @@ -455,7 +455,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n[13], n[15]}, DirectedEdge{n[14], n[15]}, DirectedEdge{n[15], n[16]}, DirectedEdge{n[16], n[17]}}); - std::unordered_map cost_map = { + std::map cost_map = { {n[0], 1.0f}, {n[1], 3.0f}, {n[2], 5.0f}, diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc index e783c2b2f3..5679d53f08 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc @@ -18,7 +18,7 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraph g = DiGraph::create(); std::vector n = add_nodes(g, 4); - std::unordered_map cost_map = { + std::map cost_map = { {n.at(0), 1.0f}, {n.at(1), 5.0f}, {n.at(2), 2.0f}, {n.at(3), 3.0f}}; SeriesParallelDecomposition sp = naive_stratum_sync_sp_ization(g); @@ -43,7 +43,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(3), n.at(4)}, }); - std::unordered_map cost_map = {{n.at(0), 1.0f}, + std::map cost_map = {{n.at(0), 1.0f}, {n.at(1), 2.0f}, {n.at(2), 3.0f}, {n.at(3), 4.0f}, @@ -73,7 +73,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(4)}, }); - std::unordered_map cost_map = {{n.at(0), 2.0f}, + std::map cost_map = {{n.at(0), 2.0f}, {n.at(1), 3.0f}, {n.at(2), 5.0f}, {n.at(3), 7.0f}, @@ -105,7 +105,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(5)}, }); - std::unordered_map cost_map = {{n.at(0), 1.0f}, + std::map cost_map = {{n.at(0), 1.0f}, {n.at(1), 1.0f}, {n.at(2), 10.0f}, {n.at(3), 1.0f}, diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc index 4a3dece443..df7e980db1 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc @@ -7,7 +7,7 @@ #include "utils/graph/node/node.dtg.h" #include "utils/graph/series_parallel/sp_ization/node_role.dtg.h" #include -#include +#include using namespace FlexFlow; @@ -25,7 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; add_edges(g, edges); - std::unordered_map node_roles = { + std::map node_roles = { {n.at(0), NodeRole::PURE}, {n.at(1), NodeRole::DUMMY}, {n.at(2), NodeRole::DUMMY}, @@ -37,9 +37,9 @@ TEST_SUITE(FF_TEST_SUITE) { contract_out_nodes_of_given_role(g, NodeRole::DUMMY, node_roles); CHECK(get_nodes(result) == - std::unordered_set{n.at(0), n.at(3), n.at(4)}); + std::set{n.at(0), n.at(3), n.at(4)}); CHECK(get_edges(result) == - std::unordered_set{DirectedEdge{n.at(0), n.at(4)}, + std::set{DirectedEdge{n.at(0), n.at(4)}, DirectedEdge{n.at(0), n.at(3)}, DirectedEdge{n.at(3), n.at(4)}}); } @@ -52,7 +52,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(1), n.at(3)}}); - std::unordered_map node_roles = { + std::map node_roles = { {n.at(0), NodeRole::PURE}, {n.at(1), NodeRole::PURE}, {n.at(2), NodeRole::PURE}, diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc index 85c4d66f6d..6ebbe5ab42 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc @@ -1,6 +1,6 @@ #include "utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.h" #include "test/utils/rapidcheck.h" -#include "utils/containers/generate_unordered_map.h" +#include "utils/containers/generate_map.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_initial_nodes.h" #include "utils/graph/digraph/algorithms/get_terminal_nodes.h" @@ -15,7 +15,7 @@ using namespace FlexFlow; -static std::pair> +static std::pair> generate_random_2_terminal_weighted_dag(int max_num_nodes = 10, int max_num_edges = 20) { assert(max_num_nodes >= 2); @@ -45,8 +45,8 @@ static std::pair> } } - std::unordered_map cost_map = - generate_unordered_map(get_nodes(g), [](Node const &) { + std::map cost_map = + generate_map(get_nodes(g), [](Node const &) { return static_cast(*rc::gen::inRange(1, 101)); }); diff --git a/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc b/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc index 20b3eaa74a..bb03895540 100644 --- a/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc +++ b/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc @@ -1,5 +1,5 @@ #include "utils/graph/undirected/algorithms/get_connected_components.h" -#include "utils/fmt/unordered_set.h" +#include "utils/fmt/set.h" #include "utils/graph/algorithms.h" #include "utils/graph/instances/hashmap_undirected_graph.h" #include "utils/graph/undirected/algorithms/make_undirected_edge.h" @@ -15,12 +15,12 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("disjoint nodes") { std::vector n = add_nodes(g, 3); - std::unordered_set> correct = { + std::set> correct = { {n.at(0)}, {n.at(1)}, {n.at(2)}, }; - std::unordered_set> result = + std::set> result = get_connected_components(g); CHECK(correct == result); @@ -36,10 +36,10 @@ TEST_SUITE(FF_TEST_SUITE) { make_undirected_edge(n.at(3), n.at(0)), }); - std::unordered_set> correct = { + std::set> correct = { {n.at(0), n.at(1), n.at(2), n.at(3)}, }; - std::unordered_set> result = + std::set> result = get_connected_components(g); CHECK(correct == result); @@ -53,11 +53,11 @@ TEST_SUITE(FF_TEST_SUITE) { make_undirected_edge(n.at(2), n.at(1)), }); - std::unordered_set> correct = { + std::set> correct = { {n.at(0), n.at(1), n.at(2)}, {n.at(3)}, }; - std::unordered_set> result = + std::set> result = get_connected_components(g); CHECK(correct == result); @@ -73,20 +73,20 @@ TEST_SUITE(FF_TEST_SUITE) { make_undirected_edge(n.at(3), n.at(4)), }); - std::unordered_set> correct = { + std::set> correct = { {n.at(0), n.at(1), n.at(2)}, {n.at(3), n.at(4)}, {n.at(5)}, }; - std::unordered_set> result = + std::set> result = get_connected_components(g); CHECK(correct == result); } SUBCASE("empty graph") { - std::unordered_set> correct = {}; - std::unordered_set> result = + std::set> correct = {}; + std::set> result = get_connected_components(g); CHECK(correct == result); diff --git a/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc b/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc index 898c8aa154..df4d28dfd5 100644 --- a/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc +++ b/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc @@ -26,8 +26,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("query_nodes") { SUBCASE("query_all") { - std::unordered_set result = g.query_nodes(node_query_all()); - std::unordered_set correct = std::unordered_set{ + std::set result = g.query_nodes(node_query_all()); + std::set correct = std::set{ n.at(0), n.at(1), n.at(2), @@ -43,8 +43,8 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n.at(0), n.at(2)}), }; - std::unordered_set result = g.query_nodes(query); - std::unordered_set correct = std::unordered_set{ + std::set result = g.query_nodes(query); + std::set correct = std::set{ n.at(0), n.at(2), }; @@ -55,10 +55,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("query_edges") { SUBCASE("query_all") { - std::unordered_set result = + std::set result = g.query_edges(undirected_edge_query_all()); - std::unordered_set correct = { + std::set correct = { e.at(0), e.at(1), e.at(2), @@ -74,9 +74,9 @@ TEST_SUITE(FF_TEST_SUITE) { query_set::match_values_in(std::set{n.at(0), n.at(1)}), }; - std::unordered_set result = g.query_edges(query); - std::unordered_set correct = - std::unordered_set{ + std::set result = g.query_edges(query); + std::set correct = + std::set{ e.at(0), }; @@ -88,35 +88,35 @@ TEST_SUITE(FF_TEST_SUITE) { g.remove_node_unsafe(n.at(0)); CHECK(g.query_nodes(node_query_all()) == - std::unordered_set{n.at(1), n.at(2), n.at(3), n.at(4)}); + std::set{n.at(1), n.at(2), n.at(3), n.at(4)}); // removing a node also removes its adjacent edges CHECK(g.query_edges(undirected_edge_query_all()) == - std::unordered_set{e.at(2), e.at(3), e.at(4)}); + std::set{e.at(2), e.at(3), e.at(4)}); g.remove_node_unsafe(n.at(1)); CHECK(g.query_nodes(node_query_all()) == - std::unordered_set{n.at(2), n.at(3), n.at(4)}); + std::set{n.at(2), n.at(3), n.at(4)}); CHECK(g.query_edges(undirected_edge_query_all()) == - std::unordered_set{e.at(3)}); + std::set{e.at(3)}); } SUBCASE("remove_edge") { g.remove_edge(e.at(0)); CHECK(g.query_edges(undirected_edge_query_all()) == - std::unordered_set{ + std::set{ e.at(1), e.at(2), e.at(3), e.at(4)}); CHECK(g.query_nodes(node_query_all()) == - std::unordered_set{ + std::set{ n.at(0), n.at(1), n.at(2), n.at(3), n.at(4)}); g.remove_edge(e.at(1)); g.remove_edge(e.at(3)); CHECK(g.query_edges(undirected_edge_query_all()) == - std::unordered_set{e.at(2), e.at(4)}); + std::set{e.at(2), e.at(4)}); } } } diff --git a/lib/utils/test/src/utils/graph/views/views.cc b/lib/utils/test/src/utils/graph/views/views.cc index 925c0b1f74..d363264cc1 100644 --- a/lib/utils/test/src/utils/graph/views/views.cc +++ b/lib/utils/test/src/utils/graph/views/views.cc @@ -1,8 +1,8 @@ #include "utils/graph/views/views.h" #include "utils/containers/set_union.h" -#include "utils/containers/unordered_set_of.h" -#include "utils/fmt/unordered_map.h" -#include "utils/fmt/unordered_set.h" +#include "utils/containers/set_of.h" +#include "utils/fmt/map.h" +#include "utils/fmt/set.h" #include "utils/graph/algorithms.h" #include "utils/graph/digraph/algorithms/get_edges.h" #include "utils/graph/instances/adjacency_digraph.h" @@ -26,25 +26,25 @@ TEST_SUITE(FF_TEST_SUITE) { make_undirected_edge(n.at(1), n.at(3)), make_undirected_edge(n.at(2), n.at(3)), make_undirected_edge(n.at(2), n.at(4))}); - std::unordered_set sub_nodes = {n.at(0), n.at(1), n.at(3)}; + std::set sub_nodes = {n.at(0), n.at(1), n.at(3)}; UndirectedGraphView view = view_subgraph(g, sub_nodes); SUBCASE("get_nodes") { - std::unordered_set expected = {n.at(0), n.at(1), n.at(3)}; + std::set expected = {n.at(0), n.at(1), n.at(3)}; - std::unordered_set result = get_nodes(view); + std::set result = get_nodes(view); CHECK(result == expected); } SUBCASE("get_edges") { - std::unordered_set expected = { + std::set expected = { make_undirected_edge(n.at(0), n.at(3)), make_undirected_edge(n.at(1), n.at(1)), make_undirected_edge(n.at(1), n.at(3)), }; - std::unordered_set result = get_edges(view); + std::set result = get_edges(view); CHECK(result == expected); } @@ -62,26 +62,26 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(2), n.at(3)}, DirectedEdge{n.at(3), n.at(2)}, DirectedEdge{n.at(2), n.at(4)}}); - std::unordered_set sub_nodes = {n.at(0), n.at(1), n.at(3)}; + std::set sub_nodes = {n.at(0), n.at(1), n.at(3)}; DiGraphView view = view_subgraph(g, sub_nodes); SUBCASE("get_nodes") { - std::unordered_set expected = {n.at(0), n.at(1), n.at(3)}; + std::set expected = {n.at(0), n.at(1), n.at(3)}; - std::unordered_set result = get_nodes(view); + std::set result = get_nodes(view); CHECK(result == expected); } SUBCASE("get_edges") { - std::unordered_set expected = { + std::set expected = { DirectedEdge{n.at(0), n.at(3)}, DirectedEdge{n.at(3), n.at(0)}, DirectedEdge{n.at(1), n.at(1)}, DirectedEdge{n.at(1), n.at(3)}, }; - std::unordered_set result = get_edges(view); + std::set result = get_edges(view); CHECK(result == expected); } @@ -99,20 +99,20 @@ TEST_SUITE(FF_TEST_SUITE) { UndirectedGraphView view = as_undirected(g); SUBCASE("get_nodes") { - std::unordered_set expected = unordered_set_of(n); + std::set expected = set_of(n); - std::unordered_set result = get_nodes(view); + std::set result = get_nodes(view); CHECK(result == expected); } SUBCASE("get_edges") { - std::unordered_set expected = { + std::set expected = { make_undirected_edge(n.at(0), n.at(1)), make_undirected_edge(n.at(1), n.at(2)), make_undirected_edge(n.at(2), n.at(0))}; - std::unordered_set result = get_edges(view); + std::set result = get_edges(view); CHECK(result == expected); } @@ -130,15 +130,15 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraphView view = as_digraph(g); SUBCASE("get_nodes") { - std::unordered_set expected = unordered_set_of(n); + std::set expected = set_of(n); - std::unordered_set result = get_nodes(view); + std::set result = get_nodes(view); CHECK(result == expected); } SUBCASE("get_edges") { - std::unordered_set expected = { + std::set expected = { DirectedEdge{n.at(0), n.at(0)}, DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(1), n.at(0)}, @@ -147,7 +147,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(2), n.at(0)}, DirectedEdge{n.at(0), n.at(2)}}; - std::unordered_set result = get_edges(view); + std::set result = get_edges(view); CHECK(result == expected); } diff --git a/lib/utils/test/src/utils/many_to_one/many_to_one.cc b/lib/utils/test/src/utils/many_to_one/many_to_one.cc index ce219676e1..1674a1c277 100644 --- a/lib/utils/test/src/utils/many_to_one/many_to_one.cc +++ b/lib/utils/test/src/utils/many_to_one/many_to_one.cc @@ -1,6 +1,6 @@ #include "utils/many_to_one/many_to_one.h" #include "test/utils/doctest/fmt/multiset.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/containers/multiset_of.h" #include @@ -39,17 +39,17 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("at_r") { - std::unordered_set result = m.at_r("two"); + nonempty_set result = m.at_r("two"); - std::unordered_set correct = {2, 20}; + nonempty_set correct = {2, 20}; CHECK(result == correct); } SUBCASE("left_values") { - std::unordered_set result = m.left_values(); + std::set result = m.left_values(); - std::unordered_set correct = { + std::set correct = { 1, 10, 100, @@ -61,9 +61,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("right_values") { - std::unordered_set result = m.right_values(); + std::set result = m.right_values(); - std::unordered_set correct = {"one", "two"}; + std::set correct = {"one", "two"}; CHECK(result == correct); } @@ -141,7 +141,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("many_to_one_from_unstructured_relation") { SUBCASE("relation is many-to-one") { - std::unordered_set> input = { + std::set> input = { {1, "odd"}, {2, "even"}, {3, "odd"}, @@ -158,7 +158,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("relation is one-to-one") { - std::unordered_set> input = { + std::set> input = { {1, "one"}, {2, "two"}, {3, "three"}, @@ -176,7 +176,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("relation is not many-to-one") { - std::unordered_set> input = { + std::set> input = { {1, "one"}, {1, "ODD"}, {2, "two"}, diff --git a/lib/utils/test/src/utils/nonnegative_int/nonnegative_int.cc b/lib/utils/test/src/utils/nonnegative_int/nonnegative_int.cc index 8c5ecd3e2c..cc4e1dee29 100644 --- a/lib/utils/test/src/utils/nonnegative_int/nonnegative_int.cc +++ b/lib/utils/test/src/utils/nonnegative_int/nonnegative_int.cc @@ -365,7 +365,7 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(hash_fn(nn_int_1a) != hash_fn(nn_int_2)); } SUBCASE("Unordered set works with nonnegative_int") { - std::unordered_set<::FlexFlow::nonnegative_int> nonnegative_int_set; + std::set<::FlexFlow::nonnegative_int> nonnegative_int_set; nonnegative_int_set.insert(nn_int_1a); nonnegative_int_set.insert(nn_int_1b); nonnegative_int_set.insert(nn_int_2); diff --git a/lib/utils/test/src/utils/one_to_many/one_to_many.cc b/lib/utils/test/src/utils/one_to_many/one_to_many.cc index de149ce609..f6ffabc698 100644 --- a/lib/utils/test/src/utils/one_to_many/one_to_many.cc +++ b/lib/utils/test/src/utils/one_to_many/one_to_many.cc @@ -1,11 +1,11 @@ #include "utils/one_to_many/one_to_many.h" #include "test/utils/doctest/fmt/multiset.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/set.h" #include "utils/containers/multiset_of.h" #include "utils/one_to_many/one_to_many_from_l_to_r_mapping.h" #include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; @@ -139,9 +139,9 @@ TEST_SUITE(FF_TEST_SUITE) { {2, {"two"}}, }; - std::unordered_set> result = + std::set> result = unstructured_relation_from_one_to_many(input); - std::unordered_set> correct = { + std::set> correct = { {1, "one"}, {1, "ONE"}, {2, "two"}, @@ -152,7 +152,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("one_to_many_from_unstructured_relation") { SUBCASE("relation is one-to-many") { - std::unordered_set> input = { + std::set> input = { {1, "one"}, {1, "ONE"}, {2, "two"}, @@ -169,7 +169,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("relation is one-to-one") { - std::unordered_set> input = { + std::set> input = { {1, "one"}, {2, "two"}, }; @@ -185,7 +185,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("relation is not one-to-many") { - std::unordered_set> input = { + std::set> input = { {1, "one"}, {1, "ONE"}, {2, "two"}, diff --git a/lib/utils/test/src/utils/orthotope/dim_coord.cc b/lib/utils/test/src/utils/orthotope/dim_coord.cc index cc2bfb63ea..0766f1067a 100644 --- a/lib/utils/test/src/utils/orthotope/dim_coord.cc +++ b/lib/utils/test/src/utils/orthotope/dim_coord.cc @@ -1,5 +1,5 @@ #include "utils/orthotope/dim_coord.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/orthotope/dim_ordering.h" #include @@ -14,7 +14,7 @@ TEST_SUITE(FF_TEST_SUITE) { }}; SUBCASE("lifted dims are a superset of coord dims") { - std::unordered_set lifted_dims = {1, 3, 6, 7}; + std::set lifted_dims = {1, 3, 6, 7}; DimCoord result = lift_dim_coord(coord, lifted_dims); DimCoord correct = DimCoord{{ @@ -28,7 +28,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("lifted dims are the same as coord dims") { - std::unordered_set lifted_dims = {1, 3, 7}; + std::set lifted_dims = {1, 3, 7}; DimCoord result = lift_dim_coord(coord, lifted_dims); DimCoord correct = coord; @@ -37,13 +37,13 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("lifted dims are a subset of coord dims") { - std::unordered_set lifted_dims = {1, 7}; + std::set lifted_dims = {1, 7}; CHECK_THROWS(lift_dim_coord(coord, lifted_dims)); } SUBCASE("lifted dims are overlapping with coord dims") { - std::unordered_set lifted_dims = {1, 2, 7}; + std::set lifted_dims = {1, 2, 7}; CHECK_THROWS(lift_dim_coord(coord, lifted_dims)); } @@ -104,10 +104,10 @@ TEST_SUITE(FF_TEST_SUITE) { {7, 2_p}, }}; - std::unordered_set> result = + std::set> result = get_coords_in_dim_domain(dim_domain); - std::unordered_set> correct = { + std::set> correct = { DimCoord{{ {7, 0_n}, }}, @@ -125,7 +125,7 @@ TEST_SUITE(FF_TEST_SUITE) { {2, 3_p}, }}; - std::unordered_set> result = + std::set> result = get_coords_in_dim_domain(dim_domain); auto mk_dim_coord = [](nonnegative_int dim7, nonnegative_int dim2) { @@ -135,7 +135,7 @@ TEST_SUITE(FF_TEST_SUITE) { }}; }; - std::unordered_set> correct = { + std::set> correct = { mk_dim_coord(0_n, 0_n), mk_dim_coord(0_n, 1_n), mk_dim_coord(0_n, 2_n), @@ -153,7 +153,7 @@ TEST_SUITE(FF_TEST_SUITE) { {2, 3_p}, }}; - std::unordered_set> result = + std::set> result = get_coords_in_dim_domain(dim_domain); auto mk_dim_coord = [](nonnegative_int dim7, nonnegative_int dim2) { @@ -163,7 +163,7 @@ TEST_SUITE(FF_TEST_SUITE) { }}; }; - std::unordered_set> correct = { + std::set> correct = { mk_dim_coord(0_n, 0_n), mk_dim_coord(0_n, 1_n), mk_dim_coord(0_n, 2_n), @@ -175,10 +175,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("zero-dimensional dim domain") { DimDomain dim_domain = DimDomain{{}}; - std::unordered_set> result = + std::set> result = get_coords_in_dim_domain(dim_domain); - std::unordered_set> correct = { + std::set> correct = { DimCoord{{}}, }; diff --git a/lib/utils/test/src/utils/orthotope/dim_domain.cc b/lib/utils/test/src/utils/orthotope/dim_domain.cc index 2bc10e6858..5969e9796b 100644 --- a/lib/utils/test/src/utils/orthotope/dim_domain.cc +++ b/lib/utils/test/src/utils/orthotope/dim_domain.cc @@ -1,5 +1,5 @@ #include "utils/orthotope/dim_domain.h" -#include "test/utils/doctest/fmt/unordered_set.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/orthotope/dim_ordering.h" #include @@ -13,8 +13,8 @@ TEST_SUITE(FF_TEST_SUITE) { {1, 3_p}, }}; - std::unordered_set result = get_domain_dims(domain); - std::unordered_set correct = { + std::set result = get_domain_dims(domain); + std::set correct = { 3, 7, 1, @@ -31,7 +31,7 @@ TEST_SUITE(FF_TEST_SUITE) { }}; SUBCASE("allowed is a strict subset of the dims") { - std::unordered_set allowed = { + std::set allowed = { 3, 1, }; @@ -46,7 +46,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("allowed is the same as dims") { - std::unordered_set allowed = { + std::set allowed = { 3, 7, 1, @@ -59,7 +59,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("allowed is empty") { - std::unordered_set allowed = {}; + std::set allowed = {}; DimDomain result = restrict_domain_to_dims(domain, allowed); DimDomain correct = DimDomain{{}}; @@ -68,7 +68,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("allowed is mutually exclusive with dims") { - std::unordered_set allowed = { + std::set allowed = { 6, 8, }; @@ -80,7 +80,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("allowed is overlapping with dims") { - std::unordered_set allowed = { + std::set allowed = { 6, 8, 7, @@ -95,7 +95,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("allowed is a superset of dims") { - std::unordered_set allowed = { + std::set allowed = { 6, 8, 7, @@ -135,7 +135,7 @@ TEST_SUITE(FF_TEST_SUITE) { 2_p, }}; - std::unordered_set dims = {3, 7, 1}; + std::set dims = {3, 7, 1}; DimDomain result = dim_domain_from_orthotope( orthotope, dims, make_default_dim_ordering()); diff --git a/lib/utils/test/src/utils/orthotope/dim_projection.cc b/lib/utils/test/src/utils/orthotope/dim_projection.cc index de37af61da..c5a435b2b6 100644 --- a/lib/utils/test/src/utils/orthotope/dim_projection.cc +++ b/lib/utils/test/src/utils/orthotope/dim_projection.cc @@ -67,14 +67,14 @@ TEST_SUITE(FF_TEST_SUITE) { project_dims(projection, /*onto=*/0, /*from=*/ - std::unordered_set{ + std::set{ shard1_idx, discard_copy_idx, }); project_dims(projection, /*onto=*/1, /*from=*/ - std::unordered_set{ + std::set{ shard0_idx, sum_idx, }); @@ -103,13 +103,13 @@ TEST_SUITE(FF_TEST_SUITE) { project_dims(projection, /*onto=*/0, /*from=*/ - std::unordered_set{ + std::set{ shard1_idx, }); project_dims(projection, /*onto=*/1, /*from=*/ - std::unordered_set{ + std::set{ shard0_idx, sum_idx, }); @@ -138,14 +138,14 @@ TEST_SUITE(FF_TEST_SUITE) { project_dims(projection, /*onto=*/0, /*from=*/ - std::unordered_set{ + std::set{ shard1_idx, discard_copy_idx, }); project_dims(projection, /*onto=*/1, /*from=*/ - std::unordered_set{ + std::set{ sum_idx, }); diff --git a/lib/utils/test/src/utils/positive_int/positive_int.cc b/lib/utils/test/src/utils/positive_int/positive_int.cc index 88454bbfbd..1b19acb052 100644 --- a/lib/utils/test/src/utils/positive_int/positive_int.cc +++ b/lib/utils/test/src/utils/positive_int/positive_int.cc @@ -469,8 +469,8 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(hash_fn(nn_int_1a) != hash_fn(nn_int_2)); } - SUBCASE("unordered_set works with positive_int") { - std::unordered_set<::FlexFlow::positive_int> positive_int_set; + SUBCASE("set works with positive_int") { + std::set<::FlexFlow::positive_int> positive_int_set; positive_int_set.insert(nn_int_1a); positive_int_set.insert(nn_int_1b); positive_int_set.insert(nn_int_2); From e4fbc04644c0bb38522b1abe8f5161f097d9e6e7 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 12 Jun 2026 20:27:38 -0700 Subject: [PATCH 26/35] Address PR comments --- flake.lock | 27 ++- .../computation_graph_instance.h | 12 +- .../cost_estimator/local_cost_estimator.h | 8 +- .../local_task_argument_accessor.h | 8 +- .../per_device_op_state_initialization.h | 6 +- .../include/local-execution/task_execution.h | 4 +- .../computation_graph_instance.cc | 12 +- .../cost_estimator/local_cost_estimator.cc | 4 +- .../local_task_argument_accessor.cc | 4 +- .../per_device_op_state_initialization.cc | 4 +- .../src/local-execution/task_execution.cc | 4 +- .../computation_graph_instance.cc | 10 +- .../cost_estimator/local_cost_estimator.cc | 4 +- .../local_task_argument_accessor.cc | 2 +- ..._space_to_parallel_tensor_space_mapping.cc | 21 -- .../include/pcg/device_in_node_idx_t.dtg.toml | 22 ++ lib/pcg/include/pcg/node_idx_t.dtg.toml | 21 ++ .../v1_mapped_parallel_computation_graph.cc | 1 - .../include/realm-execution/address_space.h | 14 ++ ...ce_specific_managed_per_device_ff_handle.h | 6 +- .../realm-execution/device_specific_ptr.h | 12 +- ...buted_per_device_op_state_initialization.h | 3 +- .../realm-execution/fmt/realm_processor.h | 34 +++ .../fmt/realm_processor_kind.h | 83 ++++++++ .../realm-execution/instance_allocation.h | 5 +- .../include/realm-execution/pcg_instance.h | 2 +- .../include/realm-execution/processor_query.h | 12 ++ .../include/realm-execution/realm_context.h | 20 +- .../include/realm-execution/realm_manager.h | 2 +- .../serializable_device_specific_ptr.dtg.toml | 4 +- .../src/realm-execution/address_space.cc | 17 ++ ...e_specific_managed_per_device_ff_handle.cc | 4 +- .../realm-execution/distributed_ff_handle.cc | 2 +- ...uted_per_device_op_state_initialization.cc | 9 +- .../realm-execution/fmt/realm_processor.cc | 9 + .../fmt/realm_processor_kind.cc | 9 + .../realm-execution/instance_allocation.cc | 9 +- .../src/realm-execution/pcg_instance.cc | 10 +- .../src/realm-execution/processor_kind.cc | 3 +- .../src/realm-execution/processor_query.cc | 13 ++ .../src/realm-execution/realm_context.cc | 193 ++++++++++++------ .../tasks/impl/ff_handle_init_task.cc | 2 +- .../src/realm-execution/tasks/impl/op_task.cc | 6 +- .../impl/per_device_op_state_init_task.cc | 8 +- .../fmt/realm_processor_kind.cc | 13 ++ .../test/src/realm-execution/test_e2e.cc | 2 - .../output_expr_to_result_sub_pcg_mapping.cc | 4 +- .../src/substitutions/pcg_pattern_match.cc | 4 +- .../include/task-spec/device_specific.h | 10 +- .../dynamic_graph/dynamic_node_attrs.dtg.toml | 6 +- .../dynamic_graph/dynamic_node_mapping.h | 6 +- .../task-spec/dynamic_graph/machine_slicing.h | 4 +- .../parallel_tensor_mapping.dtg.toml | 4 +- .../serializable_dynamic_node_attrs.dtg.toml | 6 +- ...t.dtg.toml => global_device_id_t.dtg.toml} | 2 +- .../include/task-spec/global_device_id_t.h | 15 ++ .../task-spec/local_device_id_t.dtg.toml | 23 +++ .../include/task-spec/per_device_op_state.h | 2 +- .../itask_argument_accessor.h | 4 +- .../task_argument_accessor.h | 2 +- .../task-spec/dynamic_graph/copy_insertion.cc | 6 +- .../dynamic_graph/dynamic_node_mapping.cc | 12 +- .../dynamic_open_dataflow_graph.cc | 9 + .../dynamic_graph/machine_slicing.cc | 10 +- .../serializable_dynamic_node_attrs.cc | 4 +- .../dynamic_graph/shard_expansion.cc | 18 +- .../src/task-spec/global_device_id_t.cc | 25 +++ .../src/task-spec/per_device_op_state.cc | 2 +- .../test/src/task-spec/device_specific.cc | 4 +- .../task-spec/dynamic_graph/copy_insertion.cc | 18 +- .../dynamic_graph/machine_slicing.cc | 10 +- .../dynamic_graph/shard_expansion.cc | 26 +-- .../binary_merge_disjoint_bidicts.h | 38 ++++ .../algorithms/merge_disjoint_bidicts.h | 33 +-- .../binary_merge_disjoint_bidicts.cc | 13 ++ .../algorithms/merge_disjoint_bidicts.cc | 10 + ...ts.cc => binary_merge_disjoint_bidicts.cc} | 12 +- 77 files changed, 724 insertions(+), 298 deletions(-) create mode 100644 lib/pcg/include/pcg/device_in_node_idx_t.dtg.toml create mode 100644 lib/pcg/include/pcg/node_idx_t.dtg.toml create mode 100644 lib/realm-execution/include/realm-execution/address_space.h create mode 100644 lib/realm-execution/include/realm-execution/fmt/realm_processor.h create mode 100644 lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h create mode 100644 lib/realm-execution/include/realm-execution/processor_query.h create mode 100644 lib/realm-execution/src/realm-execution/address_space.cc create mode 100644 lib/realm-execution/src/realm-execution/fmt/realm_processor.cc create mode 100644 lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc create mode 100644 lib/realm-execution/src/realm-execution/processor_query.cc create mode 100644 lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc rename lib/task-spec/include/task-spec/{device_id_t.dtg.toml => global_device_id_t.dtg.toml} (91%) create mode 100644 lib/task-spec/include/task-spec/global_device_id_t.h create mode 100644 lib/task-spec/include/task-spec/local_device_id_t.dtg.toml create mode 100644 lib/task-spec/src/task-spec/global_device_id_t.cc create mode 100644 lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h create mode 100644 lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc rename lib/utils/test/src/utils/bidict/algorithms/{merge_disjoint_bidicts.cc => binary_merge_disjoint_bidicts.cc} (72%) diff --git a/flake.lock b/flake.lock index e52833e4ed..71bb407e27 100644 --- a/flake.lock +++ b/flake.lock @@ -63,14 +63,15 @@ ], "nixpkgs": [ "nixpkgs" - ] + ], + "python38-nixpkgs": "python38-nixpkgs" }, "locked": { - "lastModified": 1778104328, - "narHash": "sha256-bn0G8xDqBrVjp5htw1i3u8fPdPMVtoZXFzX7hJ6m9YY=", + "lastModified": 1780709897, + "narHash": "sha256-55VJWpNnt/tMEP8tadRQSI66cX1BgJS9UNy/EkaQt3A=", "ref": "refs/heads/master", - "rev": "535ac756b2674dc10051b37be978b9d2cb9f817d", - "revCount": 156, + "rev": "18d04112d94ee0a1c173a30923aac9fa2b505539", + "revCount": 160, "type": "git", "url": "https://git.sr.ht/~lockshaw/proj" }, @@ -79,6 +80,22 @@ "url": "https://git.sr.ht/~lockshaw/proj" } }, + "python38-nixpkgs": { + "locked": { + "lastModified": 1645131486, + "narHash": "sha256-AuWJe0TiqHD6L7Vzcar4QOblMGubKHsuVs3YFMCRYl0=", + "owner": "nixos", + "repo": "nixpkgs", + "rev": "7592790b9e02f7f99ddcb1bd33fd44ff8df6a9a7", + "type": "github" + }, + "original": { + "owner": "nixos", + "repo": "nixpkgs", + "rev": "7592790b9e02f7f99ddcb1bd33fd44ff8df6a9a7", + "type": "github" + } + }, "root": { "inputs": { "flake-utils": "flake-utils", diff --git a/lib/local-execution/include/local-execution/computation_graph_instance.h b/lib/local-execution/include/local-execution/computation_graph_instance.h index aa05cbd582..e6c7ded9c1 100644 --- a/lib/local-execution/include/local-execution/computation_graph_instance.h +++ b/lib/local-execution/include/local-execution/computation_graph_instance.h @@ -7,7 +7,7 @@ #include "kernels/profiling_settings.dtg.h" #include "local-execution/loss_config.dtg.h" #include "pcg/computation_graph.dtg.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" @@ -49,31 +49,31 @@ ComputationGraphInstance create_computation_graph_instance( Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, - device_id_t device_idx); + global_device_id_t global_device_id); std::unordered_map> perform_all_passes_for_computation_graph_instance( ComputationGraphInstance &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx); + global_device_id_t global_device_id); std::unordered_map> perform_forward_pass_for_computation_graph_instance( ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx); + global_device_id_t global_device_id); std::unordered_map> perform_backward_pass_for_computation_graph_instance( ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx); + global_device_id_t global_device_id); void perform_update_pass_for_computation_graph_instance( ComputationGraphInstance &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx); + global_device_id_t global_device_id); } // namespace FlexFlow diff --git a/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h b/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h index a978f3c996..15aa267eec 100644 --- a/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h +++ b/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h @@ -5,7 +5,7 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" #include "kernels/profiling_settings.dtg.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "pcg/machine_interconnect_specification.dtg.h" namespace FlexFlow { @@ -16,7 +16,7 @@ struct LocalCostEstimator : public ICostEstimator { Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, - device_id_t device_idx); + global_device_id_t device_idx); LocalCostEstimator(LocalCostEstimator const &) = delete; LocalCostEstimator(LocalCostEstimator &&) = delete; @@ -31,7 +31,7 @@ struct LocalCostEstimator : public ICostEstimator { Allocator allocator; ProfilingSettings profiling_settings; device_handle_t device_handle; - device_id_t device_idx; + global_device_id_t device_idx; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(LocalCostEstimator); @@ -40,7 +40,7 @@ CostEstimator get_local_cost_estimator( Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, - device_id_t device_idx); + global_device_id_t device_idx); } // namespace FlexFlow diff --git a/lib/local-execution/include/local-execution/local_task_argument_accessor.h b/lib/local-execution/include/local-execution/local_task_argument_accessor.h index a7df8814d0..38b1023cd2 100644 --- a/lib/local-execution/include/local-execution/local_task_argument_accessor.h +++ b/lib/local-execution/include/local-execution/local_task_argument_accessor.h @@ -2,7 +2,7 @@ #define _FLEXFLOW_LIB_LOCAL_EXECUTION_INCLUDE_LOCAL_EXECUTION_LOCAL_TASK_ARGUMENT_ACCESSOR_H #include "kernels/accessor.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" #include "task-spec/task_argument_accessor/itask_argument_accessor.h" #include "task-spec/task_argument_accessor/task_tensor_parameter.dtg.h" @@ -21,7 +21,7 @@ struct LocalTaskArgumentAccessor : public ITaskArgumentAccessor { std::optional const &loss_attrs, std::optional const &per_device_op_state, std::optional const &optimizer_attrs, - device_id_t device_idx); + global_device_id_t device_idx); LocalTaskArgumentAccessor(LocalTaskArgumentAccessor const &) = delete; LocalTaskArgumentAccessor(LocalTaskArgumentAccessor &&) = delete; @@ -41,7 +41,7 @@ struct LocalTaskArgumentAccessor : public ITaskArgumentAccessor { Allocator get_allocator() const override; - device_id_t get_device_idx() const override; + global_device_id_t get_device_idx() const override; private: Allocator allocator; @@ -56,7 +56,7 @@ struct LocalTaskArgumentAccessor : public ITaskArgumentAccessor { std::optional per_device_op_state; std::optional optimizer_attrs; - device_id_t device_idx; + global_device_id_t device_idx; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(LocalTaskArgumentAccessor); diff --git a/lib/local-execution/include/local-execution/per_device_op_state_initialization.h b/lib/local-execution/include/local-execution/per_device_op_state_initialization.h index 174e51f307..90a1ffc688 100644 --- a/lib/local-execution/include/local-execution/per_device_op_state_initialization.h +++ b/lib/local-execution/include/local-execution/per_device_op_state_initialization.h @@ -4,7 +4,7 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" #include "kernels/profiling_settings.dtg.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" @@ -18,7 +18,7 @@ DynamicNodeInvocation ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, OptimizerAttrs const &optimizer_attrs, - device_id_t device_idx); + global_device_id_t device_idx); /** * @brief Initialize all operators and save the per-device op state @@ -29,7 +29,7 @@ DynamicOpenDataflowGraph perform_per_device_op_state_initialization( ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, OptimizerAttrs const &optimizer_attrs, - device_id_t device_idx); + global_device_id_t device_idx); } // namespace FlexFlow diff --git a/lib/local-execution/include/local-execution/task_execution.h b/lib/local-execution/include/local-execution/task_execution.h index 3bb0c6b92b..83f69a96e3 100644 --- a/lib/local-execution/include/local-execution/task_execution.h +++ b/lib/local-execution/include/local-execution/task_execution.h @@ -16,7 +16,7 @@ TaskArgumentAccessor make_task_argument_accessor_for_invocation( device_handle_t const &ff_handle, std::optional const &per_device_op_state, std::optional const &optimizer_attrs, - device_id_t device_idx); + global_device_id_t device_idx); std::optional execute_dynamic_node_invocation( DynamicNodeInvocation const &invocation, @@ -25,7 +25,7 @@ std::optional execute_dynamic_node_invocation( device_handle_t const &ff_handle, std::optional const &per_device_op_state, std::optional const &optimizer_attrs, - device_id_t device_idx); + global_device_id_t device_idx); } // namespace FlexFlow diff --git a/lib/local-execution/src/local-execution/computation_graph_instance.cc b/lib/local-execution/src/local-execution/computation_graph_instance.cc index d2781472b0..2bc15d1187 100644 --- a/lib/local-execution/src/local-execution/computation_graph_instance.cc +++ b/lib/local-execution/src/local-execution/computation_graph_instance.cc @@ -67,7 +67,7 @@ ComputationGraphInstance create_computation_graph_instance( Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, - device_id_t device_idx) { + global_device_id_t device_idx) { DynamicOpenDataflowGraph dg = make_dynamic_open_dataflow_graph_from_cg(cg); dg = perform_pass_expansion(dg); @@ -116,7 +116,7 @@ static std::unordered_map> OptimizerAttrs const &optimizer_attrs, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx) { + global_device_id_t device_idx) { return unordered_map_from_pairs( transform(invocations, [&](DynamicNodeInvocation const &invocation) { std::optional timing = execute_dynamic_node_invocation( @@ -141,7 +141,7 @@ std::unordered_map> ComputationGraphInstance &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx) { + global_device_id_t device_idx) { std::vector execution_order = instance.get_execution_order(); std::unordered_map> @@ -161,7 +161,7 @@ std::unordered_map> ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx) { + global_device_id_t device_idx) { std::vector execution_order = filter(instance.get_execution_order(), [](DynamicNodeInvocation const &invocation) { @@ -184,7 +184,7 @@ std::unordered_map> ComputationGraphInstance const &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx) { + global_device_id_t device_idx) { std::vector execution_order = filter(instance.get_execution_order(), [](DynamicNodeInvocation const &invocation) { @@ -206,7 +206,7 @@ void perform_update_pass_for_computation_graph_instance( ComputationGraphInstance &instance, ProfilingSettings const &profiling_settings, device_handle_t const &ff_handle, - device_id_t device_idx) { + global_device_id_t device_idx) { std::vector execution_order = filter(instance.get_execution_order(), [](DynamicNodeInvocation const &invocation) { diff --git a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc index 27a1c39677..f151ab1d17 100644 --- a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -31,7 +31,7 @@ LocalCostEstimator::LocalCostEstimator( Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, - device_id_t device_idx) + global_device_id_t device_idx) : interconnect_specification(interconnect_specification), allocator(allocator), profiling_settings(profiling_settings), device_handle(device_handle), device_idx(device_idx) {} @@ -181,7 +181,7 @@ CostEstimator get_local_cost_estimator( Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, - device_id_t device_idx) { + global_device_id_t device_idx) { return CostEstimator::create(interconnect_specification, allocator, profiling_settings, diff --git a/lib/local-execution/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/src/local-execution/local_task_argument_accessor.cc index 5fbe207e6c..a5f553db0b 100644 --- a/lib/local-execution/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/src/local-execution/local_task_argument_accessor.cc @@ -16,7 +16,7 @@ LocalTaskArgumentAccessor::LocalTaskArgumentAccessor( std::optional const &loss_attrs, std::optional const &per_device_op_state, std::optional const &optimizer_attrs, - device_id_t device_idx) + global_device_id_t device_idx) : allocator(allocator), tensor_slots_backing(tensor_slots_backing), profiling_settings(profiling_settings), ff_handle(ff_handle), op_attrs(op_attrs), loss_attrs(loss_attrs), @@ -105,7 +105,7 @@ Allocator LocalTaskArgumentAccessor::get_allocator() const { return this->allocator; } -device_id_t LocalTaskArgumentAccessor::get_device_idx() const { +global_device_id_t LocalTaskArgumentAccessor::get_device_idx() const { return this->device_idx; } diff --git a/lib/local-execution/src/local-execution/per_device_op_state_initialization.cc b/lib/local-execution/src/local-execution/per_device_op_state_initialization.cc index bf72843daf..571e3dd1c1 100644 --- a/lib/local-execution/src/local-execution/per_device_op_state_initialization.cc +++ b/lib/local-execution/src/local-execution/per_device_op_state_initialization.cc @@ -24,7 +24,7 @@ DynamicNodeInvocation ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, OptimizerAttrs const &optimizer_attrs, - device_id_t device_idx) { + global_device_id_t device_idx) { if (!i.node_attrs.op_attrs.has_value() || !i.node_attrs.op_attrs.value().is_pcg_op()) { return i; @@ -61,7 +61,7 @@ DynamicOpenDataflowGraph perform_per_device_op_state_initialization( ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, OptimizerAttrs const &optimizer_attrs, - device_id_t device_idx) { + global_device_id_t device_idx) { ASSERT(no_nodes_are_initialized(dg)); DynamicOpenDataflowGraph result = transform_dynamic_invocation_set( diff --git a/lib/local-execution/src/local-execution/task_execution.cc b/lib/local-execution/src/local-execution/task_execution.cc index 0b1bc3513a..48b9b89d94 100644 --- a/lib/local-execution/src/local-execution/task_execution.cc +++ b/lib/local-execution/src/local-execution/task_execution.cc @@ -49,7 +49,7 @@ TaskArgumentAccessor make_task_argument_accessor_for_invocation( device_handle_t const &ff_handle, std::optional const &per_device_op_state, std::optional const &optimizer_attrs, - device_id_t device_idx) { + global_device_id_t device_idx) { auto make_param = [&](DynamicTensorSlot const &slot) { return make_task_tensor_parameter_from_dynamic_slot(slot, optimizer_attrs); }; @@ -88,7 +88,7 @@ std::optional execute_dynamic_node_invocation( device_handle_t const &ff_handle, std::optional const &per_device_op_state, std::optional const &optimizer_attrs, - device_id_t device_idx) { + global_device_id_t device_idx) { TaskArgumentAccessor arg_accessor = make_task_argument_accessor_for_invocation( /*invocation=*/invocation, diff --git a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc index a4049a609f..69091b5d2e 100644 --- a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc +++ b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc @@ -139,7 +139,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*nesterov=*/false, /*weight_decay=*/0.001}}; device_handle_t ff_handle = cpu_make_device_handle_t(); - device_id_t device_idx = device_id_t{ + global_device_id_t global_device_id = global_device_id_t{ /*coord=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, @@ -164,7 +164,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*allocator=*/allocator, /*profiling_settings=*/ProfilingSettings{0, 0}, /*device_handle=*/ff_handle, - /*device_idx=*/device_idx); + /*global_device_id=*/global_device_id); // begin training loop int num_epochs = 5; @@ -175,7 +175,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*instance=*/computation_graph_instance, /*profiling_settings=*/ProfilingSettings{0, 0}, /*ff_handle=*/ff_handle, - /*device_idx=*/device_idx); + /*global_device_id=*/global_device_id); loss_values.push_back(copy_tensor_accessor_r( computation_graph_instance.get_loss_tensor_accessor().value(), allocator)); @@ -312,7 +312,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { /*weight_decay=*/0.001, }, }; - device_id_t device_idx = device_id_t{ + global_device_id_t device_idx = global_device_id_t{ /*coord=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, @@ -433,7 +433,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { }, }; - device_id_t device_idx = device_id_t{ + global_device_id_t device_idx = global_device_id_t{ /*coord=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, diff --git a/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc index 34cfbb4e3d..8d1506b5fc 100644 --- a/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -18,7 +18,7 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("LocalCostEstimator") { Allocator allocator = create_local_cpu_memory_allocator(); device_handle_t ff_handle = cpu_make_device_handle_t(); - device_id_t device_idx = device_id_t{ + global_device_id_t device_idx = global_device_id_t{ /*coord=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, @@ -93,7 +93,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { Allocator allocator = create_local_cuda_memory_allocator(); - device_id_t device_idx = device_id_t{ + global_device_id_t device_idx = global_device_id_t{ /*coord=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, diff --git a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc index 9594e13336..dba4464611 100644 --- a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc @@ -51,7 +51,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - device_id_t device_idx = device_id_t{ + global_device_id_t device_idx = global_device_id_t{ /*coord=*/MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc index 2a161838cd..b502c456e6 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.cc @@ -11,27 +11,6 @@ ParallelTensorSpaceToParallelTensorSpaceMapping ParallelTensorDimDegrees const &l_degrees, ParallelTensorDimDegrees const &r_degrees) { - // TODO(@lockshaw)(#pr): - // { - // std::unordered_set - // l_dims = - // unordered_set_of(get_nontrivial_parallel_tensor_dim_indices(l_degrees)); - // std::unordered_set - // projection_input_dims = input_dims_of_projection(projection); - // - // ASSERT(l_dims == projection_input_dims); - // } - // - // { - // std::unordered_set - // r_dims = - // unordered_set_of(get_nontrivial_parallel_tensor_dim_indices(r_degrees)); - // std::unordered_set - // projection_output_dims = output_dims_of_projection(projection); - // - // ASSERT(r_dims == projection_output_dims); - // } - return ParallelTensorSpaceToParallelTensorSpaceMapping{ dim_domain_mapping_from_projection( /*projection=*/projection, diff --git a/lib/pcg/include/pcg/device_in_node_idx_t.dtg.toml b/lib/pcg/include/pcg/device_in_node_idx_t.dtg.toml new file mode 100644 index 0000000000..c7b73e7724 --- /dev/null +++ b/lib/pcg/include/pcg/device_in_node_idx_t.dtg.toml @@ -0,0 +1,22 @@ +namespace = "FlexFlow" +name = "device_in_node_idx_t" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", + "rapidcheck", +] + +includes = [ + "utils/nonnegative_int/nonnegative_int.h", +] + +src_includes = [] + +[[fields]] +name = "raw" +type = "::FlexFlow::nonnegative_int" + diff --git a/lib/pcg/include/pcg/node_idx_t.dtg.toml b/lib/pcg/include/pcg/node_idx_t.dtg.toml new file mode 100644 index 0000000000..df592635b0 --- /dev/null +++ b/lib/pcg/include/pcg/node_idx_t.dtg.toml @@ -0,0 +1,21 @@ +namespace = "FlexFlow" +name = "node_idx_t" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", + "rapidcheck", +] + +includes = [ + "utils/nonnegative_int/nonnegative_int.h", +] + +src_includes = [] + +[[fields]] +name = "raw" +type = "::FlexFlow::nonnegative_int" diff --git a/lib/pcg/test/src/pcg/file_format/v1/v1_mapped_parallel_computation_graph.cc b/lib/pcg/test/src/pcg/file_format/v1/v1_mapped_parallel_computation_graph.cc index 0c43ad06a0..ac264191a8 100644 --- a/lib/pcg/test/src/pcg/file_format/v1/v1_mapped_parallel_computation_graph.cc +++ b/lib/pcg/test/src/pcg/file_format/v1/v1_mapped_parallel_computation_graph.cc @@ -36,7 +36,6 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate coord = MachineSpaceCoordinate{ /*node_idx=*/0_n, /*device_idx=*/0_n, - /*device_type=*/DeviceType::GPU, }; OperatorAtomicTaskShardBinding binding = OperatorAtomicTaskShardBinding{ diff --git a/lib/realm-execution/include/realm-execution/address_space.h b/lib/realm-execution/include/realm-execution/address_space.h new file mode 100644 index 0000000000..3d00bc2bd4 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/address_space.h @@ -0,0 +1,14 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_ADDRESS_SPACE_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_ADDRESS_SPACE_H + +#include "pcg/node_idx_t.dtg.h" +#include "realm-execution/realm.h" + +namespace FlexFlow { + +node_idx_t node_idx_from_realm_address_space(Realm::AddressSpace); +Realm::AddressSpace realm_address_space_from_node_idx(node_idx_t); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h b/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h index e5ab8b597e..a7e265f777 100644 --- a/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h +++ b/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h @@ -4,17 +4,17 @@ #include "kernels/device_handle_t.dtg.h" #include "kernels/managed_per_device_ff_handle.h" #include "realm-execution/device_specific_ptr.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include namespace FlexFlow { DeviceSpecificPtr make_device_specific_managed_ff_handle( - device_id_t const &, std::optional const &); + global_device_id_t const &, std::optional const &); device_handle_t device_handle_t_from_device_specific_managed_ff_handle( - DeviceSpecificPtr const &, device_id_t); + DeviceSpecificPtr const &, global_device_id_t); } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/device_specific_ptr.h b/lib/realm-execution/include/realm-execution/device_specific_ptr.h index e8de4cfdad..f1126f64c1 100644 --- a/lib/realm-execution/include/realm-execution/device_specific_ptr.h +++ b/lib/realm-execution/include/realm-execution/device_specific_ptr.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_DEVICE_SPECIFIC_PTR_H #define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_DEVICE_SPECIFIC_PTR_H -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include #include @@ -18,7 +18,7 @@ namespace FlexFlow { * transfer the pointers back-and-forth between workers and the controller * task. To prevent accidentally accessing one of these pointers on the wrong * device (as the pointer is only valid in the memory where it was created), we - * wrap them with \ref DeviceSpecificPtr, which holds the \ref device_id_t + * wrap them with \ref DeviceSpecificPtr, which holds the \ref global_device_id_t * where the pointer was created, and any attempt to interact with the raw * pointer value (i.e., \ref DeviceSpecificPtr::get) checks that the current * device matches the original device, and throws a readable error message if @@ -32,15 +32,15 @@ template struct DeviceSpecificPtr { public: DeviceSpecificPtr() = delete; - explicit DeviceSpecificPtr(device_id_t device_idx, std::optional ptr) + explicit DeviceSpecificPtr(global_device_id_t device_idx, std::optional ptr) : device_idx(device_idx), ptr(ptr) {} - std::optional get(device_id_t device_idx) const { + std::optional get(global_device_id_t device_idx) const { ASSERT(this->device_idx == device_idx); return this->ptr; } - device_id_t get_device_idx() const { + global_device_id_t get_device_idx() const { return this->device_idx; } @@ -49,7 +49,7 @@ struct DeviceSpecificPtr { } private: - device_id_t device_idx; + global_device_id_t device_idx; std::optional ptr; }; diff --git a/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h b/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h index f160146b96..5d52f8caaf 100644 --- a/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h +++ b/lib/realm-execution/include/realm-execution/distributed_per_device_op_state_initialization.h @@ -25,8 +25,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( ProfilingSettings const &profiling_settings, DistributedFfHandle const &device_handle, OptimizerAttrs const &optimizer_attrs, - Realm::Event precondition, - DeviceType device_type); + Realm::Event precondition); } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/fmt/realm_processor.h b/lib/realm-execution/include/realm-execution/fmt/realm_processor.h new file mode 100644 index 0000000000..e2fadb6b18 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/fmt/realm_processor.h @@ -0,0 +1,34 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_FMT_REALM_PROCESSOR_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_FMT_REALM_PROCESSOR_H + +#include "realm-execution/realm.h" +#include +#include + +namespace fmt { + +template +struct formatter< + ::FlexFlow::Realm::Processor, + Char, + std::enable_if_t::value>> + : formatter<::std::string> { + template + auto format(::FlexFlow::Realm::Processor const &m, FormatContext &ctx) + -> decltype(ctx.out()) { + + std::string result = fmt::format("", m.id); + + return formatter::format(result, ctx); + } +}; + +} // namespace fmt + +namespace FlexFlow { + +std::ostream &operator<<(std::ostream &s, ::FlexFlow::Realm::Processor const &m); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h b/lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h new file mode 100644 index 0000000000..d75ab067cc --- /dev/null +++ b/lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h @@ -0,0 +1,83 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_FMT_REALM_PROCESSOR_KIND_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_FMT_REALM_PROCESSOR_KIND_H + +#include "realm-execution/realm.h" +#include +#include +#include + +namespace fmt { + +template +struct formatter< + ::FlexFlow::Realm::Processor::Kind, + Char + > + : formatter<::std::string> { + template + auto format(::FlexFlow::Realm::Processor::Kind const &m, FormatContext &ctx) + -> decltype(ctx.out()) { + + std::string result; + switch (m) { + case ::FlexFlow::Realm::Processor::Kind::NO_KIND: { + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::TOC_PROC: { + // Throughput core + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::LOC_PROC: { + // Latency core + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::UTIL_PROC: { + // Utility core + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::IO_PROC: { + // I/O core + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::PROC_GROUP: { + // Processor group + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::PROC_SET: { + // Set of Processors for OpenMP/Kokkos etc. + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::OMP_PROC: { + // OpenMP (or similar) thread pool + result = ""; + break; + } + case ::FlexFlow::Realm::Processor::Kind::PY_PROC: { + // Python interpreter + result = ""; + break; + } + default: + PANIC("Unknown Processor::Kind {}", static_cast(m)); + }; + + return formatter::format(result, ctx); + } +}; + +} // namespace fmt + +namespace FlexFlow { + +std::ostream &operator<<(std::ostream &, ::FlexFlow::Realm::Processor::Kind const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/instance_allocation.h b/lib/realm-execution/include/realm-execution/instance_allocation.h index 1f2a6d9134..cc4af624a6 100644 --- a/lib/realm-execution/include/realm-execution/instance_allocation.h +++ b/lib/realm-execution/include/realm-execution/instance_allocation.h @@ -12,7 +12,7 @@ namespace FlexFlow { * on the device represented by \p device_coord. */ std::pair - perform_instance_allocation_for_value(device_id_t const &device_id, + perform_instance_allocation_for_value(global_device_id_t const &device_id, DynamicValueAttrs const &value, RealmContext &ctx); @@ -27,8 +27,7 @@ TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, std::unordered_map const &preallocated, - RealmContext &ctx, - DeviceType device_type); + RealmContext &ctx); /** * @brief Destroys all of the instances held in \p instances. diff --git a/lib/realm-execution/include/realm-execution/pcg_instance.h b/lib/realm-execution/include/realm-execution/pcg_instance.h index 1e17856999..6c82e4c61c 100644 --- a/lib/realm-execution/include/realm-execution/pcg_instance.h +++ b/lib/realm-execution/include/realm-execution/pcg_instance.h @@ -11,7 +11,7 @@ #include "realm-execution/per_device_op_state_backing.dtg.h" #include "realm-execution/realm_context.h" #include "realm-execution/tensor_instance_backing.dtg.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" diff --git a/lib/realm-execution/include/realm-execution/processor_query.h b/lib/realm-execution/include/realm-execution/processor_query.h new file mode 100644 index 0000000000..693d488390 --- /dev/null +++ b/lib/realm-execution/include/realm-execution/processor_query.h @@ -0,0 +1,12 @@ +#ifndef _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_PROCESSOR_QUERY_H +#define _FLEXFLOW_LIB_REALM_EXECUTION_INCLUDE_REALM_EXECUTION_PROCESSOR_QUERY_H + +#include "realm-execution/realm.h" + +namespace FlexFlow { + +std::set processor_set_from_query(Realm::Machine::ProcessorQuery const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/realm-execution/include/realm-execution/realm_context.h b/lib/realm-execution/include/realm-execution/realm_context.h index 0f2bed7de4..e06ac522d7 100644 --- a/lib/realm-execution/include/realm-execution/realm_context.h +++ b/lib/realm-execution/include/realm-execution/realm_context.h @@ -9,9 +9,10 @@ #include "pcg/machine_space_coordinate.dtg.h" #include "realm-execution/realm.h" #include "realm-execution/tasks/task_id_t.dtg.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include #include +#include "task-spec/local_device_id_t.dtg.h" namespace FlexFlow { @@ -32,8 +33,11 @@ struct RealmContext { /** \name Device mapping */ ///\{ - Realm::Processor map_device_coord_to_processor(device_id_t const &) const; - device_id_t map_processor_to_device_coord(Realm::Processor) const; + Realm::Processor processor_from_global_device_id(global_device_id_t const &); + global_device_id_t global_device_id_from_processor(Realm::Processor); + Realm::Processor processor_from_local_device_id(local_device_id_t const &) const; + local_device_id_t local_device_id_from_processor(Realm::Processor) const; + static Realm::Memory get_nearest_memory(Realm::Processor); ///\} @@ -41,7 +45,7 @@ struct RealmContext { ///\{ Realm::Processor get_current_processor() const; Allocator &get_current_device_allocator(); - device_id_t get_current_device_idx() const; + global_device_id_t get_current_global_device_id() const; ///\} /** \name Task creation */ @@ -109,15 +113,15 @@ struct RealmContext { */ Realm::Runtime get_runtime(); - void discover_machine_topology(); + bidict const &get_global_machine_topology(); -public: +private: Realm::Runtime runtime; Realm::Processor processor; Allocator allocator; std::vector outstanding_events; - std::optional> processors = - std::nullopt; + bidict local_machine_topology; + std::optional> cached_global_machine_topology = std::nullopt; }; } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/realm_manager.h b/lib/realm-execution/include/realm-execution/realm_manager.h index ed01f4a022..0aa81c0e74 100644 --- a/lib/realm-execution/include/realm-execution/realm_manager.h +++ b/lib/realm-execution/include/realm-execution/realm_manager.h @@ -6,7 +6,7 @@ #include "realm-execution/realm.h" #include "realm-execution/realm_context.h" #include "realm-execution/tasks/impl/controller_task.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" namespace FlexFlow { diff --git a/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml b/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml index 6a5675772e..492993edd2 100644 --- a/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml +++ b/lib/realm-execution/include/realm-execution/tasks/serializer/serializable_device_specific_ptr.dtg.toml @@ -9,7 +9,7 @@ features = [ ] includes = [ - "pcg/device_id_t.dtg.h", + "task-spec/global_device_id_t.dtg.h", "", "", ] @@ -21,7 +21,7 @@ src_includes = [ [[fields]] name = "device_idx" -type = "::FlexFlow::device_id_t" +type = "::FlexFlow::global_device_id_t" [[fields]] name = "ptr" diff --git a/lib/realm-execution/src/realm-execution/address_space.cc b/lib/realm-execution/src/realm-execution/address_space.cc new file mode 100644 index 0000000000..2886f16312 --- /dev/null +++ b/lib/realm-execution/src/realm-execution/address_space.cc @@ -0,0 +1,17 @@ +#include "realm-execution/address_space.h" + +namespace FlexFlow { + +node_idx_t node_idx_from_realm_address_space(Realm::AddressSpace address_space) { + return node_idx_t{ + nonnegative_int{ + static_cast(address_space), + }, + }; +} + +Realm::AddressSpace realm_address_space_from_node_idx(node_idx_t node_idx) { + return static_cast(node_idx.raw.unwrap_nonnegative()); +}; + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/device_specific_managed_per_device_ff_handle.cc b/lib/realm-execution/src/realm-execution/device_specific_managed_per_device_ff_handle.cc index c16f17d168..bfe845e6ab 100644 --- a/lib/realm-execution/src/realm-execution/device_specific_managed_per_device_ff_handle.cc +++ b/lib/realm-execution/src/realm-execution/device_specific_managed_per_device_ff_handle.cc @@ -8,14 +8,14 @@ namespace FlexFlow { DeviceSpecificPtr make_device_specific_managed_ff_handle( - device_id_t const &device_id, + global_device_id_t const &device_id, std::optional const &managed_handle) { return DeviceSpecificPtr{device_id, managed_handle}; } device_handle_t device_handle_t_from_device_specific_managed_ff_handle( DeviceSpecificPtr const &device_specific, - device_id_t device_idx) { + global_device_id_t device_idx) { return device_handle_t_from_managed_ff_handle_ptr( device_specific.get(device_idx)); } diff --git a/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc b/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc index 2fdc31d2c5..2e5b832f2f 100644 --- a/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc +++ b/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc @@ -32,7 +32,7 @@ DistributedFfHandle proc.kind() == Realm::Processor::TOC_PROC) { handles.insert({proc, make_device_specific_managed_ff_handle( - ctx.get_current_device_idx(), std::nullopt)}); + ctx.get_current_global_device_id(), std::nullopt)}); } } diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index 015af2b1ae..a4a4240ff3 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -22,8 +22,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( ProfilingSettings const &profiling_settings, DistributedFfHandle const &device_handle, OptimizerAttrs const &optimizer_attrs, - Realm::Event precondition, - DeviceType device_type) { + Realm::Event precondition) { // Initialize all operators and save the per-device op state ASSERT(no_nodes_are_initialized(dg)); @@ -32,15 +31,15 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( DeviceSpecificPtr *> device_state_map; for (DynamicNodeInvocation const &invocation : dg.invocations) { - Realm::Processor target_proc = ctx.map_device_coord_to_processor( - assert_unwrap(invocation.node_attrs.device_coord)); + Realm::Processor target_proc = ctx.processor_from_global_device_id( + assert_unwrap(invocation.node_attrs.device_id)); TensorInstanceBacking tensor_backing = subset_tensor_instance_backing_for_invocation(tensor_instance_backing, invocation); DeviceSpecificPtr *device_state_ptr = - new DeviceSpecificPtr{ctx.get_current_device_idx(), + new DeviceSpecificPtr{ctx.get_current_global_device_id(), std::nullopt}; std::optional completion_event = diff --git a/lib/realm-execution/src/realm-execution/fmt/realm_processor.cc b/lib/realm-execution/src/realm-execution/fmt/realm_processor.cc new file mode 100644 index 0000000000..a4afe3cd1b --- /dev/null +++ b/lib/realm-execution/src/realm-execution/fmt/realm_processor.cc @@ -0,0 +1,9 @@ +#include "realm-execution/fmt/realm_processor.h" + +namespace FlexFlow { + +std::ostream &operator<<(std::ostream &s, ::FlexFlow::Realm::Processor const &m) { + return s << fmt::to_string(m); +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc b/lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc new file mode 100644 index 0000000000..2552ec7ef5 --- /dev/null +++ b/lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc @@ -0,0 +1,9 @@ +#include "realm-execution/fmt/realm_processor_kind.h" + +namespace FlexFlow { + +std::ostream &operator<<(std::ostream &s, ::FlexFlow::Realm::Processor::Kind const &k) { + return (s << fmt::to_string(k)); +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index 37bf0ca03d..8829cda275 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -22,14 +22,14 @@ namespace FlexFlow { std::pair - perform_instance_allocation_for_value(device_id_t const &device_id, + perform_instance_allocation_for_value(global_device_id_t const &device_id, DynamicValueAttrs const &value, RealmContext &ctx) { ASSERT(value.accessor == std::nullopt); TensorShape shape = get_piece_shape(value.parallel_tensor_shape.value()); - Realm::Processor proc = ctx.map_device_coord_to_processor(device_id); + Realm::Processor proc = ctx.processor_from_global_device_id(device_id); Realm::Memory memory = ctx.get_nearest_memory(proc); return ctx.create_instance(memory, shape, Realm::ProfilingRequestSet()); } @@ -38,8 +38,7 @@ TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, std::unordered_map const &preallocated, - RealmContext &ctx, - DeviceType device_type) { + RealmContext &ctx) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); for (DynamicValueAttrs const &v : keys(preallocated)) { @@ -53,7 +52,7 @@ TensorInstanceBacking perform_instance_allocation( NOT_IMPLEMENTED(); } else { if (!contains_key(result.backing, v)) { - device_id_t device = assert_unwrap(n.device_coord); + global_device_id_t device = assert_unwrap(n.device_id); result.backing.insert(std::pair{ v, perform_instance_allocation_for_value(device, v, ctx)}); } diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 62747fda7a..e8a8494a55 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -117,11 +117,10 @@ PCGInstance create_pcg_instance( dg = perform_update_insertion(dg, optimizer_attrs); dg = perform_copy_insertion(dg); - debug_print_dynamic_open_dataflow_graph_as_dot(dg); dg = perform_shard_expansion(dg); TensorInstanceBacking tensor_instance_backing = - perform_instance_allocation(dg, inputs, ctx, device_type); + perform_instance_allocation(dg, inputs, ctx); logit_grad_value = transform(logit_grad_value, [&](DynamicValueAttrs const &lgv) { @@ -154,8 +153,7 @@ PCGInstance create_pcg_instance( profiling_settings, device_handle, optimizer_attrs, - ctx.get_outstanding_events(), - device_type); + ctx.get_outstanding_events()); // Compute the topological ordering of the graph auto [kwarg_graph, node_map] = @@ -199,8 +197,8 @@ static Realm::Event spawn_dynamic_node_invocation( invocation); auto spawn_task = [&]() { - Realm::Processor target_proc = ctx.map_device_coord_to_processor( - assert_unwrap(invocation.node_attrs.device_coord)); + Realm::Processor target_proc = ctx.processor_from_global_device_id( + assert_unwrap(invocation.node_attrs.device_id)); return spawn_op_task(ctx, target_proc, invocation, diff --git a/lib/realm-execution/src/realm-execution/processor_kind.cc b/lib/realm-execution/src/realm-execution/processor_kind.cc index 3a40de0a42..29a9b05274 100644 --- a/lib/realm-execution/src/realm-execution/processor_kind.cc +++ b/lib/realm-execution/src/realm-execution/processor_kind.cc @@ -1,4 +1,5 @@ #include "realm-execution/processor_kind.h" +#include "realm-execution/fmt/realm_processor_kind.h" #include namespace FlexFlow { @@ -11,7 +12,7 @@ DeviceType case Realm::Processor::Kind::TOC_PROC: return DeviceType::GPU; default: - PANIC("Unhandled Realm::Processor::Kind", processor_kind); + PANIC("Unhandled Realm::Processor::Kind", fmt::to_string(processor_kind)); } } diff --git a/lib/realm-execution/src/realm-execution/processor_query.cc b/lib/realm-execution/src/realm-execution/processor_query.cc new file mode 100644 index 0000000000..b5dd3e12da --- /dev/null +++ b/lib/realm-execution/src/realm-execution/processor_query.cc @@ -0,0 +1,13 @@ +#include "realm-execution/processor_query.h" + +namespace FlexFlow { + +std::set processor_set_from_query(Realm::Machine::ProcessorQuery const &pq) { + std::set result; + for (Realm::Processor p : pq) { + result.insert(p); + } + return result; +} + +} // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/realm_context.cc b/lib/realm-execution/src/realm-execution/realm_context.cc index 4cd7bdc221..07d60aa73c 100644 --- a/lib/realm-execution/src/realm-execution/realm_context.cc +++ b/lib/realm-execution/src/realm-execution/realm_context.cc @@ -15,17 +15,116 @@ #include "utils/one_to_many/one_to_many.h" #include "utils/optional.h" #include "utils/positive_int/positive_int.h" +#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "realm-execution/address_space.h" +#include "utils/bidict/algorithms/bidict_from_enumerating.h" +#include "utils/bidict/algorithms/transform_values.h" +#include "utils/containers/group_by.h" +#include "task-spec/global_device_id_t.h" +#include "utils/containers/are_all_same.h" +#include "utils/containers/set_of.h" +#include "realm-execution/processor_query.h" +#include "realm-execution/fmt/realm_processor.h" +#include "realm-execution/fmt/realm_processor_kind.h" namespace FlexFlow { + +bidict + build_local_machine_topology(std::set const &local_procs) +{ + { + bool procs_are_local = are_all_same( + transform(local_procs, + [&](Realm::Processor p) -> Realm::AddressSpace { + return p.address_space(); + })); + ASSERT(procs_are_local); + } + + OneToMany by_proc_kind = + group_by(local_procs, + [](Realm::Processor p) -> Realm::Processor::Kind { + return p.kind(); + }); + + auto local_machine_topology_for_proc_kind = [&](Realm::Processor::Kind k) + -> bidict + { + if (!contains(by_proc_kind.left_values(), k)) { + return {}; + } + + bidict enumerated = + bidict_from_enumerating(set_of(by_proc_kind.at_l(k).unwrap_as_unordered_set())).reversed(); + + bidict result = + transform_values( + enumerated, + [&](nonnegative_int idx) -> local_device_id_t { + return local_device_id_t{ + /*idx=*/device_in_node_idx_t{idx}, + /*device_type=*/device_type_from_processor_kind(k), + }; + }); + + return result; + }; + + return binary_merge_disjoint_bidicts( + local_machine_topology_for_proc_kind(Realm::Processor::Kind::LOC_PROC), + local_machine_topology_for_proc_kind(Realm::Processor::Kind::TOC_PROC)); +} + +static bidict + build_global_machine_topology(std::set const &global_procs) +{ + OneToMany by_node_idx = + group_by(global_procs, + [](Realm::Processor p) -> node_idx_t { + return node_idx_from_realm_address_space(p.address_space()); + }); + + auto build_global_machine_topology_for_node = [&](node_idx_t const &node_idx) + -> bidict + { + std::set procs_for_node = set_of(by_node_idx.at_l(node_idx).unwrap_as_unordered_set()); + + bidict + local_topology_for_node = build_local_machine_topology(procs_for_node); + + return transform_values( + local_topology_for_node, + [&](local_device_id_t const &local_device_id) -> global_device_id_t { + return global_device_id_from_local(local_device_id, node_idx); + }); + }; + + return merge_disjoint_bidicts( + transform( + set_of(by_node_idx.left_values()), + build_global_machine_topology_for_node)); +} + +static bidict discover_local_machine_topology(Realm::Processor local_processor) { + Realm::Machine::ProcessorQuery pq(Realm::Machine::get_machine()); + pq.same_address_space_as(local_processor); + + return build_local_machine_topology(processor_set_from_query(pq)); +} + +static bidict discover_global_machine_topology() { + Realm::Machine::ProcessorQuery pq(Realm::Machine::get_machine()); + + return build_global_machine_topology(processor_set_from_query(pq)); +} + RealmContext::RealmContext(Realm::Processor processor) : processor(processor), allocator(get_realm_allocator( - processor, RealmContext::get_nearest_memory(processor))) { - if (processor != Realm::Processor::NO_PROC) { - this->discover_machine_topology(); - } -} + processor, RealmContext::get_nearest_memory(processor))), + local_machine_topology(discover_local_machine_topology(processor)) +{ } RealmContext::~RealmContext() { if (!this->outstanding_events.empty()) { @@ -34,23 +133,28 @@ RealmContext::~RealmContext() { } } -static std::tuple - convert_machine_space_coordinate(MachineSpaceCoordinate const &device_coord, - DeviceType device_type) { - Realm::AddressSpace as = int{device_coord.node_idx}; - Realm::Processor::Kind kind = processor_kind_from_device_type(device_type); - nonnegative_int proc_in_node = device_coord.device_idx; - return std::tuple{as, kind, proc_in_node}; +Realm::Processor RealmContext::processor_from_global_device_id( + global_device_id_t const &global_device_id) { + + return this->get_global_machine_topology().at_r(global_device_id); } -Realm::Processor RealmContext::map_device_coord_to_processor( - device_id_t const &device_id) const { - return assert_unwrap(this->processors).at_r(device_id); +global_device_id_t RealmContext::global_device_id_from_processor( + Realm::Processor processor) { + + return this->get_global_machine_topology().at_l(processor); } -device_id_t - RealmContext::map_processor_to_device_coord(Realm::Processor p) const { - return assert_unwrap(this->processors).at_l(p); +Realm::Processor RealmContext::processor_from_local_device_id( + local_device_id_t const &local_device_id) const { + + return this->local_machine_topology.at_r(local_device_id); +} + +local_device_id_t RealmContext::local_device_id_from_processor( + Realm::Processor processor) const { + + return this->local_machine_topology.at_l(processor); } Realm::Memory RealmContext::get_nearest_memory(Realm::Processor proc) { @@ -74,9 +178,12 @@ Allocator &RealmContext::get_current_device_allocator() { return this->allocator; } -device_id_t RealmContext::get_current_device_idx() const { +global_device_id_t RealmContext::get_current_global_device_id() const { Realm::Processor proc = this->get_current_processor(); - return this->map_processor_to_device_coord(proc); + + return global_device_id_from_local( + this->local_device_id_from_processor(proc), + node_idx_from_realm_address_space(proc.address_space())); } Realm::Event @@ -290,50 +397,12 @@ Realm::Event RealmContext::merge_outstanding_events() { return result; } -void RealmContext::discover_machine_topology() { - if (this->processors.has_value()) { - return; - } - - std::unordered_map, nonnegative_int> - next_device_idx; - - auto fresh_device_id = [&](nonnegative_int node_idx, - DeviceType device_type) -> device_id_t { - std::pair key = - std::pair{node_idx, device_type}; - if (!contains_key(next_device_idx, key)) { - next_device_idx.insert({key, 0_n}); - } - - nonnegative_int device_idx = next_device_idx.at(key); - next_device_idx.at(key)++; - - return device_id_t{ - MachineSpaceCoordinate{node_idx, device_idx}, - device_type, - }; - }; - - bidict procs; - Realm::Machine::ProcessorQuery pq(Realm::Machine::get_machine()); - for (Realm::Processor proc : pq) { - Realm::AddressSpace as = proc.address_space(); - Realm::Processor::Kind kind = proc.kind(); - - nonnegative_int node_idx = nonnegative_int{static_cast(as)}; - - if (kind != Realm::Processor::LOC_PROC && - kind != Realm::Processor::TOC_PROC) { - continue; - } - - DeviceType device_type = device_type_from_processor_kind(kind); - device_id_t coord = fresh_device_id(node_idx, device_type); - procs.equate_strict(proc, coord); +bidict const &RealmContext::get_global_machine_topology() { + if (!this->cached_global_machine_topology.has_value()) { + this->cached_global_machine_topology = discover_global_machine_topology(); } - this->processors = procs; + return assert_unwrap(this->cached_global_machine_topology); } Realm::Runtime RealmContext::get_runtime() { diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/ff_handle_init_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/ff_handle_init_task.cc index f47b957f32..42e69a967b 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/ff_handle_init_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/ff_handle_init_task.cc @@ -39,7 +39,7 @@ void ff_handle_init_task_body(void const *args, RealmContext ctx{proc}; DeviceSpecificPtr managed_handle = make_device_specific_managed_ff_handle( - ctx.get_current_device_idx(), + ctx.get_current_global_device_id(), make_ff_handle_for_processor(proc, task_args.workSpaceSize, task_args.allowTensorOpMathConversion)); diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/op_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/op_task.cc index 02b0dd7f54..9026caba96 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/op_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/op_task.cc @@ -27,7 +27,7 @@ void op_task_body(void const *args, RealmContext ctx{proc}; device_handle_t device_handle = device_handle_t_from_device_specific_managed_ff_handle( - task_args.device_handle, ctx.get_current_device_idx()); + task_args.device_handle, ctx.get_current_global_device_id()); // Patch the invocation to include the provided instances auto map_instance_to_accessor = [&](DynamicValueAttrs const &value) { @@ -53,11 +53,11 @@ void op_task_body(void const *args, /*per_device_op_state=*/ transform(and_then(task_args.device_state, [&](DeviceSpecificPtr const &d) { - return d.get(ctx.get_current_device_idx()); + return d.get(ctx.get_current_global_device_id()); }), [](PerDeviceOpState *ptr) { return *ptr; }), /*optimizer_attrs=*/task_args.optimizer_attrs, - /*device_idx=*/ctx.get_current_device_idx()); + /*global_device_id=*/ctx.get_current_global_device_id()); } Realm::Event spawn_op_task( diff --git a/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc b/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc index c05932ff61..3466890d26 100644 --- a/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc +++ b/lib/realm-execution/src/realm-execution/tasks/impl/per_device_op_state_init_task.cc @@ -31,7 +31,7 @@ void per_device_op_state_init_task_body(void const *args, RealmContext ctx{proc}; device_handle_t device_handle = device_handle_t_from_device_specific_managed_ff_handle( - task_args.device_handle, ctx.get_current_device_idx()); + task_args.device_handle, ctx.get_current_global_device_id()); // Patch the invocation to include the provided instances auto map_instance_to_accessor = [&](DynamicValueAttrs const &value) { @@ -55,16 +55,16 @@ void per_device_op_state_init_task_body(void const *args, task_args.profiling_settings, device_handle, task_args.optimizer_attrs, - ctx.get_current_device_idx()); + ctx.get_current_global_device_id()); DeviceSpecificPerDeviceOpState result_state = assert_unwrap(result_invocation.node_attrs.per_device_op_state); // Important: to make sure this doesn't get deallocated, we intentionally leak // the allocation here PerDeviceOpState *result_state_ptr = new PerDeviceOpState{get_per_device_op_state_from_device_specific( - result_state, ctx.get_current_device_idx())}; + result_state, ctx.get_current_global_device_id())}; DeviceSpecificPtr result_device_specific{ - ctx.get_current_device_idx(), result_state_ptr}; + ctx.get_current_global_device_id(), result_state_ptr}; spawn_per_device_op_state_init_return_task(ctx, task_args.origin_proc, result_device_specific, diff --git a/lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc b/lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc new file mode 100644 index 0000000000..75430da876 --- /dev/null +++ b/lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc @@ -0,0 +1,13 @@ +#include +#include "realm-execution/fmt/realm_processor_kind.h" + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("fmt::to_string(Realm::Processor::Kind)") { + std::string result = fmt::to_string(::FlexFlow::Realm::Processor::Kind::TOC_PROC); + std::string correct = ""; + + CHECK(result == correct); + } +} diff --git a/lib/realm-execution/test/src/realm-execution/test_e2e.cc b/lib/realm-execution/test/src/realm-execution/test_e2e.cc index f698d2a07f..57871fb148 100644 --- a/lib/realm-execution/test/src/realm-execution/test_e2e.cc +++ b/lib/realm-execution/test/src/realm-execution/test_e2e.cc @@ -224,8 +224,6 @@ TEST_SUITE(FF_TEST_SUITE) { RealmManager manager = RealmManager{&fake_argc, &fake_argv}; (void)manager.start_controller([](RealmContext &ctx) { - ASSERT(ctx.processors.has_value()); - E2ETrainingConfig cfg = create_e2e_test_case(); Allocator allocator = ctx.get_current_device_allocator(); diff --git a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc index 2ad5b54a17..7dc6a2cc5e 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc @@ -2,9 +2,9 @@ #include "substitutions/output_graph/output_graph_expr.h" #include "substitutions/sub_parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_pairs.h" -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" #include "utils/containers/values.h" #include "utils/containers/zip_values_strict.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" namespace FlexFlow { @@ -26,7 +26,7 @@ bidict mapping_for_layer = bidict_from_pairs(values( zip_values_strict(layer_outputs, output_graph_expr_outputs))); - result = merge_disjoint_bidicts(result, mapping_for_layer); + result = binary_merge_disjoint_bidicts(result, mapping_for_layer); } return result; diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index 498fd6c1bf..46d284c4a0 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -5,12 +5,12 @@ #include "utils/bidict/algorithms/bidict_from_keys_and_values.h" #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/bidict/algorithms/exhaustive_relational_join.h" -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" #include "utils/bidict/algorithms/transform_values.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/map_values.h" #include "utils/containers/values.h" #include "utils/containers/zip.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" namespace FlexFlow { @@ -34,7 +34,7 @@ bidict exhaustive_relational_join(pattern_node_outputs.reversed(), matched_layer_output_tensors); - result = merge_disjoint_bidicts(result, mapping); + result = binary_merge_disjoint_bidicts(result, mapping); } return result; diff --git a/lib/task-spec/include/task-spec/device_specific.h b/lib/task-spec/include/task-spec/device_specific.h index 49a2555411..bc09cfc7a7 100644 --- a/lib/task-spec/include/task-spec/device_specific.h +++ b/lib/task-spec/include/task-spec/device_specific.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LOCAL_EXECUTION_DEVICE_SPECIFIC_H #define _FLEXFLOW_LOCAL_EXECUTION_DEVICE_SPECIFIC_H -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "utils/hash/tuple.h" #include @@ -12,7 +12,7 @@ struct DeviceSpecific { DeviceSpecific() = delete; template - static DeviceSpecific create(device_id_t const &device_idx, + static DeviceSpecific create(global_device_id_t const &device_idx, Args &&...args) { return DeviceSpecific(std::make_shared(std::forward(args)...), device_idx); @@ -26,18 +26,18 @@ struct DeviceSpecific { return this->tie() != other.tie(); } - T const *get(device_id_t const &curr_device_idx) const { + T const *get(global_device_id_t const &curr_device_idx) const { ASSERT(curr_device_idx == this->device_idx); return (T const *)this->ptr.get(); } private: - DeviceSpecific(std::shared_ptr ptr, device_id_t const &device_idx) + DeviceSpecific(std::shared_ptr ptr, global_device_id_t const &device_idx) : ptr(ptr), device_idx(device_idx) {} private: std::shared_ptr ptr; - device_id_t device_idx; + global_device_id_t device_idx; private: std::tuple tie() const { diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml index 7655bd24e1..a25c9fd08b 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml @@ -13,7 +13,7 @@ includes = [ "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h", "task-spec/dynamic_graph/training_operation_attrs.dtg.h", "task-spec/device_specific_per_device_op_state.dtg.h", - "task-spec/device_id_t.dtg.h", + "task-spec/global_device_id_t.dtg.h", "task-spec/dynamic_graph/dynamic_node_mapping.dtg.h", ] @@ -26,8 +26,8 @@ name = "task_type" type = "std::optional<::FlexFlow::DynamicTaskType>" [[fields]] -name = "device_coord" -type = "std::optional<::FlexFlow::device_id_t>" +name = "device_id" +type = "std::optional<::FlexFlow::global_device_id_t>" docstring = ''' \brief The device on which this task should execute. diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h index 6a52773a71..c5a4976db7 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h @@ -1,16 +1,16 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_MAPPING_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_MAPPING_H -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_mapping.dtg.h" namespace FlexFlow { -bidict +bidict dynamic_node_mapping_bindings_for_slot_name(DynamicNodeMapping const &, TensorSlotName const &); -std::unordered_set +std::unordered_set target_devices_of_dynamic_node_mapping(DynamicNodeMapping const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h b/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h index b40fabe7bf..a9cb86c185 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/machine_slicing.h @@ -7,11 +7,11 @@ namespace FlexFlow { std::unordered_set perform_machine_slicing_for_invocation(DynamicNodeInvocation const &, - device_id_t const &); + global_device_id_t const &); DynamicOpenDataflowGraph perform_machine_slicing(DynamicOpenDataflowGraph const &, - device_id_t const &); + global_device_id_t const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml index ba9785bc92..2be195f7c1 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml @@ -11,7 +11,7 @@ features = [ includes = [ "utils/bidict/bidict.h", "op-attrs/parallel_tensor_space_coordinate.dtg.h", - "task-spec/device_id_t.dtg.h", + "task-spec/global_device_id_t.dtg.h", ] src_includes = [ @@ -19,4 +19,4 @@ src_includes = [ [[fields]] name = "raw" -type = "::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::device_id_t>" +type = "::FlexFlow::bidict<::FlexFlow::ParallelTensorSpaceCoordinate, ::FlexFlow::global_device_id_t>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml index 43fb8490ec..6638350bfd 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_node_attrs.dtg.toml @@ -11,7 +11,7 @@ features = [ includes = [ "", "task-spec/dynamic_graph/dynamic_task_type.dtg.h", - "task-spec/device_id_t.dtg.h", + "task-spec/global_device_id_t.dtg.h", "task-spec/dynamic_graph/dynamic_node_mapping.dtg.h", "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h", "task-spec/dynamic_graph/training_operation_attrs.dtg.h", @@ -27,8 +27,8 @@ name = "task_type" type = "std::optional<::FlexFlow::DynamicTaskType>" [[fields]] -name = "device_coord" -type = "std::optional<::FlexFlow::device_id_t>" +name = "device_id" +type = "std::optional<::FlexFlow::global_device_id_t>" [[fields]] name = "mapping" diff --git a/lib/task-spec/include/task-spec/device_id_t.dtg.toml b/lib/task-spec/include/task-spec/global_device_id_t.dtg.toml similarity index 91% rename from lib/task-spec/include/task-spec/device_id_t.dtg.toml rename to lib/task-spec/include/task-spec/global_device_id_t.dtg.toml index f6d29ee1a7..5644c33c17 100644 --- a/lib/task-spec/include/task-spec/device_id_t.dtg.toml +++ b/lib/task-spec/include/task-spec/global_device_id_t.dtg.toml @@ -1,5 +1,5 @@ namespace = "FlexFlow" -name = "device_id_t" +name = "global_device_id_t" type = "struct" features = [ "eq", diff --git a/lib/task-spec/include/task-spec/global_device_id_t.h b/lib/task-spec/include/task-spec/global_device_id_t.h new file mode 100644 index 0000000000..2f48046ff2 --- /dev/null +++ b/lib/task-spec/include/task-spec/global_device_id_t.h @@ -0,0 +1,15 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_GLOBAL_DEVICE_ID_T_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_GLOBAL_DEVICE_ID_T_H + +#include "task-spec/global_device_id_t.dtg.h" +#include "task-spec/local_device_id_t.dtg.h" +#include "pcg/node_idx_t.dtg.h" + +namespace FlexFlow { + +global_device_id_t global_device_id_from_local(local_device_id_t const &, node_idx_t); +local_device_id_t local_device_id_from_global(global_device_id_t const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/local_device_id_t.dtg.toml b/lib/task-spec/include/task-spec/local_device_id_t.dtg.toml new file mode 100644 index 0000000000..cb6ae47621 --- /dev/null +++ b/lib/task-spec/include/task-spec/local_device_id_t.dtg.toml @@ -0,0 +1,23 @@ +namespace = "FlexFlow" +name = "local_device_id_t" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "pcg/device_in_node_idx_t.dtg.h", + "pcg/device_type.dtg.h", +] + +[[fields]] +name = "idx" +type = "::FlexFlow::device_in_node_idx_t" + +[[fields]] +name = "device_type" +type = "::FlexFlow::DeviceType" diff --git a/lib/task-spec/include/task-spec/per_device_op_state.h b/lib/task-spec/include/task-spec/per_device_op_state.h index 8783f902e4..3d9f90c0a6 100644 --- a/lib/task-spec/include/task-spec/per_device_op_state.h +++ b/lib/task-spec/include/task-spec/per_device_op_state.h @@ -9,7 +9,7 @@ namespace FlexFlow { PerDeviceOpState get_per_device_op_state_from_device_specific( - DeviceSpecificPerDeviceOpState const &, device_id_t device_idx); + DeviceSpecificPerDeviceOpState const &, global_device_id_t device_idx); } diff --git a/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h b/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h index 165631889f..776157e644 100644 --- a/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h +++ b/lib/task-spec/include/task-spec/task_argument_accessor/itask_argument_accessor.h @@ -9,7 +9,7 @@ #include "op-attrs/tensor_slot_name.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/concrete_arg_spec.h" -#include "task-spec/device_id_t.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/ops/arg_slot_id_t.dtg.h" #include "task-spec/per_device_op_state.dtg.h" #include "task-spec/privilege_tensor_accessor.h" @@ -37,7 +37,7 @@ struct ITaskArgumentAccessor { virtual OptimizerAttrs get_optimizer_attrs() const = 0; virtual Allocator get_allocator() const = 0; - virtual device_id_t get_device_idx() const = 0; + virtual global_device_id_t get_device_idx() const = 0; }; CHECK_RC_COPY_VIRTUAL_COMPLIANT(ITaskArgumentAccessor); diff --git a/lib/task-spec/include/task-spec/task_argument_accessor/task_argument_accessor.h b/lib/task-spec/include/task-spec/task_argument_accessor/task_argument_accessor.h index 0c63643e6e..9ad56bf69d 100644 --- a/lib/task-spec/include/task-spec/task_argument_accessor/task_argument_accessor.h +++ b/lib/task-spec/include/task-spec/task_argument_accessor/task_argument_accessor.h @@ -58,7 +58,7 @@ struct TaskArgumentAccessor { return this->ptr->get_allocator(); } - device_id_t get_device_idx() const { + global_device_id_t get_device_idx() const { return this->ptr->get_device_idx(); } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index f611735e26..3487b2979e 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -61,16 +61,16 @@ bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { static std::pair filter_mapping_to_avoid_degenerate_copies(DynamicValueAttrs const &input, DynamicValueAttrs const &output) { - std::unordered_set> + std::unordered_set> input_mapping = unordered_set_of(assert_unwrap(input.mapping).raw); - std::unordered_set> + std::unordered_set> output_mapping = unordered_set_of(assert_unwrap(output.mapping).raw); // Exclude the point shared between the input and output mappings, because // those will not result in actual copies once shard expansion is performed std::unordered_set< - std::pair> + std::pair> remove = set_intersection(input_mapping, output_mapping); DynamicValueAttrs filtered_input = input; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc index 9e56dddcf9..4e21da9c2a 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc @@ -4,24 +4,24 @@ namespace FlexFlow { -bidict +bidict dynamic_node_mapping_bindings_for_slot_name( DynamicNodeMapping const &mapping, TensorSlotName const &slot_name) { bidict coord_bindings = get_tensor_bindings_for_slot_name(mapping.op_task_group, slot_name); return transform_values( - coord_bindings, [&](MachineSpaceCoordinate const &coord) -> device_id_t { - return device_id_t{coord, mapping.device_type}; + coord_bindings, [&](MachineSpaceCoordinate const &coord) -> global_device_id_t { + return global_device_id_t{coord, mapping.device_type}; }); } -std::unordered_set +std::unordered_set target_devices_of_dynamic_node_mapping(DynamicNodeMapping const &mapping) { return transform(mapping.op_task_group.get_shard_bindings().left_values(), - [&](MachineSpaceCoordinate const &c) -> device_id_t { - return device_id_t{ + [&](MachineSpaceCoordinate const &c) -> global_device_id_t { + return global_device_id_t{ /*coord=*/c, /*device_type=*/mapping.device_type, }; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index 07c49664a5..edb70c2f3a 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -59,6 +59,15 @@ void check_dynamic_open_dataflow_graph_is_valid( ASSERT(values_produced_multiple_times.size() == 0, keys(values_produced_multiple_times)); + // since DynamicOpenDataflowGraph contains a set of invocations rather than a + // LabelledOpenKwargDataflowGraph, some properties guaranteed by + // LabelledOpenKwargDataflowGraph (e.g., the graph is acyclic, all tensors + // originate from another operator's output unless they're a designated graph + // input, etc.) are not automatically guaranteed. Since + // LabelledOpenKwargDataflowGraph guarantees these properties, the easiest + // way to check them is to try to convert the DynamicOpenDataflowGraph into a + // LabelledOpenKwargDataflowGraph, and if a value is returned without an + // assertion we know the properties hold. labelled_open_kwarg_dataflow_graph_from_dynamic_open_dataflow_graph(g); } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc index e46f8b1510..4871a3ee4e 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc @@ -6,11 +6,11 @@ namespace FlexFlow { std::unordered_set perform_machine_slicing_for_invocation( DynamicNodeInvocation const &invocation, - device_id_t const &device_coord) { + global_device_id_t const &device_id) { - ASSERT(invocation.node_attrs.device_coord.has_value()); + ASSERT(invocation.node_attrs.device_id.has_value()); - if (invocation.node_attrs.device_coord.value() == device_coord) { + if (invocation.node_attrs.device_id.value() == device_id) { return {invocation}; } else { return {}; @@ -19,12 +19,12 @@ std::unordered_set DynamicOpenDataflowGraph perform_machine_slicing(DynamicOpenDataflowGraph const &g, - device_id_t const &device_coord) { + global_device_id_t const &device_id) { DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( g, [&](DynamicNodeInvocation const &invocation) -> std::unordered_set { - return perform_machine_slicing_for_invocation(invocation, device_coord); + return perform_machine_slicing_for_invocation(invocation, device_id); }); return result; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc index d613194d14..7ad4686c73 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_node_attrs.cc @@ -7,7 +7,7 @@ SerializableDynamicNodeAttrs dynamic_node_attrs_to_serializable(DynamicNodeAttrs const &attrs) { return SerializableDynamicNodeAttrs{ /*task_type=*/attrs.task_type, - /*device_coord=*/attrs.device_coord, + /*device_id=*/attrs.device_id, /*mapping=*/attrs.mapping, /*op_attrs=*/attrs.op_attrs, /*layer_guid=*/attrs.layer_guid, @@ -18,7 +18,7 @@ DynamicNodeAttrs dynamic_node_attrs_from_serializable( SerializableDynamicNodeAttrs const &attrs) { return DynamicNodeAttrs{ /*task_type=*/attrs.task_type, - /*device_coord=*/attrs.device_coord, + /*device_id=*/attrs.device_id, /*mapping=*/attrs.mapping, /*op_attrs=*/attrs.op_attrs, /*layer_guid=*/attrs.layer_guid, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index 117819a639..15f78944c8 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -13,7 +13,7 @@ namespace FlexFlow { bool node_is_shard_expanded(DynamicNodeAttrs const &n) { - return n.device_coord.has_value(); + return n.device_id.has_value(); } bool value_is_shard_expanded(DynamicValueAttrs const &n) { @@ -42,9 +42,9 @@ bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { slot_is_shard_expanded); } -static bidict +static bidict restrict_tensor_mapping_keys_to_coord( - bidict const &mapping, + bidict const &mapping, ParallelTensorSpaceCoordinate const ¶llel_tensor_coord) { return filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { return p == parallel_tensor_coord; @@ -53,7 +53,7 @@ static bidict static DynamicNodeInvocation shard_invocation_for_binding( DynamicNodeInvocation const &i, - device_id_t const &device_coord, + global_device_id_t const &device_id, OperatorAtomicTaskShardBinding const &binding) { auto shard_expand_value_attrs = [&](DynamicTensorSlot const &s, @@ -76,7 +76,7 @@ static DynamicNodeInvocation shard_invocation_for_binding( DynamicNodeAttrs expanded_node_attrs = [&]() { DynamicNodeAttrs result = i.node_attrs; - result.device_coord = device_coord; + result.device_id = device_id; return result; }(); @@ -91,7 +91,7 @@ static std::unordered_set perform_shard_expansion_for_copy(DynamicNodeInvocation const &i) { auto [input_slot, input] = get_only(i.inputs); auto [output_slot, output] = get_only(i.outputs); - bidict input_mapping = + bidict input_mapping = assert_unwrap(input.mapping).raw; require_same(input_mapping.left_values(), assert_unwrap(output.mapping).raw.left_values()); @@ -105,7 +105,7 @@ static std::unordered_set // because we expect this to align with the most efficient way to issue // copies in Realm, although the current Realm backend uses a // centralized controller and thus issues copies all from a single node. - device_id_t machine_coord = input_mapping.at_l(p); + global_device_id_t machine_coord = input_mapping.at_l(p); return shard_invocation_for_binding(i, machine_coord, @@ -125,11 +125,11 @@ std::unordered_set DynamicNodeMapping mapping = assert_unwrap(i.node_attrs.mapping); - std::unordered_set shard_machine_coords = + std::unordered_set shard_machine_coords = target_devices_of_dynamic_node_mapping(mapping); return transform( - shard_machine_coords, [&](device_id_t const &c) -> DynamicNodeInvocation { + shard_machine_coords, [&](global_device_id_t const &c) -> DynamicNodeInvocation { OperatorAtomicTaskShardBinding slot_bindings = mapping.op_task_group.get_shard_bindings().at_l(c.coord); diff --git a/lib/task-spec/src/task-spec/global_device_id_t.cc b/lib/task-spec/src/task-spec/global_device_id_t.cc new file mode 100644 index 0000000000..e32a5f5f8b --- /dev/null +++ b/lib/task-spec/src/task-spec/global_device_id_t.cc @@ -0,0 +1,25 @@ +#include "task-spec/global_device_id_t.h" + +namespace FlexFlow { + +global_device_id_t global_device_id_from_local( + local_device_id_t const &local_device_id, + node_idx_t node_idx) +{ + return global_device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/node_idx.raw, + /*device_idx=*/local_device_id.idx.raw, + }, + /*device_type=*/local_device_id.device_type, + }; +} + +local_device_id_t local_device_id_from_global(global_device_id_t const &global_device_id) { + return local_device_id_t{ + /*idx=*/device_in_node_idx_t{global_device_id.coord.device_idx}, + /*device_type=*/global_device_id.device_type, + }; +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/per_device_op_state.cc b/lib/task-spec/src/task-spec/per_device_op_state.cc index 438cd8886c..db4c3f95d8 100644 --- a/lib/task-spec/src/task-spec/per_device_op_state.cc +++ b/lib/task-spec/src/task-spec/per_device_op_state.cc @@ -5,7 +5,7 @@ namespace FlexFlow { PerDeviceOpState get_per_device_op_state_from_device_specific( DeviceSpecificPerDeviceOpState const &device_specific, - device_id_t device_idx) { + global_device_id_t device_idx) { return device_specific.visit( [&](auto const &x) { return PerDeviceOpState{*(x.get(device_idx))}; }); } diff --git a/lib/task-spec/test/src/task-spec/device_specific.cc b/lib/task-spec/test/src/task-spec/device_specific.cc index 6a42a9b570..983ba4943d 100644 --- a/lib/task-spec/test/src/task-spec/device_specific.cc +++ b/lib/task-spec/test/src/task-spec/device_specific.cc @@ -7,12 +7,12 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("DeviceSpecific") { DeviceSpecific device_specific1 = DeviceSpecific::create( - device_id_t{MachineSpaceCoordinate{0_n, 0_n}, DeviceType::GPU}, + global_device_id_t{MachineSpaceCoordinate{0_n, 0_n}, DeviceType::GPU}, "hello world"); DeviceSpecific device_specific2 = DeviceSpecific::create( - device_id_t{MachineSpaceCoordinate{0_n, 1_n}, DeviceType::GPU}, + global_device_id_t{MachineSpaceCoordinate{0_n, 1_n}, DeviceType::GPU}, "hello world"); std::string result1 = fmt::to_string(device_specific1); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index 58a8a36fcd..97304ebbf0 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -65,8 +65,8 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - auto mk_device_id = [&](nonnegative_int device_idx) -> device_id_t { - return device_id_t{ + auto mk_device_id = [&](nonnegative_int device_idx) -> global_device_id_t { + return global_device_id_t{ mk_machine_coord(device_idx), DeviceType::GPU, }; @@ -209,7 +209,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*src_slot=*/TensorSlotName::OUTPUT, /*mapping=*/ ParallelTensorMapping{ - bidict{ + bidict{ {mk_ptensor_coord(0_n), mk_device_id(0_n)}, {mk_ptensor_coord(1_n), mk_device_id(1_n)}, }, @@ -220,7 +220,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*src_slot=*/TensorSlotName::OUTPUT, /*mapping=*/ ParallelTensorMapping{ - bidict{ + bidict{ {mk_ptensor_coord(0_n), mk_device_id(0_n)}, {mk_ptensor_coord(1_n), mk_device_id(1_n)}, }, @@ -231,7 +231,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*src_slot=*/TensorSlotName::OUTPUT, /*mapping=*/ ParallelTensorMapping{ - bidict{ + bidict{ {mk_ptensor_coord(0_n), mk_device_id(0_n)}, {mk_ptensor_coord(1_n), mk_device_id(2_n)}, }, @@ -242,7 +242,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*src_slot=*/TensorSlotName::OUTPUT, /*mapping=*/ ParallelTensorMapping{ - bidict{ + bidict{ {mk_ptensor_coord(0_n), mk_device_id(0_n)}, {mk_ptensor_coord(1_n), mk_device_id(2_n)}, }, @@ -399,7 +399,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*src_slot=*/TensorSlotName::OUTPUT, /*mapping=*/ ParallelTensorMapping{ - bidict{ + bidict{ {mk_ptensor_coord(0_n), mk_device_id(0_n)}, {mk_ptensor_coord(1_n), mk_device_id(1_n)}, }, @@ -410,7 +410,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*src_slot=*/TensorSlotName::OUTPUT, /*mapping=*/ ParallelTensorMapping{ - bidict{ + bidict{ {mk_ptensor_coord(0_n), mk_device_id(0_n)}, {mk_ptensor_coord(1_n), mk_device_id(1_n)}, }, @@ -421,7 +421,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*src_slot=*/TensorSlotName::OUTPUT, /*mapping=*/ ParallelTensorMapping{ - bidict{ + bidict{ {mk_ptensor_coord(0_n), mk_device_id(0_n)}, {mk_ptensor_coord(1_n), mk_device_id(1_n)}, }, diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc index 6bc31999a6..9abb802af2 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc @@ -7,8 +7,8 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("perform_machine_slicing_for_invocation") { auto mk_device_id = [](nonnegative_int node_idx, - nonnegative_int device_idx) -> device_id_t { - return device_id_t{ + nonnegative_int device_idx) -> global_device_id_t { + return global_device_id_t{ MachineSpaceCoordinate{ /*node_idx=*/node_idx, /*device_idx=*/device_idx, @@ -33,9 +33,9 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - device_id_t mc1 = mk_device_id(0_n, 0_n); - device_id_t mc2 = mk_device_id(2_n, 0_n); - device_id_t mc3 = mk_device_id(4_n, 0_n); + global_device_id_t mc1 = mk_device_id(0_n, 0_n); + global_device_id_t mc2 = mk_device_id(2_n, 0_n); + global_device_id_t mc3 = mk_device_id(4_n, 0_n); ParallelTensorSpaceCoordinate mc1_input_coord = mk_pt_coord(0_n, 0_n, 0_n, 0_n); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index 65870192d7..cc51c838b9 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -44,14 +44,14 @@ TEST_SUITE(FF_TEST_SUITE) { }; DeviceType device_type = DeviceType::GPU; - auto mk_device_id = [&](MachineSpaceCoordinate const &c) -> device_id_t { - return device_id_t{c, device_type}; + auto mk_device_id = [&](MachineSpaceCoordinate const &c) -> global_device_id_t { + return global_device_id_t{c, device_type}; }; auto mk_value = [&](size_t src_node_id, TensorSlotName src_slot_name, - bidict tensor_binding, + bidict tensor_binding, std::optional const &shard_coord) -> DynamicValueAttrs { if (shard_coord.has_value()) { @@ -155,7 +155,7 @@ TEST_SUITE(FF_TEST_SUITE) { TensorSlotName use_slot_name, std::optional const &shard_coord) -> DynamicValueAttrs { - bidict tensor_binding = + bidict tensor_binding = dynamic_node_mapping_bindings_for_slot_name(node_mapping, use_slot_name); return mk_value( @@ -212,7 +212,7 @@ TEST_SUITE(FF_TEST_SUITE) { perform_shard_expansion_for_invocation(input); auto mk_invocation_shard = - [&](device_id_t const &device_coord, + [&](global_device_id_t const &device_coord, ParallelTensorSpaceCoordinate const &input_shard_coord, ParallelTensorSpaceCoordinate const &weight_shard_coord, ParallelTensorSpaceCoordinate const &output_1_shard_coord, @@ -283,19 +283,19 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("for copy operator") { - device_id_t dev1 = mk_device_id(mk_machine_coord(0_n, 0_n)); - device_id_t dev2 = mk_device_id(mk_machine_coord(1_n, 0_n)); - device_id_t dev3 = mk_device_id(mk_machine_coord(2_n, 0_n)); - device_id_t dev4 = mk_device_id(mk_machine_coord(3_n, 0_n)); + global_device_id_t dev1 = mk_device_id(mk_machine_coord(0_n, 0_n)); + global_device_id_t dev2 = mk_device_id(mk_machine_coord(1_n, 0_n)); + global_device_id_t dev3 = mk_device_id(mk_machine_coord(2_n, 0_n)); + global_device_id_t dev4 = mk_device_id(mk_machine_coord(3_n, 0_n)); ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); - bidict src_binding{ + bidict src_binding{ {pt1, dev1}, {pt2, dev2}, }; - bidict dst_binding{ + bidict dst_binding{ {pt1, dev3}, {pt2, dev4}, }; @@ -331,7 +331,7 @@ TEST_SUITE(FF_TEST_SUITE) { perform_shard_expansion_for_invocation(input); auto mk_invocation_shard = - [&](device_id_t const &device_coord, + [&](global_device_id_t const &device_id, ParallelTensorSpaceCoordinate const &tensor_shard_coord) -> DynamicNodeInvocation { DynamicNodeInvocation result = input; @@ -343,7 +343,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; // See perform_shard_expansion_for_copy in shard_expansion.cc for explanation of the choice of device placement. - result.node_attrs.device_coord = device_coord; + result.node_attrs.device_id = device_id; result.outputs = { { mk_slot(TensorSlotName::OUTPUT), diff --git a/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h b/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h new file mode 100644 index 0000000000..607cb08229 --- /dev/null +++ b/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h @@ -0,0 +1,38 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BINARY_MERGE_DISJOINT_BIDICTS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BINARY_MERGE_DISJOINT_BIDICTS_H + +#include "utils/bidict/bidict.h" +#include +#include "utils/containers/are_disjoint.h" +#include "utils/bidict/algorithms/right_entries.h" +#include "utils/bidict/algorithms/left_entries.h" + +namespace FlexFlow { + +template +bidict binary_merge_disjoint_bidicts(bidict const &lhs, + bidict const &rhs) { + ASSERT( + are_disjoint(left_entries(lhs), left_entries(rhs)), + "Left entries of {} and {} are non-disjoint", lhs, rhs + ); + + ASSERT( + are_disjoint(right_entries(lhs), right_entries(rhs)), + "Right entries of {} and {} are non-disjoint", lhs, rhs + ); + + bidict result; + for (auto const &kv : lhs) { + result.equate_strict(kv.first, kv.second); + } + for (auto const &kv : rhs) { + result.equate_strict(kv.first, kv.second); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h index 97e7334c26..1ce7cb3a13 100644 --- a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h +++ b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h @@ -1,35 +1,20 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H -#include "utils/bidict/algorithms/left_entries.h" -#include "utils/bidict/algorithms/right_entries.h" #include "utils/bidict/bidict.h" -#include "utils/containers/are_disjoint.h" -#include "utils/exception.h" +#include "utils/containers/foldl.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" namespace FlexFlow { template -bidict merge_disjoint_bidicts(bidict const &lhs, - bidict const &rhs) { - if (!are_disjoint(left_entries(lhs), left_entries(rhs))) { - throw mk_runtime_error( - fmt::format("Left entries of {} and {} are non-disjoint", lhs, rhs)); - } - if (!are_disjoint(right_entries(lhs), right_entries(rhs))) { - throw mk_runtime_error( - fmt::format("Right entries of {} and {} are non-disjoint", lhs, rhs)); - } - - bidict result; - for (auto const &kv : lhs) { - result.equate(kv.first, kv.second); - } - for (auto const &kv : rhs) { - result.equate(kv.first, kv.second); - } - - return result; +bidict merge_disjoint_bidicts(std::set> const &bidicts) { + return foldl( + bidicts, + bidict{}, + [](bidict const &accum, bidict const &x) -> bidict { + return binary_merge_disjoint_bidicts(accum, x); + }); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc new file mode 100644 index 0000000000..d3deb887b1 --- /dev/null +++ b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc @@ -0,0 +1,13 @@ +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; + +template + bidict binary_merge_disjoint_bidicts(bidict const &, + bidict const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc index 754b8d2e90..e5ed21bb40 100644 --- a/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc +++ b/lib/utils/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc @@ -1 +1,11 @@ #include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using K = ordered_value_type<0>; +using V = ordered_value_type<1>; + +template bidict merge_disjoint_bidicts(std::set> const &); + +} // namespace FlexFlow diff --git a/lib/utils/test/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc b/lib/utils/test/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc similarity index 72% rename from lib/utils/test/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc rename to lib/utils/test/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc index 0a1babd9f9..8a3371b8d8 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/merge_disjoint_bidicts.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc @@ -1,17 +1,17 @@ -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("merge_disjoint_bidicts") { + TEST_CASE("binary_merge_disjoint_bidicts") { SUBCASE("disjoint keys and values") { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{3, "three"}, {4, "four"}}; - bidict result = merge_disjoint_bidicts(bd1, bd2); + bidict result = binary_merge_disjoint_bidicts(bd1, bd2); bidict correct = { {1, "one"}, {2, "two"}, {3, "three"}, {4, "four"}}; @@ -22,21 +22,21 @@ TEST_SUITE(FF_TEST_SUITE) { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{2, "three"}, {3, "four"}}; - CHECK_THROWS(merge_disjoint_bidicts(bd1, bd2)); + CHECK_THROWS(binary_merge_disjoint_bidicts(bd1, bd2)); } SUBCASE("overlapping key, same associated value") { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{2, "two"}, {3, "three"}}; - CHECK_THROWS(merge_disjoint_bidicts(bd1, bd2)); + CHECK_THROWS(binary_merge_disjoint_bidicts(bd1, bd2)); } SUBCASE("overlapping values") { bidict bd1 = {{1, "one"}, {2, "two"}}; bidict bd2 = {{3, "two"}, {4, "four"}}; - CHECK_THROWS(merge_disjoint_bidicts(bd1, bd2)); + CHECK_THROWS(binary_merge_disjoint_bidicts(bd1, bd2)); } } } From 9ecf61a3581b1e8d7e42c71006bcba2db9b9ba5e Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 12 Jun 2026 20:55:50 -0700 Subject: [PATCH 27/35] Format and fix checks --- .../unstructured_device_mapping.h | 18 --- .../machine_mapping/allowed_machine_views.cc | 4 +- .../machine_mapping_mutation_set.cc | 4 +- .../compiler/machine_mapping/machine_view.cc | 4 +- .../start_invariant_machine_view.cc | 3 +- .../src/compiler/mcmc/mcmc_over_mapped_pcg.cc | 8 +- .../task_graph_simulator/pcg_task_graph.cc | 5 +- .../task_graph_simulator/task_simulator.cc | 6 +- .../unity_algorithm/unity_algorithm.cc | 3 +- .../machine_mapping/allowed_machine_views.cc | 2 +- .../compiler/machine_mapping/machine_view.cc | 47 ++++--- .../start_invariant_machine_view.cc | 33 ++--- .../src/compiler/mcmc/mcmc_over_mapped_pcg.cc | 2 +- .../task_graph_simulator/task_simulator.cc | 21 +-- .../unity_algorithm/graph_optimize_state.cc | 1 - .../unity_algorithm/unity_algorithm.cc | 2 +- .../computation_graph_instance.h | 2 +- .../cost_estimator/local_cost_estimator.h | 2 +- .../local_task_argument_accessor.h | 2 +- .../per_device_op_state_initialization.h | 2 +- .../cost_estimator/local_cost_estimator.cc | 4 +- .../computation_graph_instance.cc | 31 +++-- .../cost_estimator/local_cost_estimator.cc | 34 +++-- .../local_task_argument_accessor.cc | 10 +- ...ce_specific_managed_per_device_ff_handle.h | 3 +- .../realm-execution/device_specific_ptr.h | 3 +- .../realm-execution/fmt/realm_processor.h | 11 +- .../fmt/realm_processor_kind.h | 10 +- .../include/realm-execution/pcg_instance.h | 2 +- .../include/realm-execution/processor_query.h | 3 +- .../include/realm-execution/realm_context.h | 11 +- .../src/realm-execution/address_space.cc | 9 +- ...uted_per_device_op_state_initialization.cc | 4 +- .../realm-execution/fmt/realm_processor.cc | 3 +- .../fmt/realm_processor_kind.cc | 3 +- .../src/realm-execution/processor_query.cc | 3 +- .../src/realm-execution/realm_context.cc | 127 +++++++++--------- .../fmt/realm_processor_kind.cc | 5 +- .../output_expr_to_result_sub_pcg_mapping.cc | 2 +- .../src/substitutions/pcg_pattern_match.cc | 2 +- .../dynamic_graph/dynamic_node_mapping.h | 2 +- .../include/task-spec/global_device_id_t.h | 5 +- .../task-spec/dynamic_graph/copy_insertion.cc | 9 +- .../dynamic_graph/dynamic_node_mapping.cc | 3 +- .../dynamic_graph/shard_expansion.cc | 16 ++- .../src/task-spec/global_device_id_t.cc | 24 ++-- .../test/src/task-spec/device_specific.cc | 6 +- .../dynamic_graph/shard_expansion.cc | 12 +- .../binary_merge_disjoint_bidicts.h | 24 ++-- .../algorithms/merge_disjoint_bidicts.h | 12 +- lib/utils/include/utils/variant.h | 2 +- .../binary_merge_disjoint_bidicts.cc | 5 +- .../include/test/utils/rapidcheck/doctest.h | 2 +- .../non_normal_sp_decomposition.cc | 2 +- 54 files changed, 274 insertions(+), 301 deletions(-) delete mode 100644 lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.h diff --git a/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.h b/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.h deleted file mode 100644 index 8c1333fabc..0000000000 --- a/lib/compiler/include/compiler/machine_mapping/unstructured_device_mapping.h +++ /dev/null @@ -1,18 +0,0 @@ -#ifndef _FLEXFLOW_COMPILER_MACHINE_MAPPING_UNSTRUCTURED_DEVICE_MAPPING_H -#define _FLEXFLOW_COMPILER_MACHINE_MAPPING_UNSTRUCTURED_DEVICE_MAPPING_H - -#include "compiler/machine_mapping/machine_mapping.dtg.h" -#include "compiler/machine_mapping/unstructured_device_mapping.dtg.h" -#include "pcg/machine_compute_specification.dtg.h" -#include "pcg/parallel_computation_graph/parallel_computation_graph.dtg.h" - -namespace FlexFlow { - -UnstructuredDeviceMapping get_unstructured_device_mapping( - MachineMapping const &machine_mapping, - MachineComputeSpecification const &machine_spec, - ParallelComputationGraph const &pcg); - -} // namespace FlexFlow - -#endif diff --git a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc index 9194c2e982..4cd4abb056 100644 --- a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc @@ -91,13 +91,11 @@ static std::unordered_set auto get_candidate_starts = [](MachineComputeResourceSlice const &slice) -> std::unordered_set { - std::unordered_set result; for (nonnegative_int node_idx : nonnegative_range(slice.num_nodes)) { for (nonnegative_int device_idx : nonnegative_range(slice.num_gpus_per_node)) { - result.insert( - MachineSpaceCoordinate{node_idx, device_idx}); + result.insert(MachineSpaceCoordinate{node_idx, device_idx}); } } return result; diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc index d6cdca97d1..41fbce8a78 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc @@ -17,8 +17,8 @@ std::optional for (parallel_layer_guid_t layer : layers) { OperatorTaskSpace task = get_operator_task_space(pcg, layer); std::unordered_set allowed_machine_views = - get_allowed_machine_views( - compute_slice_from_specification(resources), task); + get_allowed_machine_views(compute_slice_from_specification(resources), + task); if (allowed_machine_views.empty()) { return std::nullopt; } diff --git a/lib/compiler/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/src/compiler/machine_mapping/machine_view.cc index 5c38a66901..baf40faf12 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_view.cc @@ -125,8 +125,8 @@ MachineSpaceCoordinate compute_index(machine_view.start.node_idx, inter_dimension_indices); nonnegative_int device_idx = compute_index(machine_view.start.device_idx, intra_dimension_indices); - MachineSpaceCoordinate ms_coord = MachineSpaceCoordinate{ - node_idx, device_idx}; + MachineSpaceCoordinate ms_coord = + MachineSpaceCoordinate{node_idx, device_idx}; return ms_coord; } diff --git a/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc b/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc index 4a2d66acc1..cd9309d886 100644 --- a/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/start_invariant_machine_view.cc @@ -54,8 +54,7 @@ MachineSpaceOffset get_machine_space_offset( StartInvariantMachineView const &start_inv_machine_view, TaskSpaceCoordinate const &coord) { - MachineSpaceCoordinate dummy_start = - MachineSpaceCoordinate{0_n, 0_n}; + MachineSpaceCoordinate dummy_start = MachineSpaceCoordinate{0_n, 0_n}; MachineView mv = machine_view_from_start_invariant(start_inv_machine_view, dummy_start); diff --git a/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc b/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc index 0d2c1e4455..0dd1607d9b 100644 --- a/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc +++ b/lib/compiler/src/compiler/mcmc/mcmc_over_mapped_pcg.cc @@ -21,8 +21,8 @@ SearchResult MCMCOverMappedPCGConfig const &search_config) { MachineComputeSpecification compute_spec = machine_spec.compute_specification; std::vector substitutions = get_substitution_set(compute_spec); - MachineMapping random_mapping = assert_unwrap( - get_random_mapping(pcg, compute_spec)); + MachineMapping random_mapping = + assert_unwrap(get_random_mapping(pcg, compute_spec)); SearchResult starting_state = SearchResult{pcg, random_mapping}; auto sampler = [&](SearchResult mapped_pcg) -> std::optional { @@ -42,8 +42,8 @@ SearchResult mapped_pcg, random_substitution, match); }); } else { - MachineMapping new_machine_mapping = assert_unwrap(get_random_mutation( - mapped_pcg, compute_spec)); + MachineMapping new_machine_mapping = + assert_unwrap(get_random_mutation(mapped_pcg, compute_spec)); return SearchResult{mapped_pcg.pcg, new_machine_mapping}; } }; diff --git a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc index 058f5b72bb..d3666a2451 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc @@ -24,7 +24,8 @@ PCGTaskGraph DiGraph digraph = DiGraph::create(); bidict node_to_task; bidict node_to_layer; - std::unordered_map> node_to_devices; + std::unordered_map> + node_to_devices; for (parallel_layer_guid_t const &layer : get_parallel_layers(pcg)) { MachineView mv = machine_mapping.machine_views.at(layer); @@ -35,7 +36,7 @@ PCGTaskGraph node_to_layer.equate(node, layer); node_to_devices[node] = get_machine_space_coordinates(get_operator_task_space(pcg, layer), - machine_mapping.machine_views.at(layer)); + machine_mapping.machine_views.at(layer)); } for (ParallelComputationGraphEdge const &edge : get_edges(pcg)) { diff --git a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc index 28a6a3efae..843470a107 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc @@ -52,13 +52,15 @@ milliseconds_t task_simulator_estimate_forward_pass_time( } assert(current_task.is_operator()); - auto get_devices = [&](Node const &n) -> std::unordered_set { + auto get_devices = + [&](Node const &n) -> std::unordered_set { return task_graph.node_to_devices.at(n); }; std::unordered_set devices_occupied = set_union(transform(in_progress_tasks, get_devices)); - std::unordered_set required_devices = get_devices(task); + std::unordered_set required_devices = + get_devices(task); return set_intersection(devices_occupied, required_devices).empty(); }; diff --git a/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc b/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc index 240ad0b1c4..dd95f9ea26 100644 --- a/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc +++ b/lib/compiler/src/compiler/unity_algorithm/unity_algorithm.cc @@ -67,8 +67,7 @@ SearchResult graph_optimize(ParallelComputationGraph &pcg, OperatorTaskSpace op_task_space = get_operator_task_space_for_runtime_only_op_cost_estimate_key(key); - return get_allowed_machine_views( - resources, op_task_space); + return get_allowed_machine_views(resources, op_task_space); }, }; diff --git a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc index 6a867b16f3..c465f88c51 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc @@ -1,11 +1,11 @@ #include "compiler/machine_mapping/allowed_machine_views.h" -#include #include "utils/containers/extend.h" #include "utils/containers/range.h" #include "utils/containers/transform.h" #include "utils/containers/unordered_set_of.h" #include "utils/containers/zip.h" #include "utils/fmt/unordered_set.h" +#include #include using namespace FlexFlow; diff --git a/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc index 29fdfff3f6..ce4c9b4b25 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/machine_view.cc @@ -162,8 +162,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 0_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/2_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/2_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -171,8 +171,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/4_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/4_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -180,8 +180,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/2_n, /*device_idx=*/2_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/2_n, /*device_idx=*/2_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -189,8 +189,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/2_n, /*device_idx=*/4_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/2_n, /*device_idx=*/4_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -237,8 +237,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 0_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/0_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/0_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -246,8 +246,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/4_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/4_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -255,8 +255,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/1_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/1_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -264,8 +264,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/5_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/5_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -302,8 +302,7 @@ TEST_SUITE(FF_TEST_SUITE) { }}, }; MachineView mv = MachineView{ - MachineSpaceCoordinate{ - /*node_idx=*/0_n, /*device_idx=*/1_n}, + MachineSpaceCoordinate{/*node_idx=*/0_n, /*device_idx=*/1_n}, {MachineViewDimension{stride_t{1_p}, MachineSpecificationDimension::INTER_NODE}, MachineViewDimension{stride_t{2_p}, @@ -313,8 +312,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n, 0_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/0_n, /*device_idx=*/3_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/0_n, /*device_idx=*/3_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -322,8 +321,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n, 1_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/5_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/5_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); @@ -331,8 +330,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n, 1_n}); - MachineSpaceCoordinate correct = MachineSpaceCoordinate{ - /*node_idx=*/1_n, /*device_idx=*/7_n}; + MachineSpaceCoordinate correct = + MachineSpaceCoordinate{/*node_idx=*/1_n, /*device_idx=*/7_n}; MachineSpaceCoordinate result = get_machine_space_coordinate(task, mv, coord); CHECK(correct == result); diff --git a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc index e3b08e3805..be04954f18 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc @@ -36,8 +36,7 @@ TEST_SUITE(FF_TEST_SUITE) { } TEST_CASE("StartInvariantMachineView - conversions") { - MachineSpaceCoordinate start = - MachineSpaceCoordinate{1_n, 2_n}; + MachineSpaceCoordinate start = MachineSpaceCoordinate{1_n, 2_n}; std::vector dimensions = { MachineViewDimension{stride_t{2_p}, MachineSpecificationDimension::INTER_NODE}, @@ -45,8 +44,7 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpecificationDimension::INTRA_NODE}}; MachineView mv = MachineView{start, dimensions}; - StartInvariantMachineView simv = - StartInvariantMachineView{dimensions}; + StartInvariantMachineView simv = StartInvariantMachineView{dimensions}; SUBCASE("start_invariant_from_machine_view") { StartInvariantMachineView result = start_invariant_from_machine_view(mv); @@ -93,9 +91,9 @@ TEST_SUITE(FF_TEST_SUITE) { 3_ge2, }}, }; - StartInvariantMachineView simv = StartInvariantMachineView{ - {MachineViewDimension{stride_t{2_p}, - MachineSpecificationDimension::INTRA_NODE}}}; + StartInvariantMachineView simv = + StartInvariantMachineView{{MachineViewDimension{ + stride_t{2_p}, MachineSpecificationDimension::INTRA_NODE}}}; MachineComputeSpecification ms = MachineComputeSpecification{ /*num_nodes=*/1_p, /*num_cpus_per_node=*/6_p, @@ -105,8 +103,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_machine_space_offset") { SUBCASE("Task with TaskSpaceCoordinate = (0,)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n}); - MachineSpaceOffset correct = - MachineSpaceOffset{0, 0}; + MachineSpaceOffset correct = MachineSpaceOffset{0, 0}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -114,8 +111,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n}); - MachineSpaceOffset correct = - MachineSpaceOffset{0, 2}; + MachineSpaceOffset correct = MachineSpaceOffset{0, 2}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -123,8 +119,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (2,)") { TaskSpaceCoordinate coord = make_task_space_coordinate({2_n}); - MachineSpaceOffset correct = - MachineSpaceOffset{0, 4}; + MachineSpaceOffset correct = MachineSpaceOffset{0, 4}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -178,8 +173,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_machine_space_offset") { SUBCASE("Task with TaskSpaceCoordinate = (0,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 0_n}); - MachineSpaceOffset correct = - MachineSpaceOffset{0, 0}; + MachineSpaceOffset correct = MachineSpaceOffset{0, 0}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -187,8 +181,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (0,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({0_n, 1_n}); - MachineSpaceOffset correct = - MachineSpaceOffset{0, 2}; + MachineSpaceOffset correct = MachineSpaceOffset{0, 2}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -196,8 +189,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,0)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 0_n}); - MachineSpaceOffset correct = - MachineSpaceOffset{1, 0}; + MachineSpaceOffset correct = MachineSpaceOffset{1, 0}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); @@ -205,8 +197,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("Task with TaskSpaceCoordinate = (1,1)") { TaskSpaceCoordinate coord = make_task_space_coordinate({1_n, 1_n}); - MachineSpaceOffset correct = - MachineSpaceOffset{1, 2}; + MachineSpaceOffset correct = MachineSpaceOffset{1, 2}; MachineSpaceOffset result = get_machine_space_offset(task, simv, coord); CHECK(correct == result); diff --git a/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc b/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc index 0fe4277ab6..f006c21ba3 100644 --- a/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc +++ b/lib/compiler/test/src/compiler/mcmc/mcmc_over_mapped_pcg.cc @@ -1,6 +1,5 @@ #include "compiler/mcmc/mcmc_over_mapped_pcg.h" #include "compiler/task_graph_simulator/task_simulator.h" -#include #include "internal/runtime_only_cost_estimator_for_test.h" #include "op-attrs/parallel_tensor_dims.h" #include "op-attrs/parallel_tensor_shape.dtg.h" @@ -10,6 +9,7 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph_builder.h" #include "pcg/pcg_from_computation_graph.h" #include "utils/integer_conversions.h" +#include using namespace FlexFlow; diff --git a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc index a7f8886846..b46619ab20 100644 --- a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc @@ -65,10 +65,8 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector dims = {}; ParallelComputationGraph pcg = b.pcg; - MachineView mv1 = - MachineView{MachineSpaceCoordinate{0_n, 0_n}, dims}; - MachineView mv2 = - MachineView{MachineSpaceCoordinate{0_n, 1_n}, dims}; + MachineView mv1 = MachineView{MachineSpaceCoordinate{0_n, 0_n}, dims}; + MachineView mv2 = MachineView{MachineSpaceCoordinate{0_n, 1_n}, dims}; MachineMapping device_mapping = MachineMapping{{ {layer0, mv1}, @@ -146,14 +144,10 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector dims = {}; SUBCASE("all different devices") { - MachineView mv0 = MachineView{ - MachineSpaceCoordinate{0_n, 0_n}, dims}; - MachineView mv1 = MachineView{ - MachineSpaceCoordinate{0_n, 1_n}, dims}; - MachineView mv2 = MachineView{ - MachineSpaceCoordinate{1_n, 0_n}, dims}; - MachineView mv3 = MachineView{ - MachineSpaceCoordinate{1_n, 1_n}, dims}; + MachineView mv0 = MachineView{MachineSpaceCoordinate{0_n, 0_n}, dims}; + MachineView mv1 = MachineView{MachineSpaceCoordinate{0_n, 1_n}, dims}; + MachineView mv2 = MachineView{MachineSpaceCoordinate{1_n, 0_n}, dims}; + MachineView mv3 = MachineView{MachineSpaceCoordinate{1_n, 1_n}, dims}; MachineMapping device_mapping = MachineMapping{{ {layer0, mv0}, @@ -206,8 +200,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("all the same device") { - MachineView mv = MachineView{ - MachineSpaceCoordinate{0_n, 0_n}, dims}; + MachineView mv = MachineView{MachineSpaceCoordinate{0_n, 0_n}, dims}; MachineMapping device_mapping = MachineMapping{{ {layer0, mv}, {layer1, mv}, diff --git a/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc b/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc index cc7ca9425f..a9ee3f34e8 100644 --- a/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc +++ b/lib/compiler/test/src/compiler/unity_algorithm/graph_optimize_state.cc @@ -3,7 +3,6 @@ #include "compiler/machine_mapping/machine_mapping.h" #include "compiler/machine_mapping/machine_view.dtg.h" #include "compiler/machine_mapping/machine_view.h" -#include #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_builder.h" #include "test/utils/doctest/check_without_stringify.h" diff --git a/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc b/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc index 0d4123d381..65a1860034 100644 --- a/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc +++ b/lib/compiler/test/src/compiler/unity_algorithm/unity_algorithm.cc @@ -1,6 +1,5 @@ #include "compiler/unity_algorithm/unity_algorithm.h" #include "compiler/cost_estimator/runtime_only_cost_estimator_from_cost_estimator.h" -#include #include "internal/cost_estimator_for_test.h" #include "op-attrs/parallel_tensor_dims.h" #include "op-attrs/parallel_tensor_shape.dtg.h" @@ -10,6 +9,7 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph_builder.h" #include "pcg/pcg_from_computation_graph.h" #include "utils/integer_conversions.h" +#include using namespace FlexFlow; diff --git a/lib/local-execution/include/local-execution/computation_graph_instance.h b/lib/local-execution/include/local-execution/computation_graph_instance.h index e6c7ded9c1..23467f8da5 100644 --- a/lib/local-execution/include/local-execution/computation_graph_instance.h +++ b/lib/local-execution/include/local-execution/computation_graph_instance.h @@ -7,12 +7,12 @@ #include "kernels/profiling_settings.dtg.h" #include "local-execution/loss_config.dtg.h" #include "pcg/computation_graph.dtg.h" -#include "task-spec/global_device_id_t.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "utils/units/milliseconds_t.h" #include #include diff --git a/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h b/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h index 15aa267eec..8c4af368b0 100644 --- a/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h +++ b/lib/local-execution/include/local-execution/cost_estimator/local_cost_estimator.h @@ -5,8 +5,8 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" #include "kernels/profiling_settings.dtg.h" -#include "task-spec/global_device_id_t.dtg.h" #include "pcg/machine_interconnect_specification.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" namespace FlexFlow { diff --git a/lib/local-execution/include/local-execution/local_task_argument_accessor.h b/lib/local-execution/include/local-execution/local_task_argument_accessor.h index 38b1023cd2..e2d093967c 100644 --- a/lib/local-execution/include/local-execution/local_task_argument_accessor.h +++ b/lib/local-execution/include/local-execution/local_task_argument_accessor.h @@ -2,8 +2,8 @@ #define _FLEXFLOW_LIB_LOCAL_EXECUTION_INCLUDE_LOCAL_EXECUTION_LOCAL_TASK_ARGUMENT_ACCESSOR_H #include "kernels/accessor.h" -#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/task_argument_accessor/itask_argument_accessor.h" #include "task-spec/task_argument_accessor/task_tensor_parameter.dtg.h" #include diff --git a/lib/local-execution/include/local-execution/per_device_op_state_initialization.h b/lib/local-execution/include/local-execution/per_device_op_state_initialization.h index 90a1ffc688..4c546992e4 100644 --- a/lib/local-execution/include/local-execution/per_device_op_state_initialization.h +++ b/lib/local-execution/include/local-execution/per_device_op_state_initialization.h @@ -4,9 +4,9 @@ #include "kernels/allocation.h" #include "kernels/device_handle_t.dtg.h" #include "kernels/profiling_settings.dtg.h" -#include "task-spec/global_device_id_t.dtg.h" #include "pcg/optimizer_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" namespace FlexFlow { diff --git a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc index f151ab1d17..6c74264dff 100644 --- a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -103,8 +103,8 @@ OpCostMetrics LocalCostEstimator::estimate_cost( // allocate memory std::shared_ptr tracked_allocator_ptr = - std::make_shared(create_local_allocator_for_device_type( - this->device_idx.device_type)); + std::make_shared( + create_local_allocator_for_device_type(this->device_idx.device_type)); layer_guid_t layer_guid = layer_guid_t{Node{0}}; diff --git a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc index 69091b5d2e..ae8365b127 100644 --- a/lib/local-execution/test/src/local-execution/computation_graph_instance.cc +++ b/lib/local-execution/test/src/local-execution/computation_graph_instance.cc @@ -140,14 +140,13 @@ TEST_SUITE(FF_TEST_SUITE) { /*weight_decay=*/0.001}}; device_handle_t ff_handle = cpu_make_device_handle_t(); global_device_id_t global_device_id = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - }, - /*device_type=*/DeviceType::CPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::CPU, }; - std::unordered_map input_tensors; ComputationGraphInstance computation_graph_instance = @@ -313,11 +312,11 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { }, }; global_device_id_t device_idx = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - }, - /*device_type=*/DeviceType::GPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, }; device_handle_t ff_handle = gpu_make_device_handle_t(managed_handle.raw_handle()); @@ -434,11 +433,11 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { }; global_device_id_t device_idx = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - }, - /*device_type=*/DeviceType::GPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, }; device_handle_t ff_handle = diff --git a/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc index 8d1506b5fc..3d1b691a08 100644 --- a/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/test/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -19,11 +19,11 @@ TEST_SUITE(FF_TEST_SUITE) { Allocator allocator = create_local_cpu_memory_allocator(); device_handle_t ff_handle = cpu_make_device_handle_t(); global_device_id_t device_idx = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - }, - /*device_type=*/DeviceType::CPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::CPU, }; OptimizerAttrs optimizer_attrs = OptimizerAttrs{ @@ -69,10 +69,9 @@ TEST_SUITE(FF_TEST_SUITE) { /*output_shapes=*/{{TensorSlotName::OUTPUT, output_shape}}, /*optimizer_attrs=*/optimizer_attrs, /*machine_view=*/ - make_1d_machine_view( - MachineSpaceCoordinate{0_n, 0_n}, - MachineSpecificationDimension::INTRA_NODE, - stride_t{1_p}), + make_1d_machine_view(MachineSpaceCoordinate{0_n, 0_n}, + MachineSpecificationDimension::INTRA_NODE, + stride_t{1_p}), }; OpCostMetrics result = cost_estimator.estimate_cost(op_cost_estimate_key); @@ -94,11 +93,11 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { Allocator allocator = create_local_cuda_memory_allocator(); global_device_id_t device_idx = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - }, - /*device_type=*/DeviceType::GPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, }; device_handle_t ff_handle = @@ -168,10 +167,9 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { /*output_shapes=*/{{TensorSlotName::OUTPUT, output_shape}}, /*optimizer_attrs=*/optimizer_attrs, /*machine_view=*/ - make_1d_machine_view( - MachineSpaceCoordinate{0_n, 0_n}, - MachineSpecificationDimension::INTRA_NODE, - stride_t{1_p}), + make_1d_machine_view(MachineSpaceCoordinate{0_n, 0_n}, + MachineSpecificationDimension::INTRA_NODE, + stride_t{1_p}), }; OpCostMetrics result = cost_estimator.estimate_cost(op_cost_estimate_key); diff --git a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc index dba4464611..8800711fef 100644 --- a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc @@ -52,11 +52,11 @@ TEST_SUITE(FF_TEST_SUITE) { }; global_device_id_t device_idx = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - }, - /*device_type=*/DeviceType::CPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::CPU, }; LocalTaskArgumentAccessor acc = LocalTaskArgumentAccessor{ diff --git a/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h b/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h index a7e265f777..cc08f963e7 100644 --- a/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h +++ b/lib/realm-execution/include/realm-execution/device_specific_managed_per_device_ff_handle.h @@ -11,7 +11,8 @@ namespace FlexFlow { DeviceSpecificPtr make_device_specific_managed_ff_handle( - global_device_id_t const &, std::optional const &); + global_device_id_t const &, + std::optional const &); device_handle_t device_handle_t_from_device_specific_managed_ff_handle( DeviceSpecificPtr const &, global_device_id_t); diff --git a/lib/realm-execution/include/realm-execution/device_specific_ptr.h b/lib/realm-execution/include/realm-execution/device_specific_ptr.h index f1126f64c1..fe8da0f2e6 100644 --- a/lib/realm-execution/include/realm-execution/device_specific_ptr.h +++ b/lib/realm-execution/include/realm-execution/device_specific_ptr.h @@ -32,7 +32,8 @@ template struct DeviceSpecificPtr { public: DeviceSpecificPtr() = delete; - explicit DeviceSpecificPtr(global_device_id_t device_idx, std::optional ptr) + explicit DeviceSpecificPtr(global_device_id_t device_idx, + std::optional ptr) : device_idx(device_idx), ptr(ptr) {} std::optional get(global_device_id_t device_idx) const { diff --git a/lib/realm-execution/include/realm-execution/fmt/realm_processor.h b/lib/realm-execution/include/realm-execution/fmt/realm_processor.h index e2fadb6b18..12c8787b2f 100644 --- a/lib/realm-execution/include/realm-execution/fmt/realm_processor.h +++ b/lib/realm-execution/include/realm-execution/fmt/realm_processor.h @@ -8,10 +8,10 @@ namespace fmt { template -struct formatter< - ::FlexFlow::Realm::Processor, - Char, - std::enable_if_t::value>> +struct formatter<::FlexFlow::Realm::Processor, + Char, + std::enable_if_t::value>> : formatter<::std::string> { template auto format(::FlexFlow::Realm::Processor const &m, FormatContext &ctx) @@ -27,7 +27,8 @@ struct formatter< namespace FlexFlow { -std::ostream &operator<<(std::ostream &s, ::FlexFlow::Realm::Processor const &m); +std::ostream &operator<<(std::ostream &s, + ::FlexFlow::Realm::Processor const &m); } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h b/lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h index d75ab067cc..e591eb9adf 100644 --- a/lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h +++ b/lib/realm-execution/include/realm-execution/fmt/realm_processor_kind.h @@ -3,16 +3,13 @@ #include "realm-execution/realm.h" #include -#include #include +#include namespace fmt { template -struct formatter< - ::FlexFlow::Realm::Processor::Kind, - Char - > +struct formatter<::FlexFlow::Realm::Processor::Kind, Char> : formatter<::std::string> { template auto format(::FlexFlow::Realm::Processor::Kind const &m, FormatContext &ctx) @@ -76,7 +73,8 @@ struct formatter< namespace FlexFlow { -std::ostream &operator<<(std::ostream &, ::FlexFlow::Realm::Processor::Kind const &); +std::ostream &operator<<(std::ostream &, + ::FlexFlow::Realm::Processor::Kind const &); } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/pcg_instance.h b/lib/realm-execution/include/realm-execution/pcg_instance.h index 6c82e4c61c..d9d49a780d 100644 --- a/lib/realm-execution/include/realm-execution/pcg_instance.h +++ b/lib/realm-execution/include/realm-execution/pcg_instance.h @@ -11,10 +11,10 @@ #include "realm-execution/per_device_op_state_backing.dtg.h" #include "realm-execution/realm_context.h" #include "realm-execution/tensor_instance_backing.dtg.h" -#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_accessor.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" #include "utils/units/milliseconds_t.h" #include diff --git a/lib/realm-execution/include/realm-execution/processor_query.h b/lib/realm-execution/include/realm-execution/processor_query.h index 693d488390..acc3d79440 100644 --- a/lib/realm-execution/include/realm-execution/processor_query.h +++ b/lib/realm-execution/include/realm-execution/processor_query.h @@ -5,7 +5,8 @@ namespace FlexFlow { -std::set processor_set_from_query(Realm::Machine::ProcessorQuery const &); +std::set + processor_set_from_query(Realm::Machine::ProcessorQuery const &); } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/realm_context.h b/lib/realm-execution/include/realm-execution/realm_context.h index e06ac522d7..9ed8352f97 100644 --- a/lib/realm-execution/include/realm-execution/realm_context.h +++ b/lib/realm-execution/include/realm-execution/realm_context.h @@ -10,9 +10,9 @@ #include "realm-execution/realm.h" #include "realm-execution/tasks/task_id_t.dtg.h" #include "task-spec/global_device_id_t.dtg.h" +#include "task-spec/local_device_id_t.dtg.h" #include #include -#include "task-spec/local_device_id_t.dtg.h" namespace FlexFlow { @@ -35,7 +35,8 @@ struct RealmContext { ///\{ Realm::Processor processor_from_global_device_id(global_device_id_t const &); global_device_id_t global_device_id_from_processor(Realm::Processor); - Realm::Processor processor_from_local_device_id(local_device_id_t const &) const; + Realm::Processor + processor_from_local_device_id(local_device_id_t const &) const; local_device_id_t local_device_id_from_processor(Realm::Processor) const; static Realm::Memory get_nearest_memory(Realm::Processor); @@ -113,7 +114,8 @@ struct RealmContext { */ Realm::Runtime get_runtime(); - bidict const &get_global_machine_topology(); + bidict const & + get_global_machine_topology(); private: Realm::Runtime runtime; @@ -121,7 +123,8 @@ struct RealmContext { Allocator allocator; std::vector outstanding_events; bidict local_machine_topology; - std::optional> cached_global_machine_topology = std::nullopt; + std::optional> + cached_global_machine_topology = std::nullopt; }; } // namespace FlexFlow diff --git a/lib/realm-execution/src/realm-execution/address_space.cc b/lib/realm-execution/src/realm-execution/address_space.cc index 2886f16312..b1f229cce8 100644 --- a/lib/realm-execution/src/realm-execution/address_space.cc +++ b/lib/realm-execution/src/realm-execution/address_space.cc @@ -2,11 +2,12 @@ namespace FlexFlow { -node_idx_t node_idx_from_realm_address_space(Realm::AddressSpace address_space) { +node_idx_t + node_idx_from_realm_address_space(Realm::AddressSpace address_space) { return node_idx_t{ - nonnegative_int{ - static_cast(address_space), - }, + nonnegative_int{ + static_cast(address_space), + }, }; } diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index a4a4240ff3..7c47701cc9 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -39,8 +39,8 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( invocation); DeviceSpecificPtr *device_state_ptr = - new DeviceSpecificPtr{ctx.get_current_global_device_id(), - std::nullopt}; + new DeviceSpecificPtr{ + ctx.get_current_global_device_id(), std::nullopt}; std::optional completion_event = spawn_per_device_op_state_init_task(ctx, diff --git a/lib/realm-execution/src/realm-execution/fmt/realm_processor.cc b/lib/realm-execution/src/realm-execution/fmt/realm_processor.cc index a4afe3cd1b..efbe36d5cb 100644 --- a/lib/realm-execution/src/realm-execution/fmt/realm_processor.cc +++ b/lib/realm-execution/src/realm-execution/fmt/realm_processor.cc @@ -2,7 +2,8 @@ namespace FlexFlow { -std::ostream &operator<<(std::ostream &s, ::FlexFlow::Realm::Processor const &m) { +std::ostream &operator<<(std::ostream &s, + ::FlexFlow::Realm::Processor const &m) { return s << fmt::to_string(m); } diff --git a/lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc b/lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc index 2552ec7ef5..b631ef5027 100644 --- a/lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc +++ b/lib/realm-execution/src/realm-execution/fmt/realm_processor_kind.cc @@ -2,7 +2,8 @@ namespace FlexFlow { -std::ostream &operator<<(std::ostream &s, ::FlexFlow::Realm::Processor::Kind const &k) { +std::ostream &operator<<(std::ostream &s, + ::FlexFlow::Realm::Processor::Kind const &k) { return (s << fmt::to_string(k)); } diff --git a/lib/realm-execution/src/realm-execution/processor_query.cc b/lib/realm-execution/src/realm-execution/processor_query.cc index b5dd3e12da..f122552e52 100644 --- a/lib/realm-execution/src/realm-execution/processor_query.cc +++ b/lib/realm-execution/src/realm-execution/processor_query.cc @@ -2,7 +2,8 @@ namespace FlexFlow { -std::set processor_set_from_query(Realm::Machine::ProcessorQuery const &pq) { +std::set + processor_set_from_query(Realm::Machine::ProcessorQuery const &pq) { std::set result; for (Realm::Processor p : pq) { result.insert(p); diff --git a/lib/realm-execution/src/realm-execution/realm_context.cc b/lib/realm-execution/src/realm-execution/realm_context.cc index 07d60aa73c..869e71d616 100644 --- a/lib/realm-execution/src/realm-execution/realm_context.cc +++ b/lib/realm-execution/src/realm-execution/realm_context.cc @@ -4,67 +4,62 @@ #include "op-attrs/datatype.h" #include "op-attrs/parallel_tensor_shape.h" #include "op-attrs/tensor_dims.dtg.h" +#include "realm-execution/address_space.h" +#include "realm-execution/fmt/realm_processor.h" +#include "realm-execution/fmt/realm_processor_kind.h" #include "realm-execution/processor_kind.h" +#include "realm-execution/processor_query.h" #include "realm-execution/realm_allocator.h" #include "realm-execution/tasks/task_id_t.dtg.h" #include "realm-execution/tasks/task_id_t.h" +#include "task-spec/global_device_id_t.h" +#include "utils/bidict/algorithms/bidict_from_enumerating.h" +#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" +#include "utils/bidict/algorithms/transform_values.h" +#include "utils/containers/are_all_same.h" #include "utils/containers/contains_key.h" +#include "utils/containers/group_by.h" +#include "utils/containers/set_of.h" #include "utils/containers/transform.h" #include "utils/exception.h" #include "utils/nonnegative_int/nonnegative_int.h" #include "utils/one_to_many/one_to_many.h" #include "utils/optional.h" #include "utils/positive_int/positive_int.h" -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" -#include "realm-execution/address_space.h" -#include "utils/bidict/algorithms/bidict_from_enumerating.h" -#include "utils/bidict/algorithms/transform_values.h" -#include "utils/containers/group_by.h" -#include "task-spec/global_device_id_t.h" -#include "utils/containers/are_all_same.h" -#include "utils/containers/set_of.h" -#include "realm-execution/processor_query.h" -#include "realm-execution/fmt/realm_processor.h" -#include "realm-execution/fmt/realm_processor_kind.h" namespace FlexFlow { - -bidict - build_local_machine_topology(std::set const &local_procs) -{ +bidict build_local_machine_topology( + std::set const &local_procs) { { bool procs_are_local = are_all_same( - transform(local_procs, - [&](Realm::Processor p) -> Realm::AddressSpace { - return p.address_space(); - })); + transform(local_procs, [&](Realm::Processor p) -> Realm::AddressSpace { + return p.address_space(); + })); ASSERT(procs_are_local); } - OneToMany by_proc_kind = - group_by(local_procs, - [](Realm::Processor p) -> Realm::Processor::Kind { - return p.kind(); - }); + OneToMany by_proc_kind = + group_by(local_procs, [](Realm::Processor p) -> Realm::Processor::Kind { + return p.kind(); + }); auto local_machine_topology_for_proc_kind = [&](Realm::Processor::Kind k) - -> bidict - { + -> bidict { if (!contains(by_proc_kind.left_values(), k)) { return {}; } - bidict enumerated = - bidict_from_enumerating(set_of(by_proc_kind.at_l(k).unwrap_as_unordered_set())).reversed(); + bidict enumerated = + bidict_from_enumerating( + set_of(by_proc_kind.at_l(k).unwrap_as_unordered_set())) + .reversed(); - bidict result = - transform_values( - enumerated, - [&](nonnegative_int idx) -> local_device_id_t { + bidict result = transform_values( + enumerated, [&](nonnegative_int idx) -> local_device_id_t { return local_device_id_t{ - /*idx=*/device_in_node_idx_t{idx}, - /*device_type=*/device_type_from_processor_kind(k), + /*idx=*/device_in_node_idx_t{idx}, + /*device_type=*/device_type_from_processor_kind(k), }; }); @@ -72,48 +67,48 @@ bidict }; return binary_merge_disjoint_bidicts( - local_machine_topology_for_proc_kind(Realm::Processor::Kind::LOC_PROC), - local_machine_topology_for_proc_kind(Realm::Processor::Kind::TOC_PROC)); + local_machine_topology_for_proc_kind(Realm::Processor::Kind::LOC_PROC), + local_machine_topology_for_proc_kind(Realm::Processor::Kind::TOC_PROC)); } static bidict - build_global_machine_topology(std::set const &global_procs) -{ - OneToMany by_node_idx = - group_by(global_procs, - [](Realm::Processor p) -> node_idx_t { - return node_idx_from_realm_address_space(p.address_space()); - }); - - auto build_global_machine_topology_for_node = [&](node_idx_t const &node_idx) - -> bidict - { - std::set procs_for_node = set_of(by_node_idx.at_l(node_idx).unwrap_as_unordered_set()); + build_global_machine_topology( + std::set const &global_procs) { + OneToMany by_node_idx = + group_by(global_procs, [](Realm::Processor p) -> node_idx_t { + return node_idx_from_realm_address_space(p.address_space()); + }); + + auto build_global_machine_topology_for_node = [&](node_idx_t const &node_idx) + -> bidict { + std::set procs_for_node = + set_of(by_node_idx.at_l(node_idx).unwrap_as_unordered_set()); - bidict - local_topology_for_node = build_local_machine_topology(procs_for_node); + bidict local_topology_for_node = + build_local_machine_topology(procs_for_node); return transform_values( - local_topology_for_node, - [&](local_device_id_t const &local_device_id) -> global_device_id_t { - return global_device_id_from_local(local_device_id, node_idx); - }); + local_topology_for_node, + [&](local_device_id_t const &local_device_id) -> global_device_id_t { + return global_device_id_from_local(local_device_id, node_idx); + }); }; return merge_disjoint_bidicts( - transform( - set_of(by_node_idx.left_values()), - build_global_machine_topology_for_node)); + transform(set_of(by_node_idx.left_values()), + build_global_machine_topology_for_node)); } -static bidict discover_local_machine_topology(Realm::Processor local_processor) { +static bidict + discover_local_machine_topology(Realm::Processor local_processor) { Realm::Machine::ProcessorQuery pq(Realm::Machine::get_machine()); pq.same_address_space_as(local_processor); return build_local_machine_topology(processor_set_from_query(pq)); } -static bidict discover_global_machine_topology() { +static bidict + discover_global_machine_topology() { Realm::Machine::ProcessorQuery pq(Realm::Machine::get_machine()); return build_global_machine_topology(processor_set_from_query(pq)); @@ -123,8 +118,7 @@ RealmContext::RealmContext(Realm::Processor processor) : processor(processor), allocator(get_realm_allocator( processor, RealmContext::get_nearest_memory(processor))), - local_machine_topology(discover_local_machine_topology(processor)) -{ } + local_machine_topology(discover_local_machine_topology(processor)) {} RealmContext::~RealmContext() { if (!this->outstanding_events.empty()) { @@ -139,8 +133,8 @@ Realm::Processor RealmContext::processor_from_global_device_id( return this->get_global_machine_topology().at_r(global_device_id); } -global_device_id_t RealmContext::global_device_id_from_processor( - Realm::Processor processor) { +global_device_id_t + RealmContext::global_device_id_from_processor(Realm::Processor processor) { return this->get_global_machine_topology().at_l(processor); } @@ -182,8 +176,8 @@ global_device_id_t RealmContext::get_current_global_device_id() const { Realm::Processor proc = this->get_current_processor(); return global_device_id_from_local( - this->local_device_id_from_processor(proc), - node_idx_from_realm_address_space(proc.address_space())); + this->local_device_id_from_processor(proc), + node_idx_from_realm_address_space(proc.address_space())); } Realm::Event @@ -397,7 +391,8 @@ Realm::Event RealmContext::merge_outstanding_events() { return result; } -bidict const &RealmContext::get_global_machine_topology() { +bidict const & + RealmContext::get_global_machine_topology() { if (!this->cached_global_machine_topology.has_value()) { this->cached_global_machine_topology = discover_global_machine_topology(); } diff --git a/lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc b/lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc index 75430da876..3897be92ff 100644 --- a/lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc +++ b/lib/realm-execution/test/src/realm-execution/fmt/realm_processor_kind.cc @@ -1,11 +1,12 @@ -#include #include "realm-execution/fmt/realm_processor_kind.h" +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("fmt::to_string(Realm::Processor::Kind)") { - std::string result = fmt::to_string(::FlexFlow::Realm::Processor::Kind::TOC_PROC); + std::string result = + fmt::to_string(::FlexFlow::Realm::Processor::Kind::TOC_PROC); std::string correct = ""; CHECK(result == correct); diff --git a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc index 7dc6a2cc5e..4374a951f8 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/output_expr_to_result_sub_pcg_mapping.cc @@ -2,9 +2,9 @@ #include "substitutions/output_graph/output_graph_expr.h" #include "substitutions/sub_parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_pairs.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include "utils/containers/values.h" #include "utils/containers/zip_values_strict.h" -#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" namespace FlexFlow { diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index 46d284c4a0..85a0493e33 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -4,13 +4,13 @@ #include "substitutions/unlabelled/unlabelled_graph_pattern.h" #include "utils/bidict/algorithms/bidict_from_keys_and_values.h" #include "utils/bidict/algorithms/bidict_from_map.h" +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include "utils/bidict/algorithms/exhaustive_relational_join.h" #include "utils/bidict/algorithms/transform_values.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/map_values.h" #include "utils/containers/values.h" #include "utils/containers/zip.h" -#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" namespace FlexFlow { diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h index c5a4976db7..3c5b27c344 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_MAPPING_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_MAPPING_H -#include "task-spec/global_device_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_mapping.dtg.h" +#include "task-spec/global_device_id_t.dtg.h" namespace FlexFlow { diff --git a/lib/task-spec/include/task-spec/global_device_id_t.h b/lib/task-spec/include/task-spec/global_device_id_t.h index 2f48046ff2..428f3f7a3c 100644 --- a/lib/task-spec/include/task-spec/global_device_id_t.h +++ b/lib/task-spec/include/task-spec/global_device_id_t.h @@ -1,13 +1,14 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_GLOBAL_DEVICE_ID_T_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_GLOBAL_DEVICE_ID_T_H +#include "pcg/node_idx_t.dtg.h" #include "task-spec/global_device_id_t.dtg.h" #include "task-spec/local_device_id_t.dtg.h" -#include "pcg/node_idx_t.dtg.h" namespace FlexFlow { -global_device_id_t global_device_id_from_local(local_device_id_t const &, node_idx_t); +global_device_id_t global_device_id_from_local(local_device_id_t const &, + node_idx_t); local_device_id_t local_device_id_from_global(global_device_id_t const &); } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 3487b2979e..2c00d7412a 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -61,10 +61,12 @@ bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { static std::pair filter_mapping_to_avoid_degenerate_copies(DynamicValueAttrs const &input, DynamicValueAttrs const &output) { - std::unordered_set> + std::unordered_set< + std::pair> input_mapping = unordered_set_of(assert_unwrap(input.mapping).raw); - std::unordered_set> + std::unordered_set< + std::pair> output_mapping = unordered_set_of(assert_unwrap(output.mapping).raw); // Exclude the point shared between the input and output mappings, because @@ -214,8 +216,7 @@ std::unordered_map [&](InternalDynamicSlotSite const &s) -> ParallelTensorMapping { return ParallelTensorMapping{ dynamic_node_mapping_bindings_for_slot_name( - assert_unwrap(i.node_attrs.mapping), - s.slot_name.slot_name), + assert_unwrap(i.node_attrs.mapping), s.slot_name.slot_name), }; }); }; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc index 4e21da9c2a..b2a6e71af8 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc @@ -11,7 +11,8 @@ bidict get_tensor_bindings_for_slot_name(mapping.op_task_group, slot_name); return transform_values( - coord_bindings, [&](MachineSpaceCoordinate const &coord) -> global_device_id_t { + coord_bindings, + [&](MachineSpaceCoordinate const &coord) -> global_device_id_t { return global_device_id_t{coord, mapping.device_type}; }); } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index 15f78944c8..887c53998f 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -44,7 +44,8 @@ bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { static bidict restrict_tensor_mapping_keys_to_coord( - bidict const &mapping, + bidict const + &mapping, ParallelTensorSpaceCoordinate const ¶llel_tensor_coord) { return filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { return p == parallel_tensor_coord; @@ -128,13 +129,14 @@ std::unordered_set std::unordered_set shard_machine_coords = target_devices_of_dynamic_node_mapping(mapping); - return transform( - shard_machine_coords, [&](global_device_id_t const &c) -> DynamicNodeInvocation { - OperatorAtomicTaskShardBinding slot_bindings = - mapping.op_task_group.get_shard_bindings().at_l(c.coord); + return transform(shard_machine_coords, + [&](global_device_id_t const &c) -> DynamicNodeInvocation { + OperatorAtomicTaskShardBinding slot_bindings = + mapping.op_task_group.get_shard_bindings().at_l( + c.coord); - return shard_invocation_for_binding(i, c, slot_bindings); - }); + return shard_invocation_for_binding(i, c, slot_bindings); + }); } DynamicOpenDataflowGraph diff --git a/lib/task-spec/src/task-spec/global_device_id_t.cc b/lib/task-spec/src/task-spec/global_device_id_t.cc index e32a5f5f8b..af0a52293a 100644 --- a/lib/task-spec/src/task-spec/global_device_id_t.cc +++ b/lib/task-spec/src/task-spec/global_device_id_t.cc @@ -2,23 +2,23 @@ namespace FlexFlow { -global_device_id_t global_device_id_from_local( - local_device_id_t const &local_device_id, - node_idx_t node_idx) -{ +global_device_id_t + global_device_id_from_local(local_device_id_t const &local_device_id, + node_idx_t node_idx) { return global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/node_idx.raw, - /*device_idx=*/local_device_id.idx.raw, - }, - /*device_type=*/local_device_id.device_type, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/node_idx.raw, + /*device_idx=*/local_device_id.idx.raw, + }, + /*device_type=*/local_device_id.device_type, }; } -local_device_id_t local_device_id_from_global(global_device_id_t const &global_device_id) { +local_device_id_t + local_device_id_from_global(global_device_id_t const &global_device_id) { return local_device_id_t{ - /*idx=*/device_in_node_idx_t{global_device_id.coord.device_idx}, - /*device_type=*/global_device_id.device_type, + /*idx=*/device_in_node_idx_t{global_device_id.coord.device_idx}, + /*device_type=*/global_device_id.device_type, }; } diff --git a/lib/task-spec/test/src/task-spec/device_specific.cc b/lib/task-spec/test/src/task-spec/device_specific.cc index 983ba4943d..2d4082bb1a 100644 --- a/lib/task-spec/test/src/task-spec/device_specific.cc +++ b/lib/task-spec/test/src/task-spec/device_specific.cc @@ -7,12 +7,14 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("DeviceSpecific") { DeviceSpecific device_specific1 = DeviceSpecific::create( - global_device_id_t{MachineSpaceCoordinate{0_n, 0_n}, DeviceType::GPU}, + global_device_id_t{MachineSpaceCoordinate{0_n, 0_n}, + DeviceType::GPU}, "hello world"); DeviceSpecific device_specific2 = DeviceSpecific::create( - global_device_id_t{MachineSpaceCoordinate{0_n, 1_n}, DeviceType::GPU}, + global_device_id_t{MachineSpaceCoordinate{0_n, 1_n}, + DeviceType::GPU}, "hello world"); std::string result1 = fmt::to_string(device_specific1); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index cc51c838b9..0c50416e7c 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -44,14 +44,16 @@ TEST_SUITE(FF_TEST_SUITE) { }; DeviceType device_type = DeviceType::GPU; - auto mk_device_id = [&](MachineSpaceCoordinate const &c) -> global_device_id_t { + auto mk_device_id = + [&](MachineSpaceCoordinate const &c) -> global_device_id_t { return global_device_id_t{c, device_type}; }; auto mk_value = [&](size_t src_node_id, TensorSlotName src_slot_name, - bidict tensor_binding, + bidict + tensor_binding, std::optional const &shard_coord) -> DynamicValueAttrs { if (shard_coord.has_value()) { @@ -155,9 +157,9 @@ TEST_SUITE(FF_TEST_SUITE) { TensorSlotName use_slot_name, std::optional const &shard_coord) -> DynamicValueAttrs { - bidict tensor_binding = - dynamic_node_mapping_bindings_for_slot_name(node_mapping, - use_slot_name); + bidict + tensor_binding = dynamic_node_mapping_bindings_for_slot_name( + node_mapping, use_slot_name); return mk_value( src_node_id, src_slot_name, tensor_binding, shard_coord); }; diff --git a/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h b/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h index 607cb08229..bcc59c1e6b 100644 --- a/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h +++ b/lib/utils/include/utils/bidict/algorithms/binary_merge_disjoint_bidicts.h @@ -1,26 +1,26 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BINARY_MERGE_DISJOINT_BIDICTS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BINARY_MERGE_DISJOINT_BIDICTS_H +#include "utils/bidict/algorithms/left_entries.h" +#include "utils/bidict/algorithms/right_entries.h" #include "utils/bidict/bidict.h" -#include #include "utils/containers/are_disjoint.h" -#include "utils/bidict/algorithms/right_entries.h" -#include "utils/bidict/algorithms/left_entries.h" +#include namespace FlexFlow { template bidict binary_merge_disjoint_bidicts(bidict const &lhs, bidict const &rhs) { - ASSERT( - are_disjoint(left_entries(lhs), left_entries(rhs)), - "Left entries of {} and {} are non-disjoint", lhs, rhs - ); - - ASSERT( - are_disjoint(right_entries(lhs), right_entries(rhs)), - "Right entries of {} and {} are non-disjoint", lhs, rhs - ); + ASSERT(are_disjoint(left_entries(lhs), left_entries(rhs)), + "Left entries of {} and {} are non-disjoint", + lhs, + rhs); + + ASSERT(are_disjoint(right_entries(lhs), right_entries(rhs)), + "Right entries of {} and {} are non-disjoint", + lhs, + rhs); bidict result; for (auto const &kv : lhs) { diff --git a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h index 1ce7cb3a13..add49b3089 100644 --- a/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h +++ b/lib/utils/include/utils/bidict/algorithms/merge_disjoint_bidicts.h @@ -1,20 +1,20 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_MERGE_DISJOINT_BIDICTS_H +#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include "utils/bidict/bidict.h" #include "utils/containers/foldl.h" -#include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" namespace FlexFlow { template bidict merge_disjoint_bidicts(std::set> const &bidicts) { return foldl( - bidicts, - bidict{}, - [](bidict const &accum, bidict const &x) -> bidict { - return binary_merge_disjoint_bidicts(accum, x); - }); + bidicts, + bidict{}, + [](bidict const &accum, bidict const &x) -> bidict { + return binary_merge_disjoint_bidicts(accum, x); + }); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/variant.h b/lib/utils/include/utils/variant.h index ba689d57be..3c63a28c11 100644 --- a/lib/utils/include/utils/variant.h +++ b/lib/utils/include/utils/variant.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_UTILS_VARIANT_H #define _FLEXFLOW_UTILS_VARIANT_H -#include #include "utils/type_traits.h" #include +#include #include #include diff --git a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc index d3deb887b1..d0d2aa3b8f 100644 --- a/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc +++ b/lib/utils/src/utils/bidict/algorithms/binary_merge_disjoint_bidicts.cc @@ -6,8 +6,7 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = ordered_value_type<1>; -template - bidict binary_merge_disjoint_bidicts(bidict const &, - bidict const &); +template bidict binary_merge_disjoint_bidicts(bidict const &, + bidict const &); } // namespace FlexFlow diff --git a/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h b/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h index 15ea9fc663..bca939efc2 100644 --- a/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h +++ b/lib/utils/test/common/include/test/utils/rapidcheck/doctest.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_RAPIDCHECK_DOCTEST_H #define _FLEXFLOW_UTILS_TEST_COMMON_INCLUDE_TEST_UTILS_RAPIDCHECK_DOCTEST_H -#include #include +#include namespace FlexFlow { diff --git a/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc b/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc index 4d91d13c05..eb0268e3bd 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc @@ -1,6 +1,6 @@ #include "utils/graph/series_parallel/non_normal_sp_decomposition.h" -#include #include "utils/graph/series_parallel/series_parallel_decomposition.dtg.h" +#include using namespace ::FlexFlow; From fb7280a276eba6e87b62f87a4843081dbe2abe4a Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 12 Jun 2026 21:13:11 -0700 Subject: [PATCH 28/35] Minor fix in run-model --- bin/run-model/src/run-model/main.cc | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/bin/run-model/src/run-model/main.cc b/bin/run-model/src/run-model/main.cc index 46e1186b6b..a6f1d026ef 100644 --- a/bin/run-model/src/run-model/main.cc +++ b/bin/run-model/src/run-model/main.cc @@ -100,7 +100,8 @@ int main(int argc, char **argv) { /*loss=*/std::nullopt, /*input_tensors=*/input_tensors, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle); + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::GPU); // begin training loop int num_epochs = 5; From 9bab7b50cb441754e5446f75078a4a0d724c6faa Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 29 May 2026 16:58:26 -0700 Subject: [PATCH 29/35] Add repeat_until_converged utility function --- .../utils/containers/repeat_until_converged.h | 20 +++++++++++++++++ .../containers/repeat_until_converged.cc | 11 ++++++++++ .../containers/repeat_until_converged.cc | 22 +++++++++++++++++++ 3 files changed, 53 insertions(+) create mode 100644 lib/utils/include/utils/containers/repeat_until_converged.h create mode 100644 lib/utils/src/utils/containers/repeat_until_converged.cc create mode 100644 lib/utils/test/src/utils/containers/repeat_until_converged.cc diff --git a/lib/utils/include/utils/containers/repeat_until_converged.h b/lib/utils/include/utils/containers/repeat_until_converged.h new file mode 100644 index 0000000000..a1e7da77a1 --- /dev/null +++ b/lib/utils/include/utils/containers/repeat_until_converged.h @@ -0,0 +1,20 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REPEAT_UNTIL_CONVERGED_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REPEAT_UNTIL_CONVERGED_H + +namespace FlexFlow { + +template +T repeat_until_converged(T const &initial, F &&f) { + T previous = initial; + T current = initial; + do { + previous = current; + current = f(current); + } while (current != previous); + + return current; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/src/utils/containers/repeat_until_converged.cc b/lib/utils/src/utils/containers/repeat_until_converged.cc new file mode 100644 index 0000000000..85dc756ab7 --- /dev/null +++ b/lib/utils/src/utils/containers/repeat_until_converged.cc @@ -0,0 +1,11 @@ +#include "utils/containers/repeat_until_converged.h" +#include "utils/archetypes/value_type.h" + +namespace FlexFlow { + +using T = value_type<0>; +using F = std::function; + +template T repeat_until_converged(T const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/test/src/utils/containers/repeat_until_converged.cc b/lib/utils/test/src/utils/containers/repeat_until_converged.cc new file mode 100644 index 0000000000..23e9ad754b --- /dev/null +++ b/lib/utils/test/src/utils/containers/repeat_until_converged.cc @@ -0,0 +1,22 @@ +#include "utils/containers/repeat_until_converged.h" +#include + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("repeat_until_converged") { + SUBCASE("standard case") { + int result = repeat_until_converged(500, [](int x) { return x / 2; }); + int correct = 0; + + CHECK(result == correct); + } + + SUBCASE("value is already converged") { + int result = repeat_until_converged(500, [](int x) { return x; }); + int correct = 500; + + CHECK(result == correct); + } + } +} From a2709acab11238d6ce6f2a81e82d1936dbc63c3e Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 26 Jun 2026 18:11:24 -0700 Subject: [PATCH 30/35] Fix 'task-spec'-side issues in failing tests --- .proj.toml | 1 + flake.lock | 12 +- flake.nix | 3 +- .../pcg_task_graph.dtg.toml | 7 +- ...l_tensor_space_to_machine_space_mapping.cc | 8 +- .../task_graph_simulator/task_simulator.cc | 2 +- .../include/op-attrs/initializer_attrs.h | 7 +- .../src/op-attrs/initializer_attrs.cc | 10 +- ...sk_space_to_operator_task_space_mapping.cc | 6 +- lib/op-attrs/src/op-attrs/ops/transpose.cc | 8 +- .../file_format/v1/v1_computation_graph.cc | 4 +- .../mapped_operator_task_group.cc | 2 +- .../mapped_parallel_computation_graph.cc | 2 +- ...uted_per_device_op_state_initialization.cc | 8 +- .../realm-execution/instance_allocation.cc | 6 +- .../src/realm-execution/pcg_instance.cc | 8 +- .../src/realm-execution/realm_context.cc | 6 +- .../test/src/realm-execution/test_e2e.cc | 2 +- .../src/realm-execution/test_op_replicate.cc | 24 +- .../evaluate_substitution_output.cc | 12 +- .../src/substitutions/pcg_pattern.cc | 4 +- .../src/substitutions/pcg_pattern_match.cc | 4 +- .../task-spec/dynamic_graph/copy_insertion.h | 57 +- .../dynamic_external_value_id_t.dtg.toml | 23 + .../dynamic_graph/dynamic_graph_edge.dtg.toml | 9 +- .../dynamic_internal_value_id_t.dtg.toml | 22 + .../dynamic_invocation_id_t.dtg.toml | 20 + .../dynamic_graph/dynamic_node_attrs.dtg.toml | 2 +- .../dynamic_graph/dynamic_node_invocation.h | 12 +- ...mic_node_invocation_sharding_info.dtg.toml | 14 +- .../dynamic_graph/dynamic_node_mapping.h | 9 +- .../dynamic_open_dataflow_graph.dtg.toml | 1 + .../dynamic_open_dataflow_graph.h | 42 +- .../dynamic_graph/dynamic_slot_site.dtg.toml | 2 + .../dynamic_graph/dynamic_slot_site.h | 13 - .../dynamic_graph/dynamic_tensor_slot.h | 2 + .../dynamic_value_attrs.dtg.toml | 4 + ...dynamic_value_attrs_sharding_info.dtg.toml | 4 +- .../dynamic_value_copy_info.dtg.toml | 28 + .../dynamic_graph/dynamic_value_id_t.dtg.toml | 23 + .../external_dynamic_slot_site.dtg.toml | 8 +- .../internal_dynamic_slot_site.dtg.toml | 8 +- .../loss_insertion_result.dtg.toml | 1 + ...amic_open_dataflow_graph_from_mapped_pcg.h | 3 +- .../parallel_tensor_mapping.dtg.toml | 1 + .../dynamic_graph/parallel_tensor_mapping.h | 18 + .../task-spec/dynamic_graph/pass_expansion.h | 38 +- ...zable_dynamic_open_dataflow_graph.dtg.toml | 22 + ...serializable_dynamic_open_dataflow_graph.h | 16 + .../serializable_dynamic_value_attrs.dtg.toml | 4 + .../task-spec/dynamic_graph/shard_expansion.h | 22 +- .../dynamic_graph/training_op_type.dtg.toml | 1 + .../task-spec/dynamic_graph/copy_insertion.cc | 444 +-- .../dynamic_graph/dynamic_graph_edge.cc | 2 +- .../dynamic_graph/dynamic_node_invocation.cc | 38 +- .../dynamic_graph/dynamic_node_mapping.cc | 29 +- .../dynamic_open_dataflow_graph.cc | 261 +- .../dynamic_graph/dynamic_slot_site.cc | 26 - .../dynamic_graph/dynamic_tensor_slot.cc | 7 + .../task-spec/dynamic_graph/loss_insertion.cc | 2 + ...ake_dynamic_open_dataflow_graph_from_cg.cc | 2 + ...mic_open_dataflow_graph_from_mapped_pcg.cc | 20 +- .../dynamic_graph/parallel_tensor_mapping.cc | 23 + .../task-spec/dynamic_graph/pass_expansion.cc | 324 ++- ...erializable_dynamic_open_dataflow_graph.cc | 23 + .../serializable_dynamic_value_attrs.cc | 2 + .../dynamic_graph/shard_expansion.cc | 431 +-- .../task-spec/dynamic_graph/copy_insertion.cc | 2371 ++++++++++++----- .../dynamic_open_dataflow_graph.cc | 96 +- .../dynamic_graph/machine_slicing.cc | 3 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 1137 ++++++-- .../task-spec/dynamic_graph/pass_expansion.cc | 653 ++++- .../dynamic_graph/shard_expansion.cc | 694 +++-- .../dynamic_graph/update_insertion.cc | 10 +- .../bidict_from_unstructured_relation.h | 46 + ...ansform_keys.h => bidict_transform_keys.h} | 6 +- ...orm_values.h => bidict_transform_values.h} | 6 +- .../utils/binary_relation/binary_relation.h | 196 ++ .../binary_relation_from_map.h | 19 + .../binary_relation_transform_left.h | 28 + .../binary_relation_transform_left2.h | 28 + .../binary_relation_transform_right.h | 28 + .../binary_relation_transform_right2.h | 28 + .../binary_relation/filter_binary_relation.h | 21 + .../require_binary_relation_is_left_unique.h | 22 + .../require_binary_relation_is_right_unique.h | 22 + .../transform_binary_relation.h | 29 + lib/utils/include/utils/containers/at_idx.h | 11 + lib/utils/include/utils/containers/filtrans.h | 32 + .../utils/containers/get_element_counts.h | 26 +- .../kwarg_dataflow_output.dtg.toml | 4 + .../algorithms/is_isomorphic_under.h | 6 +- ...arg_dataflow_graphs_are_isomorphic_under.h | 4 +- .../multidigraph/algorithms/get_edge_counts.h | 3 +- ...arg_dataflow_graphs_are_isomorphic_under.h | 4 +- ...g_dataflow_graph_by_materializing_inputs.h | 1 + .../include/utils/nonempty_set/nonempty_set.h | 1 + ...sform_keys.cc => bidict_transform_keys.cc} | 4 +- ...m_values.cc => bidict_transform_values.cc} | 4 +- .../utils/binary_relation/binary_relation.cc | 34 + .../binary_relation_from_map.cc | 11 + .../binary_relation_transform_left.cc | 14 + .../binary_relation_transform_left2.cc | 14 + .../binary_relation_transform_right.cc | 14 + .../binary_relation_transform_right2.cc | 14 + .../binary_relation/filter_binary_relation.cc | 12 + .../require_binary_relation_is_left_unique.cc | 13 + ...require_binary_relation_is_right_unique.cc | 12 + lib/utils/src/utils/containers/at_idx.cc | 5 + lib/utils/src/utils/containers/filtrans.cc | 10 + .../utils/containers/get_element_counts.cc | 8 +- .../algorithms/get_imm_dominators_map.cc | 4 +- .../digraph/algorithms/transitive_closure.cc | 4 +- .../algorithms/transitive_reduction.cc | 4 +- ...aph_data_from_kwarg_dataflow_graph_data.cc | 4 +- ...ataflow_graph_from_kwarg_dataflow_graph.cc | 4 +- ...omorphism_between_kwarg_dataflow_graphs.cc | 4 +- ..._outgoing_kwarg_dataflow_edges_for_node.cc | 4 +- .../algorithms/kwarg_dataflow_graph_as_dot.cc | 4 +- .../kwarg_dataflow_graphs_are_isomorphic.cc | 4 +- ...belled_kwarg_dataflow_graph_view_as_dot.cc | 4 +- ...een_labelled_open_kwarg_dataflow_graphs.cc | 6 +- ...d_open_kwarg_dataflow_graph_view_as_dot.cc | 6 +- ...rg_dataflow_graphs_are_isomorphic_under.cc | 6 +- ...led_open_kwarg_dataflow_graph_input_ids.cc | 6 +- .../algorithms/get_edge_counts.cc | 2 +- .../algorithms/is_isomorphic_under.cc | 6 +- ...isms_between_open_kwarg_dataflow_graphs.cc | 6 +- ...arg_dataflow_graph_input_id_permutation.cc | 6 +- .../get_all_open_kwarg_dataflow_edges.cc | 12 +- .../get_open_kwarg_dataflow_graph_subgraph.cc | 6 +- .../open_kwarg_dataflow_graph_as_dot.cc | 6 +- ...ute_open_kwarg_dataflow_graph_input_ids.cc | 6 +- ...hism_between_open_kwarg_dataflow_graphs.cc | 6 +- ..._dataflow_graph_by_materializing_inputs.cc | 6 +- lib/utils/src/utils/orthotope/dim_coord.cc | 4 +- .../src/utils/orthotope/dim_domain_mapping.cc | 14 +- .../src/utils/orthotope/dim_projection.cc | 12 +- .../src/utils/orthotope/down_projection.cc | 12 +- .../orthotope/minimal_dim_domain_mapping.cc | 14 +- ...sform_keys.cc => bidict_transform_keys.cc} | 6 +- ...m_values.cc => bidict_transform_values.cc} | 6 +- .../utils/binary_relation/binary_relation.cc | 125 + .../binary_relation/filter_binary_relation.cc | 47 + .../require_binary_relation_is_left_unique.cc | 52 + ...require_binary_relation_is_right_unique.cc | 52 + .../utils/containers/get_element_counts.cc | 4 +- lib/utils/test/src/utils/fmt/unordered_map.cc | 4 +- .../get_inverse_line_graph.cc | 20 +- .../series_parallel/parallel_reduction.cc | 20 +- 150 files changed, 6323 insertions(+), 2108 deletions(-) create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_external_value_id_t.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_internal_value_id_t.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.toml delete mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_copy_info.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_id_t.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.dtg.toml create mode 100644 lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h delete mode 100644 lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc create mode 100644 lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc rename lib/utils/include/utils/bidict/algorithms/{transform_keys.h => bidict_transform_keys.h} (58%) rename lib/utils/include/utils/bidict/algorithms/{transform_values.h => bidict_transform_values.h} (58%) create mode 100644 lib/utils/include/utils/binary_relation/binary_relation.h create mode 100644 lib/utils/include/utils/binary_relation/binary_relation_from_map.h create mode 100644 lib/utils/include/utils/binary_relation/binary_relation_transform_left.h create mode 100644 lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h create mode 100644 lib/utils/include/utils/binary_relation/binary_relation_transform_right.h create mode 100644 lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h create mode 100644 lib/utils/include/utils/binary_relation/filter_binary_relation.h create mode 100644 lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h create mode 100644 lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h create mode 100644 lib/utils/include/utils/binary_relation/transform_binary_relation.h rename lib/utils/src/utils/bidict/algorithms/{transform_keys.cc => bidict_transform_keys.cc} (63%) rename lib/utils/src/utils/bidict/algorithms/{transform_values.cc => bidict_transform_values.cc} (62%) create mode 100644 lib/utils/src/utils/binary_relation/binary_relation.cc create mode 100644 lib/utils/src/utils/binary_relation/binary_relation_from_map.cc create mode 100644 lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc create mode 100644 lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc create mode 100644 lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc create mode 100644 lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc create mode 100644 lib/utils/src/utils/binary_relation/filter_binary_relation.cc create mode 100644 lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc create mode 100644 lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc rename lib/utils/test/src/utils/bidict/algorithms/{transform_keys.cc => bidict_transform_keys.cc} (65%) rename lib/utils/test/src/utils/bidict/algorithms/{transform_values.cc => bidict_transform_values.cc} (62%) create mode 100644 lib/utils/test/src/utils/binary_relation/binary_relation.cc create mode 100644 lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc create mode 100644 lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc create mode 100644 lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc diff --git a/.proj.toml b/.proj.toml index 522e11c369..b40a92fea9 100644 --- a/.proj.toml +++ b/.proj.toml @@ -3,6 +3,7 @@ testsuite_macro = "FF_TEST_SUITE" namespace_name = "FlexFlow" header_extension = ".h" doxygen = true +build_tool = "ninja" cuda_launch_cmd = [ "nixGL", "--", diff --git a/flake.lock b/flake.lock index 71bb407e27..c06618c181 100644 --- a/flake.lock +++ b/flake.lock @@ -67,17 +67,17 @@ "python38-nixpkgs": "python38-nixpkgs" }, "locked": { - "lastModified": 1780709897, - "narHash": "sha256-55VJWpNnt/tMEP8tadRQSI66cX1BgJS9UNy/EkaQt3A=", + "lastModified": 1781917763, + "narHash": "sha256-Qd8CW+G/orxU9Ne6fNqKAbJL404WeV7/kdY/9TNJ5Lw=", "ref": "refs/heads/master", - "rev": "18d04112d94ee0a1c173a30923aac9fa2b505539", - "revCount": 160, + "rev": "f2bdddab299b98fb67612e3c04136949b2339d74", + "revCount": 161, "type": "git", - "url": "https://git.sr.ht/~lockshaw/proj" + "url": "file:///home/lockshaw/x/ff/proj/proj" }, "original": { "type": "git", - "url": "https://git.sr.ht/~lockshaw/proj" + "url": "file:///home/lockshaw/x/ff/proj/proj" } }, "python38-nixpkgs": { diff --git a/flake.nix b/flake.nix index ad71cbefb4..ccb0402498 100644 --- a/flake.nix +++ b/flake.nix @@ -18,7 +18,8 @@ flake-utils.url = "github:numtide/flake-utils"; proj-repo = { - url = "git+https://git.sr.ht/~lockshaw/proj"; + # url = "git+https://git.sr.ht/~lockshaw/proj"; + url = "git+file:///home/lockshaw/x/ff/proj/proj"; inputs.nixpkgs.follows = "nixpkgs"; inputs.flake-utils.follows = "flake-utils"; }; diff --git a/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml b/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml index ea76f1be05..8b99628b76 100644 --- a/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml +++ b/lib/compiler/include/compiler/task_graph_simulator/pcg_task_graph.dtg.toml @@ -2,14 +2,13 @@ namespace = "FlexFlow" name = "PCGTaskGraph" type = "struct" -features = [ -] +features = [] includes = [ "utils/graph/digraph/digraph_view.h", "compiler/task_graph_simulator/pcg_task.dtg.h", "", - "" + "", "utils/many_to_one/many_to_one.h", "pcg/machine_space_coordinate.dtg.h", ] @@ -18,7 +17,7 @@ src_includes = [ "utils/fmt/set.h", "utils/hash/set.h", "utils/fmt/map.h", - "utils/hash/map.h" + "utils/hash/map.h", ] [[fields]] diff --git a/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc b/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc index 270e615ff1..01c9a2e765 100644 --- a/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc +++ b/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc @@ -4,8 +4,8 @@ #include "op-attrs/parallel_tensor_space_coordinate.h" #include "op-attrs/task_space_coordinate.h" #include "utils/bidict/algorithms/exhaustive_relational_join.h" -#include "utils/bidict/algorithms/transform_keys.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include namespace FlexFlow { @@ -19,8 +19,8 @@ ParallelTensorSpaceToMachineSpaceMapping ptensor_machine_map_from_composition( op_task_to_parallel_tensor_space_mapping)); bidict - pt_to_op_coord_map = transform_keys( - transform_values(op_task_to_parallel_tensor_space_mapping.raw_mapping + pt_to_op_coord_map = bidict_transform_keys( + bidict_transform_values(op_task_to_parallel_tensor_space_mapping.raw_mapping .coord_mapping.reversed(), task_space_coordinate_from_dim_coord), parallel_tensor_space_coord_from_dim_coord); diff --git a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc index 0af2f18aee..77812b4174 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc @@ -53,7 +53,7 @@ milliseconds_t task_simulator_estimate_forward_pass_time( assert(current_task.is_operator()); auto get_devices = - [&](Node const &n) -> std::unordered_set { + [&](Node const &n) -> std::set { return task_graph.node_to_devices.at(n); }; diff --git a/lib/op-attrs/include/op-attrs/initializer_attrs.h b/lib/op-attrs/include/op-attrs/initializer_attrs.h index 71d6da5363..4fba656cec 100644 --- a/lib/op-attrs/include/op-attrs/initializer_attrs.h +++ b/lib/op-attrs/include/op-attrs/initializer_attrs.h @@ -7,7 +7,12 @@ namespace FlexFlow { InitializerAttrs make_zero_initializer(); -InitializerAttrs make_kaiming_uniform(TensorDims const &); +InitializerAttrs make_kaiming_uniform( + TensorDims const &, + float a = 0.0, + KaimingInitializerMode mode = KaimingInitializerMode::FAN_IN, + KaimingInitializerNonlinearity nonlinearity = KaimingInitializerNonlinearity::LEAKY_RELU, + int seed = 0); } // namespace FlexFlow diff --git a/lib/op-attrs/src/op-attrs/initializer_attrs.cc b/lib/op-attrs/src/op-attrs/initializer_attrs.cc index b24b28a339..c0dc827ea3 100644 --- a/lib/op-attrs/src/op-attrs/initializer_attrs.cc +++ b/lib/op-attrs/src/op-attrs/initializer_attrs.cc @@ -46,11 +46,11 @@ static float // from pytorch: // see // https://github.com/pytorch/pytorch/blob/bd019c0bb485904a99fb38589444b1461ab1e486/torch/nn/init.py#L456-L518 -InitializerAttrs kaiming_uniform(TensorDims const &dims, - float a, - KaimingInitializerMode mode, - KaimingInitializerNonlinearity nonlinearity, - int seed) { +InitializerAttrs make_kaiming_uniform(TensorDims const &dims, + float a, + KaimingInitializerMode mode, + KaimingInitializerNonlinearity nonlinearity, + int seed) { positive_int fan = calculate_fan_for_mode(dims, mode); float gain = gain_for_nonlinearity(nonlinearity, a); diff --git a/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc b/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc index 605578acdc..3d70c366a1 100644 --- a/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc +++ b/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc @@ -1,8 +1,8 @@ #include "op-attrs/operator_task_space_to_operator_task_space_mapping.h" #include "op-attrs/operator_task_space.h" #include "op-attrs/task_space_coordinate.h" -#include "utils/bidict/algorithms/transform_keys.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/orthotope/minimal_dim_domain.h" #include "utils/orthotope/minimal_dim_domain_mapping.h" @@ -40,7 +40,7 @@ OperatorTaskSpace op_mapping_get_dst_space( bidict op_to_op_get_coord_mapping( OperatorTaskSpaceToOperatorTaskSpaceMapping const &mapping) { - return transform_values(transform_keys(mapping.raw_mapping.coord_mapping, + return bidict_transform_values(bidict_transform_keys(mapping.raw_mapping.coord_mapping, task_space_coordinate_from_dim_coord), task_space_coordinate_from_dim_coord); } diff --git a/lib/op-attrs/src/op-attrs/ops/transpose.cc b/lib/op-attrs/src/op-attrs/ops/transpose.cc index a13a5d1724..4aa58aa2d3 100644 --- a/lib/op-attrs/src/op-attrs/ops/transpose.cc +++ b/lib/op-attrs/src/op-attrs/ops/transpose.cc @@ -5,8 +5,8 @@ #include "op-attrs/parallel_tensor_shape.h" #include "op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.dtg.h" #include "op-attrs/parallel_tensor_space_to_parallel_tensor_space_mapping.h" -#include "utils/bidict/algorithms/transform_keys.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" namespace FlexFlow { @@ -51,8 +51,8 @@ static ParallelTensorSpaceToParallelTensorSpaceMapping EqProjection inp_to_out = EqProjection{ - transform_keys( - transform_values(attrs.permutation.as_bidict(), ff_dim_to_pt_dim), + bidict_transform_keys( + bidict_transform_values(attrs.permutation.as_bidict(), ff_dim_to_pt_dim), ff_dim_to_pt_dim), }; diff --git a/lib/pcg/src/pcg/file_format/v1/v1_computation_graph.cc b/lib/pcg/src/pcg/file_format/v1/v1_computation_graph.cc index e52b5708e5..b949a16434 100644 --- a/lib/pcg/src/pcg/file_format/v1/v1_computation_graph.cc +++ b/lib/pcg/src/pcg/file_format/v1/v1_computation_graph.cc @@ -1,6 +1,6 @@ #include "pcg/file_format/v1/v1_computation_graph.h" #include "pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h" #include "utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h" @@ -21,7 +21,7 @@ std::pair> TensorAttrs, TensorSlotName>(cg.raw_graph); V1ComputationGraph v1_cg = V1ComputationGraph{raw.first}; - bidict v1_node_ids = transform_values( + bidict v1_node_ids = bidict_transform_values( raw.second, [](Node const &n) { return layer_guid_t{n}; }); return {v1_cg, v1_node_ids}; diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc index 88ccb50647..0c9bdb3029 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc @@ -3,7 +3,7 @@ #include "op-attrs/operator_task_space.h" #include "op-attrs/parallel_tensor_space_coordinate.h" #include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/bidict/generate_bidict.h" #include "utils/containers/are_all_distinct.h" #include "utils/containers/require_all_same.h" diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 7571daecfa..38d99f561b 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -3,7 +3,7 @@ #include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_attrs.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_map.h" -#include "utils/bidict/algorithms/transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" #include "utils/containers/transform.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h" diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index 9077a24f01..797951311b 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -34,14 +34,14 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( for (DynamicNodeInvocation const &invocation : dg.invocations) { // Nodes mapped to multiple devices are always parallel operators and don't // have any initialization to perform anyway - std::optional device_coord = - maybe_get_only(assert_unwrap(invocation.node_attrs.device_coords)); - if (!device_coord.has_value()) { + std::optional device_id = + maybe_get_only(assert_unwrap(invocation.node_attrs.device_ids)); + if (!device_id.has_value()) { continue; } Realm::Processor target_proc = ctx.processor_from_global_device_id( - assert_unwrap(device_coord)); + assert_unwrap(device_id)); TensorInstanceBacking tensor_backing = subset_tensor_instance_backing_for_invocation(tensor_instance_backing, diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index a3336f570a..a4782c85ac 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -52,10 +52,10 @@ TensorInstanceBacking perform_instance_allocation( NOT_IMPLEMENTED(); } else { if (!contains_key(result.backing, v)) { - MachineSpaceCoordinate device_coord = - assert_unwrap(v.mapping).at_l(assert_unwrap(v.shard_coord)); + global_device_id_t device_id = + assert_unwrap(v.mapping).raw.at_l(assert_unwrap(v.shard_coord)); result.backing.insert(std::pair{ - v, perform_instance_allocation_for_value(device, v, ctx)}); + v, perform_instance_allocation_for_value(device_id, v, ctx)}); } return result.backing.at(v); } diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index faefadd73c..2b60580b73 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -243,10 +243,14 @@ static Realm::Event spawn_dynamic_node_invocation( // chain reductions sequentially to avoid write races on dst Realm::Event result = precondition; - for (auto const &[p, m] : assert_unwrap(output_grad.mapping)) { + for (auto const &[p, d] : assert_unwrap(output_grad.mapping).raw) { DynamicValueAttrs replica_key = output_grad; replica_key.mapping = - bidict{{p, m}}; + ParallelTensorMapping{ + bidict{ + {p, d}, + }, + }; replica_key.shard_coord = p; Realm::RegionInstance src_inst = diff --git a/lib/realm-execution/src/realm-execution/realm_context.cc b/lib/realm-execution/src/realm-execution/realm_context.cc index 8740022834..56b1d11c65 100644 --- a/lib/realm-execution/src/realm-execution/realm_context.cc +++ b/lib/realm-execution/src/realm-execution/realm_context.cc @@ -16,7 +16,7 @@ #include "task-spec/global_device_id_t.h" #include "utils/bidict/algorithms/bidict_from_enumerating.h" #include "utils/bidict/algorithms/merge_disjoint_bidicts.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/containers/are_all_same.h" #include "utils/containers/contains_key.h" #include "utils/containers/group_by.h" @@ -56,7 +56,7 @@ bidict build_local_machine_topology( set_of(by_proc_kind.at_l(k).unwrap_as_unordered_set())) .reversed(); - bidict result = transform_values( + bidict result = bidict_transform_values( enumerated, [&](nonnegative_int idx) -> local_device_id_t { return local_device_id_t{ /*idx=*/device_in_node_idx_t{idx}, @@ -88,7 +88,7 @@ static bidict bidict local_topology_for_node = build_local_machine_topology(procs_for_node); - return transform_values( + return bidict_transform_values( local_topology_for_node, [&](local_device_id_t const &local_device_id) -> global_device_id_t { return global_device_id_from_local(local_device_id, node_idx); diff --git a/lib/realm-execution/test/src/realm-execution/test_e2e.cc b/lib/realm-execution/test/src/realm-execution/test_e2e.cc index d9c85dbf6d..d7f777e686 100644 --- a/lib/realm-execution/test/src/realm-execution/test_e2e.cc +++ b/lib/realm-execution/test/src/realm-execution/test_e2e.cc @@ -148,7 +148,7 @@ static E2ETrainingConfig create_e2e_test_case() { MachineSpaceCoordinate cpu1{0_n, 1_n}; ParallelTensorSpaceCoordinate tensor_coord0{0_n, 0_n, FFOrdered{0_n}}; - std::unordered_map mapping = { + std::map mapping = { {inputs_layer.parallel_layer, MappedOperatorTaskGroup{ {{cpu0, diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index fac4b86871..708fb73d73 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -133,8 +133,8 @@ MappedParallelComputationGraph parallel_tensor_guid_t t_relu_1 = require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); - MachineSpaceCoordinate cpu0{0_n, 0_n, device_type}; - MachineSpaceCoordinate cpu1{0_n, 1_n, device_type}; + MachineSpaceCoordinate mc0{0_n, 0_n}; + MachineSpaceCoordinate mc1{0_n, 1_n}; ParallelTensorSpaceCoordinate tensor_coord0{ /*sum_component=*/0_n, @@ -154,7 +154,7 @@ MappedParallelComputationGraph MappedOperatorTaskGroup{ { { - cpu0, + mc0, OperatorAtomicTaskShardBinding{{ {TensorSlotName::OUTPUT, tensor_coord0}, }}, @@ -167,7 +167,7 @@ MappedParallelComputationGraph MappedOperatorTaskGroup{ { { - cpu0, + mc0, OperatorAtomicTaskShardBinding{{ {TensorSlotName::OUTPUT, tensor_coord0}, }}, @@ -180,7 +180,7 @@ MappedParallelComputationGraph MappedOperatorTaskGroup{ { { - cpu0, + mc0, OperatorAtomicTaskShardBinding{{ {TensorSlotName::LHS_INPUT, tensor_coord0}, {TensorSlotName::RHS_INPUT, tensor_coord0}, @@ -195,14 +195,14 @@ MappedParallelComputationGraph MappedOperatorTaskGroup{ { { - cpu0, + mc0, OperatorAtomicTaskShardBinding{{ {TensorSlotName::INPUT, tensor_coord0}, {TensorSlotName::OUTPUT, tensor_coord0}, }}, }, { - cpu1, + mc1, OperatorAtomicTaskShardBinding{{ {TensorSlotName::INPUT, tensor_coord0}, {TensorSlotName::OUTPUT, tensor_coord1}, @@ -216,14 +216,14 @@ MappedParallelComputationGraph MappedOperatorTaskGroup{ { { - cpu0, + mc0, OperatorAtomicTaskShardBinding{{ {TensorSlotName::INPUT, tensor_coord0}, {TensorSlotName::OUTPUT, tensor_coord0}, }}, }, { - cpu1, + mc1, OperatorAtomicTaskShardBinding{{ {TensorSlotName::INPUT, tensor_coord1}, {TensorSlotName::OUTPUT, tensor_coord1}, @@ -276,7 +276,8 @@ TEST_SUITE(FF_TEST_SUITE) { /*loss=*/std::nullopt, /*input_tensors=*/input_tensors, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle); + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::CPU); // begin training loop int num_epochs = 1; @@ -331,7 +332,8 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { /*loss=*/std::nullopt, /*input_tensors=*/input_tensors, /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle); + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::GPU); // begin training loop int num_epochs = 1; diff --git a/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc b/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc index 28bfac0f69..b308bc0bde 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc @@ -2,8 +2,8 @@ #include "substitutions/apply_substitution/perform_shape_inference.h" #include "substitutions/output_graph/output_operator_attrs_assignment.h" #include "substitutions/sub_parallel_computation_graph.h" -#include "utils/bidict/algorithms/transform_keys.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/bidict/generate_bidict.h" #include "utils/containers/map_keys.h" #include "utils/containers/map_values.h" @@ -70,8 +70,8 @@ std::pair }); bidict result_input_map = - transform_keys( - transform_values(new_input_id_permutation, + bidict_transform_keys( + bidict_transform_values(new_input_id_permutation, [](KwargDataflowGraphInput const &i) { return OutputGraphExprInput{i}; }), @@ -80,8 +80,8 @@ std::pair }); bidict result_node_map = - transform_keys( - transform_values( + bidict_transform_keys( + bidict_transform_values( new_node_id_permutation, [](Node const &n) { return OutputGraphExprNode{n}; }), [](NewNode const &n) { return parallel_layer_guid_t{n.raw_node}; }); diff --git a/lib/substitutions/src/substitutions/pcg_pattern.cc b/lib/substitutions/src/substitutions/pcg_pattern.cc index d55bc2f9b5..c8b521bf1d 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern.cc @@ -5,7 +5,7 @@ #include "substitutions/tensor_pattern/satisfies_pattern.h" #include "substitutions/unlabelled/find_pattern_matches.h" #include "substitutions/unlabelled/pattern_value.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/containers/map_values.h" #include "utils/containers/transform.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_outputs_for_node.h" @@ -60,7 +60,7 @@ std::vector auto pcg_match_from_unlabelled_match = [](UnlabelledKwargDataflowGraphPatternMatch const &m) { return PCGPatternMatch{ - transform_values( + bidict_transform_values( m.node_assignment, [](Node const &n) { return parallel_layer_guid_t{n}; }), map_values( diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index 9f4b207b30..0c0d4b11bb 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -6,7 +6,7 @@ #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include "utils/bidict/algorithms/exhaustive_relational_join.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/map_values.h" #include "utils/containers/values.h" @@ -43,7 +43,7 @@ bidict UnlabelledKwargDataflowGraphPatternMatch get_unlabelled_pattern_match(PCGPatternMatch const &match) { return UnlabelledKwargDataflowGraphPatternMatch{ - transform_values( + bidict_transform_values( match.node_assignment, [](parallel_layer_guid_t const &l) { return l.raw_graph_node; }), map_values(match.input_assignment, diff --git a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h index da51750cd7..fa8776f5ef 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h @@ -4,24 +4,61 @@ #include "task-spec/dynamic_graph/dynamic_node_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/dynamic_graph/dynamic_value_copy_info.dtg.h" +#include "task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" namespace FlexFlow { bool node_is_copy(DynamicNodeAttrs const &n); bool value_is_mapped(DynamicValueAttrs const &); -bool no_part_of_graph_is_copy_inserted(DynamicOpenDataflowGraph const &); -bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &); +void require_node_is_ready_for_copy_insertion(DynamicNodeAttrs const &); +void require_value_is_ready_for_copy_insertion(DynamicValueAttrs const &); +void require_invocation_is_ready_for_copy_insertion(DynamicNodeInvocation const &); +void require_graph_is_ready_for_copy_insertion(DynamicOpenDataflowGraph const &); -std::set copies_for_invocation_inputs( - DynamicNodeInvocation const &i, - std::map const - &unmapped_value_to_mapped_source_value); +void require_value_is_copy_inserted(DynamicValueAttrs const &); +void require_invocation_is_fully_copy_inserted(DynamicNodeInvocation const &); +void require_graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &); -std::set perform_copy_insertion_for_invocation( - DynamicNodeInvocation const &i, - std::map const - &unmapped_value_to_mapped_source_value); +std::map + get_mappings_for_invocation( + DynamicNodeInvocation const &, + std::map const &); + +DynamicNodeInvocation apply_mappings_for_invocation( + dynamic_invocation_id_t const &, + DynamicNodeInvocation const &, + std::map const &); + +DynamicNodeInvocation make_copy_invocation(DynamicValueCopyInfo const &); + +std::set copies_for_value( + DynamicValueAttrs const &value_attrs, + DynamicSlotSite const &src_site, + std::set const &dst_sites, + std::map const &mappings); + +std::set copies_for_internal_value( + DynamicValueAttrs const &value_attrs, + InternalDynamicSlotSite const &src_site, + ParallelTensorMapping const &src_site_mapping, + std::map const &sink_site_mappings); + +std::set infer_all_copies_in_graph(DynamicOpenDataflowGraph const &); + +std::map + resolve_tensor_mappings(DynamicOpenDataflowGraph const &); + +std::map + resolve_partial_tensor_mappings_from_node_mappings( + DynamicOpenDataflowGraph const &); + +std::map + resolve_missing_tensor_mappings_from_adjacent_values( + DynamicOpenDataflowGraph const &, + std::map const &); DynamicOpenDataflowGraph perform_copy_insertion(DynamicOpenDataflowGraph const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_external_value_id_t.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_external_value_id_t.dtg.toml new file mode 100644 index 0000000000..72686cfff5 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_external_value_id_t.dtg.toml @@ -0,0 +1,23 @@ +namespace = "FlexFlow" +name = "dynamic_external_value_id_t" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", + "rapidcheck", +] + +includes = [ + "utils/nonnegative_int/nonnegative_int.h" +] + +src_includes = [ +] + +[[fields]] +name = "idx" +type = "::FlexFlow::nonnegative_int" + diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml index 578bfe7585..0c63997814 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_graph_edge.dtg.toml @@ -3,14 +3,15 @@ name = "DynamicGraphEdge" type = "struct" features = [ "eq", + "ord", "hash", "fmt", ] includes = [ - "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h", - "task-spec/dynamic_graph/dynamic_slot_site.h", + "task-spec/dynamic_graph/dynamic_slot_site.dtg.h", + "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h", ] src_includes = [ @@ -21,8 +22,8 @@ name = "src" type = "::FlexFlow::DynamicSlotSite" [[fields]] -name = "dst_node" -type = "::FlexFlow::DynamicNodeInvocation" +name = "dst_node_id" +type = "::FlexFlow::dynamic_invocation_id_t" [[fields]] name = "dst_slot" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_internal_value_id_t.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_internal_value_id_t.dtg.toml new file mode 100644 index 0000000000..32c04e1c8a --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_internal_value_id_t.dtg.toml @@ -0,0 +1,22 @@ +namespace = "FlexFlow" +name = "dynamic_internal_value_id_t" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", + "rapidcheck", +] + +includes = [ + "utils/nonnegative_int/nonnegative_int.h", +] + +src_includes = [ +] + +[[fields]] +name = "idx" +type = "::FlexFlow::nonnegative_int" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.toml new file mode 100644 index 0000000000..f600824ab3 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.toml @@ -0,0 +1,20 @@ +namespace = "FlexFlow" +name = "dynamic_invocation_id_t" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "utils/nonnegative_int/nonnegative_int.h", +] + +src_includes = [] + +[[fields]] +name = "idx" +type = "::FlexFlow::nonnegative_int" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml index 1d7b05b1fa..9725ce4542 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_attrs.dtg.toml @@ -29,7 +29,7 @@ type = "std::optional<::FlexFlow::DynamicTaskType>" [[fields]] name = "device_ids" -type = "std::optional<::FlexFlow::nonempty_set<::FlexFlow::global_device_id_t>+" +type = "std::optional<::FlexFlow::nonempty_set<::FlexFlow::global_device_id_t>>" docstring = ''' \brief The devices on which this task should execute. diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h index 81b549e9bc..3ba9d69c3a 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h @@ -19,14 +19,20 @@ void require_invocation_fully_satisfies(DynamicNodeInvocation const &, std::function const &require_value_condition, std::function const &require_slot_condition); -std::unordered_map +std::map get_slot_map_for_direction(DynamicNodeInvocation const &, TensorDirection); TrainingOpType dynamic_node_invocation_get_op_type(DynamicNodeInvocation const &); -std::unordered_set - get_dynamic_slot_sites_for_invocation(DynamicNodeInvocation const &); +std::set + get_incoming_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &, DynamicNodeInvocation const &); + +std::set + get_output_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &, DynamicNodeInvocation const &); + +std::set + get_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &, DynamicNodeInvocation const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml index 00d98a2e6c..851dbd44fe 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.toml @@ -10,21 +10,17 @@ features = [ ] includes = [ - "pcg/machine_space_coordinate.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h", "task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.h", "utils/nonempty_set/nonempty_set.h", -] - -src_includes = [ - "utils/hash/map.h", - "utils/fmt/map.h", + "task-spec/global_device_id_t.dtg.h", + "utils/binary_relation/binary_relation.h", ] [[fields]] -name = "device_coords" -type = "::FlexFlow::nonempty_set<::FlexFlow::MachineSpaceCoordinate>" +name = "device_ids" +type = "::FlexFlow::nonempty_set<::FlexFlow::global_device_id_t>" [[fields]] name = "value_sharding" -type = "std::map<::FlexFlow::DynamicTensorSlot, ::FlexFlow::DynamicValueAttrsShardingInfo>" +type = "::FlexFlow::BinaryRelation<::FlexFlow::DynamicTensorSlot, ::FlexFlow::DynamicValueAttrsShardingInfo>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h index 3c5b27c344..2e41f223f6 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h @@ -6,11 +6,18 @@ namespace FlexFlow { +bidict + dynamic_node_mapping_get_shard_bindings(DynamicNodeMapping const &); + +OperatorAtomicTaskShardBinding + dynamic_node_mapping_get_shard_binding_for_device(DynamicNodeMapping const &, + global_device_id_t const &); + bidict dynamic_node_mapping_bindings_for_slot_name(DynamicNodeMapping const &, TensorSlotName const &); -std::unordered_set +std::set target_devices_of_dynamic_node_mapping(DynamicNodeMapping const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml index 8f5876beba..1042b89ff7 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.toml @@ -3,6 +3,7 @@ name = "DynamicOpenDataflowGraph" type = "struct" features = [ "eq", + "ord", "fmt", ] diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h index 29c485dc11..82bfe59a15 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h @@ -6,6 +6,8 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h" +#include "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h" +#include "task-spec/dynamic_graph/dynamic_value_id_t.dtg.h" namespace FlexFlow { @@ -30,9 +32,7 @@ bool no_part_of_dynamic_graph_satisfies( void require_full_dynamic_graph_satisfies( DynamicOpenDataflowGraph const &, - std::function const &, - std::function const &, - std::function const &); + std::function const &); std::multiset get_dynamic_nodes(DynamicOpenDataflowGraph const &); @@ -43,28 +43,52 @@ std::multiset std::set get_dynamic_invocation_set(DynamicOpenDataflowGraph const &); -std::unordered_set +std::set + dynamic_graph_get_internal_values(DynamicOpenDataflowGraph const &); +std::set + dynamic_graph_get_external_values(DynamicOpenDataflowGraph const &); + +dynamic_invocation_id_t dynamic_graph_get_id_for_invocation(DynamicOpenDataflowGraph const &, + DynamicNodeInvocation const &); +DynamicNodeInvocation dynamic_graph_get_invocation_for_id(DynamicOpenDataflowGraph const &, + dynamic_invocation_id_t const &); + +dynamic_value_id_t dynamic_graph_get_id_for_value(DynamicOpenDataflowGraph const &, + DynamicValueAttrs const &); +DynamicValueAttrs dynamic_graph_get_value_for_id(DynamicOpenDataflowGraph const &, + dynamic_value_id_t const &); + +std::set get_dynamic_graph_edges(DynamicOpenDataflowGraph const &); -std::unordered_set +std::set get_dynamic_graph_edges_incoming_to_invocation( DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); -std::unordered_set +std::set get_dynamic_graph_edges_outgoing_from_invocation( DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); -std::unordered_set +std::set get_internal_dynamic_slot_sites(DynamicOpenDataflowGraph const &); -std::unordered_set +std::set get_dynamic_slot_sites(DynamicOpenDataflowGraph const &); +DynamicSlotSite + dynamic_graph_find_source_of_slot_site(DynamicOpenDataflowGraph const &, + InternalDynamicSlotSite const &); +std::set + dynamic_graph_find_sinks_of_slot_site(DynamicOpenDataflowGraph const &, + InternalDynamicSlotSite const &); + DynamicSlotSite dynamic_graph_find_source_of_value(DynamicOpenDataflowGraph const &, DynamicValueAttrs const &); -std::unordered_set +std::set dynamic_graph_find_sinks_of_value(DynamicOpenDataflowGraph const &, DynamicValueAttrs const &); +DynamicValueAttrs dynamic_value_attrs_for_slot_site(DynamicOpenDataflowGraph const &, DynamicSlotSite const &); + std::optional find_output_value_attrs(DynamicOpenDataflowGraph const &, dynamic_tensor_guid_t, diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml index fa56b0f105..a3a5db6ea1 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.dtg.toml @@ -3,8 +3,10 @@ name = "DynamicSlotSite" type = "variant" features = [ "eq", + "ord", "hash", "fmt", + "json", ] includes = [ diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h deleted file mode 100644 index 4340cdcf39..0000000000 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_slot_site.h +++ /dev/null @@ -1,13 +0,0 @@ -#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_SLOT_SITE_H -#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_SLOT_SITE_H - -#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" -#include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" - -namespace FlexFlow { - -DynamicValueAttrs dynamic_value_attrs_for_slot_site(DynamicSlotSite const &); - -} // namespace FlexFlow - -#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.h index 129f1a2eae..d72d1b9dc1 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_tensor_slot.h @@ -8,6 +8,8 @@ namespace FlexFlow { DynamicTensorSlot decide_tensor_slot_role(DynamicTensorSlot const &, DynamicTensorRole); +DynamicTensorSlot slot_without_task_shard(DynamicTensorSlot const &); + } // namespace FlexFlow #endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml index ad719db3f8..2c10f23f22 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml @@ -42,6 +42,10 @@ docstring = ''' For a \ref DynamicOpenDataflowGraph originating form a \ref MappedParallelComputationGraph, this field is filled in by \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h. ''' +[[fields]] +name = "create_grad" +type = "std::optional" + [[fields]] name = "shard_coord" type = "std::optional<::FlexFlow::ParallelTensorSpaceCoordinate>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml index 63ad868e68..8598914f00 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.toml @@ -11,7 +11,7 @@ features = [ includes = [ "op-attrs/parallel_tensor_space_coordinate.dtg.h", - "pcg/machine_space_coordinate.dtg.h", + "task-spec/global_device_id_t.dtg.h", ] src_includes = [] @@ -22,4 +22,4 @@ type = "::FlexFlow::ParallelTensorSpaceCoordinate" [[fields]] name = "mapping" -type = "::FlexFlow::MachineSpaceCoordinate" +type = "::FlexFlow::global_device_id_t" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_copy_info.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_copy_info.dtg.toml new file mode 100644 index 0000000000..934e01bad9 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_copy_info.dtg.toml @@ -0,0 +1,28 @@ +namespace = "FlexFlow" +name = "DynamicValueCopyInfo" +type = "struct" +features = [ + "eq", + "ord", + "hash", + "fmt", +] + +includes = [ + "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h", + "task-spec/dynamic_graph/parallel_tensor_mapping.dtg.h", +] + +src_includes = [] + +[[fields]] +name = "value_attrs" +type = "::FlexFlow::DynamicValueAttrs" + +[[fields]] +name = "src_mapping" +type = "::FlexFlow::ParallelTensorMapping" + +[[fields]] +name = "dst_mapping" +type = "::FlexFlow::ParallelTensorMapping" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_id_t.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_id_t.dtg.toml new file mode 100644 index 0000000000..6d5159dcb1 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_id_t.dtg.toml @@ -0,0 +1,23 @@ +namespace = "FlexFlow" +name = "dynamic_value_id_t" +type = "variant" +features = [ + "eq", + "ord", + "hash", + "json", + "fmt", +] + +includes = [ + "task-spec/dynamic_graph/dynamic_internal_value_id_t.dtg.h", + "task-spec/dynamic_graph/dynamic_external_value_id_t.dtg.h", +] + +[[values]] +type = "::FlexFlow::dynamic_internal_value_id_t" +key = "internal" + +[[values]] +type = "::FlexFlow::dynamic_external_value_id_t" +key = "external" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml index bdcc93251c..fac58715de 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/external_dynamic_slot_site.dtg.toml @@ -3,17 +3,19 @@ name = "ExternalDynamicSlotSite" type = "struct" features = [ "eq", + "ord", "hash", "fmt", + "json", ] includes = [ - "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h", + "task-spec/dynamic_graph/dynamic_external_value_id_t.dtg.h", ] src_includes = [ ] [[fields]] -name = "value" -type = "::FlexFlow::DynamicValueAttrs" +name = "value_id" +type = "::FlexFlow::dynamic_external_value_id_t" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml index bf29bf2b2f..b0164a6086 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.toml @@ -3,22 +3,24 @@ name = "InternalDynamicSlotSite" type = "struct" features = [ "eq", + "ord", "hash", "fmt", + "json", ] includes = [ - "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h", "pcg/tensor_direction.dtg.h", "task-spec/dynamic_graph/dynamic_tensor_slot.h", + "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h", ] src_includes = [ ] [[fields]] -name = "invocation" -type = "::FlexFlow::DynamicNodeInvocation" +name = "invocation_id" +type = "::FlexFlow::dynamic_invocation_id_t" [[fields]] name = "direction" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion_result.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion_result.dtg.toml index 4c2c316d1d..8bc54b1151 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion_result.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/loss_insertion_result.dtg.toml @@ -3,6 +3,7 @@ name = "LossInsertionResult" type = "struct" features = [ "eq", + "ord", "fmt", ] diff --git a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h index 3b09778a5f..cae4d37229 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h @@ -8,7 +8,8 @@ namespace FlexFlow { DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( - MappedParallelLayerInvocationInfo const &); + MappedParallelLayerInvocationInfo const &, + DeviceType device_type); DynamicNodeInvocation build_replicate_invocation(MappedParallelLayerInvocationInfo const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml index 2be195f7c1..734edcad1e 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.dtg.toml @@ -3,6 +3,7 @@ name = "ParallelTensorMapping" type = "struct" features = [ "eq", + "ord", "hash", "json", "fmt", diff --git a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h new file mode 100644 index 0000000000..14bd7600b1 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h @@ -0,0 +1,18 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_PARALLEL_TENSOR_MAPPING_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_PARALLEL_TENSOR_MAPPING_H + +#include "task-spec/dynamic_graph/parallel_tensor_mapping.dtg.h" + +namespace FlexFlow { + +global_device_id_t pt_mapping_get_device_for_coord(ParallelTensorMapping const &, + ParallelTensorSpaceCoordinate const &); +ParallelTensorSpaceCoordinate pt_mapping_get_coord_for_device(ParallelTensorMapping const &, + global_device_id_t const &); + +std::set pt_mapping_get_coord_set(ParallelTensorMapping const &); +std::set pt_mapping_get_device_set(ParallelTensorMapping const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h b/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h index ad07b2941f..8aeefca69c 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h @@ -3,20 +3,40 @@ #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h" namespace FlexFlow { -bool node_is_pass_expanded(DynamicNodeAttrs const &); -bool value_is_pass_expanded(DynamicValueAttrs const &); -bool slot_is_pass_expanded(DynamicTensorSlot const &); +void require_node_might_be_pass_expanded(DynamicNodeAttrs const &); +void require_node_might_not_be_pass_expanded(DynamicNodeAttrs const &); -bool node_is_ready_for_pass_expansion(DynamicNodeAttrs const &); -bool value_is_ready_for_pass_expansion(DynamicValueAttrs const &); -bool slot_is_ready_for_pass_expansion(DynamicTensorSlot const &); +void require_slot_is_not_pass_expanded(DynamicTensorSlot const &); -bool no_part_of_graph_is_pass_expanded(DynamicOpenDataflowGraph const &); -bool graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &); -bool graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &); +void require_value_is_pass_expanded(DynamicValueAttrs const &); +void require_value_is_not_pass_expanded(DynamicValueAttrs const &); + +void require_invocation_is_fully_pass_expanded(DynamicNodeInvocation const &); +void require_invocation_is_ready_for_pass_expansion(DynamicNodeInvocation const &); + +void require_graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &); +void require_graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &); + +std::set + determine_intermediate_values_needed_to_compute_gradients_of_value( + DynamicOpenDataflowGraph const &, + DynamicValueAttrs const &); + +std::set + determine_intermediate_values_needed_for_gradient_computation( + DynamicOpenDataflowGraph const &); + +std::set + determine_invocations_needed_in_backward_pass_for_gradient_computation( + DynamicOpenDataflowGraph const &); + +DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &, FwbTensorType); +DynamicValueAttrs pass_expand_value(DynamicValueAttrs const &, FwbTensorType); +DynamicNodeAttrs pass_expand_node(DynamicNodeAttrs const &, DynamicTaskType); DynamicNodeInvocation perform_fwd_pass_expansion_for_invocation(DynamicNodeInvocation const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.dtg.toml new file mode 100644 index 0000000000..d5e92e9b3a --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.dtg.toml @@ -0,0 +1,22 @@ +namespace = "FlexFlow" +name = "SerializableDynamicOpenDataflowGraph" +type = "struct" +features = [ + "eq", + "ord", + "fmt", + "json", +] + +includes = [ + "task-spec/dynamic_graph/serializable_dynamic_node_invocation.dtg.h", + "", +] + +src_includes = [ + "utils/fmt/set.h", +] + +[[fields]] +name = "invocations" +type = "std::set<::FlexFlow::SerializableDynamicNodeInvocation>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h new file mode 100644 index 0000000000..d1c38a6694 --- /dev/null +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h @@ -0,0 +1,16 @@ +#ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_SERIALIZABLE_DYNAMIC_OPEN_DATAFLOW_GRAPH_H +#define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_SERIALIZABLE_DYNAMIC_OPEN_DATAFLOW_GRAPH_H + +#include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" + +namespace FlexFlow { + +SerializableDynamicOpenDataflowGraph + dynamic_open_dataflow_graph_to_serializable(DynamicOpenDataflowGraph const &); +DynamicOpenDataflowGraph dynamic_open_dataflow_graph_from_serializable( + SerializableDynamicOpenDataflowGraph const &); + +} // namespace FlexFlow + +#endif diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml index 67d1a171c8..a72c674fda 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_value_attrs.dtg.toml @@ -33,6 +33,10 @@ type = "::FlexFlow::dynamic_tensor_guid_t" name = "parallel_tensor_shape" type = "std::optional<::FlexFlow::ParallelTensorShape>" +[[fields]] +name = "create_grad" +type = "std::optional" + [[fields]] name = "shard_coord" type = "std::optional<::FlexFlow::ParallelTensorSpaceCoordinate>" diff --git a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h index 4a93575b51..68253369d7 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h @@ -9,17 +9,15 @@ namespace FlexFlow { -[[nodiscard]] bool node_is_shard_expanded(DynamicNodeAttrs const &); -[[nodiscard]] bool value_is_shard_expanded(DynamicValueAttrs const &); -[[nodiscard]] bool invocation_is_fully_shard_expanded(DynamicNodeInvocation const &); +void require_node_is_shard_expanded(DynamicNodeAttrs const &); +void require_value_is_shard_expanded(DynamicValueAttrs const &); +void require_invocation_is_fully_shard_expanded(DynamicNodeInvocation const &); +void require_graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &); -[[nodiscard]] bool node_is_ready_for_shard_expansion(DynamicNodeAttrs const &); -[[nodiscard]] bool value_is_ready_for_shard_expansion(DynamicValueAttrs const &); -[[nodiscard]] bool invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &); - -[[nodiscard]] bool no_part_of_graph_is_shard_expanded(DynamicOpenDataflowGraph const &); -[[nodiscard]] bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &); -[[nodiscard]] bool graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &); +void require_node_is_ready_for_shard_expansion(DynamicNodeAttrs const &); +void require_value_is_ready_for_shard_expansion(DynamicValueAttrs const &); +void require_invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &); +void require_graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &); [[nodiscard]] DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( DynamicNodeAttrs const &, @@ -37,10 +35,10 @@ namespace FlexFlow { generate_shard_expansion_for_invocation(DynamicNodeInvocation const &); [[nodiscard]] std::set - perform_shard_expansion_for_invocation(DynamicNodeInvocation const &); + perform_shard_expansion_for_invocation(DynamicNodeInvocation const &); [[nodiscard]] DynamicOpenDataflowGraph - perform_shard_expansion(DynamicOpenDataflowGraph const &); + perform_shard_expansion(DynamicOpenDataflowGraph const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml index 7744f3775f..d1f44bed25 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/training_op_type.dtg.toml @@ -1,6 +1,7 @@ namespace = "FlexFlow" name = "TrainingOpType" type = "variant" +#include "task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h" features = [ "eq", "hash", diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 827e9f3dc1..544f249243 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -9,7 +9,7 @@ #include "task-spec/dynamic_graph/dynamic_node_invocation.h" #include "task-spec/dynamic_graph/dynamic_node_mapping.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" -#include "task-spec/dynamic_graph/dynamic_slot_site.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" #include "task-spec/dynamic_graph/dynamic_task_type.h" #include "task-spec/dynamic_graph/dynamic_tensor_slot.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" @@ -32,6 +32,9 @@ #include "task-spec/dynamic_graph/training_operation_attrs.h" #include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" #include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" +#include "utils/overload.h" +#include "utils/containers/binary_merge_disjoint_maps.h" +#include "utils/containers/get_only.h" namespace FlexFlow { @@ -43,102 +46,63 @@ bool value_is_mapped(DynamicValueAttrs const &n) { return n.mapping.has_value(); } -bool no_part_of_graph_is_copy_inserted(DynamicOpenDataflowGraph const &g) { - auto slot_is_mapped = [](DynamicTensorSlot const &) -> bool { return false; }; - for (DynamicNodeInvocation const &i : g.invocations) { - if (node_is_copy(i.node_attrs)) { - return false; - } - for (auto const &[slot, value] : i.inputs) { - if (value_is_mapped(value)) { - return false; - } - } - for (auto const &[slot, value] : i.outputs) { - if (value_is_mapped(value)) { - return false; - } - } - } - return true; -} - -bool graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { - auto node_is_any = [](DynamicNodeAttrs const &) -> bool { return true; }; - auto slot_is_mapped = [](DynamicTensorSlot const &) -> bool { return true; }; - - return full_dynamic_graph_satisfies( - g, node_is_any, value_is_mapped, slot_is_mapped); -} - void require_node_is_ready_for_copy_insertion(DynamicNodeAttrs const &node_attrs) { + ASSERT(node_attrs.op_attrs.has_value()); ASSERT(node_attrs.mapping.has_value()); } -void require_graph_is_ready_for_copy_insertion(DynamicOpenDataflowGraph const &g) { - auto require_slot_is_ready_for_copy_insertion = [](DynamicTensorSlot const &slot) -> void { - return; - }; +void require_value_is_ready_for_copy_insertion(DynamicValueAttrs const &v) { + ASSERT(!v.mapping.has_value(), v); +} - auto require_value_is_ready_for_copy_insertion = [](DynamicValueAttrs const &value_attrs) -> void { +void require_invocation_is_ready_for_copy_insertion(DynamicNodeInvocation const &i) { + auto require_slot_is_ready_for_copy_insertion = [](DynamicTensorSlot const &) { return; }; - require_full_dynamic_graph_satisfies( - g, - require_node_is_ready_for_copy_insertion, - require_value_is_ready_for_copy_insertion, - require_slot_is_ready_for_copy_insertion); + require_invocation_fully_satisfies(i, + require_node_is_ready_for_copy_insertion, + require_value_is_ready_for_copy_insertion, + require_slot_is_ready_for_copy_insertion); } -static DynamicValueAttrs map_dynamic_value_attrs_for_task_group( - DynamicTensorSlot const &slot, - DynamicValueAttrs const &value, - MappedOperatorTaskGroup const &mapping) { - DynamicValueAttrs result = value; - result.mapping = get_tensor_bindings_for_slot_name(mapping, slot.slot_name); - return result; +void require_graph_is_ready_for_copy_insertion(DynamicOpenDataflowGraph const &g) { + require_full_dynamic_graph_satisfies(g, require_invocation_is_ready_for_copy_insertion); } -static std::pair - filter_mapping_to_avoid_degenerate_copies(DynamicValueAttrs const &input, - DynamicValueAttrs const &output) { - std::set< - std::pair> - input_mapping = unstructured_relation_from_bidict(assert_unwrap(input.mapping).raw); - - std::set< - std::pair> - output_mapping = unstructured_relation_from_bidict(assert_unwrap(output.mapping).raw); - - // Exclude the point shared between the input and output mappings, because - // those will not result in actual copies once shard expansion is performed - std::set< - std::pair> - remove = set_intersection(input_mapping, output_mapping); - - DynamicValueAttrs filtered_input = input; - filtered_input.mapping = ParallelTensorMapping{ - bidict_from_unstructured_relation(set_difference(input_mapping, remove)), +void require_value_is_copy_inserted(DynamicValueAttrs const &v) { + ASSERT(v.mapping.has_value()); +} + +void require_invocation_is_fully_copy_inserted(DynamicNodeInvocation const &i) { + auto require_node_is_copy_inserted = [](DynamicNodeAttrs const &) { + return; }; - DynamicValueAttrs filtered_output = output; - filtered_output.mapping = ParallelTensorMapping{ - bidict_from_unstructured_relation(set_difference(output_mapping, remove)), + auto require_slot_is_copy_inserted = [](DynamicTensorSlot const &) { + return; }; - return std::pair{filtered_input, filtered_output}; + require_invocation_fully_satisfies(i, + require_node_is_copy_inserted, + require_value_is_copy_inserted, + require_slot_is_copy_inserted); +} + +void require_graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { + require_full_dynamic_graph_satisfies(g, require_invocation_is_fully_copy_inserted); } -std::unordered_map - get_mappings_for_invocation( - DynamicNodeInvocation const &i, - std::unordered_map const + +std::map + get_mappings_for_invocation_id( + dynamic_invocation_id_t const &i, + std::map const &mappings) { return filtermap_keys(mappings, [&](InternalDynamicSlotSite const &s) -> std::optional { - if (s.invocation == i) { + if (s.invocation_id == i) { return s.slot_name; } else { return std::nullopt; @@ -147,25 +111,29 @@ std::unordered_map } DynamicNodeInvocation apply_mappings_for_invocation( + dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i, - std::unordered_map const + std::map const &all_mappings) { - std::unordered_map i_mappings = - get_mappings_for_invocation(i, all_mappings); - std::unordered_map + require_invocation_is_ready_for_copy_insertion(i); + + std::map i_mappings = + get_mappings_for_invocation_id(id, all_mappings); + + std::map i_input_mappings = restrict_keys(i_mappings, keys(i.inputs)); - std::unordered_map + std::map i_output_mappings = restrict_keys(i_mappings, keys(i.outputs)); auto apply_mapping = [&](DynamicValueAttrs const &v, ParallelTensorMapping const &mapping) -> DynamicValueAttrs { - return dynamic_value_attrs_with_mapping(v, mapping); + return decide_dynamic_value_attrs_mapping(v, mapping); }; - return DynamicNodeInvocation{ + DynamicNodeInvocation result = DynamicNodeInvocation{ /*inputs=*/ zip_values_strict_with(i.inputs, i_input_mappings, apply_mapping), /*node_attrs=*/ @@ -173,164 +141,282 @@ DynamicNodeInvocation apply_mappings_for_invocation( /*outputs=*/ zip_values_strict_with(i.outputs, i_output_mappings, apply_mapping), }; + + require_invocation_is_fully_copy_inserted(result); + + return result; +} + +DynamicNodeInvocation make_copy_invocation(DynamicValueCopyInfo const ©_info) { + DynamicNodeInvocation result = DynamicNodeInvocation{ + /*inputs=*/{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + decide_dynamic_value_attrs_mapping(copy_info.value_attrs, copy_info.src_mapping), + }, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + decide_dynamic_value_attrs_mapping(copy_info.value_attrs, copy_info.dst_mapping), + }, + }, + }; + + require_invocation_is_fully_copy_inserted(result); + + return result; } -std::unordered_set copies_for_value( - DynamicOpenDataflowGraph const &g, +std::set copies_for_value( DynamicValueAttrs const &v, - std::unordered_map const - &mappings) { - InternalDynamicSlotSite src = ({ - DynamicSlotSite found = dynamic_graph_find_source_of_value(g, v); + DynamicSlotSite const &src_site, + std::set const &dst_sites, + std::map const &all_mappings) { - if (found.is_external()) { + require_value_is_ready_for_copy_insertion(v); + + return src_site.visit>(overload { + [&](ExternalDynamicSlotSite const &) -> std::set { return {}; - } + }, + [&](InternalDynamicSlotSite const &s) -> std::set { + ParallelTensorMapping src_mapping = all_mappings.at(s); + std::map sink_site_mappings = + restrict_keys(all_mappings, dst_sites); - found.require_internal(); + return copies_for_internal_value(v, s, src_mapping, sink_site_mappings); + } }); +} - std::unordered_set sinks = - dynamic_graph_find_sinks_of_value(g, DynamicValueAttrs{v}); - - ParallelTensorMapping src_mapping = mappings.at(src); +std::set copies_for_internal_value( + DynamicValueAttrs const &v, + InternalDynamicSlotSite const &src_site, + ParallelTensorMapping const &src_mapping, + std::map const &sink_site_mappings) { - std::unordered_map - mappings_for_sinks = generate_map( - sinks, - [&](InternalDynamicSlotSite const &s) -> ParallelTensorMapping { - return mappings.at(s); - }); + require_value_is_ready_for_copy_insertion(v); - std::unordered_set sink_mapping_set = - unordered_set_of(values(mappings_for_sinks)); + std::set sink_mapping_set = + set_of(values(sink_site_mappings)); - std::unordered_set required_copies = - set_difference(sink_mapping_set, std::unordered_set{src_mapping}); + std::set required_copies = + set_difference(sink_mapping_set, std::set{src_mapping}); auto make_copy_to = - [&](ParallelTensorMapping const &sink_mapping) -> DynamicNodeInvocation { - return DynamicNodeInvocation{ - /*inputs=*/{ - { - DynamicTensorSlot{ - TensorSlotName::INPUT, - src.slot_name.slot_tensor_role, - }, - dynamic_value_attrs_with_mapping(v, src_mapping), - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/transform( - src.slot_name.slot_tensor_role, - dynamic_task_type_from_tensor_role_for_copy), - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - DynamicTensorSlot{ - TensorSlotName::OUTPUT, - src.slot_name.slot_tensor_role, - }, - dynamic_value_attrs_with_mapping(v, sink_mapping), - }, - }, + [&](ParallelTensorMapping const &sink_mapping) -> DynamicValueCopyInfo { + return DynamicValueCopyInfo{ + /*value_attrs=*/v, + /*src_mapping=*/src_mapping, + /*sink_mapping=*/sink_mapping, }; }; return transform(required_copies, make_copy_to); } -std::unordered_map - resolve_tensor_mappings_from_node_mappings( +std::map + resolve_tensor_mappings(DynamicOpenDataflowGraph const &g) +{ + require_graph_is_ready_for_copy_insertion(g); + + std::map resolved_from_node_mappings = + resolve_partial_tensor_mappings_from_node_mappings(g); + + std::map resolved_from_adjacent_values = + resolve_missing_tensor_mappings_from_adjacent_values(g, resolved_from_node_mappings); + + std::map result = + binary_merge_disjoint_maps(resolved_from_node_mappings, resolved_from_adjacent_values); + + { + std::set all_internal_slot_sites = + get_internal_dynamic_slot_sites(g); + std::set resolved_slot_sites = + keys(result); + ASSERT(resolved_slot_sites == all_internal_slot_sites); + } + + return result; +} + +std::map + resolve_partial_tensor_mappings_from_node_mappings( DynamicOpenDataflowGraph const &g) { - auto get_mappings_for_invocation = [&](DynamicNodeInvocation const &i) - -> std::unordered_map { + require_graph_is_ready_for_copy_insertion(g); + + auto slots_to_map_for_replicate = [](dynamic_invocation_id_t const &invocation_id, + DynamicNodeInvocation const &invocation) + -> std::set + { + TrainingOpType op_type = dynamic_node_invocation_get_op_type(invocation); + + ASSERT(op_type == TrainingOpType{OperatorType::REPLICATE}); + + std::set slot_sites = + (invocation.node_attrs.task_type == DynamicTaskType::BWD) + ? get_incoming_dynamic_slot_sites_for_invocation(invocation_id, invocation) + : get_output_dynamic_slot_sites_for_invocation(invocation_id, invocation); + + { + InternalDynamicSlotSite slot_site = get_only(slot_sites); + ASSERT(slot_site.slot_name.slot_name == TensorSlotName::OUTPUT); + }; + + return slot_sites; + }; + + auto get_mappings_for_invocation = [&](DynamicNodeInvocation const &invocation) + -> std::map { + + TrainingOpType op_type = dynamic_node_invocation_get_op_type(invocation); + dynamic_invocation_id_t invocation_id = dynamic_graph_get_id_for_invocation(g, invocation); + + std::set slot_sites_to_resolve = [&] { + if (op_type == TrainingOpType{OperatorType::REPLICATE}) { + return slots_to_map_for_replicate(invocation_id, invocation); + } else { + return get_dynamic_slot_sites_for_invocation(invocation_id, invocation); + } + }(); + return generate_map( - get_dynamic_slot_sites_for_invocation(i), + slot_sites_to_resolve, [&](InternalDynamicSlotSite const &s) -> ParallelTensorMapping { return ParallelTensorMapping{ dynamic_node_mapping_bindings_for_slot_name( - assert_unwrap(i.node_attrs.mapping), s.slot_name.slot_name), + assert_unwrap(invocation.node_attrs.mapping), s.slot_name.slot_name), }; }); }; - std::unordered_map result = + std::map result = merge_disjoint_maps(transform(get_dynamic_invocation_set(g), get_mappings_for_invocation)); return result; } -std::set perform_copy_insertion_for_invocation( - DynamicNodeInvocation const &i, - std::map const - &unmapped_value_to_mapped_source_value) { +std::map + resolve_missing_tensor_mappings_from_adjacent_values( + DynamicOpenDataflowGraph const &g, + std::map const &resolved_mappings) +{ + require_graph_is_ready_for_copy_insertion(g); + + std::set all_internal_slot_sites = + get_internal_dynamic_slot_sites(g); + + std::set missing_mappings = + set_minus(all_internal_slot_sites, keys(resolved_mappings)); + + auto get_mapping_for_slot_site_from_adjacent_values = + [&](InternalDynamicSlotSite const &slot_site) + -> ParallelTensorMapping { + + DynamicNodeInvocation invocation = dynamic_graph_get_invocation_for_id(g, slot_site.invocation_id); - MappedOperatorTaskGroup mapping = assert_unwrap(i.node_attrs.mapping); + TrainingOpType op_type = dynamic_node_invocation_get_op_type(invocation); + std::optional task_type = invocation.node_attrs.task_type; - auto map_tensor = [&](DynamicTensorSlot const &slot, - DynamicValueAttrs const &value) { - return map_dynamic_value_attrs_for_task_group(slot, value, mapping); + TrainingOpType replicate_op_type = TrainingOpType{OperatorType::REPLICATE}; + + if (op_type == replicate_op_type && task_type == DynamicTaskType::BWD) { + ASSERT(slot_site.direction == TensorDirection::OUTPUT); + + InternalDynamicSlotSite slot_site_sink = + get_only(dynamic_graph_find_sinks_of_slot_site(g, slot_site)); + + ASSERT(contains_key(resolved_mappings, slot_site_sink)); + + return resolved_mappings.at(slot_site_sink); + } else if ( + op_type == replicate_op_type + && (task_type == std::nullopt || task_type == DynamicTaskType::FWD) + ) { + ASSERT(slot_site.direction == TensorDirection::INCOMING); + + InternalDynamicSlotSite slot_site_src = + dynamic_graph_find_source_of_slot_site(g, slot_site).require_internal(); + + ASSERT(contains_key(resolved_mappings, slot_site_src)); + + return resolved_mappings.at(slot_site_src); + } else { + PANIC("Unhandled case"); + } }; - DynamicNodeInvocation mapped_i = [&] { - std::map mapped_inputs = - map_values2(i.inputs, map_tensor); - std::map mapped_outputs = - map_values2(i.outputs, map_tensor); + return generate_map(missing_mappings, + get_mapping_for_slot_site_from_adjacent_values); +} + +std::set infer_all_copies_in_graph(DynamicOpenDataflowGraph const &g) +{ + std::map + fully_resolved_tensor_mappings = + resolve_tensor_mappings(g); + + std::set all_copies = + flatmap(set_of(get_dynamic_values(g)), + [&](DynamicValueAttrs const &v) + -> std::set { - DynamicNodeInvocation r = i; - r.inputs = mapped_inputs; - r.outputs = mapped_outputs; - return r; - }(); + DynamicSlotSite src_site = dynamic_graph_find_source_of_value(g, v); + std::set sinks = dynamic_graph_find_sinks_of_value(g, v); - std::set result = set_union( - copies_for_invocation_inputs(i, unmapped_value_to_mapped_source_value), - std::set{ - mapped_i, - }); + return copies_for_value(v, src_site, sinks, fully_resolved_tensor_mappings); + }); - return result; + return all_copies; } DynamicOpenDataflowGraph perform_copy_insertion(DynamicOpenDataflowGraph const &g) { - ASSERT(no_part_of_graph_is_copy_inserted(g)); require_graph_is_ready_for_copy_insertion(g); std::map fully_resolved_tensor_mappings = - resolve_tensor_mappings_from_node_mappings(g); + resolve_tensor_mappings(g); - std::set all_copies = - flatmap(set_of(get_dynamic_values(g)), - [&](DynamicValueAttrs const &v) - -> std::set { - return copies_for_value(g, v, fully_resolved_tensor_mappings); - }); + std::set all_copies = infer_all_copies_in_graph(g); + + std::set all_copy_invocations = transform(all_copies, make_copy_invocation); std::set mapped_invocations = transform( get_dynamic_invocation_set(g), [&](DynamicNodeInvocation const &i) -> DynamicNodeInvocation { - return apply_mappings_for_invocation(i, fully_resolved_tensor_mappings); + dynamic_invocation_id_t id = dynamic_graph_get_id_for_invocation(g, i); + + return apply_mappings_for_invocation(id, i, fully_resolved_tensor_mappings); }); DynamicOpenDataflowGraph result = dynamic_open_dataflow_graph_from_invocation_set( - set_union(all_copies, mapped_invocations)); + set_union(all_copy_invocations, mapped_invocations)); - ASSERT(graph_is_fully_copy_inserted(result)); + require_graph_is_fully_copy_inserted(result); return result; } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc index 788c010084..1ddd4b3a51 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_graph_edge.cc @@ -12,7 +12,7 @@ DynamicGraphEdge return DynamicGraphEdge{ /*src=*/src, - /*dst_node=*/dst.invocation, + /*dst_node=*/dst.invocation_id, /*dst_slot=*/dst.slot_name, }; } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc index 5e1350d136..d10224a947 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc @@ -3,7 +3,6 @@ #include "task-spec/dynamic_graph/training_operation_attrs.h" #include "utils/containers/are_disjoint.h" #include "utils/containers/set_union.h" -#include "utils/containers/unordered_set_of.h" #include "utils/optional.h" #include "utils/containers/values.h" #include "utils/containers/keys.h" @@ -38,7 +37,7 @@ void require_invocation_fully_satisfies(DynamicNodeInvocation const &i, } } -std::unordered_map +std::map get_slot_map_for_direction(DynamicNodeInvocation const &invocation, TensorDirection direction) { switch (direction) { @@ -59,34 +58,49 @@ TrainingOpType return training_op_attrs_get_op_type(training_op_attrs); } -std::unordered_set - get_dynamic_slot_sites_for_invocation(DynamicNodeInvocation const &i) { +std::set + get_incoming_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { - std::unordered_set input_slots = - transform(unordered_set_of(i.inputs), + std::set incoming_slots = + transform(set_of(i.inputs), [&](std::pair const &p) -> InternalDynamicSlotSite { return InternalDynamicSlotSite{ - /*invocation=*/i, + /*invocation=*/id, /*direction=*/TensorDirection::INCOMING, /*slot_name=*/p.first, }; }); - std::unordered_set output_slots = - transform(unordered_set_of(i.outputs), + return incoming_slots; +} + +std::set + get_output_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { + + std::set output_slots = + transform(set_of(i.outputs), [&](std::pair const &p) -> InternalDynamicSlotSite { return InternalDynamicSlotSite{ - /*invocation=*/i, + /*invocation=*/id, /*direction=*/TensorDirection::OUTPUT, /*slot_name=*/p.first, }; }); - ASSERT(are_disjoint(input_slots, output_slots)); + return output_slots; +} + +std::set + get_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { + + std::set incoming_slots = get_incoming_dynamic_slot_sites_for_invocation(id, i); + std::set output_slots = get_output_dynamic_slot_sites_for_invocation(id, i); + + ASSERT(are_disjoint(incoming_slots, output_slots)); - return set_union(input_slots, output_slots); + return set_union(incoming_slots, output_slots); } } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc index b2a6e71af8..610d6609c3 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc @@ -1,23 +1,46 @@ #include "task-spec/dynamic_graph/dynamic_node_mapping.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/containers/transform.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" namespace FlexFlow { +bidict + dynamic_node_mapping_get_shard_bindings(DynamicNodeMapping const &m) +{ + return bidict_transform_keys( + m.op_task_group.get_shard_bindings(), + [&](MachineSpaceCoordinate const &mc) -> global_device_id_t { + return global_device_id_t{ + /*coord=*/mc, + /*device_type=*/m.device_type, + }; + }); +} + +OperatorAtomicTaskShardBinding + dynamic_node_mapping_get_shard_binding_for_device(DynamicNodeMapping const &mapping, + global_device_id_t const &device_id) +{ + ASSERT(device_id.device_type == mapping.device_type); + + return mapping.op_task_group.get_shard_bindings().at_l(device_id.coord); +} + bidict dynamic_node_mapping_bindings_for_slot_name( DynamicNodeMapping const &mapping, TensorSlotName const &slot_name) { bidict coord_bindings = get_tensor_bindings_for_slot_name(mapping.op_task_group, slot_name); - return transform_values( + return bidict_transform_values( coord_bindings, [&](MachineSpaceCoordinate const &coord) -> global_device_id_t { return global_device_id_t{coord, mapping.device_type}; }); } -std::unordered_set +std::set target_devices_of_dynamic_node_mapping(DynamicNodeMapping const &mapping) { return transform(mapping.op_task_group.get_shard_bindings().left_values(), diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index cab244e688..af6c288288 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -4,7 +4,7 @@ #include "op-attrs/pcg_operator_attrs.h" #include "task-spec/dynamic_graph/dynamic_graph_edge.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.h" -#include "task-spec/dynamic_graph/dynamic_slot_site.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" #include "task-spec/dynamic_graph/serializable_dynamic_node_attrs.h" #include "task-spec/dynamic_graph/serializable_dynamic_value_attrs.h" #include "utils/containers/all_of.h" @@ -31,6 +31,8 @@ #include "utils/many_to_one/many_to_one.h" #include "utils/containers/require_all_of.h" #include "utils/containers/multiset_of.h" +#include "utils/containers/repeat.h" +#include "utils/containers/at_idx.h" namespace FlexFlow { @@ -42,7 +44,7 @@ DynamicOpenDataflowGraph make_empty_dynamic_open_dataflow_graph() { void check_dynamic_open_dataflow_graph_is_valid( DynamicOpenDataflowGraph const &g) { - std::unordered_map> + std::map> invocations_by_value_produced; for (DynamicNodeInvocation const &i : g.invocations) { @@ -51,7 +53,7 @@ void check_dynamic_open_dataflow_graph_is_valid( } } - std::unordered_map> + std::map> values_produced_multiple_times = filter_values( invocations_by_value_produced, [](std::vector const &producers) -> bool { @@ -103,16 +105,11 @@ bool no_part_of_dynamic_graph_satisfies( void require_full_dynamic_graph_satisfies( DynamicOpenDataflowGraph const &g, - std::function const &node_condition, - std::function const &value_condition, - std::function const &slot_condition) + std::function const &invocation_condition) { - require_all_of(get_dynamic_nodes(g), node_condition); - require_all_of(get_dynamic_values(g), value_condition); - require_all_of(get_dynamic_tensor_slots(g), slot_condition); + require_all_of(g.invocations, invocation_condition); } - std::multiset get_dynamic_nodes(DynamicOpenDataflowGraph const &g) { return transform(multiset_of(g.invocations), @@ -145,44 +142,147 @@ std::set return g.invocations; } -std::unordered_set +std::set + dynamic_graph_get_internal_values(DynamicOpenDataflowGraph const &g) +{ + std::set internal_slot_sites = + get_internal_dynamic_slot_sites(g); + + std::set internal_values = + filtrans(internal_slot_sites, + [&](InternalDynamicSlotSite const &s) + -> std::optional { + if (s.direction == TensorDirection::OUTPUT) { + return dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{s}); + } else { + return std::nullopt; + } + }); + + return internal_values; +} + +std::set + dynamic_graph_get_external_values(DynamicOpenDataflowGraph const &g) +{ + std::set all_values = + set_of(get_dynamic_values(g)); + + std::set internal_values = + dynamic_graph_get_internal_values(g); + + return set_minus(all_values, internal_values); +} + +dynamic_invocation_id_t dynamic_graph_get_id_for_invocation( + DynamicOpenDataflowGraph const &g, + DynamicNodeInvocation const &invocation) +{ + return dynamic_invocation_id_t{ + nonnegative_int{assert_unwrap(index_of(g.invocations, invocation))}, + }; +} + +DynamicNodeInvocation dynamic_graph_get_invocation_for_id(DynamicOpenDataflowGraph const &g, + dynamic_invocation_id_t const &id) +{ + return at_idx(g.invocations, id.idx); +} + +dynamic_value_id_t dynamic_graph_get_id_for_value(DynamicOpenDataflowGraph const &g, + DynamicValueAttrs const &value) +{ + auto idx_in_set = [](std::set const &s, + DynamicValueAttrs const &v) + -> nonnegative_int + { + return nonnegative_int{assert_unwrap(index_of(s, v))}; + }; + + { + std::set internal_values = dynamic_graph_get_internal_values(g); + if (contains(internal_values, value)) { + return dynamic_value_id_t{ + dynamic_internal_value_id_t{ + idx_in_set(internal_values, value), + }, + }; + } + } + + { + std::set external_values = dynamic_graph_get_external_values(g); + if (contains(external_values, value)) { + return dynamic_value_id_t{ + dynamic_external_value_id_t{ + idx_in_set(external_values, value), + }, + }; + } + } + + PANIC("Could not find id for value {}", value); +} + +DynamicValueAttrs dynamic_graph_get_value_for_id(DynamicOpenDataflowGraph const &g, + dynamic_value_id_t const &id) +{ + return id.visit(overload { + [&](dynamic_internal_value_id_t const &internal_id) -> DynamicValueAttrs { + std::set internal_values = + dynamic_graph_get_internal_values(g); + + return at_idx(internal_values, internal_id.idx); + }, + [&](dynamic_external_value_id_t const &external_id) -> DynamicValueAttrs { + std::set external_values = + dynamic_graph_get_external_values(g); + + return at_idx(external_values, external_id.idx); + } + }); +} + +std::set get_dynamic_graph_edges(DynamicOpenDataflowGraph const &g) { return flatmap(get_dynamic_invocation_set(g), [&](DynamicNodeInvocation const &i) - -> std::unordered_set { + -> std::set { return get_dynamic_graph_edges_incoming_to_invocation(g, i); }); } -std::unordered_set +std::set get_dynamic_graph_edges_incoming_to_invocation( DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { - return transform(unordered_set_of(i.inputs), - [&](std::pair const &p) - -> DynamicGraphEdge { - DynamicSlotSite src = - dynamic_graph_find_source_of_value(g, p.second); - - InternalDynamicSlotSite dst = InternalDynamicSlotSite{ - /*invocation=*/i, - /*direction=*/TensorDirection::INCOMING, - /*slot_name=*/p.first, - }; - - return dynamic_graph_edge_from_slot_sites(src, dst); - }); + return transform( + set_of(i.inputs), + [&](std::pair const &p) + -> DynamicGraphEdge { + DynamicSlotSite src = + dynamic_graph_find_source_of_value(g, p.second); + + InternalDynamicSlotSite dst = InternalDynamicSlotSite{ + /*invocation_id=*/dynamic_graph_get_id_for_invocation(g, i), + /*direction=*/TensorDirection::INCOMING, + /*slot_name=*/p.first, + }; + + return dynamic_graph_edge_from_slot_sites(src, dst); + }); } -std::unordered_set +std::set get_dynamic_graph_edges_outgoing_from_invocation( DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { return flatmap( - unordered_set_of(i.outputs), + set_of(i.outputs), [&](std::pair const &p) - -> std::unordered_set { + -> std::set { + DynamicSlotSite src = DynamicSlotSite{ InternalDynamicSlotSite{ - /*invocation=*/i, + /*invocation_id=*/dynamic_graph_get_id_for_invocation(g, i), /*direction=*/TensorDirection::OUTPUT, /*slot_name=*/p.first, }, @@ -196,41 +296,56 @@ std::unordered_set }); } -std::unordered_set +DynamicValueAttrs dynamic_value_attrs_for_slot_site(DynamicOpenDataflowGraph const &g, + DynamicSlotSite const &slot) { + return slot.visit(overload{ + + [&](ExternalDynamicSlotSite const &external_slot) -> DynamicValueAttrs { + dynamic_value_id_t value_id = dynamic_value_id_t{external_slot.value_id}; + + return dynamic_graph_get_value_for_id(g, value_id); + }, + + [&](InternalDynamicSlotSite const &internal_slot) -> DynamicValueAttrs { + DynamicNodeInvocation invocation = dynamic_graph_get_invocation_for_id(g, internal_slot.invocation_id); + switch (internal_slot.direction) { + case TensorDirection::INCOMING: + return invocation.inputs.at(internal_slot.slot_name); + case TensorDirection::OUTPUT: + return invocation.outputs.at(internal_slot.slot_name); + default: + PANIC("Unexpected direction {}", internal_slot.direction); + } + }}); +} + +std::set get_internal_dynamic_slot_sites(DynamicOpenDataflowGraph const &g) { return flatmap(get_dynamic_invocation_set(g), - [](DynamicNodeInvocation const &i) - -> std::unordered_set { - return get_dynamic_slot_sites_for_invocation(i); + [&](DynamicNodeInvocation const &i) + -> std::set { + + dynamic_invocation_id_t id = dynamic_graph_get_id_for_invocation(g, i); + + return get_dynamic_slot_sites_for_invocation(id, i); }); } -std::unordered_set +std::set get_dynamic_slot_sites(DynamicOpenDataflowGraph const &g) { - std::unordered_set internal_slot_sites = - get_internal_dynamic_slot_sites(g); - std::unordered_set internal_values = - filtrans(internal_slot_sites, - [&](InternalDynamicSlotSite const &s) - -> std::optional { - if (s.direction == TensorDirection::OUTPUT) { - return dynamic_value_attrs_for_slot_site(DynamicSlotSite{s}); - } else { - return std::nullopt; - } - }); - - std::unordered_set all_values = - unordered_set_of(get_dynamic_values(g)); + std::set internal_slot_sites = + get_internal_dynamic_slot_sites(g); - std::unordered_set external_values = - set_minus(all_values, internal_values); + std::set external_values = + dynamic_graph_get_external_values(g); - std::unordered_set external_slot_sites = transform( + std::set external_slot_sites = transform( external_values, - [](DynamicValueAttrs const &external_value) -> ExternalDynamicSlotSite { - return ExternalDynamicSlotSite{external_value}; + [&](DynamicValueAttrs const &external_value) -> ExternalDynamicSlotSite { + dynamic_external_value_id_t value_id = + dynamic_graph_get_id_for_value(g, external_value).require_external(); + return ExternalDynamicSlotSite{value_id}; }); return set_union( @@ -244,36 +359,56 @@ std::unordered_set })); } -std::unordered_set +std::set dynamic_graph_find_sinks_of_value(DynamicOpenDataflowGraph const &g, DynamicValueAttrs const &v) { - std::unordered_set found = filter( + std::set found = filter( get_internal_dynamic_slot_sites(g), [&](InternalDynamicSlotSite const &s) -> bool { - return dynamic_value_attrs_for_slot_site(DynamicSlotSite{s}) == v && + return dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{s}) == v && s.direction == TensorDirection::INCOMING; }); return found; } +DynamicSlotSite + dynamic_graph_find_source_of_slot_site(DynamicOpenDataflowGraph const &g, + InternalDynamicSlotSite const &slot_site) +{ + DynamicValueAttrs value_attrs = dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{slot_site}); + DynamicSlotSite src_site = dynamic_graph_find_source_of_value(g, value_attrs); + return src_site; +} + +std::set + dynamic_graph_find_sinks_of_slot_site(DynamicOpenDataflowGraph const &g, + InternalDynamicSlotSite const &slot_site) +{ + DynamicValueAttrs value_attrs = dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{slot_site}); + std::set sink_sites = dynamic_graph_find_sinks_of_value(g, value_attrs); + return sink_sites; +} + DynamicSlotSite dynamic_graph_find_source_of_value(DynamicOpenDataflowGraph const &g, DynamicValueAttrs const &v) { + dynamic_value_id_t value_id = dynamic_graph_get_id_for_value(g, v); + auto is_source_of_value = [&](DynamicSlotSite const &s) -> bool { return s.visit(overload{ [&](InternalDynamicSlotSite const &internal_slot_site) -> bool { - return dynamic_value_attrs_for_slot_site(s) == v && + return dynamic_value_attrs_for_slot_site(g, s) == v && internal_slot_site.direction == TensorDirection::OUTPUT; }, [&](ExternalDynamicSlotSite const &external_slot_site) -> bool { - return external_slot_site.value == v; + return dynamic_value_id_t{external_slot_site.value_id} == value_id; }, }); }; - std::unordered_set found = + std::set found = filter(get_dynamic_slot_sites(g), is_source_of_value); return get_only(found); @@ -513,7 +648,7 @@ std::string auto render_parallel_tensor_mapping = [](ParallelTensorMapping const &mapping) -> RecordFormatter { - return mk_record_for_map(mapping.raw.as_unordered_map()); + return mk_record_for_map(mapping.raw.as_map()); }; std::function render_value_label = diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc deleted file mode 100644 index af13345abd..0000000000 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_slot_site.cc +++ /dev/null @@ -1,26 +0,0 @@ -#include "task-spec/dynamic_graph/dynamic_slot_site.h" -#include "utils/overload.h" - -namespace FlexFlow { - -DynamicValueAttrs - dynamic_value_attrs_for_slot_site(DynamicSlotSite const &slot) { - return slot.visit(overload{ - - [](ExternalDynamicSlotSite const &external_slot) -> DynamicValueAttrs { - return external_slot.value; - }, - - [](InternalDynamicSlotSite const &internal_slot) -> DynamicValueAttrs { - switch (internal_slot.direction) { - case TensorDirection::INCOMING: - return internal_slot.invocation.inputs.at(internal_slot.slot_name); - case TensorDirection::OUTPUT: - return internal_slot.invocation.outputs.at(internal_slot.slot_name); - default: - PANIC("Unexpected direction {}", internal_slot.direction); - } - }}); -} - -} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc index ef05974ccc..0d239a9069 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc @@ -12,4 +12,11 @@ DynamicTensorSlot decide_tensor_slot_role(DynamicTensorSlot const &slot, return result; } +DynamicTensorSlot slot_without_task_shard(DynamicTensorSlot const &s) { + DynamicTensorSlot result = s; + result.task_shard = std::nullopt; + return result; +} + + } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc index 24e98d90a8..bd10c4f674 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc @@ -24,6 +24,7 @@ LossInsertionResult perform_loss_insertion( DynamicValueAttrs label_value{ /*tensor_guid=*/mk_dynamic_tensor_guid_for_loss(), /*parallel_tensor_shape=*/logit_value.parallel_tensor_shape, + /*create_grad=*/false, /*shard_coord=*/logit_value.shard_coord, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -33,6 +34,7 @@ LossInsertionResult perform_loss_insertion( DynamicValueAttrs logit_grad_value{ /*tensor_guid=*/logit_value.tensor_guid, /*parallel_tensor_shape=*/logit_value.parallel_tensor_shape, + /*create_grad=*/logit_value.create_grad, /*shard_coord=*/logit_value.shard_coord, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc index 5a7f21a26b..d88f35c3cb 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc @@ -45,6 +45,7 @@ DynamicOpenDataflowGraph DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, /*parallel_tensor_shape=*/lift_to_parallel(attrs.shape), + /*create_grad=*/(attrs.create_grad == CreateGrad::YES), /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -67,6 +68,7 @@ DynamicOpenDataflowGraph DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, /*parallel_tensor_shape=*/lift_to_parallel(attrs.shape), + /*create_grad=*/(attrs.create_grad == CreateGrad::YES), /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 285b02d58d..9e11337368 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -22,12 +22,16 @@ namespace FlexFlow { DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( - MappedParallelLayerInvocationInfo const &invocation_info) + MappedParallelLayerInvocationInfo const &invocation_info, + DeviceType device_type) { DynamicNodeAttrs result_attrs{ /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/invocation_info.layer_info.mapping, + /*device_ids=*/std::nullopt, + /*mapping=*/DynamicNodeMapping{ + /*op_task_group=*/invocation_info.layer_info.mapping, + /*device_type=*/device_type, + }, /*op_attrs=*/TrainingOperationAttrs{invocation_info.layer_info.attrs.op_attrs}, /*pcg_layer_guid=*/dynamic_layer_guid_t{invocation_info.layer_info.guid}, /*per_device_op_state=*/std::nullopt, @@ -47,6 +51,7 @@ DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{tensor.guid}, /*parallel_tensor_shape=*/tensor.attrs.shape, + /*create_grad=*/(tensor.attrs.create_grad == CreateGrad::YES), /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -71,10 +76,15 @@ DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( } DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( - MappedParallelComputationGraph const &mpcg) { + MappedParallelComputationGraph const &mpcg, DeviceType device_type) { return dynamic_open_dataflow_graph_from_invocation_set( - transform(set_of(mpcg_get_invocation_set(mpcg)), make_dynamic_node_invocation_from_mapped)); + transform(mpcg_get_invocation_set(mpcg), + [&](MappedParallelLayerInvocationInfo const &mpcg_invocation) + -> DynamicNodeInvocation + { + return make_dynamic_node_invocation_from_mapped(mpcg_invocation, device_type); + })); } } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc b/lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc new file mode 100644 index 0000000000..8d8d41a59d --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc @@ -0,0 +1,23 @@ +#include "task-spec/dynamic_graph/parallel_tensor_mapping.h" + +namespace FlexFlow { + +global_device_id_t pt_mapping_get_device_for_coord(ParallelTensorMapping const &m, + ParallelTensorSpaceCoordinate const &coord) { + return m.raw.at_l(coord); +} + +ParallelTensorSpaceCoordinate pt_mapping_get_coord_for_device(ParallelTensorMapping const &m, + global_device_id_t const &device) { + return m.raw.at_r(device); +} + +std::set pt_mapping_get_coord_set(ParallelTensorMapping const &m) { + return m.raw.left_values(); +} + +std::set pt_mapping_get_device_set(ParallelTensorMapping const &m) { + return m.raw.right_values(); +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index 28342ec2e9..2a188aeb09 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -6,51 +6,183 @@ #include "utils/containers/get_only.h" #include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/transform.h" +#include "utils/containers/repeat_until_converged.h" +#include "utils/containers/flatmap.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" +#include "utils/containers/map_values.h" namespace FlexFlow { -bool node_is_pass_expanded(DynamicNodeAttrs const &n) { - return n.task_type.has_value(); +void require_node_might_be_pass_expanded(DynamicNodeAttrs const &n) { + if (assert_unwrap(n.op_attrs).is_copy()) { + return; + } + + ASSERT(n.task_type.has_value(), n); +} + +void require_node_might_not_be_pass_expanded(DynamicNodeAttrs const &n) { + if (assert_unwrap(n.op_attrs).is_copy()) { + return; + } + + ASSERT(!n.task_type.has_value(), n); } -bool slot_is_pass_expanded(DynamicTensorSlot const &s) { - return s.slot_tensor_role.has_value(); +void require_slot_is_not_pass_expanded(DynamicTensorSlot const &s) { + ASSERT(!s.slot_tensor_role.has_value(), s); } -bool value_is_pass_expanded(DynamicValueAttrs const &v) { - return v.role.has_value(); +void require_value_is_pass_expanded(DynamicValueAttrs const &v) { + ASSERT(v.role.has_value(), v); } -bool node_is_ready_for_pass_expansion(DynamicNodeAttrs const &) { - return true; +void require_value_is_not_pass_expanded(DynamicValueAttrs const &v) { + ASSERT(!v.role.has_value(), v); } -bool value_is_ready_for_pass_expansion(DynamicValueAttrs const &) { - return true; +void require_invocation_is_fully_pass_expanded(DynamicNodeInvocation const &invocation) { + auto require_slot_is_pass_expanded = [&](DynamicTensorSlot const &s) { + if (dynamic_node_invocation_get_op_type(invocation) == TrainingOpType{TrainingOnlyOpType::COPY}) { + return; + } + + ASSERT(s.slot_tensor_role.has_value(), s); + }; + + require_invocation_fully_satisfies( + invocation, + require_node_might_be_pass_expanded, + require_value_is_pass_expanded, + require_slot_is_pass_expanded); } -bool slot_is_ready_for_pass_expansion(DynamicTensorSlot const &) { - return true; +void require_invocation_is_ready_for_pass_expansion(DynamicNodeInvocation const &invocation) { + require_invocation_fully_satisfies( + invocation, + require_node_might_not_be_pass_expanded, + require_value_is_not_pass_expanded, + require_slot_is_not_pass_expanded); } -bool no_part_of_graph_is_pass_expanded(DynamicOpenDataflowGraph const &g) { - return no_part_of_dynamic_graph_satisfies( - g, node_is_pass_expanded, value_is_pass_expanded, slot_is_pass_expanded); + +void require_graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &g) { + require_full_dynamic_graph_satisfies( + g, + require_invocation_is_fully_pass_expanded); } -bool graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &g) { - return full_dynamic_graph_satisfies( - g, node_is_pass_expanded, value_is_pass_expanded, slot_is_pass_expanded); +void require_graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &g) { + require_full_dynamic_graph_satisfies( + g, + require_invocation_is_ready_for_pass_expansion); } -bool graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &g) { - return full_dynamic_graph_satisfies( - g, node_is_ready_for_pass_expansion, value_is_ready_for_pass_expansion, slot_is_ready_for_pass_expansion); +std::set + determine_intermediate_values_needed_to_compute_gradients_of_value( + DynamicOpenDataflowGraph const &g, + DynamicValueAttrs const &val) +{ + auto get_values_immediately_needed_to_compute_gradients_of_needed = + [&](std::set const &needed) + -> std::set + { + std::set additional + = flatmap(needed, + [&](DynamicValueAttrs const &v) -> std::set { + std::set sinks = + dynamic_graph_find_sinks_of_value(g, v); + + return flatmap( + sinks, + [&](InternalDynamicSlotSite const &sink) -> std::set { + DynamicNodeInvocation sink_invocation + = dynamic_graph_get_invocation_for_id(g, sink.invocation_id); + + return set_of(values(sink_invocation.outputs)); + }); + }); + + return set_union(needed, additional); + }; + + std::set result = repeat_until_converged( + std::set{val}, + get_values_immediately_needed_to_compute_gradients_of_needed); + + ASSERT(contains(result, val)); + return result; } +std::set + determine_intermediate_values_needed_for_gradient_computation( + DynamicOpenDataflowGraph const &g) +{ + auto value_is_fundamentally_required = + [&](DynamicValueAttrs const &v) -> bool { + DynamicSlotSite source = dynamic_graph_find_source_of_value(g, v); + + if (source.is_external()) { + return true; + } + + InternalDynamicSlotSite internal_source = source.require_internal(); + ASSERT(internal_source.direction == TensorDirection::OUTPUT); + + DynamicNodeInvocation source_invocation = + dynamic_graph_get_invocation_for_id(g, internal_source.invocation_id); + + TrainingOpType op_type = dynamic_node_invocation_get_op_type(source_invocation); + + TrainingOpType weight_op_type = TrainingOpType{OperatorType::WEIGHT}; + TrainingOpType input_op_type = TrainingOpType{OperatorType::INPUT}; + + if (op_type == weight_op_type) { + return true; + } else if (op_type == input_op_type) { + return assert_unwrap(v.create_grad); + } else { + return false; + } + }; + + std::set fundamentally_required_values + = filter(set_of(get_dynamic_values(g)), value_is_fundamentally_required); + + std::set required_values = + flatmap(fundamentally_required_values, + [&](DynamicValueAttrs const &fundamentally_required_value) -> std::set { + return determine_intermediate_values_needed_to_compute_gradients_of_value( + g, fundamentally_required_value); + }); + + return required_values; +} + +std::set + determine_invocations_needed_in_backward_pass_for_gradient_computation( + DynamicOpenDataflowGraph const &g) +{ + std::set required_values = + determine_intermediate_values_needed_for_gradient_computation(g); + + auto get_sink_invocations_for_value = + [&](DynamicValueAttrs const &v) -> std::set { + return + transform( + dynamic_graph_find_sinks_of_value(g, v), + [&](InternalDynamicSlotSite const &sink_site) -> dynamic_invocation_id_t { + ASSERT(sink_site.direction == TensorDirection::INCOMING); + return sink_site.invocation_id; + }); + }; + + return flatmap(required_values, get_sink_invocations_for_value); +} + DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &s, FwbTensorType tensor_type) { - ASSERT(!slot_is_pass_expanded(s)); + require_slot_is_not_pass_expanded(s); DynamicTensorSlot result = s; result.slot_tensor_role = @@ -60,7 +192,7 @@ DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &s, DynamicValueAttrs pass_expand_value(DynamicValueAttrs const &v, FwbTensorType tensor_type) { - ASSERT(!value_is_pass_expanded(v)); + require_value_is_not_pass_expanded(v); DynamicValueAttrs result = v; result.role = DynamicTensorRole{tensor_type}; @@ -69,17 +201,35 @@ DynamicValueAttrs pass_expand_value(DynamicValueAttrs const &v, DynamicNodeAttrs pass_expand_node(DynamicNodeAttrs const &n, DynamicTaskType task_type) { - ASSERT(!node_is_pass_expanded(n)); + require_node_might_not_be_pass_expanded(n); + ASSERT(task_type == DynamicTaskType::FWD || task_type == DynamicTaskType::BWD); + { + TrainingOperationAttrs op_attrs = assert_unwrap(n.op_attrs); + + if (op_attrs.is_copy()) { + return n; + } + } + DynamicNodeAttrs result = n; result.task_type = task_type; return result; } DynamicNodeInvocation perform_fwd_pass_expansion_for_invocation( - DynamicNodeInvocation const &task) { + DynamicNodeInvocation const &invocation) { + + require_invocation_is_ready_for_pass_expansion(invocation); + + TrainingOperationAttrs op_attrs = + assert_unwrap(invocation.node_attrs.op_attrs); + + auto to_fwd_value = [](DynamicValueAttrs const &v) -> DynamicValueAttrs { + return pass_expand_value(v, FwbTensorType::FORWARD); + }; auto to_fwd = [](DynamicTensorSlot const &k, DynamicValueAttrs const &v) { return std::pair{ @@ -88,19 +238,37 @@ DynamicNodeInvocation perform_fwd_pass_expansion_for_invocation( }; }; - return DynamicNodeInvocation{ - /*inputs=*/ - transform(task.inputs, to_fwd), - /*node_attrs=*/ - pass_expand_node(task.node_attrs, DynamicTaskType::FWD), - /*outputs=*/ - transform(task.outputs, to_fwd), - }; + DynamicNodeInvocation result = [&]() + -> DynamicNodeInvocation + { + if (op_attrs.is_copy()) { + return DynamicNodeInvocation{ + /*inputs=*/map_values(invocation.inputs, to_fwd_value), + /*node_attrs=*/invocation.node_attrs, + /*outputs=*/map_values(invocation.outputs, to_fwd_value), + }; + } else { + return DynamicNodeInvocation{ + /*inputs=*/ + transform(invocation.inputs, to_fwd), + /*node_attrs=*/ + pass_expand_node(invocation.node_attrs, DynamicTaskType::FWD), + /*outputs=*/ + transform(invocation.outputs, to_fwd), + }; + } + }(); + + require_invocation_is_fully_pass_expanded(result); + + return result; } DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( DynamicNodeInvocation const &invocation) { + require_invocation_is_ready_for_pass_expansion(invocation); + TrainingOperationAttrs op_attrs = assert_unwrap(invocation.node_attrs.op_attrs); @@ -111,6 +279,10 @@ DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( }; }; + auto to_grad_value = [](DynamicValueAttrs const &v) { + return pass_expand_value(v, FwbTensorType::GRADIENT); + }; + auto to_grad = [](DynamicTensorSlot const &k, DynamicValueAttrs const &v) { return std::pair{ pass_expand_slot(k, FwbTensorType::GRADIENT), @@ -118,61 +290,73 @@ DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( }; }; - if (training_op_attrs_has_op_type(op_attrs, OperatorType::REPLICATE)) { - auto [input_slot, input] = get_only(invocation.inputs); - auto [output_slot, output] = get_only(invocation.outputs); - - DynamicNodeInvocation bwd{ - /*inputs=*/{ - to_fwd(output_slot, output), - to_grad(output_slot, output), - }, - /*node_attrs=*/ - pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), - /*outputs=*/ - { - to_grad(input_slot, input), - }, - }; - - return bwd; - } else { - return DynamicNodeInvocation{ - /*inputs=*/ - merge_disjoint_maps(std::vector{ - transform(invocation.inputs, to_fwd), - transform(invocation.outputs, to_fwd), + DynamicNodeInvocation result = [&]() + -> DynamicNodeInvocation + { + if (op_attrs.is_copy()) { + return DynamicNodeInvocation{ + /*inputs=*/map_values(invocation.outputs, to_grad_value), + /*node_attrs=*/invocation.node_attrs, + /*outputs=*/map_values(invocation.inputs, to_grad_value), + }; + } else if (training_op_attrs_has_op_type(op_attrs, OperatorType::REPLICATE)) { + return DynamicNodeInvocation{ + /*inputs=*/{ transform(invocation.outputs, to_grad), - }), - /*node_attrs=*/ - pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), - /*outputs=*/ - transform(invocation.inputs, to_grad), + }, + /*node_attrs=*/ + pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), + /*outputs=*/{ + transform(invocation.inputs, to_grad), + }, + }; + } else { + return DynamicNodeInvocation{ + /*inputs=*/ + merge_disjoint_maps(std::vector{ + transform(invocation.inputs, to_fwd), + transform(invocation.outputs, to_fwd), + transform(invocation.outputs, to_grad), + }), + /*node_attrs=*/ + pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), + /*outputs=*/ + transform(invocation.inputs, to_grad), + }; }; - }; + }(); + + require_invocation_is_fully_pass_expanded(result); + + return result; } DynamicOpenDataflowGraph perform_pass_expansion(DynamicOpenDataflowGraph const &g) { - ASSERT(no_part_of_graph_is_pass_expanded(g)); - ASSERT(graph_is_ready_for_pass_expansion(g)); + require_graph_is_ready_for_pass_expansion(g); + + std::set needed_in_bwd_pass = + determine_invocations_needed_in_backward_pass_for_gradient_computation(g); DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( - g, [](DynamicNodeInvocation const &invocation) { - if (invocation.inputs.empty()) { + g, + [&](DynamicNodeInvocation const &invocation) { + dynamic_invocation_id_t invocation_id = dynamic_graph_get_id_for_invocation(g, invocation); + + if (contains(needed_in_bwd_pass, invocation_id)) { return std::set{ perform_fwd_pass_expansion_for_invocation(invocation), + perform_bwd_pass_expansion_for_invocation(invocation), }; } else { return std::set{ perform_fwd_pass_expansion_for_invocation(invocation), - perform_bwd_pass_expansion_for_invocation(invocation), }; - }; + } }); - ASSERT(graph_is_fully_pass_expanded(result)); + require_graph_is_fully_pass_expanded(result); return result; } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc new file mode 100644 index 0000000000..5118637702 --- /dev/null +++ b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc @@ -0,0 +1,23 @@ +#include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" + +namespace FlexFlow { + +SerializableDynamicOpenDataflowGraph + dynamic_open_dataflow_graph_to_serializable(DynamicOpenDataflowGraph const &g) +{ + return SerializableDynamicOpenDataflowGraph{ + /*invocations=*/transform(g.invocations, dynamic_node_invocation_to_serializable), + }; +} + +DynamicOpenDataflowGraph dynamic_open_dataflow_graph_from_serializable( + SerializableDynamicOpenDataflowGraph const &serializable) +{ + return DynamicOpenDataflowGraph{ + /*invocations=*/transform(serializable.invocations, + dynamic_node_invocation_from_serializable), + }; +} + +} // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_value_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_value_attrs.cc index b4d398c3f0..ee53d01f21 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_value_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_value_attrs.cc @@ -8,6 +8,7 @@ SerializableDynamicValueAttrs return SerializableDynamicValueAttrs{ /*tensor_guid=*/attrs.tensor_guid, /*parallel_tensor_shape=*/attrs.parallel_tensor_shape, + /*create_grad=*/attrs.create_grad, /*shard_coord=*/attrs.shard_coord, /*mapping=*/attrs.mapping, /*role=*/attrs.role, @@ -19,6 +20,7 @@ DynamicValueAttrs dynamic_value_attrs_from_serializable( return DynamicValueAttrs{ /*tensor_guid=*/attrs.tensor_guid, /*parallel_tensor_shape=*/attrs.parallel_tensor_shape, + /*create_grad=*/attrs.create_grad, /*shard_coord=*/attrs.shard_coord, /*mapping=*/attrs.mapping, /*accessor=*/std::nullopt, diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index 516ee75a43..74034d869e 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -4,7 +4,6 @@ #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" #include "utils/bidict/algorithms/bidict_filter_keys.h" #include "task-spec/dynamic_graph/shard_expansion.h" -#include "utils/bidict/algorithms/filter_keys.h" #include "utils/containers/get_only.h" #include "utils/containers/map_values2.h" #include "utils/containers/require_same.h" @@ -21,109 +20,72 @@ #include "utils/containers/map_keys.h" #include "utils/containers/require_only_key.h" #include "utils/containers/set_of.h" +#include "task-spec/dynamic_graph/parallel_tensor_mapping.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" +#include "utils/binary_relation/binary_relation_transform_right2.h" +#include "utils/binary_relation/binary_relation_from_map.h" +#include "utils/containers/are_disjoint.h" +#include "utils/binary_relation/filter_binary_relation.h" +#include "utils/containers/map_from_pairs.h" +#include "utils/containers/flatmap.h" namespace FlexFlow { -bool node_is_shard_expanded(DynamicNodeAttrs const &n) { - return n.device_coords.has_value(); +void require_node_is_shard_expanded(DynamicNodeAttrs const &n) { + ASSERT(n.device_ids.has_value()); } -bool node_is_ready_for_shard_expansion(DynamicNodeAttrs const &n) { - if (!n.op_attrs.has_value()) { - return false; - } - - if (n.op_attrs.value().is_pcg_op()) { - if (!n.mapping.has_value()) { - return false; - } - } - - return true; +void require_value_is_shard_expanded(DynamicValueAttrs const &n) { + ASSERT(n.shard_coord.has_value()); } -void require_node_is_ready_for_shard_expansion(DynamicNodeAttrs const &n) { - ASSERT(n.op_attrs.has_value()); - if (n.op_attrs.value().is_pcg_op()) { - ASSERT(n.mapping.has_value()); - } -} - - -bool invocation_is_fully_shard_expanded(DynamicNodeInvocation const &i) { - auto slot_is_shard_expanded = [](DynamicTensorSlot const &) { - return true; +void require_invocation_is_fully_shard_expanded(DynamicNodeInvocation const &i) { + auto require_slot_is_shard_expanded = [](DynamicTensorSlot const &) { + return; }; - return invocation_fully_satisfies( + return require_invocation_fully_satisfies( i, - node_is_shard_expanded, - value_is_shard_expanded, - slot_is_shard_expanded); + require_node_is_shard_expanded, + require_value_is_shard_expanded, + require_slot_is_shard_expanded); } -bool value_is_shard_expanded(DynamicValueAttrs const &n) { - return n.shard_coord.has_value() && n.mapping.has_value(); +void require_graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { + return require_full_dynamic_graph_satisfies(g, require_invocation_is_fully_shard_expanded); } -bool value_is_ready_for_shard_expansion(DynamicValueAttrs const &n) { - return true; +void require_node_is_ready_for_shard_expansion(DynamicNodeAttrs const &n) { + ASSERT(n.op_attrs.has_value()); + + if (n.op_attrs.value().is_pcg_op()) { + ASSERT(n.mapping.has_value()); + } } void require_value_is_ready_for_shard_expansion(DynamicValueAttrs const &n) { return; } -bool invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &i) { - auto slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) { - return true; - }; - - return invocation_fully_satisfies( - i, - node_is_ready_for_shard_expansion, - value_is_ready_for_shard_expansion, - slot_is_ready_for_shard_expansion); -} - void require_invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &i) { - auto require_slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) -> void { + auto require_slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) { return; }; - require_invocation_fully_satisfies( + return require_invocation_fully_satisfies( i, require_node_is_ready_for_shard_expansion, require_value_is_ready_for_shard_expansion, require_slot_is_ready_for_shard_expansion); } -bool no_part_of_graph_is_shard_expanded(DynamicOpenDataflowGraph const &g) { - auto slot_is_shard_expanded = [](DynamicTensorSlot const &) -> bool { - return false; - }; - - return no_part_of_dynamic_graph_satisfies(g, - node_is_shard_expanded, - value_is_shard_expanded, - slot_is_shard_expanded); -} - -bool graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { - auto slot_is_shard_expanded = [](DynamicTensorSlot const &) -> bool { - return true; - }; - - return full_dynamic_graph_satisfies(g, - node_is_shard_expanded, - value_is_shard_expanded, - slot_is_shard_expanded); +void require_graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &g) { + require_full_dynamic_graph_satisfies(g, require_invocation_is_ready_for_shard_expansion); } -<<<<<<< HEAD static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( DynamicNodeInvocation const &i, - MachineSpaceCoordinate const &machine_coord, + global_device_id_t const &device_id, OperatorAtomicTaskShardBinding const &binding) { auto shard_expand_value_attrs = @@ -133,32 +95,45 @@ static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( return DynamicValueAttrsShardingInfo{ /*shard_coord=*/parallel_tensor_coord, - /*mapping=*/v.mapping.value().at_l(parallel_tensor_coord), + /*mapping=*/pt_mapping_get_device_for_coord(assert_unwrap(v.mapping), parallel_tensor_coord), }; }; DynamicNodeAttrs expanded_node_attrs = [&]() { DynamicNodeAttrs result = i.node_attrs; - result.device_coords = nonempty_set{machine_coord}; + result.device_ids = nonempty_set{device_id}; return result; }(); - return DynamicNodeInvocationShardingInfo{ - /*device_coord=*/nonempty_set{machine_coord}, - /*value_sharding=*/map_values2( - binary_merge_disjoint_maps(i.inputs, i.outputs), + DynamicNodeInvocationShardingInfo result = DynamicNodeInvocationShardingInfo{ + /*device_coord=*/nonempty_set{device_id}, + /*value_sharding=*/binary_relation_transform_right2( + binary_relation_from_map( + binary_merge_disjoint_maps(i.inputs, i.outputs)), shard_expand_value_attrs), }; -======= + + { + std::set invocation_slots = set_union( + keys(i.inputs), keys(i.outputs)); + + std::set sharding_info_slots + = result.value_sharding.left_values(); + + ASSERT(invocation_slots == sharding_info_slots); + } + + return result; +} + static bidict restrict_tensor_mapping_keys_to_coord( bidict const &mapping, ParallelTensorSpaceCoordinate const ¶llel_tensor_coord) { - return filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { + return bidict_filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { return p == parallel_tensor_coord; }); ->>>>>>> device-type-agnostic-compiler-only } static DynamicNodeInvocation shard_invocation_for_binding( @@ -187,14 +162,14 @@ static DynamicNodeInvocation shard_invocation_for_binding( DynamicNodeAttrs expanded_node_attrs = [&]() { DynamicNodeAttrs result = i.node_attrs; - result.device_ids = nonempty_set{machine_coord};; + result.device_ids = nonempty_set{device_id};; return result; }(); return DynamicNodeInvocation{ - /*inputs=*/map_values2(i.inputs, shard_expand_value_attrs), - /*node_attrs=*/expanded_node_attrs, - /*outputs=*/map_values2(i.outputs, shard_expand_value_attrs), + /*inputs=*/map_values2(i.inputs, shard_expand_value_attrs), + /*node_attrs=*/expanded_node_attrs, + /*outputs=*/map_values2(i.outputs, shard_expand_value_attrs), }; } @@ -203,13 +178,16 @@ static std::set auto [input_slot, input] = get_only(i.inputs); auto [output_slot, output] = get_only(i.outputs); - bidict input_mapping = - assert_unwrap(input.mapping).raw; - require_same(input_mapping.left_values(), - assert_unwrap(output.mapping).raw.left_values()); + ParallelTensorMapping input_mapping = assert_unwrap(input.mapping); + ParallelTensorMapping output_mapping = assert_unwrap(output.mapping); + + std::set coord_set = + require_same( + pt_mapping_get_coord_set(input_mapping), + pt_mapping_get_coord_set(output_mapping)); return transform( - set_of(input_mapping.left_values()), + coord_set, [&](ParallelTensorSpaceCoordinate const &p) -> DynamicNodeInvocationShardingInfo { // The machine coord for a copy is inherently nebulous because it // doesn't strictly run in any single location. Further, Realm has the @@ -218,10 +196,10 @@ static std::set // because we expect this to align with the most efficient way to issue // copies in Realm, although the current Realm backend uses a // centralized controller and thus issues copies all from a single node. - global_device_id_t machine_coord = input_mapping.at_l(p); + global_device_id_t device_id = pt_mapping_get_device_for_coord(input_mapping, p); return invocation_sharding_info_for_binding(i, - machine_coord, + device_id, OperatorAtomicTaskShardBinding{{ {input_slot.slot_name, p}, {output_slot.slot_name, p}, @@ -237,7 +215,7 @@ static std::set generate_shard_expansion_for_fwd_replicate(DynamicNodeInvocation const &i) { ASSERT(i.node_attrs.task_type == DynamicTaskType::FWD); - MappedOperatorTaskGroup node_mapping = assert_unwrap(i.node_attrs.mapping); + DynamicNodeMapping node_mapping = assert_unwrap(i.node_attrs.mapping); DynamicTensorSlot expected_input_slot = DynamicTensorSlot{ /*slot_name=*/TensorSlotName::INPUT, @@ -255,54 +233,53 @@ static std::set DynamicValueAttrs output = require_only_key(i.outputs, expected_output_slot); - bidict - input_value_mapping = assert_unwrap(input.mapping); + ParallelTensorMapping input_value_mapping = assert_unwrap(input.mapping); - std::set input_tensor_shards = set_of(input_value_mapping.left_values()); + std::set input_tensor_shards = pt_mapping_get_coord_set(input_value_mapping); - bidict - output_value_mapping = assert_unwrap(output.mapping); + ParallelTensorMapping output_value_mapping = assert_unwrap(output.mapping); - auto get_task_shard_machine_coords_for_input_tensor_shard + auto get_task_shard_device_ids_for_input_tensor_shard = [&](ParallelTensorSpaceCoordinate const &input_tensor_shard) - -> nonempty_set + -> nonempty_set { - bidict dependent_on_input_tensor_shard + bidict dependent_on_input_tensor_shard = bidict_filter_values( - node_mapping.get_shard_bindings(), - [&](OperatorAtomicTaskShardBinding const &b) -> bool { - return ptensor_space_coord_for_slot_name(b, TensorSlotName::INPUT) == input_tensor_shard; - }); + dynamic_node_mapping_get_shard_bindings(node_mapping), + [&](OperatorAtomicTaskShardBinding const &b) -> bool { + return ptensor_space_coord_for_slot_name(b, TensorSlotName::INPUT) == input_tensor_shard; + }); - return nonempty_set(set_of(dependent_on_input_tensor_shard.left_values())); + return nonempty_set(dependent_on_input_tensor_shard.left_values()); }; auto invocation_sharding_info_for_input_tensor_shard = [&](ParallelTensorSpaceCoordinate const &c) -> DynamicNodeInvocationShardingInfo { - nonempty_set task_shard_machine_coords = - get_task_shard_machine_coords_for_input_tensor_shard(c); + nonempty_set task_shard_device_ids = + get_task_shard_device_ids_for_input_tensor_shard(c); - std::map output_sharding_infos = - generate_map(task_shard_machine_coords.unwrap_as_set(), - [&](MachineSpaceCoordinate const &mc) + std::map output_sharding_infos = + generate_map(task_shard_device_ids.unwrap_as_set(), + [&](global_device_id_t const &device_id) -> DynamicValueAttrsShardingInfo { - ParallelTensorSpaceCoordinate pc = output_value_mapping.at_r(mc); + ParallelTensorSpaceCoordinate pc + = pt_mapping_get_coord_for_device(output_value_mapping, device_id); return DynamicValueAttrsShardingInfo{ /*shard_coord=*/pc, - /*mapping=*/mc, + /*mapping=*/device_id, }; }); std::map keyed_output_sharding_infos = map_keys(output_sharding_infos, - [&](MachineSpaceCoordinate const &mc) -> DynamicTensorSlot { + [&](global_device_id_t const &device_id) -> DynamicTensorSlot { return DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/mc, + /*task_shard=*/device_id.coord, }; }); @@ -314,7 +291,7 @@ static std::set DynamicValueAttrsShardingInfo input_sharding_info = DynamicValueAttrsShardingInfo{ /*shard_coord=*/c, - /*mapping=*/input_value_mapping.at_l(c), + /*mapping=*/pt_mapping_get_device_for_coord(input_value_mapping, c), }; std::map sharding_infos = @@ -328,8 +305,8 @@ static std::set }); return DynamicNodeInvocationShardingInfo{ - /*device_coords=*/task_shard_machine_coords, - /*value_sharding=*/sharding_infos, + /*device_ids=*/task_shard_device_ids, + /*value_sharding=*/binary_relation_from_map(sharding_infos), }; }; @@ -340,7 +317,7 @@ static std::set generate_shard_expansion_for_bwd_replicate(DynamicNodeInvocation const &i) { ASSERT(i.node_attrs.task_type == DynamicTaskType::BWD); - MappedOperatorTaskGroup node_mapping = assert_unwrap(i.node_attrs.mapping); + DynamicNodeMapping node_mapping = assert_unwrap(i.node_attrs.mapping); DynamicTensorSlot expected_output_grad_slot = DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, @@ -358,54 +335,53 @@ static std::set DynamicValueAttrs input_grad = require_only_key(i.outputs, expected_input_grad_slot); - bidict - output_grad_value_mapping = assert_unwrap(output_grad.mapping); + ParallelTensorMapping output_grad_value_mapping = assert_unwrap(output_grad.mapping); + ParallelTensorMapping input_grad_value_mapping = assert_unwrap(input_grad.mapping); - bidict - input_grad_value_mapping = assert_unwrap(input_grad.mapping); + std::set input_grad_tensor_shards + = pt_mapping_get_coord_set(input_grad_value_mapping); - std::set input_grad_tensor_shards = set_of(input_grad_value_mapping.left_values()); - - auto get_task_shard_machine_coords_for_input_grad_tensor_shard + auto get_task_shard_device_ids_for_input_grad_tensor_shard = [&](ParallelTensorSpaceCoordinate const &input_grad_tensor_shard) - -> nonempty_set + -> nonempty_set { - bidict produce_input_grad_tensor_shard + bidict produce_input_grad_tensor_shard = bidict_filter_values( - node_mapping.get_shard_bindings(), + dynamic_node_mapping_get_shard_bindings(node_mapping), [&](OperatorAtomicTaskShardBinding const &b) -> bool { return ptensor_space_coord_for_slot_name(b, TensorSlotName::INPUT) == input_grad_tensor_shard; }); - return nonempty_set(set_of(produce_input_grad_tensor_shard.left_values())); + return nonempty_set(produce_input_grad_tensor_shard.left_values()); }; auto invocation_sharding_info_for_input_grad_tensor_shard = [&](ParallelTensorSpaceCoordinate const &c) -> DynamicNodeInvocationShardingInfo { - nonempty_set task_shard_machine_coords = - get_task_shard_machine_coords_for_input_grad_tensor_shard(c); + nonempty_set task_shard_device_ids = + get_task_shard_device_ids_for_input_grad_tensor_shard(c); - std::map output_grad_sharding_infos = - generate_map(task_shard_machine_coords.unwrap_as_set(), - [&](MachineSpaceCoordinate const &mc) + std::map output_grad_sharding_infos = + generate_map(task_shard_device_ids.unwrap_as_set(), + [&](global_device_id_t const &device_id) -> DynamicValueAttrsShardingInfo { - ParallelTensorSpaceCoordinate pc = output_grad_value_mapping.at_r(mc); + ParallelTensorSpaceCoordinate pc + = pt_mapping_get_coord_for_device(output_grad_value_mapping, device_id); return DynamicValueAttrsShardingInfo{ /*shard_coord=*/pc, - /*mapping=*/mc, + /*mapping=*/device_id, }; }); std::map keyed_output_grad_sharding_infos = map_keys(output_grad_sharding_infos, - [&](MachineSpaceCoordinate const &mc) -> DynamicTensorSlot { + [&](global_device_id_t const &device_id) -> DynamicTensorSlot { return DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/mc, + /*task_shard=*/device_id.coord, }; }); @@ -417,7 +393,7 @@ static std::set DynamicValueAttrsShardingInfo input_grad_sharding_info = DynamicValueAttrsShardingInfo{ /*shard_coord=*/c, - /*mapping=*/input_grad_value_mapping.at_l(c), + /*mapping=*/pt_mapping_get_device_for_coord(input_grad_value_mapping, c), }; std::map sharding_infos = @@ -431,8 +407,8 @@ static std::set }); return DynamicNodeInvocationShardingInfo{ - /*device_coords=*/task_shard_machine_coords, - /*value_sharding=*/sharding_infos, + /*device_ids=*/task_shard_device_ids, + /*value_sharding=*/binary_relation_from_map(sharding_infos), }; }; @@ -454,35 +430,12 @@ std::set }); } -bool graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &g) { - auto slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) -> bool { - return false; - }; - - return full_dynamic_graph_satisfies(g, - node_is_ready_for_shard_expansion, - value_is_ready_for_shard_expansion, - slot_is_ready_for_shard_expansion); -} - - -void require_graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &g) { - auto require_slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) -> void { - return; - }; - - return require_full_dynamic_graph_satisfies(g, - require_node_is_ready_for_shard_expansion, - require_value_is_ready_for_shard_expansion, - require_slot_is_ready_for_shard_expansion); -} - DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( DynamicNodeAttrs const &node_attrs, - nonempty_set const &device_coords) + nonempty_set const &device_ids) { DynamicNodeAttrs result = node_attrs; - result.device_coords = device_coords; + result.device_ids = device_ids; return result; } @@ -494,12 +447,12 @@ DynamicValueAttrs apply_dynamic_value_attrs_sharding_info( DynamicValueAttrs result = value_attrs; result.shard_coord = value_sharding_info.shard_coord; - { - bidict value_mapping = - assert_unwrap(result.mapping); + if (result.mapping.has_value()) { + ParallelTensorMapping value_mapping = assert_unwrap(result.mapping); - MachineSpaceCoordinate from_mapping = value_mapping.at_l(value_sharding_info.shard_coord); - MachineSpaceCoordinate from_sharding_info = value_sharding_info.mapping; + global_device_id_t from_mapping = + pt_mapping_get_device_for_coord(value_mapping, value_sharding_info.shard_coord); + global_device_id_t from_sharding_info = value_sharding_info.mapping; ASSERT(from_mapping == from_sharding_info); } @@ -513,21 +466,87 @@ DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( { require_invocation_is_ready_for_shard_expansion(invocation); + { + std::set invocation_slots = set_union( + keys(invocation.inputs), keys(invocation.outputs)); + + std::set shard_info_slots_ignoring_task_shard = + transform( + invocation_sharding_info.value_sharding.left_values(), + slot_without_task_shard); + + ASSERT(invocation_slots == shard_info_slots_ignoring_task_shard, + dynamic_node_invocation_to_serializable(invocation), + invocation_sharding_info); + } + + std::set shard_labelled = + filtrans(invocation_sharding_info.value_sharding.left_values(), + [](DynamicTensorSlot const &s) -> std::optional + { + if (s.task_shard.has_value()) { + return slot_without_task_shard(s); + } else { + return std::nullopt; + } + }); + + { + std::set not_shard_labelled = + filter(invocation_sharding_info.value_sharding.left_values(), + [](DynamicTensorSlot const &s) -> bool + { + return !s.task_shard.has_value(); + }); + + ASSERT(are_disjoint(shard_labelled, not_shard_labelled)); + } + auto shard_value = [&](DynamicTensorSlot const &slot, DynamicValueAttrs const &value_attrs) - -> DynamicValueAttrs + -> std::map { - DynamicValueAttrsShardingInfo sharding_info = invocation_sharding_info.value_sharding.at(slot); - return apply_dynamic_value_attrs_sharding_info(value_attrs, sharding_info); + ASSERT(!slot.task_shard.has_value()); + + if (contains(shard_labelled, slot)) { + BinaryRelation + for_slot = filter_binary_relation(invocation_sharding_info.value_sharding, + [&](DynamicTensorSlot const &s, + DynamicValueAttrsShardingInfo const &) + -> bool + { + return slot_without_task_shard(s) == slot; + }); + + std::set> result = transform( + for_slot.unwrap_as_set(), + [&](std::pair const &p) + -> std::pair + { + return {p.first, apply_dynamic_value_attrs_sharding_info(value_attrs, p.second)}; + }); + + return map_from_pairs(result); + } else { + DynamicValueAttrsShardingInfo sharding_info + = get_only(invocation_sharding_info.value_sharding.at_l(slot)); + return { + { + slot, + apply_dynamic_value_attrs_sharding_info(value_attrs, sharding_info), + }, + }; + } }; DynamicNodeInvocation result = DynamicNodeInvocation{ - /*inputs=*/map_values2(invocation.inputs, shard_value), + /*inputs=*/flatmap(invocation.inputs, shard_value), /*node_attrs=*/apply_dynamic_node_attrs_sharding_info( - invocation.node_attrs, invocation_sharding_info.device_coords), - /*outputs=*/map_values2(invocation.outputs, shard_value), + invocation.node_attrs, invocation_sharding_info.device_ids), + /*outputs=*/flatmap(invocation.outputs, shard_value), }; - ASSERT(invocation_is_fully_shard_expanded(result)); + require_invocation_is_fully_shard_expanded(result); + return result; } @@ -536,41 +555,41 @@ std::set { require_invocation_is_ready_for_shard_expansion(i); - if (i.node_attrs.op_attrs.value().is_copy()) { - return set_of(generate_shard_expansion_for_copy(i)); - } - - if (training_op_attrs_has_op_type(i.node_attrs.op_attrs.value(), OperatorType::REPLICATE)) { - DynamicTaskType task_type = assert_unwrap(i.node_attrs.task_type); - switch (task_type) { - case DynamicTaskType::FWD: - return set_of(generate_shard_expansion_for_fwd_replicate(i)); - case DynamicTaskType::BWD: - return set_of(generate_shard_expansion_for_bwd_replicate(i)); - default: - PANIC("Unexpected task type for Replicate: {}", task_type); + std::set result = [&]() { + if (i.node_attrs.op_attrs.value().is_copy()) { + return generate_shard_expansion_for_copy(i); + } else if (training_op_attrs_has_op_type(i.node_attrs.op_attrs.value(), OperatorType::REPLICATE)) { + DynamicTaskType task_type = assert_unwrap(i.node_attrs.task_type); + switch (task_type) { + case DynamicTaskType::FWD: + return set_of(generate_shard_expansion_for_fwd_replicate(i)); + case DynamicTaskType::BWD: + return set_of(generate_shard_expansion_for_bwd_replicate(i)); + default: + PANIC("Unexpected task type for Replicate: {}", task_type); + } + } else { + DynamicNodeMapping mapping = assert_unwrap(i.node_attrs.mapping); + + std::set shard_machine_coords = + target_devices_of_dynamic_node_mapping(mapping); + + return transform(shard_machine_coords, + [&](global_device_id_t const &device_id) -> DynamicNodeInvocationShardingInfo { + OperatorAtomicTaskShardBinding slot_bindings = + dynamic_node_mapping_get_shard_binding_for_device(mapping, device_id); + + return invocation_sharding_info_for_binding(i, device_id, slot_bindings); + }); } - } - - DynamicNodeMapping mapping = assert_unwrap(i.node_attrs.mapping); - - std::set shard_machine_coords = - target_devices_of_dynamic_node_mapping(mapping); - - return transform(shard_machine_coords, - [&](global_device_id_t const &c) -> DynamicNodeInvocation { - OperatorAtomicTaskShardBinding slot_bindings = - mapping.op_task_group.get_shard_bindings().at_l( - c.coord); + }(); - return shard_invocation_for_binding(i, c, slot_bindings); - }); + return result; } DynamicOpenDataflowGraph perform_shard_expansion(DynamicOpenDataflowGraph const &g) { - ASSERT(no_part_of_graph_is_shard_expanded(g)); require_graph_is_ready_for_shard_expansion(g); DynamicOpenDataflowGraph result = @@ -578,7 +597,7 @@ DynamicOpenDataflowGraph return perform_shard_expansion_for_invocation(i); }); - ASSERT(graph_is_fully_shard_expanded(result)); + require_graph_is_fully_shard_expanded(result); return result; } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index c24488fe19..1f97f064e8 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -11,853 +11,1880 @@ #include "task-spec/dynamic_graph/dynamic_value_attrs.h" #include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" #include "op-attrs/ops/element_unary.h" +#include "utils/containers/require_only_key.h" +#include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/pass_expansion.h" using namespace ::FlexFlow; -TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("copies_for_invocation_inputs") { - auto mk_machine_coord = - [](nonnegative_int node_idx, - nonnegative_int device_idx) -> MachineSpaceCoordinate { - return MachineSpaceCoordinate{ - /*node_idx=*/node_idx, - /*device_idx=*/device_idx, - /*device_type=*/DeviceType::GPU, - }; +static DynamicValueAttrs + mk_value_attrs(size_t src_layer_guid, + TensorSlotName src_slot, + std::optional const &mapping) +{ + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{ + src_layer_guid, + }, + src_slot, + }, + }, + }, + /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/false, + /*shard_coord=*/std::nullopt, + /*mapping=*/mapping, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; +} + +static DynamicTensorSlot mk_slot(TensorSlotName slot_name) { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }; +} + +static DynamicNodeAttrs mk_node_attrs(size_t layer_guid, + PCGOperatorAttrs const &op_attrs, + DynamicNodeMapping const &mapping) { + return DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/mapping, + /*op_attrs=*/TrainingOperationAttrs{ + op_attrs, + }, + /*layer_guid=*/dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{layer_guid}, + }, + }, + /*per_device_op_state=*/std::nullopt, + }; +} + +static MachineSpaceCoordinate mk_machine_coord(nonnegative_int device_idx) { + return MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/device_idx, + }; +}; + +static global_device_id_t mk_device_id(MachineSpaceCoordinate const &mc) { + return global_device_id_t{mc, DeviceType::GPU}; +}; + +static DynamicOpenDataflowGraph mk_single_input_node_graph(MachineSpaceCoordinate const &input_device) +{ + TensorShape input_shape = TensorShape{ + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, + }, + DataType::FLOAT, + }; + + PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ + InputAttrs{ + input_shape, + }, + }; + + PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ + make_relu_attrs(), + }; + + auto mk_node_mapping = [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { + return DynamicNodeMapping{ + /*op_task_group=*/op_task_group, + /*device_type=*/DeviceType::GPU, + }; + }; + + auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/FFOrdered{ + 0_n, + 0_n, + }, + }; + }; + + DynamicValueAttrs input_op_output = + mk_value_attrs(123, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); + + MappedOperatorTaskGroup input_node_mapping = + MappedOperatorTaskGroup{ + bidict{ + { + input_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, + }, + }, + }, }; - auto mk_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { - return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }; + DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ + /*inputs=*/{}, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, + /*mapping=*/mk_node_mapping(input_node_mapping)), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, + }, + }; + + DynamicOpenDataflowGraph g + = dynamic_open_dataflow_graph_from_invocation_set( + {input_invocation}); + + return g; +} + +static DynamicOpenDataflowGraph mk_single_input_into_relu_graph(MachineSpaceCoordinate const &input_device, + MachineSpaceCoordinate const &relu1_device) +{ + TensorShape input_shape = TensorShape{ + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, + }, + DataType::FLOAT, + }; + + PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ + InputAttrs{ + input_shape, + }, + }; + + PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ + make_relu_attrs(), + }; + + auto mk_node_mapping = [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { + return DynamicNodeMapping{ + /*op_task_group=*/op_task_group, + /*device_type=*/DeviceType::GPU, + }; + }; + + auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/FFOrdered{ + 0_n, + 0_n, + }, + }; + }; + + DynamicValueAttrs input_op_output = + mk_value_attrs(123, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); + + MappedOperatorTaskGroup input_node_mapping = + MappedOperatorTaskGroup{ + bidict{ + { + input_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, + }, + }, + }, }; - auto mk_value = [](size_t src_node_id, - TensorSlotName src_slot_name) - -> DynamicValueAttrs { - return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{ - parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{src_node_id}, - src_slot_name, + DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ + /*inputs=*/{}, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, + /*mapping=*/mk_node_mapping(input_node_mapping)), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, + }, + }; + + DynamicValueAttrs relu1_op_output = + mk_value_attrs(124, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); + + MappedOperatorTaskGroup relu1_node_mapping = + MappedOperatorTaskGroup{ + bidict{ + { + relu1_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), }, }, }, - /*parallel_tensor_shape=*/std::nullopt, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }; + }, + }, }; - auto mk_pt_coord = - [](nonnegative_int idx1, - nonnegative_int idx2, - nonnegative_int idx3, - nonnegative_int idx4) -> ParallelTensorSpaceCoordinate { - return ParallelTensorSpaceCoordinate{ - /*sum_component=*/idx1, - /*discard_copy_component=*/idx2, - /*shard_components=*/ - FFOrdered{ - idx3, - idx4, + DynamicNodeInvocation relu1_invocation = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + input_op_output, + }, + }, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/124, + /*op_attrs=*/relu_attrs, + /*mapping=*/mk_node_mapping(relu1_node_mapping)), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + relu1_op_output, + }, + }, + }; + + DynamicOpenDataflowGraph g + = dynamic_open_dataflow_graph_from_invocation_set( + {input_invocation, relu1_invocation}); + + return g; +} + +struct ExampleGraphTestCase { + DynamicOpenDataflowGraph g; + dynamic_invocation_id_t input_op_id; + dynamic_invocation_id_t relu1_op_id; + dynamic_invocation_id_t replicate_op_id; + dynamic_invocation_id_t relu2_op_id; +}; + +static ExampleGraphTestCase mk_example_replicate_graph( + MachineSpaceCoordinate const &input_device, + MachineSpaceCoordinate const &relu1_device, + MachineSpaceCoordinate const &replicate_device1, + MachineSpaceCoordinate const &replicate_device2, + MachineSpaceCoordinate const &relu2_device1, + MachineSpaceCoordinate const &relu2_device2) +{ + TensorShape input_shape = TensorShape{ + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, + }, + DataType::FLOAT, + }; + + PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ + InputAttrs{ + input_shape, + }, + }; + + PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ + make_relu_attrs(), + }; + + auto mk_node_mapping = [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { + return DynamicNodeMapping{ + /*op_task_group=*/op_task_group, + /*device_type=*/DeviceType::GPU, + }; + }; + + auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/FFOrdered{ + 0_n, + 0_n, + }, + }; + }; + + DynamicValueAttrs input_op_output = + mk_value_attrs(123, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); + + MappedOperatorTaskGroup input_node_mapping = + MappedOperatorTaskGroup{ + bidict{ + { + input_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }; + }, + }, }; - size_t invocation_id = 20; + DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ + /*inputs=*/{}, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, + /*mapping=*/mk_node_mapping(input_node_mapping)), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, + }, + }; + + DynamicValueAttrs relu1_op_output = + mk_value_attrs(124, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); + + MappedOperatorTaskGroup relu1_node_mapping = + MappedOperatorTaskGroup{ + bidict{ + { + relu1_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, + }, + }, + }, + }; - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); - MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); - MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); + DynamicNodeInvocation relu1_invocation = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + input_op_output, + }, + }, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/124, + /*op_attrs=*/relu_attrs, + /*mapping=*/mk_node_mapping(relu1_node_mapping)), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + relu1_op_output, + }, + }, + }; + + DynamicValueAttrs replicate_op_output = + mk_value_attrs(125, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); + + MappedOperatorTaskGroup replicate_node_mapping = + MappedOperatorTaskGroup{ + bidict{ + { + replicate_device1, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, + }, + }, + { + replicate_device2, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }, + }, + }, + }, + }; - SUBCASE("standard operator") { - auto mk_input_shard_binding = [&](ParallelTensorSpaceCoordinate const &c) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ - { - TensorSlotName::OUTPUT, - c, - }, + DynamicNodeInvocation replicate_invocation = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + relu1_op_output, + }, + }, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/125, + /*op_attrs=*/PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, + }, + /*mapping=*/mk_node_mapping(replicate_node_mapping)), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + replicate_op_output, + }, + }, + }; + + DynamicValueAttrs relu2_op_output = + mk_value_attrs(126, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); + + MappedOperatorTaskGroup relu2_node_mapping = + MappedOperatorTaskGroup{ + bidict{ + { + relu2_device1, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, }, - }; + }, + }, + { + relu2_device2, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(1_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }, + }, + }, + }, + }; + + DynamicNodeInvocation relu2_invocation = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + replicate_op_output, + }, + }, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/125, + /*op_attrs=*/relu_attrs, + /*mapping=*/mk_node_mapping(relu2_node_mapping)), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + relu2_op_output, + }, + }, + }; + + DynamicOpenDataflowGraph g + = dynamic_open_dataflow_graph_from_invocation_set( + {input_invocation, relu1_invocation, replicate_invocation, relu2_invocation}); + + return ExampleGraphTestCase{ + /*g=*/g, + /*input_op_id=*/dynamic_graph_get_id_for_invocation(g, input_invocation), + /*relu1_op_id=*/dynamic_graph_get_id_for_invocation(g, relu1_invocation), + /*replicate_op_id=*/dynamic_graph_get_id_for_invocation(g, replicate_invocation), + /*relu2_op_id=*/dynamic_graph_get_id_for_invocation(g, relu2_invocation), + }; +}; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("resolve_tensor_mappings") { + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(1_n); + MachineSpaceCoordinate mc3 = mk_machine_coord(2_n); + + auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/FFOrdered{ + 0_n, + 0_n, + }, }; + }; - auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, - ParallelTensorSpaceCoordinate const &c2, - ParallelTensorSpaceCoordinate const &c3, - ParallelTensorSpaceCoordinate const &c4) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ - { - TensorSlotName::INPUT, - c1, - }, + + SUBCASE("dynamic graph is not pass expanded") { + auto mk_correct_mappings = + [&](ExampleGraphTestCase const &tc, + MachineSpaceCoordinate const &input_out_mc, + MachineSpaceCoordinate const &relu1_in_mc, + MachineSpaceCoordinate const &relu1_out_mc, + MachineSpaceCoordinate const &replicate_in_mc, + MachineSpaceCoordinate const &replicate_out_mc1, + MachineSpaceCoordinate const &replicate_out_mc2, + MachineSpaceCoordinate const &relu2_in_mc1, + MachineSpaceCoordinate const &relu2_in_mc2, + MachineSpaceCoordinate const &relu2_out_mc1, + MachineSpaceCoordinate const &relu2_out_mc2) + -> std::map + { + DynamicOpenDataflowGraph g = tc.g; + dynamic_invocation_id_t input_op_id = tc.input_op_id; + dynamic_invocation_id_t relu1_op_id = tc.relu1_op_id; + dynamic_invocation_id_t replicate_op_id = tc.replicate_op_id; + dynamic_invocation_id_t relu2_op_id = tc.relu2_op_id; + + auto mk_inp_slot = [](dynamic_invocation_id_t invocation_id) -> InternalDynamicSlotSite { + return InternalDynamicSlotSite{ + /*invocation_id=*/invocation_id, + /*direction=*/TensorDirection::INCOMING, + /*slot_name=*/mk_slot(TensorSlotName::INPUT), + }; + }; + + auto mk_out_slot = [](dynamic_invocation_id_t invocation_id) -> InternalDynamicSlotSite { + return InternalDynamicSlotSite{ + /*invocation_id=*/invocation_id, + /*direction=*/TensorDirection::OUTPUT, + /*slot_name=*/mk_slot(TensorSlotName::OUTPUT), + }; + }; + + auto mk_single_shard_mapping = [&](MachineSpaceCoordinate const &mc) -> ParallelTensorMapping { + return ParallelTensorMapping{ + bidict{ { - TensorSlotName::WEIGHT, - c2, + mk_pt_coord(0_n), + mk_device_id(mc), }, + }, + }; + }; + + auto mk_two_shard_mapping = [&](MachineSpaceCoordinate const &mc1, + MachineSpaceCoordinate const &mc2) -> ParallelTensorMapping { + return ParallelTensorMapping{ + bidict{ { - TensorSlotName::OUTPUT_1, - c3, + mk_pt_coord(0_n), + mk_device_id(mc1), }, { - TensorSlotName::OUTPUT_2, - c4, + mk_pt_coord(1_n), + mk_device_id(mc2), }, + }, + }; + }; + + InternalDynamicSlotSite input_op_out = mk_out_slot(input_op_id); + InternalDynamicSlotSite relu1_op_in = mk_inp_slot(relu1_op_id); + InternalDynamicSlotSite relu1_op_out = mk_out_slot(relu1_op_id); + InternalDynamicSlotSite replicate_op_in = mk_inp_slot(replicate_op_id); + InternalDynamicSlotSite replicate_op_out = mk_out_slot(replicate_op_id); + InternalDynamicSlotSite relu2_op_in = mk_inp_slot(relu2_op_id); + InternalDynamicSlotSite relu2_op_out = mk_out_slot(relu2_op_id); + + return { + { + input_op_out, + mk_single_shard_mapping(input_out_mc), + }, + { + relu1_op_in, + mk_single_shard_mapping(relu1_in_mc), + }, + { + relu1_op_out, + mk_single_shard_mapping(relu1_out_mc), }, + { + replicate_op_in, + mk_single_shard_mapping(replicate_in_mc), + }, + { + replicate_op_out, + mk_two_shard_mapping(replicate_out_mc1, replicate_out_mc2), + }, + { + relu2_op_in, + mk_two_shard_mapping(relu2_in_mc1, relu2_in_mc2), + }, + { + relu2_op_out, + mk_two_shard_mapping(relu2_out_mc1, relu2_out_mc2), + }, + }; }; - }; - ParallelTensorSpaceCoordinate mc1_input_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc1_weight_coord = - mk_pt_coord(0_n, 1_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc1_output_1_coord = - mk_pt_coord(1_n, 0_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc1_output_2_coord = - mk_pt_coord(3_n, 0_n, 0_n, 0_n); - - ParallelTensorSpaceCoordinate mc2_input_coord = - mk_pt_coord(0_n, 1_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc2_weight_coord = - mk_pt_coord(0_n, 4_n, 2_n, 0_n); - ParallelTensorSpaceCoordinate mc2_output_1_coord = - mk_pt_coord(1_n, 2_n, 0_n, 1_n); - ParallelTensorSpaceCoordinate mc2_output_2_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - - MappedOperatorTaskGroup input_mapping_same = MappedOperatorTaskGroup{ - bidict{ - { - TensorSlotName::INPUT, - mk_ptensor_coord(input_shard_idx), - }, - { - mc2, - mk_input_shard_binding(mc2_input_coord), - }, + + SUBCASE("replicate input matches tensor source") { + ExampleGraphTestCase tc + = mk_example_replicate_graph( + /*input_device=*/mc1, + /*relu1_device=*/mc1, + /*replicate_device1=*/mc1, + /*replicate_device2=*/mc2, + /*relu2_device1=*/mc1, + /*relu2_device2=*/mc2); + + std::map result = + resolve_tensor_mappings(tc.g); + + std::map correct = + mk_correct_mappings( + /*tc=*/tc, + /*input_out_mc=*/mc1, + /*relu1_in_mc=*/mc1, + /*relu1_out_mc=*/mc1, + /*replicate_in_mc=*/mc1, + /*relicate_out_mc1=*/mc1, + /*relicate_out_mc2=*/mc2, + /*relu2_in_mc1=*/mc1, + /*relu2_in_mc2=*/mc2, + /*relu2_out_mc1=*/mc1, + /*relu2_out_mc1=*/mc2); + + ASSERT(result == correct); + } + + SUBCASE("input invocation's output follows the input's mapping") { + ExampleGraphTestCase tc + = mk_example_replicate_graph( + /*input_device=*/mc3, + /*relu1_device=*/mc1, + /*replicate_device1=*/mc1, + /*replicate_device2=*/mc2, + /*relu2_device1=*/mc1, + /*relu2_device2=*/mc2); + + std::map result = + resolve_tensor_mappings(tc.g); + + std::map correct = + mk_correct_mappings( + /*tc=*/tc, + /*input_out_mc=*/mc3, + /*relu1_in_mc=*/mc1, + /*relu1_out_mc=*/mc1, + /*replicate_in_mc=*/mc1, + /*relicate_out_mc1=*/mc1, + /*relicate_out_mc2=*/mc2, + /*relu2_in_mc1=*/mc1, + /*relu2_in_mc2=*/mc2, + /*relu2_out_mc1=*/mc1, + /*relu2_out_mc1=*/mc2); + + ASSERT(result == correct); + } + + SUBCASE("src and sink can differ due to different invocation mappings") { + ExampleGraphTestCase tc + = mk_example_replicate_graph( + /*input_device=*/mc3, + /*relu1_device=*/mc1, + /*replicate_device1=*/mc2, + /*replicate_device2=*/mc3, + /*relu2_device1=*/mc1, + /*relu2_device2=*/mc3); + + std::map result = + resolve_tensor_mappings(tc.g); + + std::map correct = + mk_correct_mappings( + /*tc=*/tc, + /*input_out_mc=*/mc3, + /*relu1_in_mc=*/mc1, + /*relu1_out_mc=*/mc1, + /*replicate_in_mc=*/mc1, + /*relicate_out_mc1=*/mc2, + /*relicate_out_mc2=*/mc3, + /*relu2_in_mc1=*/mc1, + /*relu2_in_mc2=*/mc3, + /*relu2_out_mc1=*/mc1, + /*relu2_out_mc2=*/mc3); + + ASSERT(result == correct); + } + } + } + + TEST_CASE("copies_for_value") { + DynamicValueAttrs value_attrs = DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{1}, + TensorSlotName::OUTPUT, + }, }, + }, + /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + + auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shared_components=*/FFOrdered{ + idx, + 0_n, + }, }; + }; - MappedOperatorTaskGroup weight_mapping_same = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_weight_coord), - }, - }, + auto mk_device = [](nonnegative_int idx) -> global_device_id_t { + return global_device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/2_n, + /*device_idx=*/idx, + }, + /*device_type=*/DeviceType::GPU, }; }; - DynamicValueAttrs v1 = mk_value_attrs( - /*src_layer_guid=*/0, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/std::nullopt); - - DynamicValueAttrs v2 = mk_value_attrs( - /*src_layer_guid=*/1, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/std::nullopt); - - DynamicValueAttrs v3 = mk_value_attrs( - /*src_layer_guid=*/2, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/std::nullopt); - - SUBCASE("inserts copy when necessary") { - DynamicNodeMapping mapping1 = DynamicNodeMapping{ - MappedOperatorTaskGroup{ - bidict{ - { - mk_machine_coord(0_n), - mk_binding(0_n, 0_n), - }, - { - mk_machine_coord(1_n), - mk_binding(1_n, 1_n), - }, - }, - }, - DeviceType::GPU, + auto mk_slot_site = [](nonnegative_int idx) -> InternalDynamicSlotSite { + return InternalDynamicSlotSite{ + /*invocation_id=*/dynamic_invocation_id_t{idx}, + /*direction=*/TensorDirection::INCOMING, + /*slot_name=*/DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/DynamicTensorRole{FwbTensorType::FORWARD}, // could be any role + /*task_shard=*/std::nullopt, + }, }; + }; - DynamicNodeMapping mapping2 = DynamicNodeMapping{ - MappedOperatorTaskGroup{ - bidict{ - { - mk_machine_coord(0_n), - mk_binding(0_n, 0_n), - }, - { - mk_machine_coord(2_n), - mk_binding(1_n, 1_n), - }, - }, - }, - DeviceType::GPU, + SUBCASE("if src site is external, no copies no matter what") { + DynamicSlotSite src_site = DynamicSlotSite{ + ExternalDynamicSlotSite{ + dynamic_external_value_id_t{0_n}, + }, }; - DynamicNodeInvocation inv1 = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - v1, - }, - }, - /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT), - v2, - }, - }, + ParallelTensorMapping mapping1 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, + }; + + ParallelTensorMapping mapping2 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(1_n)}, + }, + }; + + ParallelTensorMapping mapping3 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(2_n)}, + }, + }; + + InternalDynamicSlotSite dst_site1 = mk_slot_site(2_n); + InternalDynamicSlotSite dst_site2 = mk_slot_site(3_n); + InternalDynamicSlotSite dst_site3 = mk_slot_site(4_n); + + std::map site_mappings = { + {dst_site1, mapping2}, + {dst_site2, mapping1}, + {dst_site3, mapping3}, + }; + + std::set result = + copies_for_value( + /*value_attrs=*/value_attrs, + /*src_site=*/DynamicSlotSite{src_site}, + /*dst_sites=*/{dst_site1, dst_site2, dst_site3}, + /*all_mappings=*/site_mappings); + + std::set correct = {}; + + CHECK(result == correct); + }; + + InternalDynamicSlotSite src_site = InternalDynamicSlotSite{ + /*invocation_id=*/dynamic_invocation_id_t{0_n}, + /*direction=*/TensorDirection::OUTPUT, + /*slot_name=*/DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + }; + + SUBCASE("if src mapping is same as dst mapping don't copy") { + ParallelTensorMapping mapping1 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, + }; + + InternalDynamicSlotSite dst_site = mk_slot_site(1_n); + + std::map site_mappings = { + {src_site, mapping1}, + {dst_site, mapping1}, + }; + + std::set result = + copies_for_value( + /*value_attrs=*/value_attrs, + /*src_site=*/DynamicSlotSite{src_site}, + /*dst_sites=*/{dst_site}, + /*all_mappings=*/site_mappings); + + std::set correct = {}; + + CHECK(result == correct); + } + + SUBCASE("if src mapping does not overlap dst mapping issue copy") { + ParallelTensorMapping mapping1 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, + }; + + ParallelTensorMapping mapping2 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(5_n)}, + }, + }; + + InternalDynamicSlotSite dst_site = mk_slot_site(1_n); + + std::map site_mappings = { + {src_site, mapping1}, + {dst_site, mapping2}, + }; + + std::set result = + copies_for_value( + /*value_attrs=*/value_attrs, + /*src_site=*/DynamicSlotSite{src_site}, + /*dst_sites=*/{dst_site}, + /*all_mappings=*/site_mappings); + + std::set correct = { + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, + }; + + CHECK(result == correct); + } + + SUBCASE("if src mapping overlaps dst mapping issue full copy") { + ParallelTensorMapping mapping1 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(1_n)}, + }, + }; + + ParallelTensorMapping mapping2 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(2_n)}, + }, + }; + + InternalDynamicSlotSite dst_site = mk_slot_site(1_n); + + std::map site_mappings = { + {src_site, mapping1}, + {dst_site, mapping2}, + }; + + std::set result = + copies_for_value( + /*value_attrs=*/value_attrs, + /*src_site=*/DynamicSlotSite{src_site}, + /*dst_sites=*/{dst_site}, + /*all_mappings=*/site_mappings); + + std::set correct = { + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, + }; + + CHECK(result == correct); + } + + SUBCASE("if src mapping overlaps multiple dst mappings issue both copies") { + ParallelTensorMapping mapping1 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(1_n)}, + }, + }; + + ParallelTensorMapping mapping2 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(2_n)}, + }, + }; + + ParallelTensorMapping mapping3 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(3_n)}, + }, + }; + + InternalDynamicSlotSite dst_site1 = mk_slot_site(1_n); + InternalDynamicSlotSite dst_site2 = mk_slot_site(2_n); + + std::map site_mappings = { + {src_site, mapping1}, + {dst_site1, mapping2}, + {dst_site2, mapping3}, + }; + + std::set result = + copies_for_value( + /*value_attrs=*/value_attrs, + /*src_site=*/DynamicSlotSite{src_site}, + /*dst_sites=*/{dst_site1, dst_site2}, + /*sink_site_mappings=*/site_mappings); + + std::set correct = { + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping3, + }, + }; + + CHECK(result == correct); + } + + SUBCASE("only copy once if multiple sinks use the same mapping") { + ParallelTensorMapping mapping1 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, + }; + + ParallelTensorMapping mapping2 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(5_n)}, + }, + }; + + InternalDynamicSlotSite dst_site1 = mk_slot_site(1_n); + InternalDynamicSlotSite dst_site2 = mk_slot_site(2_n); + InternalDynamicSlotSite dst_site3 = mk_slot_site(3_n); + + std::map site_mappings = { + {src_site, mapping1}, + {dst_site1, mapping2}, + {dst_site2, mapping2}, + {dst_site3, mapping2}, + }; + + std::set result = + copies_for_value( + /*value_attrs=*/value_attrs, + /*src_site=*/DynamicSlotSite{src_site}, + /*dst_sites=*/{dst_site1, dst_site2, dst_site3}, + /*all_mappings=*/site_mappings); + + std::set correct = { + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, + }; + + CHECK(result == correct); + } + + SUBCASE("if src mapping matches one dst mapping, still issue copies for the rest") { + ParallelTensorMapping mapping1 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, + }; + + ParallelTensorMapping mapping2 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(5_n)}, + }, + }; + + ParallelTensorMapping mapping3 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(6_n)}, + }, + }; + + ParallelTensorMapping mapping4 = ParallelTensorMapping{ + bidict{ + {mk_pt_coord(0_n), mk_device(7_n)}, + }, + }; + + InternalDynamicSlotSite dst_site1 = mk_slot_site(1_n); + InternalDynamicSlotSite dst_site2 = mk_slot_site(2_n); + InternalDynamicSlotSite dst_site3 = mk_slot_site(3_n); + InternalDynamicSlotSite dst_site4 = mk_slot_site(4_n); + InternalDynamicSlotSite dst_site5 = mk_slot_site(5_n); + + std::map site_mappings = { + {src_site, mapping1}, + {dst_site1, mapping2}, + {dst_site2, mapping1}, + {dst_site3, mapping2}, + {dst_site4, mapping4}, + {dst_site5, mapping3}, + }; + + std::set result = + copies_for_value( + /*value_attrs=*/value_attrs, + /*src_site=*/DynamicSlotSite{src_site}, + /*dst_sites=*/{dst_site1, dst_site2, dst_site3, dst_site4, dst_site5}, + /*all_mappings=*/site_mappings); + + std::set correct = { + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping3, + }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping4, + }, + }; + + CHECK(result == correct); + } + } + + TEST_CASE("perform_copy_insertion") { + + SUBCASE("standard operator") { + auto mk_ptensor_coord = + [](nonnegative_int shard_idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_components=*/ + FFOrdered{ + shard_idx, + }, + }; + }; + + auto mk_device_id = [&](nonnegative_int device_idx) -> global_device_id_t { + return global_device_id_t{ + mk_machine_coord(device_idx), + DeviceType::GPU, + }; + }; + + auto mk_pcg_layer_guid = [](size_t pcg_layer_guid) -> dynamic_layer_guid_t { + return dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{pcg_layer_guid}, + }, + }; + }; + + auto mk_node_attrs = + [](dynamic_layer_guid_t layer_guid, + std::optional const &mapping, + TrainingOperationAttrs const &op_attrs) + -> DynamicNodeAttrs { + return DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }; + }; + + auto mk_binding = [&](nonnegative_int input_shard_idx, + nonnegative_int output_shard_idx) + -> OperatorAtomicTaskShardBinding { + return OperatorAtomicTaskShardBinding{ + /*tensor_coords=*/std::map{ + { + TensorSlotName::INPUT, + mk_ptensor_coord(input_shard_idx), + }, + { + TensorSlotName::OUTPUT, + mk_ptensor_coord(output_shard_idx), + }, + }, + }; }; - DynamicNodeInvocation inv2 = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - v2, - }, - }, - /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT), - v3, - }, - }, + TrainingOperationAttrs relu_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + make_relu_attrs(), + }, }; - DynamicOpenDataflowGraph g = - dynamic_open_dataflow_graph_from_invocation_set({inv1, inv2}); + DynamicValueAttrs v1 = mk_value_attrs( + /*src_layer_guid=*/0, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/std::nullopt); - DynamicOpenDataflowGraph result = perform_copy_insertion(g); + DynamicValueAttrs v2 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/std::nullopt); - DynamicOpenDataflowGraph correct = [&] { - DynamicValueAttrs mapped_v1 = mk_value_attrs( - /*src_layer_guid=*/0, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/ - ParallelTensorMapping{ - bidict{ - {mk_ptensor_coord(0_n), mk_device_id(0_n)}, - {mk_ptensor_coord(1_n), mk_device_id(1_n)}, - }, - }); - - DynamicValueAttrs mapped_v2_placement1 = mk_value_attrs( - /*src_layer_guid=*/1, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/ - ParallelTensorMapping{ - bidict{ - {mk_ptensor_coord(0_n), mk_device_id(0_n)}, - {mk_ptensor_coord(1_n), mk_device_id(1_n)}, - }, - }); - - DynamicValueAttrs mapped_v2_placement2 = mk_value_attrs( - /*src_layer_guid=*/1, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/ - ParallelTensorMapping{ - bidict{ - {mk_ptensor_coord(0_n), mk_device_id(0_n)}, - {mk_ptensor_coord(1_n), mk_device_id(2_n)}, - }, - }); - - DynamicValueAttrs mapped_v3 = mk_value_attrs( - /*src_layer_guid=*/2, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/ - ParallelTensorMapping{ - bidict{ - {mk_ptensor_coord(0_n), mk_device_id(0_n)}, - {mk_ptensor_coord(1_n), mk_device_id(2_n)}, - }, - }); + DynamicValueAttrs v3 = mk_value_attrs( + /*src_layer_guid=*/2, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/std::nullopt); - DynamicNodeInvocation mapped_inv1 = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mapped_v1, + SUBCASE("inserts copy when necessary") { + DynamicNodeMapping mapping1 = DynamicNodeMapping{ + MappedOperatorTaskGroup{ + bidict{ + { + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(1_n), + mk_binding(1_n, 1_n), + }, }, }, - /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT), - mapped_v2_placement1, + DeviceType::GPU, + }; + + DynamicNodeMapping mapping2 = DynamicNodeMapping{ + MappedOperatorTaskGroup{ + bidict{ + { + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(2_n), + mk_binding(1_n, 1_n), + }, }, }, + DeviceType::GPU, }; - DynamicNodeInvocation inserted_copy = DynamicNodeInvocation{ + DynamicNodeInvocation inv1 = DynamicNodeInvocation{ /*inputs=*/{ { mk_slot(TensorSlotName::INPUT), - mapped_v2_placement1, + v1, }, }, /*node_attrs=*/ - mk_node_attrs(dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}), + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, relu_attrs), /*outputs=*/ { { mk_slot(TensorSlotName::OUTPUT), - mapped_v2_placement2, + v2, }, }, - }; - DynamicNodeInvocation mapped_inv2 = DynamicNodeInvocation{ + DynamicNodeInvocation inv2 = DynamicNodeInvocation{ /*inputs=*/{ { mk_slot(TensorSlotName::INPUT), - mapped_v2_placement2, + v2, }, }, /*node_attrs=*/ mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), + mk_pcg_layer_guid(2), mapping2, relu_attrs), /*outputs=*/ { { mk_slot(TensorSlotName::OUTPUT), - mapped_v3, + v3, }, }, }; - return dynamic_open_dataflow_graph_from_invocation_set( - {mapped_inv1, mapped_inv2, inserted_copy}); - }(); + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set({inv1, inv2}); - CHECK_MESSAGE( - result == correct, - check_kv("result\n", dynamic_open_dataflow_graph_as_dot(result)), - check_kv("correct\n", dynamic_open_dataflow_graph_as_dot(correct))); - } + DynamicOpenDataflowGraph result = perform_copy_insertion(g); - SUBCASE("does not insert a copy when not necessary") { - DynamicNodeMapping mapping1 = DynamicNodeMapping{ - MappedOperatorTaskGroup{ - bidict{ - { - mk_machine_coord(0_n), - mk_binding(0_n, 0_n), + DynamicOpenDataflowGraph correct = [&] { + DynamicValueAttrs mapped_v1 = mk_value_attrs( + /*src_layer_guid=*/0, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, }, - { - mk_machine_coord(1_n), - mk_binding(1_n, 1_n), + }); + + DynamicValueAttrs mapped_v2_placement1 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, }, - }, - }, - DeviceType::GPU, - }; - - DynamicNodeMapping mapping2 = DynamicNodeMapping{ - MappedOperatorTaskGroup{ - bidict{ - { - mk_machine_coord(0_n), - mk_binding(0_n, 0_n), + }); + + DynamicValueAttrs mapped_v2_placement2 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(2_n)}, }, - { - mk_machine_coord(1_n), - mk_binding(1_n, 1_n), + }); + + DynamicValueAttrs mapped_v3 = mk_value_attrs( + /*src_layer_guid=*/2, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(2_n)}, }, - }, - }, - DeviceType::GPU, - }; + }); - DynamicNodeInvocation inv1 = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - v1, + DynamicNodeInvocation mapped_inv1 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v1, + }, }, - }, - /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), - /*outputs=*/ - { + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, relu_attrs), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - v2, + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v2_placement1, + }, }, - }, - }; + }; - DynamicNodeInvocation inv2 = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - v2, + DynamicNodeInvocation inserted_copy = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v2_placement1, + }, }, - }, - /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), - /*outputs=*/ - { + /*node_attrs=*/ + mk_node_attrs(dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - v3, - }, - }, - }; - - MappedOperatorTaskGroup invocation_mapping_diff_vs_copy1 = - MappedOperatorTaskGroup{ - bidict{ { - mc2, - mk_shard_binding(mc2_input_coord, - mc2_weight_coord, - mc2_output_1_coord, - mc2_output_2_coord), + mk_slot(TensorSlotName::OUTPUT), + mapped_v2_placement2, }, }, + }; - DynamicValueAttrs graph_input1 = - mk_value(0, TensorSlotName::OUTPUT); - - DynamicValueAttrs graph_input1_use = - decide_dynamic_value_attrs_mapping( - graph_input1, - get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::INPUT)); - - DynamicValueAttrs graph_input1_use_diff_vs_copy1 = - decide_dynamic_value_attrs_mapping( - graph_input1, - get_tensor_bindings_for_slot_name(invocation_mapping_diff_vs_copy1, TensorSlotName::INPUT)); - - DynamicValueAttrs graph_input2 = - mk_value(1, TensorSlotName::OUTPUT); - - DynamicValueAttrs graph_input2_use = - decide_dynamic_value_attrs_mapping( - graph_input2, - get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::WEIGHT)); - - DynamicValueAttrs invocation_output1 = mk_value(invocation_id, - TensorSlotName::OUTPUT_1); - DynamicValueAttrs invocation_output1_src = - decide_dynamic_value_attrs_mapping( - invocation_output1, - get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT_1)); - - DynamicValueAttrs invocation_output2 = mk_value(invocation_id, - TensorSlotName::OUTPUT_2); - DynamicValueAttrs invocation_output2_src = - decide_dynamic_value_attrs_mapping( - invocation_output2, - get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT_2)); - - DynamicValueAttrs graph_input1_src_same = - decide_dynamic_value_attrs_mapping( - graph_input1, - get_tensor_bindings_for_slot_name(input_mapping_same, TensorSlotName::OUTPUT)); - - DynamicValueAttrs graph_input2_src_same = - decide_dynamic_value_attrs_mapping( - graph_input2, - get_tensor_bindings_for_slot_name(weight_mapping_same, TensorSlotName::OUTPUT)); - - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - graph_input1, - }, - { - mk_slot(TensorSlotName::WEIGHT), - graph_input2, - }, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/invocation_mapping, - /*op_attrs=*/TrainingOperationAttrs{ - PCGOperatorAttrs{ - make_relu_attrs(), - }, - }, - /*layer_guid=*/ - dynamic_layer_guid_t{parallel_layer_guid_t{Node{invocation_id}}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - { - mk_slot(TensorSlotName::OUTPUT_1), - invocation_output1, + DynamicNodeInvocation mapped_inv2 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v2_placement2, + }, }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(2), mapping2, relu_attrs), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT_2), - invocation_output2, + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v3, + }, }, - }, - }; - - auto mk_copy = [&](DynamicValueAttrs const &src, - DynamicValueAttrs const &dst) { - return DynamicNodeInvocation{ - /*inputs=*/{{mk_slot(TensorSlotName::INPUT), src}}, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{{mk_slot(TensorSlotName::OUTPUT), dst}}, - }; - }; - - SUBCASE("same mapping, no copies") { - std::map sources_same{ - {graph_input1, graph_input1_src_same}, - {graph_input2, graph_input2_src_same}, - }; - - std::set result = - copies_for_invocation_inputs(input, sources_same); + }; - std::set correct = {}; + return dynamic_open_dataflow_graph_from_invocation_set( + {mapped_inv1, mapped_inv2, inserted_copy}); + }(); - CHECK(result.size() == correct.size()); - CHECK(result == correct); + CHECK_MESSAGE( + result == correct, + check_kv("result\n", dynamic_open_dataflow_graph_as_dot(result)), + check_kv("correct\n", dynamic_open_dataflow_graph_as_dot(correct))); } - SUBCASE("copy one tensor, one point") { - MappedOperatorTaskGroup input_mapping_copy1 = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_input_shard_binding(mc1_input_coord), - }, - { - mc3, - mk_input_shard_binding(mc2_input_coord), + SUBCASE("does not insert a copy when not necessary") { + DynamicNodeMapping mapping1 = DynamicNodeMapping{ + MappedOperatorTaskGroup{ + bidict{ + { + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(1_n), + mk_binding(1_n, 1_n), + }, }, }, + DeviceType::GPU, }; - MappedOperatorTaskGroup input_mapping_copy1_diff_vs_use = + DynamicNodeMapping mapping2 = DynamicNodeMapping{ MappedOperatorTaskGroup{ bidict{ { - mc3, - mk_input_shard_binding(mc2_input_coord), + mk_machine_coord(0_n), + mk_binding(0_n, 0_n), + }, + { + mk_machine_coord(1_n), + mk_binding(1_n, 1_n), }, }, - }; - - DynamicValueAttrs graph_input1_src_copy1 = - decide_dynamic_value_attrs_mapping( - graph_input1, - get_tensor_bindings_for_slot_name(input_mapping_copy1, TensorSlotName::OUTPUT)); - - DynamicValueAttrs graph_input1_src_copy1_diff_vs_use = - decide_dynamic_value_attrs_mapping( - graph_input1, - get_tensor_bindings_for_slot_name(input_mapping_copy1_diff_vs_use, TensorSlotName::OUTPUT)); - - std::map sources_copy1{ - {graph_input1, graph_input1_src_copy1}, - {graph_input2, graph_input2_src_same}}; - - std::set result = - copies_for_invocation_inputs(input, sources_copy1); - - std::set correct = { - mk_copy(graph_input1_src_copy1_diff_vs_use, graph_input1_use_diff_vs_copy1), + }, + DeviceType::GPU, }; - CHECK(result.size() == correct.size()); - CHECK(result == correct); - } - - SUBCASE("copy two tensors, two points") { - MappedOperatorTaskGroup input_mapping_copy2 = MappedOperatorTaskGroup{ - bidict{ + DynamicNodeInvocation inv1 = DynamicNodeInvocation{ + /*inputs=*/{ { - mc3, - mk_input_shard_binding(mc1_input_coord), + mk_slot(TensorSlotName::INPUT), + v1, }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, relu_attrs), + /*outputs=*/ + { { - mc4, - mk_input_shard_binding(mc2_input_coord), + mk_slot(TensorSlotName::OUTPUT), + v2, }, }, }; - MappedOperatorTaskGroup weight_mapping_copy2 = MappedOperatorTaskGroup{ - bidict{ + + DynamicNodeInvocation inv2 = DynamicNodeInvocation{ + /*inputs=*/{ { - mc4, - mk_input_shard_binding(mc1_weight_coord), + mk_slot(TensorSlotName::INPUT), + v2, }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(2), mapping2, relu_attrs), + /*outputs=*/ + { { - mc3, - mk_input_shard_binding(mc2_weight_coord), + mk_slot(TensorSlotName::OUTPUT), + v3, }, }, }; - DynamicValueAttrs graph_input1_src_copy2 = - decide_dynamic_value_attrs_mapping( - graph_input1, - get_tensor_bindings_for_slot_name(input_mapping_copy2, TensorSlotName::OUTPUT)); + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set({inv1, inv2}); - DynamicValueAttrs graph_input2_src_copy2 = - decide_dynamic_value_attrs_mapping( - graph_input2, - get_tensor_bindings_for_slot_name(weight_mapping_copy2, TensorSlotName::OUTPUT)); + DynamicOpenDataflowGraph result = perform_copy_insertion(g); - std::map sources_copy2{ - {graph_input1, graph_input1_src_copy2}, - {graph_input2, graph_input2_src_copy2}}; + DynamicOpenDataflowGraph correct = [&] { + DynamicValueAttrs mapped_v1 = mk_value_attrs( + /*src_layer_guid=*/0, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); + + DynamicValueAttrs mapped_v2 = mk_value_attrs( + /*src_layer_guid=*/1, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); + + DynamicValueAttrs mapped_v3 = mk_value_attrs( + /*src_layer_guid=*/2, + /*src_slot=*/TensorSlotName::OUTPUT, + /*mapping=*/ + ParallelTensorMapping{ + bidict{ + {mk_ptensor_coord(0_n), mk_device_id(0_n)}, + {mk_ptensor_coord(1_n), mk_device_id(1_n)}, + }, + }); - std::set result = - copies_for_invocation_inputs(input, sources_copy2); + DynamicNodeInvocation mapped_inv1 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v1, + }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(1), mapping1, relu_attrs), + /*outputs=*/ + { + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v2, + }, + }, + }; - std::set correct = { - mk_copy(graph_input1_src_copy2, graph_input1_use), - mk_copy(graph_input2_src_copy2, graph_input2_use), - }; + DynamicNodeInvocation mapped_inv2 = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mapped_v2, + }, + }, + /*node_attrs=*/ + mk_node_attrs( + mk_pcg_layer_guid(2), mapping2, relu_attrs), + /*outputs=*/ + { + { + mk_slot(TensorSlotName::OUTPUT), + mapped_v3, + }, + }, + }; - CHECK(result.size() == correct.size()); - CHECK(result == correct); + return dynamic_open_dataflow_graph_from_invocation_set( + {mapped_inv1, mapped_inv2}); + }(); + + CHECK_MESSAGE( + result == correct, + check_kv("result\n", dynamic_open_dataflow_graph_as_dot(result)), + check_kv("correct\n", dynamic_open_dataflow_graph_as_dot(correct))); } } SUBCASE("replicate operator") { - - auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, - ParallelTensorSpaceCoordinate const &c2) - -> OperatorAtomicTaskShardBinding { - return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/{ - { - TensorSlotName::INPUT, - c1, - }, - { - TensorSlotName::OUTPUT, - c2, - }, - }, + auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/FFOrdered{ + 0_n, + 0_n, + }, }; }; - ParallelTensorSpaceCoordinate mc_input_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(1_n); + MachineSpaceCoordinate mc3 = mk_machine_coord(2_n); - ParallelTensorSpaceCoordinate mc1_output_coord = - mk_pt_coord(0_n, 0_n, 0_n, 0_n); - ParallelTensorSpaceCoordinate mc2_output_coord = - mk_pt_coord(0_n, 1_n, 0_n, 0_n); + ExampleGraphTestCase tc + = mk_example_replicate_graph( + /*input_device=*/mc3, + /*relu1_device=*/mc1, + /*replicate_device1=*/mc2, + /*replicate_device2=*/mc3, + /*relu2_device1=*/mc1, + /*relu2_device2=*/mc3); - MappedOperatorTaskGroup invocation_mapping = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_shard_binding(mc_input_coord, - mc1_output_coord), - }, - { - mc2, - mk_shard_binding(mc_input_coord, - mc2_output_coord), - }, - }, - }; + std::set copies = infer_all_copies_in_graph(tc.g); + ASSERT(copies.size() == 2); - DynamicValueAttrs graph_input_unmapped = - mk_value(0, TensorSlotName::OUTPUT); + DynamicOpenDataflowGraph result = perform_copy_insertion(tc.g); - DynamicValueAttrs invocation_output_unmapped = - mk_value(invocation_id, TensorSlotName::OUTPUT); - DynamicValueAttrs invocation_output_src_mapped = - decide_dynamic_value_attrs_mapping( - invocation_output_unmapped, - get_tensor_bindings_for_slot_name(invocation_mapping, TensorSlotName::OUTPUT)); + DynamicOpenDataflowGraph correct = [&] { + auto map_input_value = [](DynamicNodeInvocation const &invocation, + ParallelTensorMapping const &mapping) + -> DynamicNodeInvocation + { + DynamicNodeInvocation result = invocation; + DynamicTensorSlot input_slot = mk_slot(TensorSlotName::INPUT); + result.inputs = { + { + input_slot, + decide_dynamic_value_attrs_mapping( + require_only_key(invocation.inputs, input_slot), + mapping), + }, + }; + return result; + }; - DynamicNodeInvocation input = DynamicNodeInvocation{ - /*inputs=*/{ + auto map_output_value = [](DynamicNodeInvocation const &invocation, + ParallelTensorMapping const &mapping) + -> DynamicNodeInvocation + { + DynamicNodeInvocation result = invocation; + DynamicTensorSlot output_slot = mk_slot(TensorSlotName::OUTPUT); + result.outputs = { { - mk_slot(TensorSlotName::INPUT), - graph_input_unmapped, + output_slot, + decide_dynamic_value_attrs_mapping( + require_only_key(invocation.outputs, output_slot), + mapping), }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/invocation_mapping, - /*op_attrs=*/TrainingOperationAttrs{ - PCGOperatorAttrs{ - ReplicateAttrs{ - 2_p, - }, + }; + return result; + }; + + auto map_input_and_output_values = [&](DynamicNodeInvocation const &invocation, + ParallelTensorMapping const &input_mapping, + ParallelTensorMapping const &output_mapping) + -> DynamicNodeInvocation + { + return map_input_value( + map_output_value( + invocation, + output_mapping), + input_mapping); + }; + + auto mk_single_shard_mapping = [&](MachineSpaceCoordinate const &mc) -> ParallelTensorMapping { + return ParallelTensorMapping{ + bidict{ + { + mk_pt_coord(0_n), + mk_device_id(mc), }, }, - /*layer_guid=*/dynamic_layer_guid_t{ - parallel_layer_guid_t{ - Node{invocation_id}, + }; + }; + + auto mk_two_shard_mapping = [&](MachineSpaceCoordinate const &mc1, + MachineSpaceCoordinate const &mc2) -> ParallelTensorMapping { + return ParallelTensorMapping{ + bidict{ + { + mk_pt_coord(0_n), + mk_device_id(mc1), + }, + { + mk_pt_coord(1_n), + mk_device_id(mc2), }, }, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ - { - mk_slot(TensorSlotName::OUTPUT), - invocation_output_unmapped, - }, - }, - }; - - std::map unmapped_to_mapped_source_value = { - { - graph_input_unmapped, - decide_dynamic_value_attrs_mapping( - graph_input_unmapped, - bidict{ - { - mc_input_coord, - mc3, - }, - }) - }, + }; }; - std::set result = copies_for_invocation_inputs( - input, unmapped_to_mapped_source_value); - - std::set correct = {}; + DynamicNodeInvocation input_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.input_op_id); + DynamicNodeInvocation relu1_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.relu1_op_id); + DynamicNodeInvocation replicate_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.replicate_op_id); + DynamicNodeInvocation relu2_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.relu2_op_id); - nlohmann::json result_j = transform(result, dynamic_node_invocation_to_serializable); - nlohmann::json correct_j = transform(correct, dynamic_node_invocation_to_serializable); + DynamicNodeInvocation value_mapped_input_invocation = + map_output_value( + input_invocation, + mk_single_shard_mapping(mc3)); - CHECK(result_j == correct_j); - DynamicOpenDataflowGraph g = - dynamic_open_dataflow_graph_from_invocation_set({inv1, inv2}); + DynamicNodeInvocation value_mapped_relu1_invocation = + map_input_and_output_values( + relu1_invocation, + mk_single_shard_mapping(mc1), + mk_single_shard_mapping(mc1)); - DynamicOpenDataflowGraph result = perform_copy_insertion(g); + DynamicNodeInvocation value_mapped_replicate_invocation = + map_input_and_output_values( + replicate_invocation, + mk_single_shard_mapping(mc1), + mk_two_shard_mapping(mc2, mc3)); - DynamicOpenDataflowGraph correct = [&] { - DynamicValueAttrs mapped_v1 = mk_value_attrs( - /*src_layer_guid=*/0, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/ - ParallelTensorMapping{ - bidict{ - {mk_ptensor_coord(0_n), mk_device_id(0_n)}, - {mk_ptensor_coord(1_n), mk_device_id(1_n)}, - }, - }); - - DynamicValueAttrs mapped_v2 = mk_value_attrs( - /*src_layer_guid=*/1, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/ - ParallelTensorMapping{ - bidict{ - {mk_ptensor_coord(0_n), mk_device_id(0_n)}, - {mk_ptensor_coord(1_n), mk_device_id(1_n)}, - }, - }); - - DynamicValueAttrs mapped_v3 = mk_value_attrs( - /*src_layer_guid=*/2, - /*src_slot=*/TensorSlotName::OUTPUT, - /*mapping=*/ - ParallelTensorMapping{ - bidict{ - {mk_ptensor_coord(0_n), mk_device_id(0_n)}, - {mk_ptensor_coord(1_n), mk_device_id(1_n)}, - }, - }); + DynamicNodeInvocation value_mapped_relu2_invocation = + map_input_and_output_values( + relu2_invocation, + mk_two_shard_mapping(mc1, mc3), + mk_two_shard_mapping(mc1, mc3)); - DynamicNodeInvocation mapped_inv1 = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mapped_v1, - }, + DynamicNodeInvocation input_to_relu1_copy = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + require_only_key(value_mapped_input_invocation.outputs, mk_slot(TensorSlotName::OUTPUT)), }, - /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, /*op_attrs=*/std::nullopt), - /*outputs=*/ + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ { - { - mk_slot(TensorSlotName::OUTPUT), - mapped_v2, - }, + mk_slot(TensorSlotName::OUTPUT), + require_only_key(value_mapped_relu1_invocation.inputs, mk_slot(TensorSlotName::INPUT)), }, + }, }; - DynamicNodeInvocation mapped_inv2 = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mapped_v2, - }, + DynamicNodeInvocation replicate_to_relu2_copy = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + require_only_key(value_mapped_replicate_invocation.outputs, mk_slot(TensorSlotName::OUTPUT)), }, - /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, /*op_attrs=*/std::nullopt), - /*outputs=*/ + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ { - { - mk_slot(TensorSlotName::OUTPUT), - mapped_v3, - }, + mk_slot(TensorSlotName::OUTPUT), + require_only_key(value_mapped_relu2_invocation.inputs, mk_slot(TensorSlotName::INPUT)), }, + }, }; - return dynamic_open_dataflow_graph_from_invocation_set( - {mapped_inv1, mapped_inv2}); + return dynamic_open_dataflow_graph_from_invocation_set({ + value_mapped_input_invocation, + input_to_relu1_copy, + value_mapped_relu1_invocation, + value_mapped_replicate_invocation, + replicate_to_relu2_copy, + value_mapped_relu2_invocation, + }); }(); + nlohmann::json result_json = dynamic_open_dataflow_graph_to_serializable(result); + nlohmann::json correct_json = dynamic_open_dataflow_graph_to_serializable(correct); + CHECK_MESSAGE( result == correct, - check_kv("result\n", dynamic_open_dataflow_graph_as_dot(result)), - check_kv("correct\n", dynamic_open_dataflow_graph_as_dot(correct))); + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); + } + + SUBCASE("copy insertion commutes with pass expansion") { + MachineSpaceCoordinate mc1 = mk_machine_coord(0_n); + MachineSpaceCoordinate mc2 = mk_machine_coord(1_n); + MachineSpaceCoordinate mc3 = mk_machine_coord(2_n); + + SUBCASE("graph is single input node") { + DynamicOpenDataflowGraph g = mk_single_input_node_graph(mc1); + + DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = + perform_copy_insertion(perform_pass_expansion(g)); + DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = + perform_pass_expansion(perform_copy_insertion(g)); + + nlohmann::json pass_expansion_before_copy_insertion_json + = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json + = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + + CHECK_MESSAGE( + pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + check_kv("pass_expansion_before_copy_insertion_json\n", + pass_expansion_before_copy_insertion_json.dump()), + check_kv("copy_insertion_before_pass_expansion_json\n", + copy_insertion_before_pass_expansion_json.dump()), + check_kv("pass_expansion_before_copy_insertion\n", + dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + check_kv("copy_insertion_before_pass_expansion\n", + dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + } + + SUBCASE("graph is single input node followed by relu") { + SUBCASE("operations are mapped to the same device") { + DynamicOpenDataflowGraph g = mk_single_input_into_relu_graph(mc1, mc1); + + DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = + perform_copy_insertion(perform_pass_expansion(g)); + DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = + perform_pass_expansion(perform_copy_insertion(g)); + + nlohmann::json pass_expansion_before_copy_insertion_json + = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json + = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + + CHECK_MESSAGE( + pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + check_kv("pass_expansion_before_copy_insertion_json\n", + pass_expansion_before_copy_insertion_json.dump()), + check_kv("copy_insertion_before_pass_expansion_json\n", + copy_insertion_before_pass_expansion_json.dump()), + check_kv("pass_expansion_before_copy_insertion\n", + dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + check_kv("copy_insertion_before_pass_expansion\n", + dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + } + + SUBCASE("operations are mapped to different devices") { + DynamicOpenDataflowGraph g = mk_single_input_into_relu_graph(mc1, mc2); + + DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = + perform_copy_insertion(perform_pass_expansion(g)); + DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = + perform_pass_expansion(perform_copy_insertion(g)); + + nlohmann::json pass_expansion_before_copy_insertion_json + = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json + = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + + CHECK_MESSAGE( + pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + check_kv("pass_expansion_before_copy_insertion_json\n", + pass_expansion_before_copy_insertion_json.dump()), + check_kv("copy_insertion_before_pass_expansion_json\n", + copy_insertion_before_pass_expansion_json.dump()), + check_kv("pass_expansion_before_copy_insertion\n", + dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + check_kv("copy_insertion_before_pass_expansion\n", + dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + } + } + + SUBCASE("multinode graph including replicate") { + ExampleGraphTestCase tc + = mk_example_replicate_graph( + /*input_device=*/mc3, + /*relu1_device=*/mc1, + /*replicate_device1=*/mc2, + /*replicate_device2=*/mc3, + /*relu2_device1=*/mc1, + /*relu2_device2=*/mc3); + + DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = + perform_copy_insertion(perform_pass_expansion(tc.g)); + DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = + perform_pass_expansion(perform_copy_insertion(tc.g)); + + nlohmann::json pass_expansion_before_copy_insertion_json + = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json + = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + + CHECK_MESSAGE( + pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + check_kv("pass_expansion_before_copy_insertion_json\n", + pass_expansion_before_copy_insertion_json.dump()), + check_kv("copy_insertion_before_pass_expansion_json\n", + copy_insertion_before_pass_expansion_json.dump()), + check_kv("pass_expansion_before_copy_insertion\n", + dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + check_kv("copy_insertion_before_pass_expansion\n", + dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + } } } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index 2804a397c7..ce0e710b7f 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -5,6 +5,7 @@ #include "utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h" #include "utils/graph/node/algorithms.h" #include +#include using namespace ::FlexFlow; @@ -21,6 +22,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -32,6 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }; }; @@ -50,7 +53,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("correct usage") { DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ { mk_slot(TensorSlotName::INPUT), value_1, @@ -58,7 +61,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { mk_slot(TensorSlotName::OUTPUT), value_2, @@ -67,10 +70,10 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{}, + /*inputs=*/std::map{}, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { mk_slot(TensorSlotName::OUTPUT), value_3, @@ -79,7 +82,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicNodeInvocation invocation_3 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ { mk_slot(TensorSlotName::INPUT), value_1, @@ -95,10 +98,10 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{}, + std::map{}, }; - std::unordered_set invocation_set = { + std::set invocation_set = { invocation_1, invocation_2, invocation_3, @@ -112,7 +115,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("throws if multiple invocations produce the same value") { DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ { mk_slot(TensorSlotName::INPUT), value_1, @@ -120,7 +123,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { mk_slot(TensorSlotName::OUTPUT), value_2, @@ -129,10 +132,10 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{}, + /*inputs=*/std::map{}, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { mk_slot(TensorSlotName::OUTPUT), value_2, @@ -140,7 +143,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set invocation_set = { + std::set invocation_set = { invocation_1, invocation_2, }; @@ -151,7 +154,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("throws if invocations contain/create cycle") { DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ { mk_slot(TensorSlotName::INPUT), value_1, @@ -159,7 +162,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { mk_slot(TensorSlotName::OUTPUT), value_2, @@ -168,7 +171,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicNodeInvocation invocation_2 = DynamicNodeInvocation{ - /*inputs=*/std::unordered_map{ + /*inputs=*/std::map{ { mk_slot(TensorSlotName::INPUT), value_2, @@ -176,7 +179,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*node_attrs=*/node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { mk_slot(TensorSlotName::OUTPUT), value_1, @@ -184,7 +187,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set invocation_set = { + std::set invocation_set = { invocation_1, invocation_2, }; @@ -205,6 +208,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -261,7 +265,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*slot_tensor_role=*/std::nullopt, /*task_shard=*/std::nullopt, }, - value_2, + value_3, }, }, }; @@ -305,48 +309,62 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicOpenDataflowGraph g = dynamic_open_dataflow_graph_from_invocation_set( - std::unordered_set{invocation_1, invocation_2}); + invocation_set); + + dynamic_invocation_id_t invocation_1_id = + dynamic_graph_get_id_for_invocation(g, invocation_1); + dynamic_invocation_id_t invocation_2_id = + dynamic_graph_get_id_for_invocation(g, invocation_2); + dynamic_invocation_id_t invocation_3_id = + dynamic_graph_get_id_for_invocation(g, invocation_3); + + dynamic_value_id_t value_1_id = dynamic_graph_get_id_for_value(g, value_1); - std::unordered_set result = get_dynamic_slot_sites(g); + std::set result = get_dynamic_slot_sites(g); - auto mk_internal_slot_site = [](DynamicNodeInvocation const &invocation, + auto mk_internal_slot_site = [](dynamic_invocation_id_t const &invocation_id, TensorDirection direction, - TensorSlotName slot_name) { + TensorSlotName slot_name) + -> DynamicSlotSite + { return DynamicSlotSite{ InternalDynamicSlotSite{ - /*invocation=*/invocation, + /*invocation_id=*/invocation_id, /*direction=*/direction, /*slot_name=*/ DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, }, }; }; - std::unordered_set correct = { - DynamicSlotSite{ - ExternalDynamicSlotSite{ - value_1, - }, - }, + std::set correct = { DynamicSlotSite{ ExternalDynamicSlotSite{ - value_3, + value_1_id.require_external(), }, }, mk_internal_slot_site( - invocation_1, TensorDirection::INCOMING, TensorSlotName::INPUT), + invocation_1_id, TensorDirection::INCOMING, TensorSlotName::INPUT), + mk_internal_slot_site( + invocation_1_id, TensorDirection::OUTPUT, TensorSlotName::OUTPUT), mk_internal_slot_site( - invocation_1, TensorDirection::OUTPUT, TensorSlotName::OUTPUT), + invocation_2_id, TensorDirection::OUTPUT, TensorSlotName::OUTPUT), mk_internal_slot_site( - invocation_2, TensorDirection::INCOMING, TensorSlotName::INPUT), + invocation_3_id, TensorDirection::INCOMING, TensorSlotName::INPUT), mk_internal_slot_site( - invocation_2, TensorDirection::INCOMING, TensorSlotName::WEIGHT), + invocation_3_id, TensorDirection::INCOMING, TensorSlotName::WEIGHT), + mk_internal_slot_site( + invocation_3_id, TensorDirection::INCOMING, TensorSlotName::BIAS), }; - CHECK(result == correct); + nlohmann::json result_json = result; + nlohmann::json correct_json = correct; + + CHECK(result_json == correct_json); } TEST_CASE( @@ -395,11 +413,13 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot fwd_weight_output_slot1 = DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }; DynamicValueAttrs fwd_weight_output_attrs1 = DynamicValueAttrs{ /*tensor_guid=*/tensor_guid, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -410,7 +430,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*inputs=*/{}, /*node_attrs=*/fwd_weight_node_attrs, /*outputs=*/ - std::unordered_map{ + std::map{ { fwd_weight_output_slot1, fwd_weight_output_attrs1, @@ -430,11 +450,13 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicTensorSlot upd_weight_input_slot2 = DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }; DynamicValueAttrs upd_weight_input_attrs2 = DynamicValueAttrs{ /*tensor_guid=*/tensor_guid, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -445,11 +467,13 @@ TEST_SUITE(FF_TEST_SUITE) { /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/ mk_dynamic_tensor_role_opt(OptimizerSlotName::SGD_V), + /*task_shard=*/std::nullopt, }; DynamicValueAttrs upd_weight_input_attrs3 = DynamicValueAttrs{ /*tensor_guid=*/tensor_guid, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc index 6942f6658f..df602bebf3 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/machine_slicing.cc @@ -76,6 +76,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/shard_coord, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -114,7 +115,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/std::nullopt, - /*device_coords=*/nonempty_set{mc2}, + /*device_ids=*/nonempty_set{mc2}, /*mapping=*/std::nullopt, /*op_attrs=*/std::nullopt, /*layer_guid=*/ diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 2e14c88654..61fe183ac0 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -3,15 +3,20 @@ #include "utils/containers/require_only_key.h" #include "op-attrs/ops/element_unary.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" +#include "op-attrs/initializer_attrs.h" +#include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h" +#include "test/utils/doctest/check_kv.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("make_dynamic_node_invocation_from_mapped") { - SUBCASE("Replicate") { - MachineSpaceCoordinate gpu0 = MachineSpaceCoordinate{0_n, 0_n, DeviceType::GPU}; - MachineSpaceCoordinate gpu1 = MachineSpaceCoordinate{0_n, 1_n, DeviceType::GPU}; + MachineSpaceCoordinate gpu0 = MachineSpaceCoordinate{0_n, 0_n}; + MachineSpaceCoordinate gpu1 = MachineSpaceCoordinate{0_n, 1_n}; + SUBCASE("Replicate") { ParallelTensorSpaceCoordinate tensor_coord0 = ParallelTensorSpaceCoordinate{ /*sum_component=*/0_n, /*discard_copy_component=*/0_n, @@ -24,7 +29,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*shard_component=*/FFOrdered{0_n}, }; - MappedOperatorTaskGroup mapping = MappedOperatorTaskGroup{ + MappedOperatorTaskGroup mapped_op_task_group = MappedOperatorTaskGroup{ { { gpu0, @@ -100,7 +105,7 @@ TEST_SUITE(FF_TEST_SUITE) { /*op_attrs=*/op_attrs, /*name=*/std::nullopt, }, - /*mapping=*/mapping, + /*mapping=*/mapped_op_task_group, }, /*outgoing=*/{ { @@ -116,7 +121,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - DynamicNodeInvocation result = make_dynamic_node_invocation_from_mapped(input); + DynamicNodeInvocation result = make_dynamic_node_invocation_from_mapped(input, DeviceType::GPU); DynamicNodeInvocation correct = DynamicNodeInvocation{ /*inputs=*/{ @@ -129,6 +134,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, /*parallel_tensor_shape=*/input_shape, + /*create_grad=*/true, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -139,7 +145,10 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_attrs=*/DynamicNodeAttrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, - /*mapping=*/mapping, + /*mapping=*/DynamicNodeMapping{ + /*op_task_group=*/mapped_op_task_group, + /*device_type=*/DeviceType::GPU, + }, /*op_attrs=*/TrainingOperationAttrs{op_attrs}, /*layer_guid=*/dynamic_layer_guid_t{layer_guid}, /*per_device_op_state=*/std::nullopt, @@ -154,6 +163,7 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, /*parallel_tensor_shape=*/output_shape, + /*create_grad=*/true, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -163,236 +173,893 @@ TEST_SUITE(FF_TEST_SUITE) { } }; - CHECK(result == correct); + nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + + CHECK_MESSAGE(result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); } - // SUBCASE("standard op") { - // - // } + SUBCASE("standard op") { + PCGOperatorAttrs op_attrs = PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/7_p, + /*use_bias=*/true, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, + }; + + parallel_layer_guid_t layer_guid = parallel_layer_guid_t{Node{0}}; + + auto mk_tensor_guid = [](size_t node_id) + -> parallel_tensor_guid_t + { + return parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{node_id}, + TensorSlotName::OUTPUT, + }, + }; + }; + + parallel_tensor_guid_t input_tensor_guid = mk_tensor_guid(5); + parallel_tensor_guid_t weight_tensor_guid = mk_tensor_guid(6); + parallel_tensor_guid_t bias_tensor_guid = mk_tensor_guid(7); + parallel_tensor_guid_t output_tensor_guid = mk_tensor_guid(0); + + ParallelTensorShape input_shape = ParallelTensorShape{ + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{8_p, 2_p}, + ShardParallelDim{5_p, 1_p}, + }, + /*replica_dims=*/ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{1_p}, + }, + }, + /*data_type=*/DataType::FLOAT, + }; + + ParallelTensorShape weight_shape = ParallelTensorShape{ + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{5_p, 1_p}, + ShardParallelDim{7_p, 1_p}, + }, + /*replica_dims=*/ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{2_p}, + }, + }, + /*data_type=*/DataType::FLOAT, + }; + + ParallelTensorShape bias_shape = ParallelTensorShape{ + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{7_p, 1_p}, + }, + /*replica_dims=*/ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{2_p}, + }, + }, + /*data_type=*/DataType::FLOAT, + }; + + ParallelTensorShape output_shape = ParallelTensorShape{ + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{8_p, 2_p}, + ShardParallelDim{7_p, 1_p}, + }, + /*replica_dims=*/ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{1_p}, + }, + }, + /*data_type=*/DataType::FLOAT, + }; + + auto mk_2d_pt_coord = [](nonnegative_int replica_coord, nonnegative_int shard_coord) + -> ParallelTensorSpaceCoordinate + { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_compnent=*/replica_coord, + /*shard_components=*/FFOrdered{ + shard_coord, + 0_n, + }, + }; + }; + + auto mk_1d_pt_coord = [](nonnegative_int replica_coord, nonnegative_int shard_coord) + -> ParallelTensorSpaceCoordinate + { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_compnent=*/replica_coord, + /*shard_components=*/FFOrdered{ + shard_coord, + }, + }; + }; + + MappedOperatorTaskGroup mapped_op_task_group = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_2d_pt_coord(0_n, 0_n)}, + {TensorSlotName::WEIGHT, mk_2d_pt_coord(0_n, 0_n)}, + {TensorSlotName::BIAS, mk_1d_pt_coord(0_n, 0_n)}, + {TensorSlotName::OUTPUT, mk_2d_pt_coord(0_n, 0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_2d_pt_coord(0_n, 1_n)}, + {TensorSlotName::WEIGHT, mk_2d_pt_coord(1_n, 0_n)}, + {TensorSlotName::BIAS, mk_1d_pt_coord(1_n, 0_n)}, + {TensorSlotName::OUTPUT, mk_2d_pt_coord(0_n, 1_n)}, + }}, + }, + }, + }; + + MappedParallelLayerInvocationInfo input = MappedParallelLayerInvocationInfo{ + /*incoming=*/{ + { + TensorSlotName::INPUT, + ParallelTensorInfo{ + /*guid=*/input_tensor_guid, + /*attrs=*/ParallelTensorAttrs{ + /*shape=*/input_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + { + TensorSlotName::WEIGHT, + ParallelTensorInfo{ + /*guid=*/weight_tensor_guid, + /*attrs=*/ParallelTensorAttrs{ + /*shape=*/weight_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + { + TensorSlotName::BIAS, + ParallelTensorInfo{ + /*guid=*/bias_tensor_guid, + /*attrs=*/ParallelTensorAttrs{ + /*shape=*/bias_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + }, + /*layer_info=*/MappedParallelLayerInfo{ + /*guid=*/layer_guid, + /*attrs=*/ParallelLayerAttrs{ + /*op_attrs=*/op_attrs, + /*name=*/std::nullopt, + }, + /*mapping=*/mapped_op_task_group, + }, + /*outgoing=*/{ + { + TensorSlotName::OUTPUT, + ParallelTensorInfo{ + /*guid=*/output_tensor_guid, + /*attrs=*/ParallelTensorAttrs{ + /*shape=*/output_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + }, + }; + + DynamicNodeInvocation result = make_dynamic_node_invocation_from_mapped(input, DeviceType::GPU); + + DynamicNodeInvocation correct = [&]() + -> DynamicNodeInvocation + { + auto mk_slot = [](TensorSlotName slot_name) + -> DynamicTensorSlot + { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }; + }; + + auto mk_value = [](parallel_tensor_guid_t const &tensor_guid, + ParallelTensorShape const &shape) + -> DynamicValueAttrs + { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{tensor_guid}, + /*parallel_tensor_shape=*/shape, + /*create_grad=*/true, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + }; + + return DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_tensor_guid, input_shape), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(weight_tensor_guid, weight_shape), + }, + { + mk_slot(TensorSlotName::BIAS), + mk_value(bias_tensor_guid, bias_shape), + }, + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/DynamicNodeMapping{ + /*op_task_group=*/mapped_op_task_group, + /*device_type=*/DeviceType::GPU, + }, + /*op_attrs=*/TrainingOperationAttrs{op_attrs}, + /*layer_guid=*/dynamic_layer_guid_t{layer_guid}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(output_tensor_guid, output_shape), + }, + } + }; + }(); + + nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + + CHECK_MESSAGE(result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); + } } - // TEST_CASE("make_dynamic_open_dataflow_graph_from_mapped_pcg") { - // positive_int batch_size = 10_p; - // positive_int data_dim = 16_p; - // positive_int hidden_dim = 32_p; - // positive_int output_dim = 1_p; - - // auto make_layer_attrs = [](auto const &op_attrs) -> ParallelLayerAttrs { - // return ParallelLayerAttrs{ - // /*op_attrs=*/PCGOperatorAttrs{op_attrs}, - // /*name=*/std::nullopt, - // }; - // }; - - - // TensorShape output_tensor_shape = TensorShape{ - // TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - // TensorShape label_tensor_shape = TensorShape{ - // TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - // ParallelComputationGraph pcg = empty_parallel_computation_graph(); - - // TensorShape input_tensor_shape = TensorShape{ - // TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; - - // ParallelLayerAddedResult inputs_layer = - // pcg_add_input_layer(pcg, input_tensor_shape); - // parallel_tensor_guid_t t_input = - // require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); - - // ParallelLayerAddedResult inputs_layer_2 = - // pcg_add_input_layer(pcg, input_tensor_shape); - // parallel_tensor_guid_t t_input_2 = - // require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); - - // ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ - // OperatorType::EW_ADD, - // DataType::FLOAT, - // false, - // false, - // }; - - // ParallelLayerAddedResult add_operator_1 = - // add_parallel_layer(pcg, - // make_layer_attrs(add_attrs), - // { - // { - // TensorSlotName::LHS_INPUT, - // t_input, - // }, - // { - // TensorSlotName::RHS_INPUT, - // t_input_2, - // }, - // }, - // /*weights=*/{}); - - // parallel_tensor_guid_t t_add_1 = - // require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); - - // positive_int replicate_degree = 2_p; - // ReplicateAttrs repl_attrs = ReplicateAttrs{replicate_degree}; - // ParallelLayerAddedResult repl_operator_1 = - // add_parallel_layer(pcg, - // make_layer_attrs(repl_attrs), - // { - // { - // TensorSlotName::INPUT, - // t_add_1, - // }, - // }, - // /*weight=*/{}); - - // parallel_tensor_guid_t t_repl_1 = - // require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); - - // ParallelLayerAddedResult relu_operator_1 = - // add_parallel_layer(pcg, - // make_layer_attrs(make_relu_attrs()), - // /*inputs=*/ - // { - // { - // TensorSlotName::INPUT, - // t_repl_1, - // }, - // }, - // /*weights=*/{}); - - // parallel_tensor_guid_t t_relu_1 = - // require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); - - // MachineSpaceCoordinate gpu0{0_n, 0_n, DeviceType::GPU}; - // MachineSpaceCoordinate gpu1{0_n, 1_n, DeviceType::GPU}; - - // ParallelTensorSpaceCoordinate tensor_coord0{ - // /*sum_component=*/0_n, - // /*discard_copy_component=*/0_n, - // /*shard_component=*/FFOrdered{0_n}}; - // ParallelTensorSpaceCoordinate tensor_coord1{ - // /*sum_component=*/0_n, - // /*discard_copy_component=*/1_n, - // /*shard_component=*/FFOrdered{0_n}}; - - // MappedOperatorTaskGroup input_1_mapping = MappedOperatorTaskGroup{ - // { - // { - // gpu0, - // OperatorAtomicTaskShardBinding{{ - // {TensorSlotName::OUTPUT, tensor_coord0}, - // }}, - // }, - // }, - // }; - - // MappedOperatorTaskGroup input_2_mapping = MappedOperatorTaskGroup{ - // { - // { - // gpu0, - // OperatorAtomicTaskShardBinding{{ - // {TensorSlotName::OUTPUT, tensor_coord0}, - // }}, - // }, - // }, - // }; - - // MappedOperatorTaskGroup add_operator_1_mapping = MappedOperatorTaskGroup{ - // { - // { - // gpu0, - // OperatorAtomicTaskShardBinding{{ - // {TensorSlotName::LHS_INPUT, tensor_coord0}, - // {TensorSlotName::RHS_INPUT, tensor_coord0}, - // {TensorSlotName::OUTPUT, tensor_coord0}, - // }}, - // }, - // }, - // }; - - // MappedOperatorTaskGroup repl_operator_1_mapping = MappedOperatorTaskGroup{ - // { - // { - // gpu0, - // OperatorAtomicTaskShardBinding{{ - // {TensorSlotName::OUTPUT, tensor_coord0}, - // }}, - // }, - // { - // gpu1, - // OperatorAtomicTaskShardBinding{{ - // {TensorSlotName::OUTPUT, tensor_coord1}, - // }}, - // }, - // }, - // }; - - // MappedOperatorTaskGroup relu_operator_1_mapping = MappedOperatorTaskGroup{ - // { - // { - // gpu0, - // OperatorAtomicTaskShardBinding{{ - // {TensorSlotName::INPUT, tensor_coord0}, - // {TensorSlotName::OUTPUT, tensor_coord0}, - // }}, - // }, - // { - // gpu1, - // OperatorAtomicTaskShardBinding{{ - // {TensorSlotName::INPUT, tensor_coord1}, - // {TensorSlotName::OUTPUT, tensor_coord1}, - // }}, - // }, - // }, - // }; - - // MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( - // /*pcg=*/pcg, - // /*mapped_op_task_groups=*/{ - // { - // inputs_layer.parallel_layer, - // input_1_mapping, - // }, - // { - // inputs_layer_2.parallel_layer, - // input_2_mapping, - // }, - // { - // add_operator_1.parallel_layer, - // add_operator_1_mapping, - // }, - // { - // repl_operator_1.parallel_layer, - // repl_operator_1_mapping, - // }, - // { - // relu_operator_1.parallel_layer, - // relu_operator_1_mapping, - // }, - // }); - - - // DynamicOpenDataflowGraph result = make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg); - - // DynamicNodeInvocation input_1_invocation = DynamicNodeInvocation{ - // DynamicNodeAttrs{ - // /*task_type=*/std::nullopt, - // /*device_coord=*/std::nullopt, - // /*mapping=*/input_1_mapping, - // /*op_attrs=*/TrainingOperationAttrs{ - // /*pcg_layer_guid=*/ - // /*per_device_op_state=*/std::nullopt, - // }, - // }; - - // DynamicNodeInvocation input_2_invocation = - // DynamicNodeInvocation add_operator_1_invocation = - // DynamicNodeInvocation repl_operator_1_invocation = - // DynamicNodeInvocation relu_operator_1_invocation = - - // DynamicOpenDataflowGraph correct = dynamic_open_dataflow_graph_from_invocation_set( - // /*invocations=*/{ - - // }, - // }; - // } + TEST_CASE("make_dynamic_open_dataflow_graph_from_mapped_pcg") { + positive_int batch_size = 10_p; + positive_int data_dim = 16_p; + positive_int hidden_dim = 32_p; + positive_int output_dim = 1_p; + + auto make_layer_attrs = [](PCGOperatorAttrs const &op_attrs) -> ParallelLayerAttrs { + return ParallelLayerAttrs{ + /*op_attrs=*/op_attrs, + /*name=*/std::nullopt, + }; + }; + + TensorShape input_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + TensorShape weight_1_tensor_shape = TensorShape{ + TensorDims{FFOrdered{hidden_dim, data_dim}}, DataType::FLOAT}; + + TensorShape weight_2_tensor_shape = TensorShape{ + TensorDims{FFOrdered{output_dim, hidden_dim}}, DataType::FLOAT}; + + TensorShape output_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + TensorShape label_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + PCGOperatorAttrs input_op_attrs = PCGOperatorAttrs{ + InputAttrs{ + /*tensor_shape=*/input_tensor_shape + }, + }; + + PCGOperatorAttrs partition_input_op_attrs = PCGOperatorAttrs{ + RepartitionAttrs{ + /*repartition_dim=*/ff_dim_t{0_n}, + /*repartition_degree=*/2_p, + }, + }; + + PCGOperatorAttrs weight_1_op_attrs = PCGOperatorAttrs{ + WeightAttrs{ + /*tensor_shape=*/weight_1_tensor_shape, + /*initializer=*/make_kaiming_uniform(weight_1_tensor_shape.dims), + }, + }; + + PCGOperatorAttrs replicate_weight_1_op_attrs = PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, + }; + + PCGOperatorAttrs weight_2_op_attrs = PCGOperatorAttrs{ + WeightAttrs{ + /*tensor_shape=*/weight_2_tensor_shape, + /*initializer=*/make_kaiming_uniform(weight_1_tensor_shape.dims), + }, + }; + + PCGOperatorAttrs replicate_weight_2_op_attrs = PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, + }; + + PCGOperatorAttrs linear_1_op_attrs = PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/hidden_dim, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, + }; + + PCGOperatorAttrs linear_2_op_attrs = PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/output_dim, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, + }; + + ParallelLayerAddedResult input_layer = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input = + require_only_key(input_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult partition_input_layer = + add_parallel_layer(pcg, + make_layer_attrs(partition_input_op_attrs), + /*inputs=*/{ + {TensorSlotName::INPUT, t_input}, + }, + /*weights=*/{}); + parallel_tensor_guid_t t_partitioned_input = + require_only_key(partition_input_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult weight_1_layer = + add_parallel_layer(pcg, + make_layer_attrs(weight_1_op_attrs), + /*inputs=*/{}, + /*weights=*/{}); + parallel_tensor_guid_t t_weight_1 = + require_only_key(weight_1_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult replicate_weight_1_layer = + add_parallel_layer(pcg, + make_layer_attrs(replicate_weight_1_op_attrs), + /*inputs=*/{ + {TensorSlotName::INPUT, t_weight_1}, + }, + /*weights=*/{}); + parallel_tensor_guid_t t_replicated_weight_1 = + require_only_key(replicate_weight_1_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult weight_2_layer = + add_parallel_layer(pcg, + make_layer_attrs(weight_2_op_attrs), + /*inputs=*/{}, + /*weights=*/{}); + parallel_tensor_guid_t t_weight_2 = + require_only_key(weight_2_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult replicate_weight_2_layer = + add_parallel_layer(pcg, + make_layer_attrs(replicate_weight_2_op_attrs), + /*inputs=*/{ + {TensorSlotName::INPUT, t_weight_2}, + }, + /*weights=*/{}); + parallel_tensor_guid_t t_replicated_weight_2 = + require_only_key(replicate_weight_2_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult linear_1_layer = + add_parallel_layer(pcg, + make_layer_attrs(linear_1_op_attrs), + /*inputs=*/{ + {TensorSlotName::INPUT, t_partitioned_input}, + }, + /*weights=*/{ + {TensorSlotName::WEIGHT, t_replicated_weight_1}, + }); + parallel_tensor_guid_t t_hidden_activation = + require_only_key(linear_1_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult linear_2_layer = + add_parallel_layer(pcg, + make_layer_attrs(linear_2_op_attrs), + /*inputs=*/{ + {TensorSlotName::INPUT, t_hidden_activation}, + }, + /*weights=*/{ + {TensorSlotName::WEIGHT, t_replicated_weight_2}, + }); + parallel_tensor_guid_t t_output = + require_only_key(linear_2_layer.outputs, TensorSlotName::OUTPUT); + + + + MachineSpaceCoordinate gpu0 = MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }; + MachineSpaceCoordinate gpu1 = MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/1_n, + }; + + auto mk_shard_coord = [](nonnegative_int shard_idx) + -> ParallelTensorSpaceCoordinate + { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{shard_idx, 0_n}, + }; + }; + + auto mk_replica_coord = [](nonnegative_int replica_idx) + -> ParallelTensorSpaceCoordinate + { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/replica_idx, + /*shard_component=*/FFOrdered{0_n, 0_n}, + }; + }; + + ParallelTensorSpaceCoordinate nonparallel_coord = ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{0_n, 0_n}, + }; + + MappedOperatorTaskGroup input_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, nonparallel_coord}, + }}, + }, + }, + }; + + MappedOperatorTaskGroup weight_1_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, nonparallel_coord}, + }}, + }, + }, + }; + + MappedOperatorTaskGroup weight_2_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, nonparallel_coord}, + }}, + }, + }, + }; + + MappedOperatorTaskGroup partition_input_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, + }}, + }, + }, + }; + + MappedOperatorTaskGroup replicate_weight_1_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(1_n)}, + }}, + }, + }, + }; + + MappedOperatorTaskGroup replicate_weight_2_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(1_n)}, + }}, + }, + }, + }; + + MappedOperatorTaskGroup linear_1_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(0_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(0_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(1_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(1_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, + }}, + }, + }, + }; + + MappedOperatorTaskGroup linear_2_op_mapping = MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(0_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(0_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(1_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(1_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, + }}, + }, + }, + }; + + MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( + /*pcg=*/pcg, + /*mapped_op_task_groups=*/{ + { + input_layer.parallel_layer, + input_op_mapping, + }, + { + weight_1_layer.parallel_layer, + weight_1_op_mapping, + }, + { + weight_2_layer.parallel_layer, + weight_2_op_mapping, + }, + { + partition_input_layer.parallel_layer, + partition_input_op_mapping, + }, + { + replicate_weight_1_layer.parallel_layer, + replicate_weight_1_op_mapping, + }, + { + replicate_weight_2_layer.parallel_layer, + replicate_weight_2_op_mapping, + }, + { + linear_1_layer.parallel_layer, + linear_1_op_mapping, + }, + { + linear_2_layer.parallel_layer, + linear_2_op_mapping, + }, + }); + + + DynamicOpenDataflowGraph result = make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg, DeviceType::GPU); + + DynamicOpenDataflowGraph correct = [&]() + -> DynamicOpenDataflowGraph + { + auto mk_pt_shape = [](positive_int sum_degree, + positive_int discard_copy_degree, + positive_int shard_dim_size_0, + positive_int shard_degree_0, + positive_int shard_dim_size_1, + positive_int shard_degree_1) + -> ParallelTensorShape + { + return ParallelTensorShape{ + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{shard_dim_size_0, shard_degree_0}, + ShardParallelDim{shard_dim_size_1, shard_degree_1}, + }, + /*replica_dims=*/ReplicaParallelDimSet{ + /*sum_degree=*/SumDegree{sum_degree}, + /*discard_copy_degree=*/DiscardCopyDegree{discard_copy_degree}, + }, + }, + /*data_type=*/DataType::FLOAT, + }; + }; + + ParallelTensorShape input_pt_shape + = mk_pt_shape(1_p, 1_p, batch_size, 1_p, data_dim, 1_p); + ParallelTensorShape weight_1_pt_shape + = mk_pt_shape(1_p, 1_p, hidden_dim, 1_p, data_dim, 1_p); + ParallelTensorShape weight_2_pt_shape + = mk_pt_shape(1_p, 1_p, output_dim, 1_p, hidden_dim, 1_p); + ParallelTensorShape partitioned_input_pt_shape + = mk_pt_shape(1_p, 1_p, batch_size, 2_p, data_dim, 1_p); + ParallelTensorShape replicated_weight_1_pt_shape + = mk_pt_shape(1_p, 2_p, hidden_dim, 1_p, data_dim, 1_p); + ParallelTensorShape replicated_weight_2_pt_shape + = mk_pt_shape(1_p, 2_p, output_dim, 1_p, hidden_dim, 1_p); + ParallelTensorShape hidden_activation_pt_shape + = mk_pt_shape(1_p, 1_p, batch_size, 2_p, hidden_dim, 1_p); + ParallelTensorShape output_pt_shape + = mk_pt_shape(1_p, 1_p, batch_size, 2_p, output_dim, 1_p); + + auto mk_node_attrs = [&](parallel_layer_guid_t const &layer_guid, + PCGOperatorAttrs const &op_attrs, + MappedOperatorTaskGroup const &mapped_op_task_group) + -> DynamicNodeAttrs + { + return DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/DynamicNodeMapping{ + mapped_op_task_group, + DeviceType::GPU, + }, + /*op_attrs=*/TrainingOperationAttrs{op_attrs}, + /*pcg_layer_guid=*/dynamic_layer_guid_t{layer_guid}, + /*per_device_op_state=*/std::nullopt, + }; + }; + + auto mk_slot = [&](TensorSlotName slot_name) + -> DynamicTensorSlot + { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }; + }; + + auto mk_value = [&](parallel_tensor_guid_t const &tensor_guid, + ParallelTensorShape const &shape, + bool create_grad = true) + -> DynamicValueAttrs + { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{tensor_guid}, + /*parallel_tensor_shape=*/shape, + /*create_grad=*/create_grad, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + }; + + DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{}, + /*node=*/mk_node_attrs(input_layer.parallel_layer, + input_op_attrs, + input_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_input, input_pt_shape, /*create_grad=*/false), + }, + }, + }; + + DynamicNodeInvocation weight_1_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{}, + /*node=*/mk_node_attrs(weight_1_layer.parallel_layer, + weight_1_op_attrs, + weight_1_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_weight_1, weight_1_pt_shape), + }, + }, + }; + + DynamicNodeInvocation weight_2_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{}, + /*node=*/mk_node_attrs(weight_2_layer.parallel_layer, + weight_2_op_attrs, + weight_2_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_weight_2, weight_2_pt_shape), + }, + }, + }; + + DynamicNodeInvocation partition_input_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_input, input_pt_shape, /*create_grad=*/false), + }, + }, + /*node=*/mk_node_attrs(partition_input_layer.parallel_layer, + partition_input_op_attrs, + partition_input_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_partitioned_input, partitioned_input_pt_shape), + }, + }, + }; + + DynamicNodeInvocation replicate_weight_1_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_weight_1, weight_1_pt_shape), + }, + }, + /*node=*/mk_node_attrs(replicate_weight_1_layer.parallel_layer, + replicate_weight_1_op_attrs, + replicate_weight_1_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_replicated_weight_1, replicated_weight_1_pt_shape), + }, + }, + }; + + DynamicNodeInvocation replicate_weight_2_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_weight_2, weight_2_pt_shape), + }, + }, + /*node=*/mk_node_attrs(replicate_weight_2_layer.parallel_layer, + replicate_weight_2_op_attrs, + replicate_weight_2_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_replicated_weight_2, replicated_weight_2_pt_shape), + }, + }, + }; + + DynamicNodeInvocation linear_1_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_partitioned_input, partitioned_input_pt_shape), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(t_replicated_weight_1, replicated_weight_1_pt_shape), + }, + }, + /*node=*/mk_node_attrs(linear_1_layer.parallel_layer, + linear_1_op_attrs, + linear_1_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_hidden_activation, hidden_activation_pt_shape), + }, + }, + }; + + DynamicNodeInvocation linear_2_invocation = DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_hidden_activation, hidden_activation_pt_shape), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(t_replicated_weight_2, replicated_weight_2_pt_shape), + }, + }, + /*node=*/mk_node_attrs(linear_2_layer.parallel_layer, + linear_2_op_attrs, + linear_2_op_mapping), + /*outputs=*/std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_output, output_pt_shape), + }, + }, + }; + + return dynamic_open_dataflow_graph_from_invocation_set( + /*invocations=*/{ + input_invocation, + weight_1_invocation, + weight_2_invocation, + partition_input_invocation, + replicate_weight_1_invocation, + replicate_weight_2_invocation, + linear_1_invocation, + linear_2_invocation, + }); + }(); + + nlohmann::json result_json + = dynamic_open_dataflow_graph_to_serializable(result); + nlohmann::json correct_json + = dynamic_open_dataflow_graph_to_serializable(correct); + + CHECK_MESSAGE( + result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); + } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc index 1a90aa3f40..d9e8ffa781 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc @@ -3,10 +3,168 @@ #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" +#include "op-attrs/initializer_attrs.h" +#include "test/utils/doctest/check_kv.h" +#include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h" using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("determine_intermediate_values_needed_for_gradient_computation") { + auto mk_slot = [](TensorSlotName slot_name) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }; + }; + + auto mk_node_attrs = [](size_t layer_guid, + PCGOperatorAttrs const &op_attrs) { + return DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{ + op_attrs, + }, + /*layer_guid=*/dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{layer_guid}, + }, + }, + /*per_device_op_state=*/std::nullopt, + }; + }; + + auto mk_value_attrs = [](size_t src_layer_guid, + TensorSlotName src_slot, + bool create_grad) + -> DynamicValueAttrs + { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{ + src_layer_guid, + }, + src_slot, + }, + }, + }, + /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/create_grad, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + }; + + struct TestGraph { + DynamicOpenDataflowGraph g; + DynamicValueAttrs input_op_output; + DynamicValueAttrs relu1_op_output; + }; + + auto mk_test_graph = [&](bool input_create_grad) -> TestGraph { + TensorShape input_shape = TensorShape{ + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, + }, + DataType::FLOAT, + }; + + DynamicValueAttrs input_op_output = + mk_value_attrs(123, TensorSlotName::OUTPUT, /*create_grad=*/input_create_grad); + + DynamicValueAttrs relu1_op_output = + mk_value_attrs(124, TensorSlotName::OUTPUT, /*create_grad=*/true); + + + PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ + InputAttrs{ + input_shape, + }, + }; + + PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ + make_relu_attrs(), + }; + + DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ + /*inputs=*/{}, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, + }, + }; + + DynamicNodeInvocation relu_invocation = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + input_op_output, + }, + }, + /*node_attrs=*/mk_node_attrs( + /*layer_guid=*/124, + /*op_attrs=*/relu_attrs), + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + relu1_op_output, + }, + }, + }; + + DynamicOpenDataflowGraph g + = dynamic_open_dataflow_graph_from_invocation_set( + {input_invocation, relu_invocation}); + + return TestGraph{ + /*g=*/g, + /*input_op_output=*/input_op_output, + /*relu1_op_output=*/relu1_op_output, + }; + }; + + SUBCASE("input create_grad is false") { + TestGraph tg = mk_test_graph(/*input_create_grad=*/false); + + std::set result = + determine_intermediate_values_needed_for_gradient_computation(tg.g); + + std::set correct = {}; + + ASSERT(result == correct); + } + + SUBCASE("input create_grad is true") { + TestGraph tg = mk_test_graph(/*input_create_grad=*/true); + + std::set result = + determine_intermediate_values_needed_for_gradient_computation(tg.g); + + std::set correct = { + tg.input_op_output, + tg.relu1_op_output, + }; + + ASSERT(result == correct); + } + } + TEST_CASE("perform_fwd_pass_expansion_for_invocation") { auto mk_value_attrs = [](size_t node_id, std::optional const &tensor_role) @@ -19,6 +177,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -38,64 +197,113 @@ TEST_SUITE(FF_TEST_SUITE) { dynamic_layer_guid_t layer_guid{parallel_layer_guid_t{Node{20}}}; - TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ - PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/8_p, - /*use_bias=*/true, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, - }, - }, - }; + DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); + DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); + DynamicValueAttrs v3 = mk_value_attrs(2, std::nullopt); - DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { - DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); - DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); - DynamicValueAttrs v3 = mk_value_attrs(2, std::nullopt); + DynamicTensorRole fwd_role = DynamicTensorRole{FwbTensorType::FORWARD}; - return DynamicNodeInvocation{ - /*inputs=*/{ - {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, - {mk_slot(TensorSlotName::WEIGHT, std::nullopt), v2}, - {mk_slot(TensorSlotName::BIAS, std::nullopt), v1}, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/op_attrs, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v3}, + DynamicValueAttrs v1_fwd = mk_value_attrs(0, fwd_role); + DynamicValueAttrs v2_fwd = mk_value_attrs(1, fwd_role); + DynamicValueAttrs v3_fwd = mk_value_attrs(2, fwd_role); + + SUBCASE("standard operator") { + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/8_p, + /*use_bias=*/true, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }, }; - }(); - DynamicNodeInvocation result = - perform_fwd_pass_expansion_for_invocation(invocation); + DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, + {mk_slot(TensorSlotName::WEIGHT, std::nullopt), v2}, + {mk_slot(TensorSlotName::BIAS, std::nullopt), v1}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v3}, + }, + }; + }(); + + DynamicNodeInvocation result = + perform_fwd_pass_expansion_for_invocation(invocation); + + DynamicNodeInvocation correct = DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, fwd_role), v1_fwd}, + {mk_slot(TensorSlotName::WEIGHT, fwd_role), v2_fwd}, + {mk_slot(TensorSlotName::BIAS, fwd_role), v1_fwd}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::OUTPUT, fwd_role), v3_fwd}, + }, + }; + + ASSERT(result == correct); + } - DynamicNodeInvocation correct = [&]() -> DynamicNodeInvocation { - DynamicTensorRole fwd_role = DynamicTensorRole{FwbTensorType::FORWARD}; + SUBCASE("copy operator") { + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{CopyAttrs{}}; + + DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v2}, + }, + }; + }(); - DynamicValueAttrs v1_fwd = mk_value_attrs(0, fwd_role); - DynamicValueAttrs v2_fwd = mk_value_attrs(1, fwd_role); - DynamicValueAttrs v3_fwd = mk_value_attrs(2, fwd_role); + DynamicNodeInvocation result = + perform_fwd_pass_expansion_for_invocation(invocation); - return DynamicNodeInvocation{ + DynamicNodeInvocation correct = DynamicNodeInvocation{ /*inputs=*/{ - {mk_slot(TensorSlotName::INPUT, fwd_role), v1_fwd}, - {mk_slot(TensorSlotName::WEIGHT, fwd_role), v2_fwd}, - {mk_slot(TensorSlotName::BIAS, fwd_role), v1_fwd}, + {mk_slot(TensorSlotName::INPUT, std::nullopt), v1_fwd}, }, /*node_attrs=*/ DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, + /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*op_attrs=*/op_attrs, @@ -104,12 +312,12 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*outputs=*/ { - {mk_slot(TensorSlotName::OUTPUT, fwd_role), v3_fwd}, + {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v2_fwd}, }, }; - }(); - ASSERT(result == correct); + ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); + } } TEST_CASE("perform_bwd_pass_expansion_for_invocation") { @@ -124,6 +332,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -223,10 +432,10 @@ TEST_SUITE(FF_TEST_SUITE) { }; }(); - ASSERT(result == correct); + ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); } - SUBCASE("replicate operator optimization") { + SUBCASE("replicate operator") { TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ PCGOperatorAttrs{ ReplicateAttrs{ @@ -266,7 +475,6 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicNodeInvocation{ /*inputs=*/{ - {mk_slot(TensorSlotName::OUTPUT, fwd_role), v2_fwd}, {mk_slot(TensorSlotName::OUTPUT, grad_role), v2_grad}, }, /*node_attrs=*/ @@ -285,7 +493,62 @@ TEST_SUITE(FF_TEST_SUITE) { }; }(); - ASSERT(result == correct); + ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); + } + + SUBCASE("copy operator") { + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{CopyAttrs{}}; + + DynamicNodeInvocation invocation = [&]() -> DynamicNodeInvocation { + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, std::nullopt), v1}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v2}, + }, + }; + }(); + + DynamicNodeInvocation result = + perform_bwd_pass_expansion_for_invocation(invocation); + + DynamicNodeInvocation correct = [&]() -> DynamicNodeInvocation { + DynamicTensorRole fwd_role = DynamicTensorRole{FwbTensorType::FORWARD}; + DynamicTensorRole grad_role = + DynamicTensorRole{FwbTensorType::GRADIENT}; + + return DynamicNodeInvocation{ + /*inputs=*/{ + {mk_slot(TensorSlotName::OUTPUT, std::nullopt), v2_grad}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*pass_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::INPUT, std::nullopt), v1_grad}, + }, + }; + }(); + + ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); } } @@ -306,7 +569,8 @@ TEST_SUITE(FF_TEST_SUITE) { }; auto mk_value_attrs = - [](size_t node_id, std::optional const &tensor_type) + [](size_t node_id, + std::optional const &tensor_type) -> DynamicValueAttrs { return DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ @@ -316,6 +580,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/false, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -345,50 +610,112 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - DynamicOpenDataflowGraph input = [&]() -> DynamicOpenDataflowGraph { - DynamicNodeAttrs n1 = mk_node_attrs(10, input_op_attrs, std::nullopt); - DynamicNodeAttrs n2 = mk_node_attrs(11, relu_op_attrs, std::nullopt); + TensorShape weight_shape = TensorShape{ + TensorDims{ + FFOrdered{ + 8_p, + 6_p, + }, + }, + DataType::FLOAT, + }; - DynamicValueAttrs v1 = mk_value_attrs(0, std::nullopt); - DynamicValueAttrs v2 = mk_value_attrs(1, std::nullopt); + TrainingOperationAttrs weight_op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + WeightAttrs{ + /*tensor_shape=*/weight_shape, + /*initializer=*/make_zero_initializer(), + }, + }, + }; + + TrainingOperationAttrs linear_op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/6_p, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, + }, + }; + + DynamicOpenDataflowGraph input = [&]() -> DynamicOpenDataflowGraph { + DynamicNodeAttrs input_node = mk_node_attrs(10, input_op_attrs, std::nullopt); + DynamicNodeAttrs weight_node = mk_node_attrs(11, weight_op_attrs, std::nullopt); + DynamicNodeAttrs relu_node = mk_node_attrs(12, relu_op_attrs, std::nullopt); + DynamicNodeAttrs linear_node = mk_node_attrs(13, linear_op_attrs, std::nullopt); + + DynamicValueAttrs input_tensor = mk_value_attrs(0, std::nullopt); + DynamicValueAttrs weight_tensor = mk_value_attrs(1, std::nullopt); + DynamicValueAttrs relu_output = mk_value_attrs(2, std::nullopt); + DynamicValueAttrs linear_output = mk_value_attrs(3, std::nullopt); + + auto mk_dynamic_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }; + }; std::set invocation_set = { DynamicNodeInvocation{ /*inputs=*/std::map{}, - /*node_attrs=*/n1, + /*node_attrs=*/input_node, /*outputs=*/ std::map{ { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - v1, + mk_dynamic_slot(TensorSlotName::OUTPUT), + input_tensor, + }, + }, + }, + DynamicNodeInvocation{ + /*inputs=*/std::map{}, + /*node_attrs=*/weight_node, + /*outputs=*/ + std::map{ + { + mk_dynamic_slot(TensorSlotName::OUTPUT), + weight_tensor, }, }, }, DynamicNodeInvocation{ /*inputs=*/std::map{ { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - v1, + mk_dynamic_slot(TensorSlotName::INPUT), + input_tensor, }, }, - /*node_attrs=*/n2, + /*node_attrs=*/relu_node, /*outputs=*/ std::map{ { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - v2, + mk_dynamic_slot(TensorSlotName::OUTPUT), + relu_output, + }, + }, + }, + DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_dynamic_slot(TensorSlotName::INPUT), + relu_output, + }, + { + mk_dynamic_slot(TensorSlotName::WEIGHT), + weight_tensor, + }, + }, + /*node_attrs=*/linear_node, + /*outputs=*/ + std::map{ + { + mk_dynamic_slot(TensorSlotName::OUTPUT), + linear_output, }, }, }, @@ -400,101 +727,139 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicOpenDataflowGraph result = perform_pass_expansion(input); DynamicOpenDataflowGraph correct = [&]() -> DynamicOpenDataflowGraph { - DynamicNodeAttrs n1_fwd = + DynamicNodeAttrs input_node_fwd = mk_node_attrs(10, input_op_attrs, DynamicTaskType::FWD); - DynamicNodeAttrs n2_fwd = - mk_node_attrs(11, relu_op_attrs, DynamicTaskType::FWD); - DynamicNodeAttrs n1_bwd = - mk_node_attrs(10, input_op_attrs, DynamicTaskType::BWD); - DynamicNodeAttrs n2_bwd = - mk_node_attrs(11, relu_op_attrs, DynamicTaskType::BWD); - - DynamicValueAttrs v1_activation = + DynamicNodeAttrs weight_node_fwd = + mk_node_attrs(11, weight_op_attrs, DynamicTaskType::FWD); + DynamicNodeAttrs relu_node_fwd = + mk_node_attrs(12, relu_op_attrs, DynamicTaskType::FWD); + DynamicNodeAttrs linear_node_fwd = + mk_node_attrs(13, linear_op_attrs, DynamicTaskType::FWD); + + DynamicNodeAttrs linear_node_bwd = + mk_node_attrs(13, linear_op_attrs, DynamicTaskType::BWD); + + DynamicValueAttrs input_tensor_activation = mk_value_attrs(0, mk_dynamic_tensor_role_fwd()); - DynamicValueAttrs v1_gradient = + DynamicValueAttrs input_tensor_gradient = mk_value_attrs(0, mk_dynamic_tensor_role_bwd()); - DynamicValueAttrs v2_activation = + DynamicValueAttrs weight_tensor_activation = mk_value_attrs(1, mk_dynamic_tensor_role_fwd()); - DynamicValueAttrs v2_gradient = + DynamicValueAttrs weight_tensor_gradient = mk_value_attrs(1, mk_dynamic_tensor_role_bwd()); + DynamicValueAttrs relu_output_tensor_activation = + mk_value_attrs(2, mk_dynamic_tensor_role_fwd()); + DynamicValueAttrs relu_output_tensor_gradient = + mk_value_attrs(2, mk_dynamic_tensor_role_bwd()); + DynamicValueAttrs linear_output_tensor_activation = + mk_value_attrs(3, mk_dynamic_tensor_role_fwd()); + DynamicValueAttrs linear_output_tensor_gradient= + mk_value_attrs(3, mk_dynamic_tensor_role_bwd()); + + auto mk_fwd_slot = [&](TensorSlotName slot_name) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }; + }; + + auto mk_grad_slot = [&](TensorSlotName slot_name) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }; + }; std::set invocation_set = { DynamicNodeInvocation{ /*inputs=*/std::map{}, - /*node_attrs=*/n1_fwd, + /*node_attrs=*/input_node_fwd, /*outputs=*/ std::map{ std::pair{ - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - v1_activation, + mk_fwd_slot(TensorSlotName::OUTPUT), + input_tensor_activation, + }, + }, + }, + DynamicNodeInvocation{ + /*inputs=*/std::map{}, + /*node_attrs=*/weight_node_fwd, + /*outputs=*/ + std::map{ + std::pair{ + mk_fwd_slot(TensorSlotName::OUTPUT), + weight_tensor_activation, }, }, }, DynamicNodeInvocation{ /*inputs=*/std::map{ std::pair{ - DynamicTensorSlot{ - TensorSlotName::INPUT, - mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - v1_activation, + mk_fwd_slot(TensorSlotName::INPUT), + input_tensor_activation, }, }, - /*node_attrs=*/n2_fwd, + /*node_attrs=*/relu_node_fwd, /*outputs=*/ std::map{ std::pair{ - DynamicTensorSlot{ - TensorSlotName::OUTPUT, - mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - v2_activation, + mk_fwd_slot(TensorSlotName::OUTPUT), + relu_output_tensor_activation, }, }, }, DynamicNodeInvocation{ /*inputs=*/std::map{ std::pair{ - DynamicTensorSlot{ - TensorSlotName::INPUT, - mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - v1_activation, + mk_fwd_slot(TensorSlotName::INPUT), + relu_output_tensor_activation, }, std::pair{ - DynamicTensorSlot{ - TensorSlotName::OUTPUT, - mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - v2_activation, + mk_fwd_slot(TensorSlotName::WEIGHT), + weight_tensor_activation, }, + }, + /*node_attrs=*/linear_node_fwd, + /*outputs=*/ + std::map{ std::pair{ - DynamicTensorSlot{ - TensorSlotName::OUTPUT, - mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, - }, - v2_gradient, + mk_fwd_slot(TensorSlotName::OUTPUT), + linear_output_tensor_activation, }, }, - /*node_attrs=*/n2_bwd, + }, + DynamicNodeInvocation{ + /*inputs=*/std::map{ + std::pair{ + mk_fwd_slot(TensorSlotName::INPUT), + relu_output_tensor_activation, + }, + std::pair{ + mk_fwd_slot(TensorSlotName::WEIGHT), + weight_tensor_activation, + }, + std::pair{ + mk_fwd_slot(TensorSlotName::OUTPUT), + linear_output_tensor_activation, + }, + std::pair{ + mk_grad_slot(TensorSlotName::OUTPUT), + linear_output_tensor_gradient, + }, + }, + /*node_attrs=*/linear_node_bwd, /*outputs=*/ std::map{ std::pair{ - DynamicTensorSlot{ - TensorSlotName::INPUT, - mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, - }, - v1_gradient, + mk_grad_slot(TensorSlotName::INPUT), + relu_output_tensor_gradient, + }, + std::pair{ + mk_grad_slot(TensorSlotName::WEIGHT), + weight_tensor_gradient, }, }, }, @@ -503,9 +868,17 @@ TEST_SUITE(FF_TEST_SUITE) { return dynamic_open_dataflow_graph_from_invocation_set(invocation_set); }(); - ASSERT(get_dynamic_invocation_set(result).size() == 3); - ASSERT(get_dynamic_invocation_set(result) == - get_dynamic_invocation_set(correct)); - ASSERT(dynamic_open_dataflow_graphs_are_isomorphic(result, correct)); + CHECK(get_dynamic_invocation_set(result).size() == correct.invocations.size()); + + nlohmann::json result_json + = dynamic_open_dataflow_graph_to_serializable(result); + nlohmann::json correct_json + = dynamic_open_dataflow_graph_to_serializable(correct); + + CHECK_MESSAGE(get_dynamic_invocation_set(result) == get_dynamic_invocation_set(correct), + check_kv("result", result_json.dump()), + check_kv("correct", correct_json.dump())); + + CHECK(dynamic_open_dataflow_graphs_are_isomorphic(result, correct)); } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index 420f9bc13d..89339751ef 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -12,6 +12,9 @@ #include "utils/bidict/algorithms/bidict_filter_values.h" #include "utils/containers/map_from_pairs.h" #include "utils/containers/binary_merge_disjoint_maps.h" +#include "utils/binary_relation/binary_relation_from_map.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" +#include "test/utils/doctest/check_kv.h" using namespace ::FlexFlow; @@ -23,6 +26,13 @@ static MachineSpaceCoordinate mk_machine_coord(nonnegative_int node_idx, }; }; +static global_device_id_t mk_device_id(MachineSpaceCoordinate const &mc) { + return global_device_id_t{ + /*coord=*/mc, + /*device_type=*/DeviceType::GPU, + }; +}; + static ParallelTensorSpaceCoordinate mk_pt_coord(nonnegative_int idx1, nonnegative_int idx2, nonnegative_int idx3, @@ -50,11 +60,11 @@ DynamicTensorSlot mk_slot(TensorSlotName const &slot_name, DynamicValueAttrs mk_value(size_t src_node_id, TensorSlotName src_slot_name, - bidict const &tensor_binding, + bidict const &tensor_binding, std::optional const &shard_coord, std::optional const &role = std::nullopt) { - bidict mapping = tensor_binding; + bidict mapping = tensor_binding; if (shard_coord.has_value()) { mapping = bidict_filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { @@ -70,37 +80,427 @@ DynamicValueAttrs }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/shard_coord, - /*mapping=*/mapping, + /*mapping=*/ParallelTensorMapping{mapping}, /*accessor=*/std::nullopt, /*role=*/role, }; }; TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("apply_dynamic_node_invocation_sharding_info") { + global_device_id_t device_0 = global_device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, + }; + + global_device_id_t device_1 = global_device_id_t{ + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/2_n, + /*device_idx=*/1_n, + }, + /*device_type=*/DeviceType::GPU, + }; + + auto mk_slot = [](TensorSlotName slot_name, + std::optional const &task_shard = std::nullopt) + -> DynamicTensorSlot + { + return DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/task_shard, + }; + }; + + auto mk_value = [](size_t src_node_id, + TensorSlotName src_slot_name, + std::optional const &shard_coord = std::nullopt) + -> DynamicValueAttrs + { + return DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + /*node=*/Node{src_node_id}, + /*slot_name=*/src_slot_name, + }, + }, + }, + /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, + /*shard_coord=*/shard_coord, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }; + }; + + SUBCASE("sharding info creates additional arguments ie replicate") { + auto mk_pt_coord = [](nonnegative_int idx) + -> ParallelTensorSpaceCoordinate + { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shard_components=*/FFOrdered{ + 0_n, + 0_n, + }, + }; + }; + + size_t input_src_node_id = 234; + size_t replicate_layer_node_id = 13; + + DynamicNodeMapping node_mapping = DynamicNodeMapping{ + /*op_task_group=*/MappedOperatorTaskGroup{{ + { + device_0.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }}, + }, + { + device_1.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }}, + }, + }}, + /*device_type=*/DeviceType::GPU, + }; + + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, + }, + }; + + dynamic_layer_guid_t layer_guid = dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{replicate_layer_node_id}, + }, + }; + + DynamicNodeInvocation invocation = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, TensorSlotName::OUTPUT), + }, + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(replicate_layer_node_id, TensorSlotName::OUTPUT), + }, + }, + }; + + DynamicNodeInvocationShardingInfo invocation_sharding_info = + DynamicNodeInvocationShardingInfo{ + /*device_ids=*/nonempty_set{ + device_0, + device_1, + }, + /*value_sharding=*/{ + { + mk_slot(TensorSlotName::INPUT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(0_n), + /*mapping=*/device_0, + }, + }, + { + mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_0.coord), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(0_n), + /*mapping=*/device_0, + }, + }, + { + mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_1.coord), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(1_n), + /*mapping=*/device_1, + }, + }, + }, + }; + + DynamicNodeInvocation result = + apply_dynamic_node_invocation_sharding_info(invocation, invocation_sharding_info); + + DynamicNodeInvocation correct = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, TensorSlotName::OUTPUT, /*shrad_coord=*/mk_pt_coord(0_n)), + }, + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/nonempty_set{ + device_0, + device_1, + }, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_0.coord), + mk_value(replicate_layer_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(0_n)), + }, + { + mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_1.coord), + mk_value(replicate_layer_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(1_n)), + }, + }, + }; + + nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + + CHECK_MESSAGE( + result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump()) + ); + } + + SUBCASE("sharding info does not create additional arguments ie standard operator") { + size_t input_src_node_id = 234; + size_t weight_src_node_id = 345; + size_t linear_layer_node_id = 13; + + auto mk_pt_coord = [](nonnegative_int idx) + -> ParallelTensorSpaceCoordinate + { + return ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shard_components=*/FFOrdered{ + 0_n, + 0_n, + }, + }; + }; + + // note that the node mapping does not have to be accurate/real here. + // apply_dynamic_node_invocation_sharding_info should just function + // based on what it is given. + DynamicNodeMapping node_mapping = DynamicNodeMapping{ + /*op_task_group=*/MappedOperatorTaskGroup{{ + { + device_0.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::WEIGHT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }}, + }, + { + device_1.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(1_n), + }, + { + TensorSlotName::WEIGHT, + mk_pt_coord(2_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }}, + }, + }}, + /*device_type=*/DeviceType::GPU, + }; + + TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ + PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/8_p, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, + }, + }; + + dynamic_layer_guid_t layer_guid = dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{linear_layer_node_id}, + }, + }; + + DynamicNodeInvocation invocation = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, TensorSlotName::OUTPUT), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(weight_src_node_id, TensorSlotName::OUTPUT), + }, + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(linear_layer_node_id, TensorSlotName::OUTPUT), + }, + }, + }; + + DynamicNodeInvocationShardingInfo invocation_sharding_info = + DynamicNodeInvocationShardingInfo{ + /*device_ids=*/nonempty_set{ + device_1, + }, + /*value_sharding=*/{ + { + mk_slot(TensorSlotName::INPUT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(1_n), + /*mapping=*/device_1, + }, + }, + { + mk_slot(TensorSlotName::WEIGHT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(2_n), + /*mapping=*/device_1, + }, + }, + { + mk_slot(TensorSlotName::OUTPUT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(1_n), + /*mapping=*/device_1, + }, + }, + }, + }; + + DynamicNodeInvocation result = + apply_dynamic_node_invocation_sharding_info(invocation, invocation_sharding_info); + + DynamicNodeInvocation correct = DynamicNodeInvocation{ + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(1_n)), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(weight_src_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(2_n)), + }, + }, + /*node_attrs=*/DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/nonempty_set{device_1}, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(linear_layer_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(1_n)), + }, + }, + }; + + nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + + CHECK_MESSAGE( + result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump()) + ); + } + } + TEST_CASE("generate_shard_expansion_for_invocation") { auto mk_op_value = [&](size_t src_node_id, TensorSlotName src_slot_name, TensorSlotName use_slot_name, - MappedOperatorTaskGroup const &mapped_task_group, + DynamicNodeMapping const &node_mapping, std::optional const &shard_coord, std::optional const &role = std::nullopt) -> DynamicValueAttrs { - bidict - tensor_binding = get_tensor_bindings_for_slot_name(mapped_task_group, - use_slot_name); + + bidict + tensor_binding = dynamic_node_mapping_bindings_for_slot_name(node_mapping, + use_slot_name); return mk_value(src_node_id, src_slot_name, tensor_binding, shard_coord, role); }; auto mk_sharding_info = [&](TensorSlotName slot_name, ParallelTensorSpaceCoordinate const &shard_coord, - MappedOperatorTaskGroup const &mapped_op_task_group) + DynamicNodeMapping const &node_mapping) -> std::pair { - bidict - tensor_binding = get_tensor_bindings_for_slot_name(mapped_op_task_group, - slot_name); + bidict + tensor_binding = dynamic_node_mapping_bindings_for_slot_name(node_mapping, slot_name); + return std::pair{ mk_slot(slot_name), DynamicValueAttrsShardingInfo{ @@ -110,10 +510,12 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - SUBCASE("standard operator") { - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(2_n, 0_n); + global_device_id_t dev1 = mk_device_id(mk_machine_coord(0_n, 0_n)); + global_device_id_t dev2 = mk_device_id(mk_machine_coord(1_n, 0_n)); + global_device_id_t dev3 = mk_device_id(mk_machine_coord(2_n, 0_n)); + global_device_id_t dev4 = mk_device_id(mk_machine_coord(3_n, 0_n)); + SUBCASE("standard operator") { auto mk_shard_binding = [&](ParallelTensorSpaceCoordinate const &c1, ParallelTensorSpaceCoordinate const &c2, ParallelTensorSpaceCoordinate const &c3, @@ -150,7 +552,7 @@ TEST_SUITE(FF_TEST_SUITE) { -> DynamicValueAttrs { if (shard_coord.has_value()) { tensor_binding = - filter_keys(tensor_binding, + bidict_filter_keys(tensor_binding, [&](ParallelTensorSpaceCoordinate const &p) -> bool { return p == shard_coord.value(); }); @@ -164,6 +566,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }}, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/shard_coord, /*mapping=*/ ParallelTensorMapping{tensor_binding}, @@ -172,9 +575,6 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(2_n, 0_n); - ParallelTensorSpaceCoordinate mc1_input_coord = mk_pt_coord(0_n, 0_n, 0_n, 0_n); ParallelTensorSpaceCoordinate mc1_weight_coord = @@ -203,14 +603,14 @@ TEST_SUITE(FF_TEST_SUITE) { MappedOperatorTaskGroup{ bidict{ { - mc1, + dev1.coord, mk_shard_binding(mc1_input_coord, mc1_weight_coord, mc1_output_1_coord, mc1_output_2_coord), }, { - mc2, + dev2.coord, mk_shard_binding(mc2_input_coord, mc2_weight_coord, mc2_output_1_coord, @@ -218,7 +618,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, }, - device_type, + DeviceType::GPU, }; DynamicNodeInvocation input = DynamicNodeInvocation{ @@ -228,7 +628,7 @@ TEST_SUITE(FF_TEST_SUITE) { mk_op_value(0, TensorSlotName::OUTPUT, TensorSlotName::INPUT, - mapped_task_group, + node_mapping, std::nullopt), }, { @@ -236,7 +636,7 @@ TEST_SUITE(FF_TEST_SUITE) { mk_op_value(1, TensorSlotName::OUTPUT, TensorSlotName::WEIGHT, - mapped_task_group, + node_mapping, std::nullopt), }, }, @@ -257,7 +657,7 @@ TEST_SUITE(FF_TEST_SUITE) { mk_op_value(20, TensorSlotName::OUTPUT_1, TensorSlotName::OUTPUT_1, - mapped_task_group, + node_mapping, std::nullopt), }, { @@ -265,7 +665,7 @@ TEST_SUITE(FF_TEST_SUITE) { mk_op_value(20, TensorSlotName::OUTPUT_2, TensorSlotName::OUTPUT_2, - mapped_task_group, + node_mapping, std::nullopt), }, }, @@ -284,41 +684,32 @@ TEST_SUITE(FF_TEST_SUITE) { return DynamicNodeInvocationShardingInfo{ /*device_coord=*/nonempty_set{device_coord}, /*value_sharding=*/{ - mk_sharding_info(TensorSlotName::INPUT, input_shard_coord, mapped_task_group), - mk_sharding_info(TensorSlotName::WEIGHT, weight_shard_coord, mapped_task_group), - mk_sharding_info(TensorSlotName::OUTPUT_1, output_1_shard_coord, mapped_task_group), - mk_sharding_info(TensorSlotName::OUTPUT_2, output_2_shard_coord, mapped_task_group), + mk_sharding_info(TensorSlotName::INPUT, input_shard_coord, node_mapping), + mk_sharding_info(TensorSlotName::WEIGHT, weight_shard_coord, node_mapping), + mk_sharding_info(TensorSlotName::OUTPUT_1, output_1_shard_coord, node_mapping), + mk_sharding_info(TensorSlotName::OUTPUT_2, output_2_shard_coord, node_mapping), }, }; }; - std::unordered_set correct = { - mk_invocation_shard(mk_device_id(mc1), + std::set correct = { + mk_invocation_shard(dev1, mc1_input_coord, mc1_weight_coord, mc1_output_1_coord, mc1_output_2_coord), - mk_invocation_shard(mk_device_id(mc2), + mk_invocation_shard(dev2, mc2_input_coord, mc2_weight_coord, mc2_output_1_coord, mc2_output_2_coord), }; - nlohmann::json result_json = result; - nlohmann::json correct_json = correct; - CHECK(result.size() == correct.size()); - CHECK(result_json == correct_json); CHECK(result == correct); } SUBCASE("for copy operator") { - global_device_id_t dev1 = mk_device_id(mk_machine_coord(0_n, 0_n)); - global_device_id_t dev2 = mk_device_id(mk_machine_coord(1_n, 0_n)); - global_device_id_t dev3 = mk_device_id(mk_machine_coord(2_n, 0_n)); - global_device_id_t dev4 = mk_device_id(mk_machine_coord(3_n, 0_n)); - ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); @@ -356,54 +747,17 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::unordered_set result = - perform_shard_expansion_for_invocation(input); - - auto mk_invocation_shard = - [&](global_device_id_t const &device_id, - ParallelTensorSpaceCoordinate const &tensor_shard_coord) - -> DynamicNodeInvocation { - DynamicNodeInvocation result = input; - result.inputs = { - { - mk_slot(TensorSlotName::INPUT), - mk_value( - 0, TensorSlotName::OUTPUT, src_binding, tensor_shard_coord), - }, - }; - // See perform_shard_expansion_for_copy in shard_expansion.cc for explanation of the choice of device placement. - result.node_attrs.device_id = device_id; - result.outputs = { - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(20, - TensorSlotName::OUTPUT, - dst_binding, - tensor_shard_coord), - }, - }; - return result; - }; - - std::unordered_set correct = { - mk_invocation_shard(dev1, pt1), - mk_invocation_shard(dev2, pt2), - }; - - CHECK(result.size() == correct.size()); - CHECK(result == correct); - std::set result = generate_shard_expansion_for_invocation(input); auto mk_invocation_shard = - [&](MachineSpaceCoordinate const &device_coord, + [&](global_device_id_t const &device_coord, ParallelTensorSpaceCoordinate const &tensor_shard_coord) -> DynamicNodeInvocationShardingInfo { return DynamicNodeInvocationShardingInfo{ /*device_coord=*/nonempty_set{device_coord}, - /*value_sharding=*/std::map{ + /*value_sharding=*/BinaryRelation{ { mk_slot(TensorSlotName::INPUT), DynamicValueAttrsShardingInfo{ @@ -423,8 +777,8 @@ TEST_SUITE(FF_TEST_SUITE) { }; std::set correct = { - mk_invocation_shard(mc1, pt1), - mk_invocation_shard(mc2, pt2), + mk_invocation_shard(dev1, pt1), + mk_invocation_shard(dev2, pt2), }; CHECK(result.size() == correct.size()); @@ -432,11 +786,6 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("replicate operator") { - MachineSpaceCoordinate mc1 = mk_machine_coord(0_n, 0_n); - MachineSpaceCoordinate mc2 = mk_machine_coord(1_n, 0_n); - MachineSpaceCoordinate mc3 = mk_machine_coord(2_n, 0_n); - MachineSpaceCoordinate mc4 = mk_machine_coord(3_n, 0_n); - ParallelTensorSpaceCoordinate pt1 = mk_pt_coord(0_n, 0_n, 0_n, 0_n); ParallelTensorSpaceCoordinate pt2 = mk_pt_coord(0_n, 0_n, 0_n, 1_n); ParallelTensorSpaceCoordinate pt3 = mk_pt_coord(0_n, 1_n, 0_n, 0_n); @@ -459,38 +808,41 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - MappedOperatorTaskGroup mapped_task_group = MappedOperatorTaskGroup{ - bidict{ - { - mc1, - mk_shard_binding(pt1, pt1), - }, - { - mc2, - mk_shard_binding(pt1, pt2), - }, - { - mc3, - mk_shard_binding(pt2, pt3), - }, - { - mc4, - mk_shard_binding(pt2, pt4), - }, - }, + DynamicNodeMapping node_mapping = DynamicNodeMapping{ + /*op_task_group=*/MappedOperatorTaskGroup{ + bidict{ + { + dev1.coord, + mk_shard_binding(pt1, pt1), + }, + { + dev2.coord, + mk_shard_binding(pt1, pt2), + }, + { + dev3.coord, + mk_shard_binding(pt2, pt3), + }, + { + dev4.coord, + mk_shard_binding(pt2, pt4), + }, + }, + }, + /*device_type=*/DeviceType::GPU, }; SUBCASE("fwd") { - bidict src_binding{ - {pt1, mc1}, - {pt2, mc2}, + bidict src_binding{ + {pt1, dev1}, + {pt2, dev2}, }; - bidict dst_binding{ - {pt1, mc1}, - {pt2, mc2}, - {pt3, mc3}, - {pt4, mc4}, + bidict dst_binding{ + {pt1, dev1}, + {pt2, dev2}, + {pt3, dev3}, + {pt4, dev4}, }; DynamicNodeInvocation input = DynamicNodeInvocation{ @@ -507,8 +859,8 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/DynamicTaskType::FWD, - /*device_coords=*/std::nullopt, - /*mapping=*/mapped_task_group, + /*device_ids=*/std::nullopt, + /*mapping=*/node_mapping, /*op_attrs=*/TrainingOperationAttrs{ PCGOperatorAttrs{ ReplicateAttrs{ @@ -536,52 +888,53 @@ TEST_SUITE(FF_TEST_SUITE) { generate_shard_expansion_for_invocation(input); - auto mk_output_binding = [&](MachineSpaceCoordinate const &mc) + auto mk_output_binding = [&](global_device_id_t const &device) -> std::pair { return { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/mc, + /*task_shard=*/device.coord, }, DynamicValueAttrsShardingInfo{ - dst_binding.at_r(mc), - mc, + dst_binding.at_r(device), + device, }, }; }; auto mk_invocation_shard = - [&](nonempty_set const &device_coords, + [&](nonempty_set const &device_ids, ParallelTensorSpaceCoordinate const &input_shard_coord, - std::set const &output_task_shards) + std::set const &output_task_shards) -> DynamicNodeInvocationShardingInfo { return DynamicNodeInvocationShardingInfo{ - /*device_coords=*/device_coords, + /*device_ids=*/device_ids, /*value_sharding=*/ - binary_merge_disjoint_maps( - std::map{ - { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrsShardingInfo{ - input_shard_coord, - src_binding.at_l(input_shard_coord), + binary_relation_from_map( + binary_merge_disjoint_maps( + std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrsShardingInfo{ + input_shard_coord, + src_binding.at_l(input_shard_coord), + }, }, }, - }, - map_from_pairs(transform(output_task_shards, mk_output_binding))), + map_from_pairs(transform(output_task_shards, mk_output_binding)))), }; }; std::set correct = { - mk_invocation_shard(nonempty_set{mc1, mc2}, pt1, {mc1, mc2}), - mk_invocation_shard(nonempty_set{mc3, mc4}, pt2, {mc3, mc4}), + mk_invocation_shard(nonempty_set{dev1, dev2}, pt1, {dev1, dev2}), + mk_invocation_shard(nonempty_set{dev3, dev4}, pt2, {dev3, dev4}), }; CHECK(result.size() == correct.size()); @@ -589,16 +942,16 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("bwd") { - bidict output_grad_binding{ - {pt1, mc1}, - {pt2, mc2}, - {pt3, mc3}, - {pt4, mc4}, + bidict output_grad_binding{ + {pt1, dev1}, + {pt2, dev2}, + {pt3, dev3}, + {pt4, dev4}, }; - bidict input_grad_binding{ - {pt1, mc1}, - {pt2, mc2}, + bidict input_grad_binding{ + {pt1, dev1}, + {pt2, dev2}, }; DynamicNodeInvocation input = DynamicNodeInvocation{ @@ -615,8 +968,8 @@ TEST_SUITE(FF_TEST_SUITE) { /*node_attrs=*/ DynamicNodeAttrs{ /*task_type=*/DynamicTaskType::BWD, - /*device_coords=*/std::nullopt, - /*mapping=*/mapped_task_group, + /*device_ids=*/std::nullopt, + /*mapping=*/node_mapping, /*op_attrs=*/TrainingOperationAttrs{ PCGOperatorAttrs{ ReplicateAttrs{ @@ -643,52 +996,53 @@ TEST_SUITE(FF_TEST_SUITE) { std::set result = generate_shard_expansion_for_invocation(input); - auto mk_output_grad_binding = [&](MachineSpaceCoordinate const &mc) + auto mk_output_grad_binding = [&](global_device_id_t const &device) -> std::pair { return { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/mc, + /*task_shard=*/device.coord, }, DynamicValueAttrsShardingInfo{ - output_grad_binding.at_r(mc), - mc, + output_grad_binding.at_r(device), + device, }, }; }; auto mk_invocation_shard = - [&](nonempty_set const &device_coords, - std::set const &output_grad_task_shards, + [&](nonempty_set const &device_ids, + std::set const &output_grad_task_shards, ParallelTensorSpaceCoordinate const &input_grad_shard_coord) -> DynamicNodeInvocationShardingInfo { return DynamicNodeInvocationShardingInfo{ - /*device_coords=*/device_coords, + /*device_ids=*/device_ids, /*value_sharding=*/ - binary_merge_disjoint_maps( - std::map{ - { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrsShardingInfo{ - input_grad_shard_coord, - input_grad_binding.at_l(input_grad_shard_coord), + binary_relation_from_map( + binary_merge_disjoint_maps( + std::map{ + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrsShardingInfo{ + input_grad_shard_coord, + input_grad_binding.at_l(input_grad_shard_coord), + }, }, }, - }, - map_from_pairs(transform(output_grad_task_shards, mk_output_grad_binding))), + map_from_pairs(transform(output_grad_task_shards, mk_output_grad_binding)))), }; }; std::set correct = { - mk_invocation_shard(nonempty_set{mc1, mc2}, {mc1, mc2}, pt1), - mk_invocation_shard(nonempty_set{mc3, mc4}, {mc3, mc4}, pt2), + mk_invocation_shard(nonempty_set{dev1, dev2}, {dev1, dev2}, pt1), + mk_invocation_shard(nonempty_set{dev3, dev4}, {dev3, dev4}, pt2), }; CHECK(result.size() == correct.size()); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc index 3fb94bb98a..ae39d5afcb 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/update_insertion.cc @@ -52,15 +52,17 @@ TEST_SUITE(FF_TEST_SUITE) { /*per_device_op_state=*/std::nullopt, }, /*outputs=*/ - std::unordered_map{ + std::map{ { DynamicTensorSlot{ /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/tensor_guid, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -96,10 +98,12 @@ TEST_SUITE(FF_TEST_SUITE) { /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/ mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/tensor_guid, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -111,10 +115,12 @@ TEST_SUITE(FF_TEST_SUITE) { /*slot_name=*/TensorSlotName::OUTPUT, /*slot_tensor_role=*/ mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/tensor_guid, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, @@ -127,10 +133,12 @@ TEST_SUITE(FF_TEST_SUITE) { /*slot_tensor_role=*/ mk_dynamic_tensor_role_opt( OptimizerSlotName::SGD_V), + /*task_shard=*/std::nullopt, }, DynamicValueAttrs{ /*tensor_guid=*/tensor_guid, /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, /*shard_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*accessor=*/std::nullopt, diff --git a/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h b/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h index 3f7a3e4ff2..3125c95eb1 100644 --- a/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h @@ -2,12 +2,58 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FROM_UNSTRUCTURED_RELATION_H #include "utils/bidict/bidict.h" +#include "utils/containers/transform.h" +#include "utils/containers/get_element_counts.h" +#include +#include "utils/containers/filter_values.h" +#include "utils/containers/multiset_of.h" namespace FlexFlow { template bidict bidict_from_unstructured_relation( std::set> const &relation) { + + { + std::multiset l_values = transform( + multiset_of(relation), + [](std::pair const &p) -> L { + return p.first; + } + ); + + std::map l_value_counts = get_element_counts(l_values); + + std::map duplicated_element_counts = + filter_values(l_value_counts, + [](positive_int num_occurences) -> bool { + return num_occurences > 1; + }); + + ASSERT(duplicated_element_counts.empty(), + duplicated_element_counts); + } + + { + std::multiset r_values = transform( + multiset_of(relation), + [](std::pair const &p) -> R { + return p.second; + } + ); + + std::map r_value_counts = get_element_counts(r_values); + + std::map duplicated_element_counts = + filter_values(r_value_counts, + [](positive_int num_occurences) -> bool { + return num_occurences > 1; + }); + + ASSERT(duplicated_element_counts.empty(), + duplicated_element_counts); + } + bidict result; for (auto const &lr : relation) { result.equate_strict(lr); diff --git a/lib/utils/include/utils/bidict/algorithms/transform_keys.h b/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys.h similarity index 58% rename from lib/utils/include/utils/bidict/algorithms/transform_keys.h rename to lib/utils/include/utils/bidict/algorithms/bidict_transform_keys.h index 1d82464d17..1c4ea5b623 100644 --- a/lib/utils/include/utils/bidict/algorithms/transform_keys.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys.h @@ -1,5 +1,5 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_TRANSFORM_KEYS_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_TRANSFORM_KEYS_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_TRANSFORM_KEYS_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_TRANSFORM_KEYS_H #include "utils/bidict/bidict.h" @@ -9,7 +9,7 @@ template > -bidict transform_keys(bidict const &m, F &&f) { +bidict bidict_transform_keys(bidict const &m, F &&f) { bidict result; for (auto const &kv : m) { result.equate_strict(f(kv.first), kv.second); diff --git a/lib/utils/include/utils/bidict/algorithms/transform_values.h b/lib/utils/include/utils/bidict/algorithms/bidict_transform_values.h similarity index 58% rename from lib/utils/include/utils/bidict/algorithms/transform_values.h rename to lib/utils/include/utils/bidict/algorithms/bidict_transform_values.h index fc8655594e..82acc50676 100644 --- a/lib/utils/include/utils/bidict/algorithms/transform_values.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_transform_values.h @@ -1,5 +1,5 @@ -#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_TRANSFORM_VALUES_H -#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_TRANSFORM_VALUES_H +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_TRANSFORM_VALUES_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_TRANSFORM_VALUES_H #include "utils/bidict/bidict.h" @@ -9,7 +9,7 @@ template > -bidict transform_values(bidict const &m, F &&f) { +bidict bidict_transform_values(bidict const &m, F &&f) { bidict result; for (auto const &kv : m) { result.equate_strict({kv.first, f(kv.second)}); diff --git a/lib/utils/include/utils/binary_relation/binary_relation.h b/lib/utils/include/utils/binary_relation/binary_relation.h new file mode 100644 index 0000000000..0d09045ec2 --- /dev/null +++ b/lib/utils/include/utils/binary_relation/binary_relation.h @@ -0,0 +1,196 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_H + +#include +#include "utils/json/check_is_json_deserializable.h" +#include "utils/hash-utils.h" +#include "utils/hash/set.h" +#include "utils/hash/pair.h" +#include "utils/hash/tuple.h" +#include "utils/fmt/set.h" +#include "utils/fmt/pair.h" +#include "utils/json/check_is_json_serializable.h" +#include +#include +#include "utils/containers/multiset_of.h" +#include "utils/containers/filtrans.h" +#include "utils/containers/transform.h" + +namespace FlexFlow { + +template +struct BinaryRelation { + BinaryRelation() : raw{} {} + + BinaryRelation(std::set> const &raw) + : raw(raw) {} + + BinaryRelation(std::initializer_list> init) + : BinaryRelation(init.begin(), init.end()) {} + + template + BinaryRelation(InputIt first, InputIt last) { + for (auto it = first; it != last; it++) { + this->equate(it->first, it->second); + } + } + + bool operator==(BinaryRelation const &other) const { + return this->tie() == other.tie(); + } + + bool operator!=(BinaryRelation const &other) const { + return this->tie() != other.tie(); + } + + bool operator<(BinaryRelation const &other) const { + return this->tie() < other.tie(); + } + + bool operator<=(BinaryRelation const &other) const { + return this->tie() <= other.tie(); + } + + bool operator>(BinaryRelation const &other) const { + return this->tie() > other.tie(); + } + + bool operator>=(BinaryRelation const &other) const { + return this->tie() >= other.tie(); + } + + void equate(L const &l, R const &r) { + this->raw.insert({l, r}); + } + + void equate(std::pair const &lr) { + this->raw.insert(lr); + } + + std::set at_l(L const &l) const { + return filtrans( + this->raw, + [&](std::pair const &p) -> std::optional { + if (p.first == l) { + return p.second; + } else { + return std::nullopt; + } + }); + } + + std::set at_r(R const &r) const { + return filtrans( + this->raw, + [&](std::pair const &p) -> std::optional { + if (p.second == r) { + return p.first; + } else { + return std::nullopt; + } + }); + } + + std::multiset left_value_occurences() const { + return transform( + multiset_of(this->raw), + [&](std::pair const &p) -> L { + return p.first; + }); + } + + std::multiset right_value_occurences() const { + return transform( + multiset_of(this->raw), + [&](std::pair const &p) -> R { + return p.second; + }); + } + + std::set left_values() const { + return transform( + this->raw, + [&](std::pair const &p) -> L { + return p.first; + }); + } + + std::set right_values() const { + return transform( + this->raw, + [&](std::pair const &p) -> R { + return p.second; + }); + } + + std::size_t size() const { + return this->raw.size(); + } + + bool empty() const { + return this->raw.empty(); + } + + std::set> const &unwrap_as_set() const { + return this->raw; + } +private: + std::set> raw; + +private: + std::tuple + tie() const { + return std::tie(this->raw); + } + + friend struct std::hash>; +}; + +template +std::set> + format_as(BinaryRelation const &m) { + return m.unwrap_as_set(); +} + +template +std::ostream &operator<<(std::ostream &s, BinaryRelation const &m) { + return (s << fmt::to_string(m)); +} + +} // namespace FlexFlow + +namespace nlohmann { + +template +struct adl_serializer<::FlexFlow::BinaryRelation> { + static ::FlexFlow::BinaryRelation from_json(json const &j) { + CHECK_IS_JSON_DESERIALIZABLE(L); + CHECK_IS_JSON_DESERIALIZABLE(R); + + std::set> s = j; + + return ::FlexFlow::BinaryRelation(s.cbegin(), s.cend()); + } + + static void to_json(json &j, ::FlexFlow::BinaryRelation const &m) { + CHECK_IS_JSON_SERIALIZABLE(L); + CHECK_IS_JSON_SERIALIZABLE(R); + + j = m.unwrap_as_set(); + } +}; + +} // namespace nlohmann + +namespace std { + +template +struct hash<::FlexFlow::BinaryRelation> { + size_t operator()(::FlexFlow::BinaryRelation const &m) const { + return ::FlexFlow::get_std_hash(m.tie()); + } +}; + +} // namespace std + +#endif diff --git a/lib/utils/include/utils/binary_relation/binary_relation_from_map.h b/lib/utils/include/utils/binary_relation/binary_relation_from_map.h new file mode 100644 index 0000000000..7ddb857722 --- /dev/null +++ b/lib/utils/include/utils/binary_relation/binary_relation_from_map.h @@ -0,0 +1,19 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_FROM_MAP_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_FROM_MAP_H + +#include +#include "utils/binary_relation/binary_relation.h" +#include "utils/containers/set_of.h" + +namespace FlexFlow { + +template +BinaryRelation binary_relation_from_map(std::map const &m) { + return BinaryRelation{ + set_of(m), + }; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_left.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_left.h new file mode 100644 index 0000000000..60669bbfc1 --- /dev/null +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_left.h @@ -0,0 +1,28 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT_H + +#include +#include "utils/binary_relation/binary_relation.h" + +namespace FlexFlow { + +template < + typename L, + typename R, + typename F, + typename L2 = std::invoke_result_t +> +BinaryRelation binary_relation_transform_left(BinaryRelation const &rel, + F &&f) { + BinaryRelation result; + + for (std::pair const &p : rel.unwrap_as_set()) { + result.equate(f(p.first), p.second); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h new file mode 100644 index 0000000000..a317992edf --- /dev/null +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h @@ -0,0 +1,28 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT2_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT2_H + +#include +#include "utils/binary_relation/binary_relation.h" + +namespace FlexFlow { + +template < + typename L, + typename R, + typename F, + typename L2 = std::invoke_result_t +> +BinaryRelation binary_relation_transform_left2(BinaryRelation const &rel, + F &&f) { + BinaryRelation result; + + for (std::pair const &p : rel.unwrap_as_set()) { + result.equate(f(p.first, p.second), p.second); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_right.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_right.h new file mode 100644 index 0000000000..a2e36a7441 --- /dev/null +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_right.h @@ -0,0 +1,28 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT_H + +#include +#include "utils/binary_relation/binary_relation.h" + +namespace FlexFlow { + +template < + typename L, + typename R, + typename F, + typename R2 = std::invoke_result_t +> +BinaryRelation binary_relation_transform_right(BinaryRelation const &rel, + F &&f) { + BinaryRelation result; + + for (std::pair const &p : rel.unwrap_as_set()) { + result.equate(p.first, f(p.second)); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h new file mode 100644 index 0000000000..2ca789a8e0 --- /dev/null +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h @@ -0,0 +1,28 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT2_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT2_H + +#include +#include "utils/binary_relation/binary_relation.h" + +namespace FlexFlow { + +template < + typename L, + typename R, + typename F, + typename R2 = std::invoke_result_t +> +BinaryRelation binary_relation_transform_right2(BinaryRelation const &rel, + F &&f) { + BinaryRelation result; + + for (std::pair const &p : rel.unwrap_as_set()) { + result.equate(p.first, f(p.first, p.second)); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/filter_binary_relation.h b/lib/utils/include/utils/binary_relation/filter_binary_relation.h new file mode 100644 index 0000000000..2f8026a3ed --- /dev/null +++ b/lib/utils/include/utils/binary_relation/filter_binary_relation.h @@ -0,0 +1,21 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_FILTER_BINARY_RELATION_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_FILTER_BINARY_RELATION_H + +#include "utils/binary_relation/binary_relation.h" +#include "utils/containers/filter.h" + +namespace FlexFlow { + +template +BinaryRelation filter_binary_relation(BinaryRelation const &rel, F &&f) { + return BinaryRelation{ + filter(rel.unwrap_as_set(), + [&](std::pair const &p) -> bool { + return f(p.first, p.second); + }), + }; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h b/lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h new file mode 100644 index 0000000000..152e0db15b --- /dev/null +++ b/lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h @@ -0,0 +1,22 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_REQUIRE_BINARY_RELATION_IS_LEFT_UNIQUE_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_REQUIRE_BINARY_RELATION_IS_LEFT_UNIQUE_H + +#include "utils/binary_relation/binary_relation.h" +#include "utils/one_to_many/one_to_many.h" + +namespace FlexFlow { + +template +OneToMany require_binary_relation_is_left_unique(BinaryRelation const &rel) { + OneToMany result; + + for (std::pair const &p : rel.unwrap_as_set()) { + result.insert(p); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h b/lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h new file mode 100644 index 0000000000..5d765bbb64 --- /dev/null +++ b/lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h @@ -0,0 +1,22 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_REQUIRE_BINARY_RELATION_IS_LEFT_UNIQUE_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_REQUIRE_BINARY_RELATION_IS_LEFT_UNIQUE_H + +#include "utils/binary_relation/binary_relation.h" +#include "utils/many_to_one/many_to_one.h" + +namespace FlexFlow { + +template +ManyToOne require_binary_relation_is_right_unique(BinaryRelation const &rel) { + ManyToOne result; + + for (std::pair const &p : rel.unwrap_as_set()) { + result.insert(p); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/binary_relation/transform_binary_relation.h b/lib/utils/include/utils/binary_relation/transform_binary_relation.h new file mode 100644 index 0000000000..538c477ea1 --- /dev/null +++ b/lib/utils/include/utils/binary_relation/transform_binary_relation.h @@ -0,0 +1,29 @@ +#ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_TRANSFORM_BINARY_RELATION_H +#define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_TRANSFORM_BINARY_RELATION_H + +#include +#include "utils/binary_relation/binary_relation.h" + +namespace FlexFlow { + +template < + typename L, + typename R, + typename F, + typename L2 = std::invoke_result_t::first_type, + typename R2 = std::invoke_result_t::second_type +> +BinaryRelation binary_relation_transform_left(BinaryRelation const &rel, + F &&f) { + BinaryRelation result; + + for (std::pair const &p : rel.unwrap_as_set()) { + result.equate(f(p.first, p.second)); + } + + return result; +} + +} // namespace FlexFlow + +#endif diff --git a/lib/utils/include/utils/containers/at_idx.h b/lib/utils/include/utils/containers/at_idx.h index 2442a759ac..dae35a0dff 100644 --- a/lib/utils/include/utils/containers/at_idx.h +++ b/lib/utils/include/utils/containers/at_idx.h @@ -15,6 +15,17 @@ E at_idx(std::vector const &v, nonnegative_int idx) { return v.at(idx.unwrap_nonnegative()); } +template +E at_idx(std::set const &v, nonnegative_int idx) { + ASSERT(idx < v.size()); + + auto b = v.cbegin(); + for (int i = 0; i < idx; i++) { + b++; + }; + return *b; +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/filtrans.h b/lib/utils/include/utils/containers/filtrans.h index 9ee65dee74..76e3bfaa84 100644 --- a/lib/utils/include/utils/containers/filtrans.h +++ b/lib/utils/include/utils/containers/filtrans.h @@ -68,6 +68,38 @@ std::set filtrans(std::set const &s, F &&f) { return result; } +template >> +std::multiset filtrans(std::multiset const &s, F &&f) { + std::multiset result; + + for (In const &i : s) { + std::optional o = f(i); + if (o.has_value()) { + result.insert(o.value()); + } + } + + return result; +} + +template >> +std::unordered_multiset filtrans(std::unordered_multiset const &s, F &&f) { + std::unordered_multiset result; + + for (In const &i : s) { + std::optional o = f(i); + if (o.has_value()) { + result.insert(o.value()); + } + } + + return result; +} + } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/get_element_counts.h b/lib/utils/include/utils/containers/get_element_counts.h index e32ca7b552..121e5399d5 100644 --- a/lib/utils/include/utils/containers/get_element_counts.h +++ b/lib/utils/include/utils/containers/get_element_counts.h @@ -5,22 +5,38 @@ #include #include #include +#include +#include "utils/positive_int/positive_int.h" namespace FlexFlow { template -std::map get_element_counts(std::vector const &v) { - std::map counts; +std::map get_element_counts(std::vector const &v) { + std::map counts; for (T const &t : v) { if (!contains_key(counts, t)) { - counts[t] = 0; + counts.insert({t, 1_p}); + } else { + counts.at(t)++; } - counts.at(t)++; } return counts; } -std::map get_element_counts(std::string const &); +template +std::map get_element_counts(std::multiset const &v) { + std::map counts; + for (T const &t : v) { + if (!contains_key(counts, t)) { + counts.insert({t, 1_p}); + } else { + counts.at(t)++; + } + } + return counts; +} + +std::map get_element_counts(std::string const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.toml b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.toml index 5b537eac88..9a58e2f686 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.toml +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output.dtg.toml @@ -18,6 +18,10 @@ includes = [ "utils/nonnegative_int/nonnegative_int.h", ] +src_includes = [ + "utils/json/optional.h", +] + [[fields]] name = "node" type = "::FlexFlow::Node" diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h index df8207251f..b94f3df126 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_IS_ISOMORPHIC_UNDER_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_DATAFLOW_GRAPH_ALGORITHMS_IS_ISOMORPHIC_UNDER_H -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/get_graph_data.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/permute_input_ids.h" #include "utils/graph/labelled_open_dataflow_graph/algorithms/permute_node_ids.h" @@ -18,11 +18,11 @@ bool is_isomorphic_under( OpenDataflowGraphIsomorphism const &candidate_isomorphism) { bidict node_permutation = - transform_values(candidate_isomorphism.node_mapping, + bidict_transform_values(candidate_isomorphism.node_mapping, [](Node const &dst_node) { return NewNode{dst_node}; }) .reversed(); bidict input_permutation = - transform_values(candidate_isomorphism.input_mapping, + bidict_transform_values(candidate_isomorphism.input_mapping, [](DataflowGraphInput const &dst_input) { return NewDataflowGraphInput{dst_input}; }) diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h index 015b388fa9..d82dc86e19 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_LABELLED_OPEN_KWARG_DATAFLOW_GRAPHS_ARE_ISOMORPHIC_UNDER_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_LABELLED_OPEN_KWARG_DATAFLOW_GRAPHS_ARE_ISOMORPHIC_UNDER_H -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_input_ids.h" @@ -28,7 +28,7 @@ bool labelled_open_kwarg_dataflow_graphs_are_isomorphic_under( OpenKwargDataflowGraphIsomorphism const &candidate_isomorphism) { bidict new_node_to_old_node = - transform_values(candidate_isomorphism.node_mapping, [](Node const &n) { + bidict_transform_values(candidate_isomorphism.node_mapping, [](Node const &n) { return NewNode{n}; }).reversed(); diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h index fa18a86a8d..83391f46be 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_edge_counts.h @@ -2,10 +2,11 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_MULTIDIGRAPH_ALGORITHMS_GET_EDGE_COUNTS_H #include "utils/graph/multidigraph/multidigraph_view.h" +#include "utils/positive_int/positive_int.h" namespace FlexFlow { -std::map get_edge_counts(MultiDiGraphView const &); +std::map get_edge_counts(MultiDiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.h index 63c367a987..249e406126 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graphs_are_isomorphic_under.h @@ -1,7 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_OPEN_KWARG_DATAFLOW_GRAPHS_ARE_ISOMORPHIC_UNDER_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_OPEN_KWARG_DATAFLOW_GRAPHS_ARE_ISOMORPHIC_UNDER_H -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_isomorphism.dtg.h" @@ -17,7 +17,7 @@ bool open_kwarg_dataflow_graphs_are_isomorphic_under( OpenKwargDataflowGraphView const &dst, OpenKwargDataflowGraphIsomorphism const &isomorphism) { bidict new_node_to_old_node = - transform_values(isomorphism.node_mapping, [](Node const &n) { + bidict_transform_values(isomorphism.node_mapping, [](Node const &n) { return NewNode{n}; }).reversed(); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h index 7e1fe0e60c..19d9f10ee7 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h @@ -14,6 +14,7 @@ #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" #include "utils/overload.h" #include "utils/containers/set_of.h" +#include "utils/json/optional.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/nonempty_set/nonempty_set.h b/lib/utils/include/utils/nonempty_set/nonempty_set.h index 93d2b37def..61276f4db9 100644 --- a/lib/utils/include/utils/nonempty_set/nonempty_set.h +++ b/lib/utils/include/utils/nonempty_set/nonempty_set.h @@ -5,6 +5,7 @@ #include #include "utils/hash-utils.h" #include "utils/hash/set.h" +#include "utils/hash/tuple.h" #include "utils/fmt/set.h" #include "utils/positive_int/positive_int.h" #include "utils/containers/set_of.h" diff --git a/lib/utils/src/utils/bidict/algorithms/transform_keys.cc b/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys.cc similarity index 63% rename from lib/utils/src/utils/bidict/algorithms/transform_keys.cc rename to lib/utils/src/utils/bidict/algorithms/bidict_transform_keys.cc index a96ec0487b..1055cbdc69 100644 --- a/lib/utils/src/utils/bidict/algorithms/transform_keys.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys.cc @@ -1,4 +1,4 @@ -#include "utils/bidict/algorithms/transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" #include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -8,6 +8,6 @@ using V = ordered_value_type<1>; using K2 = ordered_value_type<2>; using F = std::function; -template bidict transform_keys(bidict const &, F &&); +template bidict bidict_transform_keys(bidict const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/transform_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_transform_values.cc similarity index 62% rename from lib/utils/src/utils/bidict/algorithms/transform_values.cc rename to lib/utils/src/utils/bidict/algorithms/bidict_transform_values.cc index d6e3d57c13..aae6140a41 100644 --- a/lib/utils/src/utils/bidict/algorithms/transform_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_transform_values.cc @@ -1,4 +1,4 @@ -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -8,6 +8,6 @@ using V = ordered_value_type<1>; using V2 = ordered_value_type<2>; using F = std::function; -template bidict transform_values(bidict const &, F &&); +template bidict bidict_transform_values(bidict const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation.cc b/lib/utils/src/utils/binary_relation/binary_relation.cc new file mode 100644 index 0000000000..0fe55f17c6 --- /dev/null +++ b/lib/utils/src/utils/binary_relation/binary_relation.cc @@ -0,0 +1,34 @@ +#include "utils/binary_relation/binary_relation.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = jsonable_ordered_value_type<0>; +using R = jsonable_ordered_value_type<1>; + +template struct BinaryRelation; + +template std::set> format_as(BinaryRelation const &); + +template std::ostream &operator<<(std::ostream &, BinaryRelation const &); + +} // namespace FlexFlow + +namespace nlohmann { + +using L = ::FlexFlow::jsonable_ordered_value_type<0>; +using R = ::FlexFlow::jsonable_ordered_value_type<1>; + +template struct adl_serializer<::FlexFlow::BinaryRelation>; + +} // namespace nlohmann + +namespace std { + +using L = ::FlexFlow::ordered_value_type<0>; +using R = ::FlexFlow::ordered_value_type<1>; + +template struct hash<::FlexFlow::BinaryRelation>; + +} // namespace std diff --git a/lib/utils/src/utils/binary_relation/binary_relation_from_map.cc b/lib/utils/src/utils/binary_relation/binary_relation_from_map.cc new file mode 100644 index 0000000000..49b8152def --- /dev/null +++ b/lib/utils/src/utils/binary_relation/binary_relation_from_map.cc @@ -0,0 +1,11 @@ +#include "utils/binary_relation/binary_relation_from_map.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template BinaryRelation binary_relation_from_map(std::map const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc new file mode 100644 index 0000000000..073265d427 --- /dev/null +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc @@ -0,0 +1,14 @@ +#include "utils/binary_relation/binary_relation_transform_left.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; +using L2 = ordered_value_type<2>; +using F = std::function; + +template + BinaryRelation binary_relation_transform_left(BinaryRelation const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc new file mode 100644 index 0000000000..a26b8d1ee4 --- /dev/null +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc @@ -0,0 +1,14 @@ +#include "utils/binary_relation/binary_relation_transform_left2.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using L2 = ordered_value_type<1>; +using R = ordered_value_type<2>; +using F = std::function; + +template + BinaryRelation binary_relation_transform_left2(BinaryRelation const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc new file mode 100644 index 0000000000..5274340885 --- /dev/null +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc @@ -0,0 +1,14 @@ +#include "utils/binary_relation/binary_relation_transform_right.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; +using R2 = ordered_value_type<2>; +using F = std::function; + +template + BinaryRelation binary_relation_transform_right(BinaryRelation const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc new file mode 100644 index 0000000000..9f76943eae --- /dev/null +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc @@ -0,0 +1,14 @@ +#include "utils/binary_relation/binary_relation_transform_right2.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; +using R2 = ordered_value_type<2>; +using F = std::function; + +template + BinaryRelation binary_relation_transform_right2(BinaryRelation const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/filter_binary_relation.cc b/lib/utils/src/utils/binary_relation/filter_binary_relation.cc new file mode 100644 index 0000000000..fd0d497b98 --- /dev/null +++ b/lib/utils/src/utils/binary_relation/filter_binary_relation.cc @@ -0,0 +1,12 @@ +#include "utils/binary_relation/filter_binary_relation.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; +using F = std::function; + +template BinaryRelation filter_binary_relation(BinaryRelation const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc b/lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc new file mode 100644 index 0000000000..09b1c1d539 --- /dev/null +++ b/lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc @@ -0,0 +1,13 @@ +#include "utils/binary_relation/require_binary_relation_is_left_unique.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template + OneToMany require_binary_relation_is_left_unique(BinaryRelation const &); + + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc b/lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc new file mode 100644 index 0000000000..aebdc9895c --- /dev/null +++ b/lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc @@ -0,0 +1,12 @@ +#include "utils/binary_relation/require_binary_relation_is_right_unique.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; + +template + ManyToOne require_binary_relation_is_right_unique(BinaryRelation const &); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/at_idx.cc b/lib/utils/src/utils/containers/at_idx.cc index 0f398b144e..c0f315d766 100644 --- a/lib/utils/src/utils/containers/at_idx.cc +++ b/lib/utils/src/utils/containers/at_idx.cc @@ -1,5 +1,6 @@ #include "utils/containers/at_idx.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -7,4 +8,8 @@ using E = value_type<0>; template E at_idx(std::vector const &, nonnegative_int); +using O_E = ordered_value_type<0>; + +template O_E at_idx(std::set const &, nonnegative_int); + } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/filtrans.cc b/lib/utils/src/utils/containers/filtrans.cc index c65a22a669..6bb5bbbb85 100644 --- a/lib/utils/src/utils/containers/filtrans.cc +++ b/lib/utils/src/utils/containers/filtrans.cc @@ -1,5 +1,6 @@ #include "utils/containers/filtrans.h" #include "utils/archetypes/value_type.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { @@ -8,5 +9,14 @@ using Out = value_type<1>; using F = std::function(In const &)>; template std::vector filtrans(std::vector const &, F &&); +template std::unordered_set filtrans(std::unordered_set const &, F &&); +template std::unordered_multiset filtrans(std::unordered_multiset const &, F &&); + +using O_In = ordered_value_type<0>; +using O_Out = ordered_value_type<0>; +using O_F = std::function(O_In const &)>; + +template std::set filtrans(std::set const &, O_F &&); +template std::multiset filtrans(std::multiset const &, O_F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/get_element_counts.cc b/lib/utils/src/utils/containers/get_element_counts.cc index 70eda44608..c640d51d35 100644 --- a/lib/utils/src/utils/containers/get_element_counts.cc +++ b/lib/utils/src/utils/containers/get_element_counts.cc @@ -1,9 +1,15 @@ #include "utils/containers/get_element_counts.h" #include "utils/containers/vector_of.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -std::map get_element_counts(std::string const &s) { +using O_T = ordered_value_type<0>; + +template std::map get_element_counts(std::vector const &); +template std::map get_element_counts(std::multiset const &); + +std::map get_element_counts(std::string const &s) { return get_element_counts(vector_of(s)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc index 21d598eca6..21dbb737e0 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc @@ -25,10 +25,10 @@ std::map> transform(vector_of(n_dominators), [&](Node const &dominator) { return vector_of(node_to_its_dominators.at(dominator)); })); - std::map dominator_counts = + std::map dominator_counts = get_element_counts(recursive_dominator_list); std::set imm_dominators = keys( - filter_values(dominator_counts, [](int count) { return count <= 1; })); + filter_values(dominator_counts, [](positive_int count) { return count <= 1; })); ASSERT(imm_dominators.size() <= 1); return maybe_get_only(imm_dominators); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc b/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc index 2081e50cf0..ce4b029f4b 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc @@ -1,6 +1,6 @@ #include "utils/graph/digraph/algorithms/transitive_closure.h" #include "utils/bidict/algorithms/bidict_from_enumerating.h" -#include "utils/bidict/algorithms/transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" #include "utils/containers/vector_of.h" #include "utils/graph/digraph/algorithms/digraph_has_edge.h" #include "utils/graph/digraph/algorithms/get_edges.h" @@ -18,7 +18,7 @@ DiGraphView transitive_closure(DiGraphView const &g) { // (i.e., 200 nodes) without optimization enabled. bidict nodes = - transform_keys(bidict_from_enumerating(get_nodes(g)), + bidict_transform_keys(bidict_from_enumerating(get_nodes(g)), [](nonnegative_int x) { return x.unwrap_nonnegative(); }); std::set edges = get_edges(g); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc b/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc index f32de6c469..2c3b7f2b0c 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc @@ -1,6 +1,6 @@ #include "utils/graph/digraph/algorithms/transitive_reduction.h" #include "utils/bidict/algorithms/bidict_from_enumerating.h" -#include "utils/bidict/algorithms/transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/set_intersection.h" #include "utils/containers/vector_of.h" @@ -41,7 +41,7 @@ DiGraph transitive_reduction(DiGraphView const &g) { // between transitive_closure and transitive_reduction bidict nodes = - transform_keys(bidict_from_enumerating(get_nodes(g)), + bidict_transform_keys(bidict_from_enumerating(get_nodes(g)), [](nonnegative_int x) { return x.unwrap_nonnegative(); }); int num_nodes = nodes.size(); diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc index 62268cc868..18ee3e677a 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc @@ -1,9 +1,9 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using SlotName = ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<0>; template DataflowGraphData dataflow_graph_data_from_kwarg_dataflow_graph_data( KwargDataflowGraphData const &, diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc index 2d10dc66fe..927c6e2d82 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc @@ -1,9 +1,9 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using SlotName = ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<0>; template DataflowGraphView dataflow_graph_from_kwarg_dataflow_graph( KwargDataflowGraphView const &, diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.cc index 06627a9b51..0a682c1190 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.cc @@ -1,9 +1,9 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using SlotName = ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<0>; template std::optional> find_isomorphism_between_kwarg_dataflow_graphs( diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_edges_for_node.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_edges_for_node.cc index 0e22b6f632..a68bac54fd 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_edges_for_node.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_edges_for_node.cc @@ -1,9 +1,9 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_kwarg_dataflow_edges_for_node.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using SlotName = ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<0>; template OneToMany> get_outgoing_kwarg_dataflow_edges_for_node( diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc index 06c2a96bf4..6f6d9bacad 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc @@ -1,9 +1,9 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using SlotName = ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<0>; template std::string kwarg_dataflow_graph_as_dot( KwargDataflowGraphView const &, diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graphs_are_isomorphic.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graphs_are_isomorphic.cc index 528d830261..d8c7e10f25 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graphs_are_isomorphic.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graphs_are_isomorphic.cc @@ -1,9 +1,9 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graphs_are_isomorphic.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using SlotName = ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<0>; template bool kwarg_dataflow_graphs_are_isomorphic( KwargDataflowGraphView const &, diff --git a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc index 1ef9fb2782..85abc54c05 100644 --- a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc +++ b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc @@ -1,12 +1,12 @@ #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { using NodeLabel = value_type<0>; using ValueLabel = value_type<1>; -using SlotName = ordered_value_type<2>; +using SlotName = jsonable_ordered_value_type<2>; template std::string labelled_kwarg_dataflow_graph_view_as_dot( LabelledKwargDataflowGraphView const &, diff --git a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.cc b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.cc index bc80881f5f..c3e01293e4 100644 --- a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.cc +++ b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.cc @@ -1,13 +1,13 @@ #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/find_isomorphism_between_labelled_open_kwarg_dataflow_graphs.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { using NodeLabel = value_type<0>; using ValueLabel = value_type<1>; -using GraphInputName = ordered_value_type<2>; -using SlotName = ordered_value_type<3>; +using GraphInputName = jsonable_ordered_value_type<2>; +using SlotName = jsonable_ordered_value_type<3>; template std::optional> find_isomorphism_between_labelled_open_kwarg_dataflow_graphs( diff --git a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc index 74187baa0b..1dcd6c9f2d 100644 --- a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc +++ b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc @@ -1,13 +1,13 @@ #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { using NodeLabel = value_type<0>; using ValueLabel = value_type<1>; -using GraphInputName = ordered_value_type<2>; -using SlotName = ordered_value_type<3>; +using GraphInputName = jsonable_ordered_value_type<2>; +using SlotName = jsonable_ordered_value_type<3>; template std::string labelled_open_kwarg_dataflow_graph_view_as_dot( LabelledOpenKwargDataflowGraphView; using ValueLabel = value_type<1>; -using GraphInputName = ordered_value_type<2>; -using SlotName = ordered_value_type<3>; +using GraphInputName = jsonable_ordered_value_type<2>; +using SlotName = jsonable_ordered_value_type<3>; template bool labelled_open_kwarg_dataflow_graphs_are_isomorphic_under( LabelledOpenKwargDataflowGraphView; using ValueLabel = value_type<1>; -using GraphInputName = ordered_value_type<2>; -using SlotName = ordered_value_type<3>; +using GraphInputName = jsonable_ordered_value_type<2>; +using SlotName = jsonable_ordered_value_type<3>; template LabelledOpenKwargDataflowGraphView +std::map get_edge_counts(MultiDiGraphView const &g) { return get_element_counts( transform(vector_of(get_edges(g)), diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc index 571b5306c9..d73a582252 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc @@ -1,5 +1,5 @@ #include "utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.h" -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/graph/node/algorithms/new_node.dtg.h" #include "utils/graph/open_dataflow_graph/algorithms/get_graph_data.h" #include "utils/graph/open_dataflow_graph/algorithms/new_dataflow_graph_input.dtg.h" @@ -14,11 +14,11 @@ bool is_isomorphic_under( OpenDataflowGraphIsomorphism const &candidate_isomorphism) { bidict node_permutation = - transform_values(candidate_isomorphism.node_mapping, + bidict_transform_values(candidate_isomorphism.node_mapping, [](Node const &dst_node) { return NewNode{dst_node}; }) .reversed(); bidict input_permutation = - transform_values(candidate_isomorphism.input_mapping, + bidict_transform_values(candidate_isomorphism.input_mapping, [](DataflowGraphInput const &dst_input) { return NewDataflowGraphInput{dst_input}; }) diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc index 3e9a67f4e6..0d5bca59c6 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.cc @@ -1,10 +1,10 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using GraphInputName = ordered_value_type<0>; -using SlotName = ordered_value_type<1>; +using GraphInputName = jsonable_ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<1>; template std::set> find_isomorphisms_between_open_kwarg_dataflow_graphs( diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.cc index 515a7bc413..12a15b55a2 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.cc @@ -1,10 +1,10 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/generate_new_kwarg_dataflow_graph_input_id_permutation.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using GraphInputName = ordered_value_type<0>; -using SlotName = ordered_value_type<1>; +using GraphInputName = jsonable_ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<1>; template bidict, KwargDataflowGraphInput> diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc index 9527e53104..291188723b 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc @@ -1,10 +1,14 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" +#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { -template -std::set> - get_all_open_kwarg_dataflow_edges( - OpenKwargDataflowGraphView const &); +using GraphInputName = ordered_value_type<0>; +using SlotName = ordered_value_type<1>; + +template + std::set> + get_all_open_kwarg_dataflow_edges( + OpenKwargDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc index 064eb36c46..8e65f18ece 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.cc @@ -1,10 +1,10 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using GraphInputName = ordered_value_type<0>; -using SlotName = ordered_value_type<1>; +using GraphInputName = jsonable_ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<1>; template OpenKwargDataflowSubgraphResult get_open_kwarg_dataflow_graph_subgraph( diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc index cc2fe65cc2..643e600ba5 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc @@ -1,11 +1,11 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/value_type.h" namespace FlexFlow { -using GraphInputName = ordered_value_type<0>; -using SlotName = ordered_value_type<1>; +using GraphInputName = jsonable_ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<1>; template std::string open_kwarg_dataflow_graph_as_dot( OpenKwargDataflowGraphView const &, diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.cc index 17ff20f99f..b3a1078910 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.cc @@ -1,10 +1,10 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/permute_open_kwarg_dataflow_graph_input_ids.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using GraphInputName = ordered_value_type<0>; -using SlotName = ordered_value_type<1>; +using GraphInputName = jsonable_ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<1>; template OpenKwargDataflowGraphView permute_open_kwarg_dataflow_graph_input_ids( diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.cc index 0d4bfebabd..88a2532017 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.cc @@ -1,10 +1,10 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using GraphInputName = ordered_value_type<0>; -using SlotName = ordered_value_type<1>; +using GraphInputName = jsonable_ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<1>; template std::optional> try_find_isomorphism_between_open_kwarg_dataflow_graphs( diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc index 491ff1b600..9309e893b7 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc @@ -1,10 +1,10 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using GraphInputName = ordered_value_type<0>; -using SlotName = ordered_value_type<1>; +using GraphInputName = jsonable_ordered_value_type<0>; +using SlotName = jsonable_ordered_value_type<1>; template std::pair>, bidict, Node>> diff --git a/lib/utils/src/utils/orthotope/dim_coord.cc b/lib/utils/src/utils/orthotope/dim_coord.cc index cf3c12712e..dbe5273929 100644 --- a/lib/utils/src/utils/orthotope/dim_coord.cc +++ b/lib/utils/src/utils/orthotope/dim_coord.cc @@ -1,9 +1,9 @@ #include "utils/orthotope/dim_coord.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using T = ordered_value_type<0>; +using T = jsonable_ordered_value_type<0>; template std::set get_coord_dims(DimCoord const &); diff --git a/lib/utils/src/utils/orthotope/dim_domain_mapping.cc b/lib/utils/src/utils/orthotope/dim_domain_mapping.cc index bd0f46e3dc..0fce76e6f3 100644 --- a/lib/utils/src/utils/orthotope/dim_domain_mapping.cc +++ b/lib/utils/src/utils/orthotope/dim_domain_mapping.cc @@ -1,9 +1,9 @@ #include "utils/orthotope/dim_domain_mapping.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" -using ::FlexFlow::ordered_value_type; -using L = ordered_value_type<0>; -using R = ordered_value_type<1>; +using ::FlexFlow::jsonable_ordered_value_type; +using L = jsonable_ordered_value_type<0>; +using R = jsonable_ordered_value_type<1>; namespace FlexFlow { @@ -32,9 +32,9 @@ template DimDomainMapping DimOrdering const &, DimOrdering const &); -using T1 = ordered_value_type<2>; -using T2 = ordered_value_type<3>; -using T3 = ordered_value_type<4>; +using T1 = jsonable_ordered_value_type<2>; +using T2 = jsonable_ordered_value_type<3>; +using T3 = jsonable_ordered_value_type<4>; template DimDomainMapping compose_dim_domain_mappings(DimDomainMapping const &, diff --git a/lib/utils/src/utils/orthotope/dim_projection.cc b/lib/utils/src/utils/orthotope/dim_projection.cc index 9c4c5e0c78..64f172ab10 100644 --- a/lib/utils/src/utils/orthotope/dim_projection.cc +++ b/lib/utils/src/utils/orthotope/dim_projection.cc @@ -1,11 +1,11 @@ #include "utils/orthotope/dim_projection.h" #include "utils/archetypes/value_type.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using L = ordered_value_type<0>; -using R = ordered_value_type<1>; +using L = jsonable_ordered_value_type<0>; +using R = jsonable_ordered_value_type<1>; template DimProjection dim_projection_identity_map(DimDomain const &, @@ -28,9 +28,9 @@ template DimCoord compute_dim_projection(DimProjection const &, DimOrdering const &, DimOrdering const &); -using T1 = ordered_value_type<2>; -using T2 = ordered_value_type<3>; -using T3 = ordered_value_type<4>; +using T1 = jsonable_ordered_value_type<2>; +using T2 = jsonable_ordered_value_type<3>; +using T3 = jsonable_ordered_value_type<4>; template DimProjection right_compose_eq_projection(DimProjection const &, diff --git a/lib/utils/src/utils/orthotope/down_projection.cc b/lib/utils/src/utils/orthotope/down_projection.cc index 271640605b..d92969165f 100644 --- a/lib/utils/src/utils/orthotope/down_projection.cc +++ b/lib/utils/src/utils/orthotope/down_projection.cc @@ -1,10 +1,10 @@ #include "utils/orthotope/down_projection.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { -using L = ordered_value_type<0>; -using R = ordered_value_type<1>; +using L = jsonable_ordered_value_type<0>; +using R = jsonable_ordered_value_type<1>; template DownProjection make_empty_down_projection(); @@ -26,9 +26,9 @@ template void project_dims(DownProjection &, template UpProjection invert_down_projection(DownProjection const &); -using T1 = ordered_value_type<2>; -using T2 = ordered_value_type<3>; -using T3 = ordered_value_type<4>; +using T1 = jsonable_ordered_value_type<2>; +using T2 = jsonable_ordered_value_type<3>; +using T3 = jsonable_ordered_value_type<4>; template DownProjection compose_down_projections(DownProjection const &, diff --git a/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc b/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc index c281d31d66..6809333df0 100644 --- a/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc +++ b/lib/utils/src/utils/orthotope/minimal_dim_domain_mapping.cc @@ -1,9 +1,9 @@ #include "utils/orthotope/minimal_dim_domain_mapping.h" -#include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" -using ::FlexFlow::ordered_value_type; -using L = ordered_value_type<0>; -using R = ordered_value_type<1>; +using ::FlexFlow::jsonable_ordered_value_type; +using L = jsonable_ordered_value_type<0>; +using R = jsonable_ordered_value_type<1>; namespace FlexFlow { @@ -40,9 +40,9 @@ template MinimalDimDomainMapping DimOrdering const &, DimOrdering const &); -using T1 = ordered_value_type<2>; -using T2 = ordered_value_type<3>; -using T3 = ordered_value_type<4>; +using T1 = jsonable_ordered_value_type<2>; +using T2 = jsonable_ordered_value_type<3>; +using T3 = jsonable_ordered_value_type<4>; template MinimalDimDomainMapping compose_minimal_dim_domain_mappings( MinimalDimDomainMapping const &, diff --git a/lib/utils/test/src/utils/bidict/algorithms/transform_keys.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_keys.cc similarity index 65% rename from lib/utils/test/src/utils/bidict/algorithms/transform_keys.cc rename to lib/utils/test/src/utils/bidict/algorithms/bidict_transform_keys.cc index 07154bc1a8..994db63aa6 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/transform_keys.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_keys.cc @@ -1,16 +1,16 @@ -#include "utils/bidict/algorithms/transform_keys.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("transform_keys(bidict, F)") { + TEST_CASE("bidict_transform_keys(bidict, F)") { bidict dict = { {1, "one"}, {2, "two"}, }; - bidict result = transform_keys(dict, [](int k) { + bidict result = bidict_transform_keys(dict, [](int k) { std::ostringstream oss; oss << k; return oss.str(); diff --git a/lib/utils/test/src/utils/bidict/algorithms/transform_values.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_values.cc similarity index 62% rename from lib/utils/test/src/utils/bidict/algorithms/transform_values.cc rename to lib/utils/test/src/utils/bidict/algorithms/bidict_transform_values.cc index 446f8ac31d..606823df87 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/transform_values.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_values.cc @@ -1,17 +1,17 @@ -#include "utils/bidict/algorithms/transform_values.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("transform_values(bidict, F)") { + TEST_CASE("bidict_transform_values(bidict, F)") { bidict dict = { {1, "one"}, {2, "two"}, }; bidict result = - transform_values(dict, [](std::string const &v) { return v + "a"; }); + bidict_transform_values(dict, [](std::string const &v) { return v + "a"; }); bidict correct = { {1, "onea"}, {2, "twoa"}, diff --git a/lib/utils/test/src/utils/binary_relation/binary_relation.cc b/lib/utils/test/src/utils/binary_relation/binary_relation.cc new file mode 100644 index 0000000000..26634dad31 --- /dev/null +++ b/lib/utils/test/src/utils/binary_relation/binary_relation.cc @@ -0,0 +1,125 @@ +#include +#include "utils/binary_relation/binary_relation.h" +#include "test/utils/doctest/fmt/set.h" +#include "test/utils/doctest/fmt/pair.h" +#include "test/utils/doctest/fmt/multiset.h" + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("BinaryRelation") { + SUBCASE("default constructor") { + BinaryRelation b; + + CHECK(b.empty()); + CHECK(b.size() == 0); + + std::set> raw = b.unwrap_as_set(); + std::set> correct_raw = {}; + + CHECK(raw == correct_raw); + } + + SUBCASE("initializer_list constuctor") { + BinaryRelation b = BinaryRelation{ + { + 2, + "even", + }, + { + 2, + "EVEN", + }, + { + 3, + "odd", + }, + { + 1, + "odd", + }, + }; + + CHECK(b.size() == 4); + CHECK(!b.empty()); + + std::set> raw = b.unwrap_as_set(); + std::set> correct_raw = { + {2, "even"}, + {2, "EVEN"}, + {3, "odd"}, + {1, "odd"}, + }; + + CHECK(raw == correct_raw); + } + + BinaryRelation empty_rel; + + BinaryRelation b = BinaryRelation{ + { + 2, + "even", + }, + { + 2, + "EVEN", + }, + { + 3, + "odd", + }, + { + 1, + "odd", + }, + }; + + SUBCASE("left_values") { + std::set result = b.left_values(); + std::set correct = {1, 2, 3}; + + CHECK(result == correct); + } + + SUBCASE("right_values") { + std::set result = b.right_values(); + std::set correct = {"odd", "even", "EVEN"}; + + CHECK(result == correct); + } + + SUBCASE("left_value_occurences") { + std::multiset result = b.left_value_occurences(); + std::multiset correct = {1, 2, 2, 3}; + + CHECK(result == correct); + } + + SUBCASE("right_value_occurences") { + std::multiset result = b.right_value_occurences(); + std::multiset correct = { + "odd", + "odd", + "even", + "EVEN", + }; + + CHECK(result == correct); + } + + SUBCASE("at_l") { + std::set result = b.at_l(2); + std::set correct = {"even", "EVEN"}; + + CHECK(result == correct); + } + + SUBCASE("at_r") { + std::set result = b.at_r("odd"); + std::set correct = {1, 3}; + + CHECK(result == correct); + } + } +} diff --git a/lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc b/lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc new file mode 100644 index 0000000000..5a470b60d8 --- /dev/null +++ b/lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc @@ -0,0 +1,47 @@ +#include +#include "utils/binary_relation/filter_binary_relation.h" + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("filter_binary_relation") { + BinaryRelation rel = BinaryRelation{ + { + 2, + "even", + }, + { + 2, + "EVEN", + }, + { + 3, + "odd", + }, + { + 1, + "odd", + }, + }; + + BinaryRelation result + = filter_binary_relation( + rel, + [](int l, std::string const &r) -> bool { + return l > 1 && r != "EVEN"; + }); + + BinaryRelation correct = BinaryRelation{ + { + 2, + "even", + }, + { + 3, + "odd", + }, + }; + + CHECK(result == correct); + } +} diff --git a/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc new file mode 100644 index 0000000000..e90f1eb633 --- /dev/null +++ b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc @@ -0,0 +1,52 @@ +#include +#include "utils/binary_relation/require_binary_relation_is_left_unique.h" + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("require_binary_relation_is_left_unique") { + SUBCASE("relation is left unique") { + BinaryRelation rel = BinaryRelation{ + { + 1, + "one", + }, + { + 2, + "two", + }, + { + 2, + "TWO", + }, + }; + + OneToMany result = require_binary_relation_is_left_unique(rel); + OneToMany correct = { + {1, {"one"}}, + {2, {"two", "TWO"}}, + }; + + CHECK(result == correct); + } + + SUBCASE("relation is not left unique") { + BinaryRelation rel = BinaryRelation{ + { + 1, + "odd", + }, + { + 2, + "even", + }, + { + 3, + "odd", + }, + }; + + CHECK_THROWS(require_binary_relation_is_left_unique(rel)); + } + } +} diff --git a/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc new file mode 100644 index 0000000000..f2951d257f --- /dev/null +++ b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc @@ -0,0 +1,52 @@ +#include +#include "utils/binary_relation/require_binary_relation_is_right_unique.h" + +using namespace ::FlexFlow; + +TEST_SUITE(FF_TEST_SUITE) { + TEST_CASE("require_binary_relation_is_right_unique") { + SUBCASE("relation is right unique") { + BinaryRelation rel = BinaryRelation{ + { + 1, + "odd", + }, + { + 2, + "even", + }, + { + 3, + "odd", + }, + }; + + ManyToOne result = require_binary_relation_is_right_unique(rel); + ManyToOne correct = { + {{1, 3}, "odd"}, + {{2}, "even"}, + }; + + CHECK(result == correct); + } + + SUBCASE("relation is not right unique") { + BinaryRelation rel = BinaryRelation{ + { + 1, + "one", + }, + { + 2, + "two", + }, + { + 2, + "TWO", + }, + }; + + CHECK_THROWS(require_binary_relation_is_right_unique(rel)); + } + } +} diff --git a/lib/utils/test/src/utils/containers/get_element_counts.cc b/lib/utils/test/src/utils/containers/get_element_counts.cc index e9bd4c55dd..0502e79dd2 100644 --- a/lib/utils/test/src/utils/containers/get_element_counts.cc +++ b/lib/utils/test/src/utils/containers/get_element_counts.cc @@ -7,8 +7,8 @@ using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("get_element_counts") { std::vector input = {1, 2, 3, 2, 3, 3, 2, 3}; - std::map result = get_element_counts(input); - std::map correct = {{1, 1}, {2, 3}, {3, 4}}; + std::map result = get_element_counts(input); + std::map correct = {{1, 1_p}, {2, 3_p}, {3, 4_p}}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/fmt/unordered_map.cc b/lib/utils/test/src/utils/fmt/unordered_map.cc index b82a13b6d2..7f43728df7 100644 --- a/lib/utils/test/src/utils/fmt/unordered_map.cc +++ b/lib/utils/test/src/utils/fmt/unordered_map.cc @@ -10,9 +10,9 @@ TEST_SUITE(FF_TEST_SUITE) { std::map input = {{0, 10}, {1, 1}, {3, 5}, {2, 8}}; std::string result = fmt::to_string(input); std::string correct = "{{0, 10}, {1, 1}, {2, 8}, {3, 5}}"; - std::map result_char_counts = + std::map result_char_counts = get_element_counts(result); - std::map correct_char_counts = + std::map correct_char_counts = get_element_counts(correct); CHECK(result_char_counts == correct_char_counts); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc index 89b24f6e95..3947aab336 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc @@ -63,14 +63,14 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector inv = get_topological_ordering(result.graph); SUBCASE("edges") { - std::map result_edges = + std::map result_edges = get_edge_counts(result.graph); - std::map correct_edges = { - {DirectedEdge{inv.at(0), inv.at(1)}, 1}, - {DirectedEdge{inv.at(1), inv.at(2)}, 1}, - {DirectedEdge{inv.at(1), inv.at(3)}, 1}, - {DirectedEdge{inv.at(2), inv.at(3)}, 1}, - {DirectedEdge{inv.at(3), inv.at(4)}, 1}, + std::map correct_edges = { + {DirectedEdge{inv.at(0), inv.at(1)}, 1_p}, + {DirectedEdge{inv.at(1), inv.at(2)}, 1_p}, + {DirectedEdge{inv.at(1), inv.at(3)}, 1_p}, + {DirectedEdge{inv.at(2), inv.at(3)}, 1_p}, + {DirectedEdge{inv.at(3), inv.at(4)}, 1_p}, }; CHECK(result_edges == correct_edges); } @@ -125,10 +125,10 @@ TEST_SUITE(FF_TEST_SUITE) { std::vector inv = get_topological_ordering(result.graph); SUBCASE("edges") { - std::map result_edges = + std::map result_edges = get_edge_counts(result.graph); - std::map correct_edges = { - {DirectedEdge{inv.at(0), inv.at(1)}, 2}, + std::map correct_edges = { + {DirectedEdge{inv.at(0), inv.at(1)}, 2_p}, }; CHECK(result_edges == correct_edges); } diff --git a/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc b/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc index a33fbf3df3..4242131d2e 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc @@ -127,9 +127,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("edge shape") { - std::map result_edges = get_edge_counts(g); - std::map correct_edges = { - {DirectedEdge{n.at(0), n.at(1)}, 1}, + std::map result_edges = get_edge_counts(g); + std::map correct_edges = { + {DirectedEdge{n.at(0), n.at(1)}, 1_p}, }; CHECK(result_edges == correct_edges); } @@ -154,7 +154,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(3), n.at(4)}, }); - std::map input_edge_counts = + std::map input_edge_counts = get_edge_counts(g); MultiDiEdge reduction_e1 = e.at(3); @@ -171,11 +171,15 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("edge shape") { - std::map result_edges = get_edge_counts(g); - std::map correct_edges = [&] { - std::map new_edge_counts = + std::map result_edges = get_edge_counts(g); + std::map correct_edges = [&] { + std::map new_edge_counts = input_edge_counts; - new_edge_counts.at(get_directed_edge(g, reduction_e1))--; + + DirectedEdge e = get_directed_edge(g, reduction_e1); + new_edge_counts.at(e) = positive_int{ + new_edge_counts.at(e).int_from_positive_int() - 1, + }; return new_edge_counts; }(); CHECK(result_edges == correct_edges); From 571119ba8a4a29b2b9f10eeca63591312037ed81 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 26 Jun 2026 18:13:05 -0700 Subject: [PATCH 31/35] Format --- bin/run-model/src/run-model/main.cc | 82 +- .../sp-ization-benchmarking/distributions.h | 12 +- .../abstracted_single_tensor_movement.h | 14 +- .../abstracted_tensor_set_movement.h | 12 +- .../machine_mapping_constraints.h | 7 +- .../machine_mapping_problem_tree.h | 3 +- .../machine_mapping/machine_mapping_result.h | 4 +- .../pareto_optimal_machine_mapping.h | 5 +- .../pcg/pcg_binary_sp_decomposition.h | 5 +- ...l_tensor_space_to_machine_space_mapping.cc | 8 +- .../abstracted_single_tensor_movement.cc | 19 +- .../abstracted_tensor_set_movement.cc | 21 +- ...racted_tensor_set_movement_across_split.cc | 17 +- .../machine_mapping/allowed_machine_views.cc | 17 +- ...substitution_and_update_machine_mapping.cc | 16 +- .../get_optimal_machine_mapping.cc | 42 +- .../machine_mapping/machine_mapping.cc | 12 +- .../machine_mapping_constraints.cc | 3 +- .../machine_mapping_mutation_set.cc | 5 +- .../get_machine_mapping_problem_tree.cc | 13 +- .../compiler/machine_mapping/machine_view.cc | 8 +- ...get_optimal_machine_mapping_with_memory.cc | 32 +- .../machine_mapping_with_memory_result.cc | 2 +- .../pareto_optimal_machine_mapping.cc | 5 +- ...el_layer_guid_oblivious_machine_mapping.cc | 26 +- .../machine_mapping/transitive_reduced_pcg.cc | 8 +- .../task_graph_simulator/pcg_task_graph.cc | 2 +- .../simulate_task_graph_execution.cc | 9 +- .../task_graph_simulator/task_simulator.cc | 13 +- .../unity_algorithm/graph_optimize_state.cc | 21 +- .../machine_mapping/allowed_machine_views.cc | 8 +- .../get_optimal_machine_mapping.cc | 10 +- .../start_invariant_machine_view.cc | 16 +- .../simulate_task_graph_execution.cc | 64 +- .../task_graph_simulator/task_simulator.cc | 2 +- .../src/internal/cost_estimator_for_test.cc | 3 +- .../runtime_only_cost_estimator_for_test.cc | 7 +- .../runtime_only_cost_estimator_for_test.h | 4 +- lib/kernels/src/kernels/accessor.cc | 24 +- .../computation_graph_instance.h | 5 +- .../local_task_argument_accessor.h | 3 +- .../local-execution/tensor_allocation.h | 3 +- .../computation_graph_instance.cc | 12 +- .../cost_estimator/local_cost_estimator.cc | 10 +- .../src/local-execution/task_execution.cc | 4 +- .../src/local-execution/tensor_allocation.cc | 6 +- .../local_task_argument_accessor.cc | 4 +- .../op-attrs/get_operator_task_space.h | 3 +- .../include/op-attrs/initializer_attrs.h | 11 +- lib/op-attrs/include/op-attrs/ops/attention.h | 3 +- .../include/op-attrs/ops/batch_norm.h | 6 +- .../include/op-attrs/ops/layer_norm.h | 3 +- lib/op-attrs/include/op-attrs/ops/linear.h | 3 +- .../include/op-attrs/parallel_tensor_shape.h | 3 +- .../op-attrs/replica_parallel_dim_set.h | 3 +- .../include/op-attrs/shape_inference.h | 6 +- .../ff_ordered/ff_ordered_from_map.cc | 3 +- .../ff_ordered/map_from_ff_ordered.cc | 3 +- .../src/op-attrs/get_incoming_tensor_roles.cc | 9 +- ...space_to_parallel_tensor_space_mappings.cc | 329 ++-- .../src/op-attrs/get_operator_task_space.cc | 3 +- .../src/op-attrs/initializer_attrs.cc | 11 +- .../src/op-attrs/operator_task_space.cc | 2 +- ...sk_space_to_operator_task_space_mapping.cc | 7 +- lib/op-attrs/src/op-attrs/ops/attention.cc | 3 +- lib/op-attrs/src/op-attrs/ops/batch_norm.cc | 6 +- lib/op-attrs/src/op-attrs/ops/layer_norm.cc | 3 +- lib/op-attrs/src/op-attrs/ops/linear.cc | 3 +- lib/op-attrs/src/op-attrs/ops/transpose.cc | 3 +- .../op-attrs/parallel_tensor_dim_degrees.cc | 30 +- .../src/op-attrs/parallel_tensor_dims.cc | 3 +- .../src/op-attrs/parallel_tensor_shape.cc | 3 +- .../parallel_tensor_space_coordinate.cc | 9 +- .../src/op-attrs/replica_parallel_dim_set.cc | 3 +- lib/op-attrs/src/op-attrs/shape_inference.cc | 629 +++---- .../src/op-attrs/tensor_dim_permutation.cc | 3 +- lib/op-attrs/src/op-attrs/tensor_dims.cc | 2 +- .../test/src/op-attrs/operator_task_space.cc | 9 +- .../op-attrs/parallel_tensor_dim_degrees.cc | 1 - lib/op-attrs/test/src/op-attrs/tensor_dims.cc | 6 +- lib/pcg/include/pcg/computation_graph.h | 9 +- .../include/pcg/computation_graph_builder.h | 4 +- .../v1/graphs/v1_kwarg_dataflow_graph.h | 5 +- .../graphs/v1_labelled_kwarg_dataflow_graph.h | 6 +- .../mapped_operator_task_group.h | 4 +- .../mapped_parallel_computation_graph.h | 2 +- .../mapped_parallel_layer_invocation_info.h | 7 +- .../parallel_computation_graph.h | 17 +- .../parallel_computation_graph_builder.h | 3 +- lib/pcg/src/pcg/computation_graph.cc | 55 +- lib/pcg/src/pcg/computation_graph_builder.cc | 62 +- .../mapped_operator_task_group.cc | 49 +- .../mapped_parallel_computation_graph.cc | 50 +- .../mapped_parallel_layer_invocation_info.cc | 23 +- lib/pcg/src/pcg/optimizer_attrs.cc | 3 +- .../parallel_computation_graph.cc | 70 +- .../parallel_computation_graph_builder.cc | 33 +- .../mapped_parallel_computation_graph.cc | 98 +- .../parallel_computation_graph.cc | 3 +- .../parallel_computation_graph_builder.cc | 9 +- .../include/realm-execution/dependency_set.h | 3 +- .../realm-execution/distributed_ff_handle.h | 6 +- .../realm-execution/instance_allocation.h | 3 +- .../include/realm-execution/pcg_instance.h | 3 +- .../include/realm-execution/realm_context.h | 2 +- .../realm-execution/distributed_ff_handle.cc | 6 +- ...uted_per_device_op_state_initialization.cc | 13 +- .../realm-execution/instance_allocation.cc | 3 +- .../src/realm-execution/pcg_instance.cc | 15 +- .../src/realm-execution/realm_context.cc | 17 +- .../serializable_tensor_instance_backing.cc | 30 +- .../test/src/realm-execution/test_e2e.cc | 6 +- .../src/realm-execution/test_op_replicate.cc | 6 +- .../perform_shape_inference.h | 4 +- .../sub_parallel_computation_graph.h | 12 +- .../substitutions/substitution_builder.h | 13 +- .../unlabelled/unlabelled_graph_pattern.h | 12 +- ...elled_kwarg_dataflow_graph_pattern_match.h | 3 +- .../apply_substitution/apply_substitution.cc | 51 +- .../evaluate_substitution_output.cc | 35 +- .../perform_shape_inference.cc | 22 +- .../materialize_operator_from_attrs_map.cc | 6 +- .../output_graph/output_graph_expr.cc | 6 +- .../output_operator_attrs_assignment.cc | 25 +- .../src/substitutions/pcg_pattern.cc | 6 +- .../src/substitutions/pcg_pattern_match.cc | 8 +- .../sub_parallel_computation_graph.cc | 12 +- .../src/substitutions/substitution.cc | 6 +- .../src/substitutions/substitution_builder.cc | 12 +- .../substitutions/unity_substitution_set.cc | 70 +- .../unlabelled/find_pattern_matches.cc | 9 +- .../unlabelled/pattern_matching.cc | 11 +- .../unlabelled/unlabelled_graph_pattern.cc | 12 +- ...lled_kwarg_dataflow_graph_pattern_match.cc | 13 +- .../evaluate_substitution_output.cc | 3 +- .../perform_shape_inference.cc | 7 +- .../test/src/substitutions/pcg_pattern.cc | 8 +- .../substitutions/unity_substitution_set.cc | 4 +- .../unlabelled/find_pattern_matches.cc | 21 +- .../unlabelled/pattern_matching.cc | 21 +- .../task-spec/dynamic_graph/copy_insertion.h | 21 +- .../dynamic_graph/dynamic_node_invocation.h | 31 +- .../dynamic_graph/dynamic_node_mapping.h | 4 +- .../dynamic_open_dataflow_graph.h | 46 +- .../dynamic_graph/dynamic_value_attrs.h | 6 +- ...amic_open_dataflow_graph_from_mapped_pcg.h | 8 +- .../dynamic_graph/parallel_tensor_mapping.h | 16 +- .../task-spec/dynamic_graph/pass_expansion.h | 21 +- ...serializable_dynamic_open_dataflow_graph.h | 5 +- .../task-spec/dynamic_graph/shard_expansion.h | 26 +- .../dynamic_graph/update_insertion.h | 6 +- .../task-spec/dynamic_graph/copy_insertion.cc | 233 +-- .../dynamic_graph/dynamic_node_invocation.cc | 53 +- .../dynamic_graph/dynamic_node_mapping.cc | 25 +- .../dynamic_open_dataflow_graph.cc | 231 ++- .../dynamic_graph/dynamic_tensor_slot.cc | 1 - .../dynamic_graph/dynamic_value_attrs.cc | 7 +- .../task-spec/dynamic_graph/loss_insertion.cc | 36 +- .../dynamic_graph/machine_slicing.cc | 7 +- ...ake_dynamic_open_dataflow_graph_from_cg.cc | 90 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 69 +- .../dynamic_graph/parallel_tensor_mapping.cc | 16 +- .../task-spec/dynamic_graph/pass_expansion.cc | 221 ++- ...erializable_dynamic_open_dataflow_graph.cc | 14 +- .../dynamic_graph/shard_expansion.cc | 569 +++--- .../dynamic_graph/update_insertion.cc | 7 +- .../task-spec/dynamic_graph/copy_insertion.cc | 1617 +++++++++-------- .../dynamic_open_dataflow_graph.cc | 36 +- ...mic_open_dataflow_graph_from_mapped_pcg.cc | 1527 ++++++++-------- .../task-spec/dynamic_graph/pass_expansion.cc | 273 +-- .../dynamic_graph/shard_expansion.cc | 986 +++++----- .../archetypes/jsonable_ordered_value_type.h | 9 +- .../bidict_from_unstructured_relation.h | 44 +- .../bidict_transform_keys_and_values.h | 3 +- .../algorithms/bidict_unordered_set_of.h | 3 +- lib/utils/include/utils/bidict/bidict.h | 36 +- .../utils/binary_relation/binary_relation.h | 88 +- .../binary_relation_from_map.h | 4 +- .../binary_relation_transform_left.h | 16 +- .../binary_relation_transform_left2.h | 16 +- .../binary_relation_transform_right.h | 16 +- .../binary_relation_transform_right2.h | 16 +- .../binary_relation/filter_binary_relation.h | 11 +- .../require_binary_relation_is_left_unique.h | 3 +- .../require_binary_relation_is_right_unique.h | 3 +- .../transform_binary_relation.h | 18 +- .../include/utils/containers/are_disjoint.h | 3 +- .../containers/binary_cartesian_product.h | 5 +- .../containers/binary_merge_disjoint_maps.h | 7 +- .../binary_merge_disjoint_unordered_maps.h | 6 +- .../utils/containers/binary_merge_maps_with.h | 7 +- .../binary_merge_maps_with_left_dominating.h | 5 +- .../binary_merge_maps_with_right_dominating.h | 5 +- .../binary_merge_unordered_maps_with.h | 4 +- .../include/utils/containers/filter_values.h | 3 +- lib/utils/include/utils/containers/filtrans.h | 3 +- lib/utils/include/utils/containers/find.h | 3 +- lib/utils/include/utils/containers/flatmap.h | 9 +- .../include/utils/containers/generate_map.h | 3 +- .../utils/containers/get_all_assignments.h | 26 +- .../utils/containers/get_element_counts.h | 6 +- lib/utils/include/utils/containers/get_only.h | 2 +- lib/utils/include/utils/containers/group_by.h | 3 +- .../include/utils/containers/invert_map.h | 5 +- .../utils/containers/invert_unordered_map.h | 1 - .../include/utils/containers/is_submapeq_of.h | 3 +- .../utils/containers/is_superseteq_of.h | 3 +- .../containers/lift_optional_through_map.h | 6 +- .../include/utils/containers/lookup_in_map.h | 6 +- .../containers/map_from_keys_and_values.h | 7 +- lib/utils/include/utils/containers/map_keys.h | 7 +- .../include/utils/containers/map_keys2.h | 3 +- .../utils/containers/map_keys_and_values.h | 6 +- .../containers/map_keys_with_value_merging.h | 7 +- .../include/utils/containers/map_values.h | 2 +- .../include/utils/containers/map_values2.h | 6 +- .../utils/containers/merge_disjoint_maps.h | 3 +- .../merge_disjoint_unordered_maps.h | 2 +- .../include/utils/containers/merge_in_map.h | 3 +- .../utils/containers/merge_maps_with.h | 8 +- .../containers/merge_unordered_maps_with.h | 5 +- ...rge_unordered_maps_with_right_dominating.h | 3 +- lib/utils/include/utils/containers/minimum.h | 4 +- .../utils/containers/require_only_key.h | 2 +- .../include/utils/containers/require_same.h | 2 +- .../utils/containers/require_two_keys.h | 5 +- .../include/utils/containers/restrict_keys.h | 7 +- .../include/utils/containers/set_difference.h | 3 +- lib/utils/include/utils/containers/set_of.h | 5 +- .../include/utils/containers/transform.h | 6 +- .../containers/try_merge_nondisjoint_maps.h | 2 +- .../utils/containers/unordered_items.h | 2 +- .../unordered_map_from_keys_and_values.h | 2 +- .../utils/containers/vector_from_idx_map.h | 2 +- .../utils/containers/without_nullopts.h | 3 +- .../utils/containers/zip_values_strict.h | 26 +- .../utils/containers/zip_values_strict_with.h | 7 +- .../utils/deduplicated_priority_queue.h | 1 - lib/utils/include/utils/disjoint_set.h | 1 - lib/utils/include/utils/dot/dot_file.h | 3 +- .../full_binary_tree/find_paths_to_leaf.h | 48 +- .../full_binary_tree/get_all_leaf_paths.h | 37 +- .../utils/full_binary_tree/get_leaves.h | 18 +- .../full_binary_tree/get_path_to_leaf_map.h | 44 +- lib/utils/include/utils/graph/algorithms.h | 31 +- .../utils/graph/dataflow_graph/algorithms.h | 3 +- .../algorithms/find_isomorphisms.h | 4 +- .../algorithms/get_incoming_edges.h | 5 +- .../algorithms/get_outgoing_edges.h | 7 +- .../algorithms/get_subgraph_incoming_edges.h | 5 +- .../algorithms/get_subgraph_outgoing_edges.h | 5 +- .../graph/dataflow_graph/dataflow_graph.h | 3 +- .../dataflow_graph/dataflow_graph_view.h | 3 +- .../graph/digraph/algorithms/contract_node.h | 3 +- .../utils/graph/digraph/algorithms/flipped.h | 3 +- .../digraph/algorithms/get_descendants.h | 3 +- .../graph/digraph/algorithms/get_dominators.h | 3 +- .../digraph/algorithms/get_dominators_map.h | 3 +- .../get_edges_from_subgraph_to_subgraph.h | 6 +- .../algorithms/get_imm_dominators_map.h | 3 +- .../digraph/algorithms/get_incoming_edges.h | 3 +- .../digraph/algorithms/get_outgoing_edges.h | 3 +- .../algorithms/get_post_dominators_map.h | 3 +- .../digraph/algorithms/get_predecessors.h | 7 +- .../algorithms/get_strict_dominators.h | 3 +- .../algorithms/get_strict_dominators_map.h | 3 +- .../algorithms/get_subgraph_outgoing_edges.h | 5 +- .../algorithms/get_subgraph_successors.h | 5 +- .../graph/digraph/algorithms/get_successors.h | 7 +- .../get_weakly_connected_components.h | 3 +- .../digraph/algorithms/transitive_reduction.h | 3 +- .../utils/graph/instances/adjacency_digraph.h | 8 +- .../graph/instances/adjacency_multidigraph.h | 11 +- .../instances/hashmap_undirected_graph.h | 3 +- .../instances/unordered_set_dataflow_graph.h | 13 +- .../unordered_set_kwarg_dataflow_graph.h | 32 +- ...ordered_set_labelled_open_dataflow_graph.h | 6 +- ...d_set_labelled_open_kwarg_dataflow_graph.h | 60 +- .../unordered_set_open_kwarg_dataflow_graph.h | 31 +- ...raph_data_from_kwarg_dataflow_graph_data.h | 61 +- ...dataflow_graph_from_kwarg_dataflow_graph.h | 4 +- .../get_all_kwarg_dataflow_outputs.h | 5 +- .../get_kwarg_dataflow_graph_subgraph.h | 9 +- .../algorithms/kwarg_dataflow_graph_as_dot.h | 6 +- .../algorithms/kwarg_dataflow_graph_data.h | 3 +- ...educed_kwarg_dataflow_edges_across_split.h | 13 +- .../view_from_kwarg_dataflow_graph_data.h | 7 +- .../i_kwarg_dataflow_graph.h | 9 +- .../kwarg_dataflow_graph.h | 9 +- ...lled_kwarg_dataflow_graph_node_label_map.h | 8 +- ...ed_kwarg_dataflow_graph_output_label_map.h | 2 +- ...t_labelled_kwarg_dataflow_graph_subgraph.h | 5 +- ...kwarg_dataflow_graph_view_with_labelling.h | 3 +- ...abelled_kwarg_dataflow_graph_view_as_dot.h | 4 +- .../i_labelled_kwarg_dataflow_graph.h | 8 +- .../labelled_kwarg_dataflow_graph.h | 8 +- .../from_labelled_open_dataflow_graph_data.h | 3 +- .../algorithms/is_isomorphic_under.h | 11 +- ..._labelled_open_kwarg_dataflow_graph_data.h | 4 +- ...ed_open_kwarg_dataflow_graph_view_as_dot.h | 4 +- ...arg_dataflow_graphs_are_isomorphic_under.h | 6 +- ...kwarg_dataflow_graph_view_with_labelling.h | 7 +- ...lled_open_kwarg_dataflow_graph_input_ids.h | 3 +- ...elled_open_kwarg_dataflow_graph_node_ids.h | 3 +- ...abelled_open_kwarg_dataflow_graph_labels.h | 3 +- .../i_labelled_open_kwarg_dataflow_graph.h | 3 +- .../labelled_open_kwarg_dataflow_graph.h | 3 +- .../algorithms/get_incoming_edges.h | 5 +- .../algorithms/get_outgoing_edges.h | 5 +- .../include/utils/graph/node/node_query.h | 3 +- .../algorithms/get_incoming_edges.h | 6 +- .../algorithms/get_subgraph.h | 3 +- .../algorithms/get_subgraph_inputs.h | 5 +- .../open_dataflow_edge_query.h | 6 +- .../open_dataflow_graph_view.h | 3 +- ...hisms_between_open_kwarg_dataflow_graphs.h | 27 +- .../get_open_kwarg_dataflow_graph_data.h | 2 +- .../get_open_kwarg_dataflow_graph_subgraph.h | 8 +- .../get_open_kwarg_dataflow_value_uses.h | 7 +- .../open_kwarg_dataflow_graph_as_dot.h | 12 +- .../open_kwarg_dataflow_graph_data.h | 4 +- ...phism_between_open_kwarg_dataflow_graphs.h | 5 +- ...g_dataflow_graph_by_materializing_inputs.h | 17 +- ...view_from_open_kwarg_dataflow_graph_data.h | 8 +- .../i_open_kwarg_dataflow_graph.h | 3 +- .../i_open_kwarg_dataflow_graph_view.h | 9 +- .../open_kwarg_dataflow_graph.h | 3 +- .../open_kwarg_dataflow_graph_view.h | 8 +- lib/utils/include/utils/graph/query_set.h | 5 +- lib/utils/include/utils/graph/render_dot.h | 12 +- .../series_parallel/digraph_generation.h | 6 +- .../graph/series_parallel/get_ancestors.h | 2 +- .../series_parallel_decomposition.h | 3 +- .../series_parallel/series_parallel_metrics.h | 3 +- .../sp_ization/escribano_algo.h | 12 +- .../series_parallel/sp_ization/node_role.h | 3 +- .../sp_ization/up_down_partition.h | 4 +- lib/utils/include/utils/graph/traversal.h | 15 +- .../algorithms/get_connected_components.h | 3 +- .../algorithms/get_neighboring_nodes.h | 3 +- .../graph/undirected/i_undirected_graph.h | 3 +- .../undirected/i_undirected_graph_view.h | 3 +- lib/utils/include/utils/graph/views/views.h | 18 +- .../include/utils/json/check_is_jsonable.h | 4 +- .../include/utils/many_to_one/many_to_one.h | 25 +- .../include/utils/nonempty_set/nonempty_set.h | 18 +- .../include/utils/one_to_many/one_to_many.h | 23 +- .../one_to_many/one_to_many_filter_values.h | 1 - .../one_to_many_from_l_to_r_mapping.h | 4 +- .../one_to_many_transform_values.h | 11 +- .../require_one_to_many_is_bijection.h | 11 +- lib/utils/include/utils/ord/unordered_map.h | 2 +- lib/utils/include/utils/orthotope/dim_coord.h | 35 +- .../include/utils/orthotope/dim_domain.h | 2 +- .../include/utils/orthotope/dim_projection.h | 15 +- .../include/utils/orthotope/down_projection.h | 3 +- .../include/utils/orthotope/eq_projection.h | 6 +- .../utils/orthotope/minimal_dim_domain.h | 10 +- .../orthotope/minimal_dim_domain_mapping.h | 45 +- lib/utils/include/utils/orthotope/orthotope.h | 3 +- .../include/utils/orthotope/up_projection.h | 14 +- .../algorithms/bidict_from_keys_and_values.cc | 6 +- .../bidict/algorithms/bidict_from_pairs.cc | 3 +- .../bidict_from_unstructured_relation.cc | 4 +- .../bidict_transform_keys_and_values.cc | 3 +- .../algorithms/bidict_unordered_set_of.cc | 3 +- lib/utils/src/utils/bidict/bidict.cc | 2 +- .../binary_relation_transform_left.cc | 4 +- .../binary_relation_transform_left2.cc | 4 +- .../binary_relation_transform_right.cc | 4 +- .../binary_relation_transform_right2.cc | 4 +- .../binary_relation/filter_binary_relation.cc | 3 +- .../require_binary_relation_is_left_unique.cc | 5 +- ...require_binary_relation_is_right_unique.cc | 4 +- .../src/utils/containers/are_disjoint.cc | 5 +- lib/utils/src/utils/containers/argmax.cc | 1 - lib/utils/src/utils/containers/argmin.cc | 1 - lib/utils/src/utils/containers/at_idx.cc | 2 +- .../containers/binary_cartesian_product.cc | 3 +- .../containers/binary_merge_disjoint_maps.cc | 7 +- .../containers/binary_merge_maps_with.cc | 7 +- .../binary_merge_maps_with_left_dominating.cc | 2 +- ...binary_merge_maps_with_right_dominating.cc | 2 +- ...rge_unordered_maps_with_left_dominating.cc | 6 +- ...ge_unordered_maps_with_right_dominating.cc | 10 +- .../utils/containers/contains_duplicates.cc | 2 +- .../src/utils/containers/contains_value.cc | 3 +- lib/utils/src/utils/containers/filter.cc | 35 +- lib/utils/src/utils/containers/filtrans.cc | 5 +- lib/utils/src/utils/containers/flatmap.cc | 4 +- .../containers/generate_unordered_map.cc | 5 +- .../utils/containers/get_all_assignments.cc | 5 +- .../utils/containers/get_element_counts.cc | 8 +- lib/utils/src/utils/containers/get_only.cc | 2 +- lib/utils/src/utils/containers/invert_map.cc | 3 +- .../utils/containers/invert_unordered_map.cc | 1 - .../src/utils/containers/is_submapeq_of.cc | 1 - .../src/utils/containers/is_subseteq_of.cc | 3 +- .../src/utils/containers/lookup_in_map.cc | 5 +- .../containers/map_from_keys_and_values.cc | 6 +- .../src/utils/containers/map_from_pairs.cc | 6 +- .../utils/containers/map_from_unordered.cc | 3 +- lib/utils/src/utils/containers/map_keys.cc | 12 +- lib/utils/src/utils/containers/map_keys2.cc | 5 +- .../utils/containers/map_keys_and_values.cc | 6 +- .../containers/map_keys_with_value_merging.cc | 4 +- lib/utils/src/utils/containers/map_values.cc | 2 +- lib/utils/src/utils/containers/map_values2.cc | 14 +- .../utils/containers/merge_disjoint_maps.cc | 2 +- .../merge_disjoint_unordered_maps.cc | 5 +- .../src/utils/containers/merge_in_map.cc | 3 +- .../src/utils/containers/merge_maps_with.cc | 6 +- .../merge_maps_with_right_dominating.cc | 2 +- .../containers/merge_unordered_maps_with.cc | 6 +- ...ge_unordered_maps_with_right_dominating.cc | 3 +- .../src/utils/containers/multiset_union.cc | 19 +- .../src/utils/containers/require_all_of.cc | 1 - .../src/utils/containers/require_only_key.cc | 2 +- .../src/utils/containers/require_two_keys.cc | 2 +- .../src/utils/containers/restrict_keys.cc | 11 +- lib/utils/src/utils/containers/set_of.cc | 1 - lib/utils/src/utils/containers/set_union.cc | 6 +- lib/utils/src/utils/containers/transform.cc | 21 +- .../src/utils/containers/transform_pairs.cc | 6 +- .../containers/try_merge_nondisjoint_maps.cc | 6 +- .../try_merge_nondisjoint_unordered_maps.cc | 7 +- .../src/utils/containers/unordered_items.cc | 3 +- .../src/utils/containers/unordered_keys.cc | 2 +- .../unordered_map_from_keys_and_values.cc | 3 +- .../containers/unordered_map_from_map.cc | 3 +- ...unstructured_exhaustive_relational_join.cc | 5 +- lib/utils/src/utils/containers/vector_of.cc | 2 +- .../src/utils/containers/without_nullopts.cc | 3 +- .../src/utils/containers/zip_values_strict.cc | 7 +- .../containers/zip_values_strict_with.cc | 7 +- .../src/utils/full_binary_tree/get_leaves.cc | 2 +- lib/utils/src/utils/graph/algorithms.cc | 33 +- .../utils/graph/dataflow_graph/algorithms.cc | 3 +- .../algorithms/find_isomorphisms.cc | 5 +- .../algorithms/get_incoming_edges.cc | 5 +- .../algorithms/get_outgoing_edges.cc | 7 +- .../algorithms/get_subgraph_incoming_edges.cc | 5 +- .../algorithms/get_subgraph_outgoing_edges.cc | 5 +- ...t_transitive_reduced_edges_across_split.cc | 11 +- .../algorithms/view_as_open_dataflow_graph.cc | 10 +- .../dataflow_graph/dataflow_graph_view.cc | 3 +- .../dataflow_graph/i_dataflow_graph_view.cc | 3 +- .../get_cbc_decomposition.cc | 9 +- .../graph/digraph/algorithms/contract_node.cc | 3 +- .../utils/graph/digraph/algorithms/flipped.cc | 3 +- .../graph/digraph/algorithms/get_ancestors.cc | 3 +- .../digraph/algorithms/get_descendants.cc | 5 +- .../digraph/algorithms/get_dominators.cc | 3 +- .../digraph/algorithms/get_dominators_map.cc | 3 +- .../get_edges_from_subgraph_to_subgraph.cc | 8 +- .../algorithms/get_imm_dominators_map.cc | 7 +- .../algorithms/get_imm_post_dominator.cc | 5 +- .../digraph/algorithms/get_incoming_edges.cc | 31 +- .../algorithms/get_lowest_common_ancestors.cc | 10 +- .../digraph/algorithms/get_outgoing_edges.cc | 30 +- .../digraph/algorithms/get_post_dominators.cc | 3 +- .../algorithms/get_post_dominators_map.cc | 3 +- .../digraph/algorithms/get_predecessors.cc | 16 +- .../algorithms/get_strict_dominators.cc | 3 +- .../algorithms/get_strict_dominators_map.cc | 3 +- .../algorithms/get_subgraph_outgoing_edges.cc | 8 +- .../algorithms/get_subgraph_successors.cc | 5 +- .../digraph/algorithms/get_successors.cc | 7 +- .../get_weakly_connected_components.cc | 3 +- .../digraph/algorithms/transitive_closure.cc | 6 +- .../algorithms/transitive_reduction.cc | 9 +- lib/utils/src/utils/graph/digraph/digraph.cc | 3 +- .../src/utils/graph/digraph/digraph_view.cc | 3 +- .../graph/instances/adjacency_digraph.cc | 3 +- .../graph/instances/adjacency_multidigraph.cc | 25 +- .../instances/unordered_set_dataflow_graph.cc | 5 +- ...aph_data_from_kwarg_dataflow_graph_data.cc | 3 +- ...ataflow_graph_from_kwarg_dataflow_graph.cc | 3 +- ..._kwarg_dataflow_subgraph_incoming_edges.cc | 3 +- ..._kwarg_dataflow_subgraph_outgoing_edges.cc | 3 +- .../algorithms/kwarg_dataflow_graph_as_dot.cc | 3 +- ...belled_kwarg_dataflow_graph_view_as_dot.cc | 3 +- ...d_open_kwarg_dataflow_graph_view_as_dot.cc | 3 +- ...warg_dataflow_graph_view_with_labelling.cc | 2 +- .../algorithms/get_incoming_edges.cc | 8 +- .../algorithms/get_outgoing_edges.cc | 8 +- .../graph/multidigraph/multidigraph_view.cc | 3 +- lib/utils/src/utils/graph/node/node_query.cc | 2 +- .../algorithms/find_isomorphisms.cc | 6 +- .../from_open_dataflow_graph_data.cc | 3 +- .../algorithms/get_incoming_edges.cc | 12 +- .../algorithms/get_subgraph.cc | 8 +- .../algorithms/is_isomorphic_under.cc | 11 +- .../i_open_dataflow_graph_view.cc | 3 +- .../open_dataflow_edge_query.cc | 6 +- .../open_dataflow_graph_view.cc | 3 +- .../unordered_set_open_dataflow_graph.cc | 3 +- .../get_all_open_kwarg_dataflow_edges.cc | 7 +- ...ming_open_kwarg_dataflow_edges_for_node.cc | 3 +- ...ing_open_kwarg_dataflow_values_for_node.cc | 3 +- .../open_kwarg_dataflow_graph_as_dot.cc | 3 +- lib/utils/src/utils/graph/render_dot.cc | 6 +- .../balanced_binary_sp_tree_from_nary.cc | 6 +- .../binary_sp_decomposition_tree.cc | 3 +- .../get_leaves.cc | 2 +- .../series_parallel/digraph_generation.cc | 6 +- .../graph/series_parallel/get_ancestors.cc | 10 +- .../get_series_parallel_decomposition.cc | 16 +- .../non_normal_sp_decomposition.cc | 20 +- .../normalize_sp_decomposition.cc | 2 +- .../series_parallel/parallel_reduction.cc | 18 +- .../series_parallel_decomposition.cc | 17 +- .../series_parallel_metrics.cc | 19 +- .../graph/series_parallel/series_reduction.cc | 2 +- .../sp_ization/escribano_algo.cc | 55 +- .../sp_ization/flexible_algo.cc | 42 +- .../sp_ization/naive_stratum_sync.cc | 32 +- .../series_parallel/sp_ization/node_role.cc | 3 +- .../sp_ization/up_down_partition.cc | 4 +- .../sp_ization/work_duplicating_sp_ization.cc | 7 +- lib/utils/src/utils/graph/traversal.cc | 18 +- .../algorithms/get_connected_components.cc | 3 +- .../algorithms/get_neighboring_nodes.cc | 2 +- .../graph/undirected/undirected_graph.cc | 3 +- .../graph/undirected/undirected_graph_view.cc | 3 +- lib/utils/src/utils/graph/views/views.cc | 19 +- .../src/utils/many_to_one/many_to_one.cc | 9 +- .../src/utils/nonempty_set/nonempty_set.cc | 2 +- .../src/utils/one_to_many/one_to_many.cc | 7 +- .../one_to_many/one_to_many_filter_values.cc | 3 +- .../one_to_many_from_l_to_r_mapping.cc | 4 +- .../require_one_to_many_is_bijection.cc | 3 +- lib/utils/src/utils/orthotope/dim_coord.cc | 13 +- .../src/utils/orthotope/dim_projection.cc | 8 +- .../src/utils/orthotope/down_projection.cc | 5 +- .../src/utils/orthotope/eq_projection.cc | 6 +- .../src/utils/orthotope/minimal_dim_domain.cc | 9 +- lib/utils/src/utils/orthotope/orthotope.cc | 2 +- .../src/utils/orthotope/up_projection.cc | 11 +- .../utils/doctest/check_without_stringify.h | 4 +- .../bidict/algorithms/bidict_filter_values.cc | 4 +- .../algorithms/bidict_filtrans_values.cc | 4 +- .../algorithms/bidict_from_enumerating.cc | 6 +- .../algorithms/bidict_transform_keys.cc | 11 +- .../algorithms/bidict_transform_values.cc | 4 +- .../utils/binary_relation/binary_relation.cc | 86 +- .../binary_relation/filter_binary_relation.cc | 60 +- .../require_binary_relation_is_left_unique.cc | 57 +- ...require_binary_relation_is_right_unique.cc | 57 +- .../binary_merge_disjoint_unordered_maps.cc | 2 +- .../binary_merge_unordered_maps_with.cc | 2 +- ...rge_unordered_maps_with_left_dominating.cc | 2 +- ...ge_unordered_maps_with_right_dominating.cc | 2 +- .../src/utils/containers/cartesian_product.cc | 21 +- .../test/src/utils/containers/contains_key.cc | 1 - .../test/src/utils/containers/enumerate.cc | 2 +- lib/utils/test/src/utils/containers/filter.cc | 2 - .../test/src/utils/containers/filter_keys.cc | 5 +- .../src/utils/containers/filtermap_keys.cc | 1 - .../src/utils/containers/filtermap_values.cc | 1 - .../test/src/utils/containers/filtrans.cc | 1 - lib/utils/test/src/utils/containers/find.cc | 1 - .../test/src/utils/containers/flatmap.cc | 8 +- .../utils/containers/get_all_assignments.cc | 9 +- .../get_all_permutations_with_repetition.cc | 16 +- .../test/src/utils/containers/group_by.cc | 1 - .../src/utils/containers/inplace_filter.cc | 2 - .../src/utils/containers/is_submapeq_of.cc | 8 +- lib/utils/test/src/utils/containers/keys.cc | 5 +- .../containers/lift_optional_through_map.cc | 5 +- .../src/utils/containers/map_from_pairs.cc | 11 +- .../test/src/utils/containers/map_keys.cc | 2 +- .../test/src/utils/containers/map_values.cc | 2 +- .../merge_disjoint_unordered_maps.cc | 2 +- .../src/utils/containers/merge_maps_with.cc | 22 +- .../containers/merge_unordered_maps_with.cc | 24 +- .../src/utils/containers/multiset_union.cc | 1 - .../test/src/utils/containers/product.cc | 1 - .../src/utils/containers/require_all_same1.cc | 3 - .../utils/containers/require_no_duplicates.cc | 2 - .../src/utils/containers/restrict_keys.cc | 3 +- .../src/utils/containers/set_intersection.cc | 1 - .../try_merge_nondisjoint_unordered_maps.cc | 7 +- .../containers/unordered_map_from_pairs.cc | 20 +- .../utils/containers/unordered_multiset_of.cc | 2 +- .../src/utils/containers/unordered_set_of.cc | 2 +- lib/utils/test/src/utils/containers/values.cc | 5 +- lib/utils/test/src/utils/fmt/unordered_map.cc | 2 +- lib/utils/test/src/utils/graph/algorithms.cc | 4 +- lib/utils/test/src/utils/graph/cow_ptr_t.cc | 2 +- ...plete_bipartite_composite_decomposition.cc | 12 +- .../graph/digraph/algorithms/contract_node.cc | 3 +- .../digraph/algorithms/get_dominators_map.cc | 3 +- .../algorithms/get_imm_dominators_map.cc | 3 +- .../algorithms/get_lowest_common_ancestors.cc | 44 +- .../algorithms/get_post_dominators_map.cc | 9 +- .../digraph/algorithms/get_predecessors.cc | 3 +- .../digraph/algorithms/get_successors.cc | 3 +- .../get_weakly_connected_components.cc | 27 +- .../get_inverse_line_graph.cc | 22 +- .../graph/instances/adjacency_digraph.cc | 3 +- .../graph/instances/adjacency_multidigraph.cc | 33 +- .../instances/unordered_set_dataflow_graph.cc | 16 +- ..._set_labelled_open_kwarg_dataflow_graph.cc | 27 +- ...unordered_set_open_kwarg_dataflow_graph.cc | 27 +- ...aph_data_from_kwarg_dataflow_graph_data.cc | 9 +- ...ataflow_graph_from_kwarg_dataflow_graph.cc | 15 +- .../get_kwarg_dataflow_graph_subgraph.cc | 4 +- .../view_from_kwarg_dataflow_graph_data.cc | 8 +- ..._labelled_kwarg_dataflow_graph_subgraph.cc | 7 +- ...warg_dataflow_graph_view_with_labelling.cc | 13 +- .../multidigraph/algorithms/add_edges.cc | 3 +- .../algorithms/get_incoming_edges.cc | 3 +- .../algorithms/get_outgoing_edges.cc | 3 +- .../get_open_dataflow_graph_inputs.cc | 3 +- .../algorithms/get_subgraph.cc | 7 +- .../algorithms/permute_node_ids.cc | 12 +- ..._dataflow_graph_by_materializing_inputs.cc | 6 +- ...ft_associative_binary_sp_tree_from_nary.cc | 3 +- ...ht_associative_binary_sp_tree_from_nary.cc | 3 +- .../series_parallel/parallel_reduction.cc | 2 +- .../graph/series_parallel/series_reduction.cc | 6 +- .../sp_ization/escribano_algo.cc | 60 +- .../sp_ization/naive_stratum_sync.cc | 26 +- .../series_parallel/sp_ization/node_role.cc | 7 +- .../sp_ization/work_duplicating_sp_ization.cc | 7 +- .../algorithms/get_connected_components.cc | 15 +- .../graph/undirected/undirected_graph.cc | 13 +- lib/utils/test/src/utils/graph/views/views.cc | 17 +- .../test/src/utils/one_to_many/one_to_many.cc | 4 +- .../test/src/utils/orthotope/dim_coord.cc | 12 +- 631 files changed, 6178 insertions(+), 6732 deletions(-) diff --git a/bin/run-model/src/run-model/main.cc b/bin/run-model/src/run-model/main.cc index 2afb546c48..49c87d5a98 100644 --- a/bin/run-model/src/run-model/main.cc +++ b/bin/run-model/src/run-model/main.cc @@ -71,47 +71,47 @@ int main(int argc, char **argv) { char **realm_argv = realm_args.data(); RealmManager manager(&realm_argc, &realm_argv); - ControllerTaskResult result = manager.start_controller([&](RealmContext - &ctx) { - MappedParallelComputationGraph mpcg = [&]() { - std::ifstream f(mapped_pcg_json); - nlohmann::json mpcg_json = nlohmann::json::parse(f); - return from_v1(mpcg_json.get()); - }(); - - // instantiate computation graph - OptimizerAttrs optimizer_attrs = - OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001}}; - - std::map input_tensors; - - DistributedFfHandle device_handle = - create_distributed_ff_handle(ctx, - /*workSpaceSize=*/1024 * 1024, - /*allowTensorOpMathConversion=*/true); - - PCGInstance pcg_instance = create_pcg_instance( - /*ctx=*/ctx, - /*mpcg=*/mpcg, - /*optimizer=*/optimizer_attrs, - /*loss=*/std::nullopt, - /*input_tensors=*/input_tensors, - /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*device_type=*/DeviceType::GPU); - - // begin training loop - int num_epochs = 5; - for (int i = 0; i < num_epochs; i++) { - perform_all_passes_for_pcg_instance( - /*instance=*/pcg_instance, - /*profiling_settings=*/ProfilingSettings{0, 1}, - /*device_handle=*/device_handle); - } - }); + ControllerTaskResult result = + manager.start_controller([&](RealmContext &ctx) { + MappedParallelComputationGraph mpcg = [&]() { + std::ifstream f(mapped_pcg_json); + nlohmann::json mpcg_json = nlohmann::json::parse(f); + return from_v1(mpcg_json.get()); + }(); + + // instantiate computation graph + OptimizerAttrs optimizer_attrs = + OptimizerAttrs{SGDOptimizerAttrs{/*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001}}; + + std::map input_tensors; + + DistributedFfHandle device_handle = + create_distributed_ff_handle(ctx, + /*workSpaceSize=*/1024 * 1024, + /*allowTensorOpMathConversion=*/true); + + PCGInstance pcg_instance = create_pcg_instance( + /*ctx=*/ctx, + /*mpcg=*/mpcg, + /*optimizer=*/optimizer_attrs, + /*loss=*/std::nullopt, + /*input_tensors=*/input_tensors, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::GPU); + + // begin training loop + int num_epochs = 5; + for (int i = 0; i < num_epochs; i++) { + perform_all_passes_for_pcg_instance( + /*instance=*/pcg_instance, + /*profiling_settings=*/ProfilingSettings{0, 1}, + /*device_handle=*/device_handle); + } + }); result.wait(); return 0; diff --git a/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h b/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h index 572c9c5f5c..640b9523b2 100644 --- a/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h +++ b/bin/sp-ization-benchmarking/include/sp-ization-benchmarking/distributions.h @@ -2,8 +2,8 @@ #define _FLEXFLOW_BIN_SP_IZATION_BENCHMARKING_INCLUDE_SP_IZATION_BENCHMARKING_DISTRIBUTIONS_H #include "utils/graph/node/node.dtg.h" -#include #include +#include #include namespace FlexFlow { @@ -55,9 +55,8 @@ struct GaussianNoise { }; template -std::map - make_cost_map(std::set const &nodes, - Dist const &distribution) { +std::map make_cost_map(std::set const &nodes, + Dist const &distribution) { std::map cost_map; for (Node const &node : nodes) { cost_map[node] = distribution(); @@ -66,9 +65,8 @@ std::map } template -std::map - add_noise_to_cost_map(std::map cost_map, - Noise const &noise) { +std::map add_noise_to_cost_map(std::map cost_map, + Noise const &noise) { std::map noisy_cost_map; for (auto const &[node, cost] : cost_map) { noisy_cost_map[node] = noise() * cost; diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h index 556a0a83e0..709be14baa 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h @@ -8,9 +8,8 @@ namespace FlexFlow { -std::set - abstracted_single_tensor_movement_get_dst_layers( - AbstractedSingleTensorMovement const &); +std::set abstracted_single_tensor_movement_get_dst_layers( + AbstractedSingleTensorMovement const &); AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( std::multiset const &); @@ -18,15 +17,12 @@ AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( AbstractedSingleTensorMovement abstracted_single_tensor_movement_from_communications( BinaryTreePath const &src_op_tree_path, - std::set const - &communications); + std::set const &communications); TensorSetMovement concretize_abstracted_single_tensor_movement( AbstractedSingleTensorMovement const &, - std::map const - &pre_machine_stencils, - std::map const - &post_machine_stencils); + std::map const &pre_machine_stencils, + std::map const &post_machine_stencils); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h index 15f6654918..0f57ca20f3 100644 --- a/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h +++ b/lib/compiler/include/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.h @@ -19,17 +19,13 @@ AbstractedTensorSetMovement abstracted_tensor_set_movement_from_single_tensor_movement( AbstractedSingleTensorMovement const &); -std::set - get_src_layers(AbstractedTensorSetMovement const &); -std::set - get_dst_layers(AbstractedTensorSetMovement const &); +std::set get_src_layers(AbstractedTensorSetMovement const &); +std::set get_dst_layers(AbstractedTensorSetMovement const &); TensorSetMovement concretize_abstracted_tensor_set_movement( AbstractedTensorSetMovement const &, - std::map const - &pre_machine_stencils, - std::map const - &post_machine_stencils); + std::map const &pre_machine_stencils, + std::map const &post_machine_stencils); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h index 5562c39f38..086ebff942 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_constraints.h @@ -10,8 +10,8 @@ namespace FlexFlow { -MachineMappingConstraints get_unconstrained_solution_for_layers( - std::set const &); +MachineMappingConstraints + get_unconstrained_solution_for_layers(std::set const &); std::set get_unconstrained_layers(MachineMappingConstraints const &); @@ -19,8 +19,7 @@ std::set std::set get_constrained_layers(MachineMappingConstraints const &); -std::set - get_all_layers(MachineMappingConstraints const &); +std::set get_all_layers(MachineMappingConstraints const &); std::optional get_machine_view_for_layer(MachineMappingConstraints const &, diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h index 8bb78a9086..7321cc744a 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_problem_tree/machine_mapping_problem_tree.h @@ -21,8 +21,7 @@ SPDecompositionTreeNodeType get_node_type(MachineMappingProblemTree const &); std::multiset get_leaves(MachineMappingProblemTree const &); -std::set - get_all_leaf_paths(MachineMappingProblemTree const &); +std::set get_all_leaf_paths(MachineMappingProblemTree const &); std::optional mm_problem_tree_get_subtree_at_path(MachineMappingProblemTree const &, diff --git a/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h b/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h index 8d5d741187..390742d9e7 100644 --- a/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h +++ b/lib/compiler/include/compiler/machine_mapping/machine_mapping_result.h @@ -12,8 +12,8 @@ namespace FlexFlow { [[nodiscard]] bool is_infeasible(MachineMappingResult const &); FeasibleMachineMappingResult require_feasible(MachineMappingResult const &); -[[nodiscard]] MachineMappingResult get_mapping_with_minimal_runtime( - std::set const &); +[[nodiscard]] MachineMappingResult + get_mapping_with_minimal_runtime(std::set const &); [[nodiscard]] MachineMappingResult series_combine(milliseconds_t comm_cost, diff --git a/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h b/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h index dcb909d59f..35e1889032 100644 --- a/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h +++ b/lib/compiler/include/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.h @@ -5,9 +5,8 @@ namespace FlexFlow { -bool is_pareto_optimal_in( - ParetoOptimalMachineMapping const &, - std::set const &); +bool is_pareto_optimal_in(ParetoOptimalMachineMapping const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h b/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h index 21ffc11af3..e289f91542 100644 --- a/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h +++ b/lib/compiler/include/compiler/series_parallel/pcg/pcg_binary_sp_decomposition.h @@ -34,9 +34,8 @@ SPDecompositionTreeNodeType get_node_type(PCGBinarySPDecomposition const &); std::set pcg_sp_tree_get_all_leaf_paths(PCGBinarySPDecomposition const &); -std::set - find_paths_to_leaf(PCGBinarySPDecomposition const &, - parallel_layer_guid_t const &); +std::set find_paths_to_leaf(PCGBinarySPDecomposition const &, + parallel_layer_guid_t const &); std::map pcg_sp_tree_get_path_to_leaf_map(PCGBinarySPDecomposition const &); diff --git a/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc b/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc index 01c9a2e765..58b135dfcf 100644 --- a/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc +++ b/lib/compiler/src/compiler/cost_estimator/parallel_tensor_space_to_machine_space_mapping.cc @@ -3,9 +3,9 @@ #include "op-attrs/parallel_tensor_dim_degrees.h" #include "op-attrs/parallel_tensor_space_coordinate.h" #include "op-attrs/task_space_coordinate.h" -#include "utils/bidict/algorithms/exhaustive_relational_join.h" #include "utils/bidict/algorithms/bidict_transform_keys.h" #include "utils/bidict/algorithms/bidict_transform_values.h" +#include "utils/bidict/algorithms/exhaustive_relational_join.h" #include namespace FlexFlow { @@ -20,9 +20,9 @@ ParallelTensorSpaceToMachineSpaceMapping ptensor_machine_map_from_composition( bidict pt_to_op_coord_map = bidict_transform_keys( - bidict_transform_values(op_task_to_parallel_tensor_space_mapping.raw_mapping - .coord_mapping.reversed(), - task_space_coordinate_from_dim_coord), + bidict_transform_values(op_task_to_parallel_tensor_space_mapping + .raw_mapping.coord_mapping.reversed(), + task_space_coordinate_from_dim_coord), parallel_tensor_space_coord_from_dim_coord); bidict op_to_ms_coord_map = diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc index a789e76789..3c163703c6 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.cc @@ -1,19 +1,18 @@ #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h" #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_communication_edge.h" #include "utils/containers/filtermap_keys.h" +#include "utils/containers/map_from_pairs.h" #include "utils/containers/map_keys_with_value_merging.h" +#include "utils/containers/merge_maps_with.h" #include "utils/containers/require_all_same1.h" #include "utils/containers/require_same.h" #include "utils/containers/transform.h" #include "utils/containers/values.h" -#include "utils/containers/merge_maps_with.h" -#include "utils/containers/map_from_pairs.h" namespace FlexFlow { -std::set - abstracted_single_tensor_movement_get_dst_layers( - AbstractedSingleTensorMovement const &m) { +std::set abstracted_single_tensor_movement_get_dst_layers( + AbstractedSingleTensorMovement const &m) { return transform( keys(m.edge_to_size), [](AbstractedSingleTensorCommunicationEdge const &e) -> BinaryTreePath { @@ -45,8 +44,7 @@ AbstractedSingleTensorMovement merge_abstracted_single_tensor_movements( AbstractedSingleTensorMovement abstracted_single_tensor_movement_from_communications( BinaryTreePath const &src_op_tree_path, - std::set const - &communications) { + std::set const &communications) { return AbstractedSingleTensorMovement{ /*src_op_tree_path=*/src_op_tree_path, @@ -61,8 +59,7 @@ AbstractedSingleTensorMovement TensorSetMovement concretize_abstracted_single_tensor_movement( AbstractedSingleTensorMovement const &abstracted, - std::map const - &pre_machine_stencils, + std::map const &pre_machine_stencils, std::map const &post_machine_stencils) { @@ -70,8 +67,8 @@ TensorSetMovement concretize_abstracted_single_tensor_movement( MachineSpaceStencil pre_machine_stencil = pre_machine_stencils.at(abstracted.src_op_tree_path); - std::map, num_bytes_t> - communication_edges = map_keys_with_value_merging( + std::map, num_bytes_t> communication_edges = + map_keys_with_value_merging( abstracted.edge_to_size, /*key_func=*/ [&](AbstractedSingleTensorCommunicationEdge const &k) { diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc index b75eed3fdf..25e961f100 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_tensor_set_movement.cc @@ -3,11 +3,11 @@ #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.dtg.h" #include "compiler/machine_mapping/abstracted_tensor_set_movement/abstracted_single_tensor_movement.h" #include "compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.h" +#include "utils/containers/binary_merge_maps_with.h" #include "utils/containers/flatmap.h" #include "utils/containers/map_keys_with_value_merging.h" #include "utils/containers/merge_maps_with.h" #include "utils/containers/transform.h" -#include "utils/containers/binary_merge_maps_with.h" namespace FlexFlow { @@ -23,8 +23,7 @@ AbstractedTensorSetMovement }; } -std::set - get_src_layers(AbstractedTensorSetMovement const &m) { +std::set get_src_layers(AbstractedTensorSetMovement const &m) { return transform( m.single_tensor_movements, [](AbstractedSingleTensorMovement const &e) -> BinaryTreePath { @@ -32,19 +31,17 @@ std::set }); } -std::set - get_dst_layers(AbstractedTensorSetMovement const &m) { - return flatmap(m.single_tensor_movements, - [](AbstractedSingleTensorMovement const &m) - -> std::set { - return abstracted_single_tensor_movement_get_dst_layers(m); - }); +std::set get_dst_layers(AbstractedTensorSetMovement const &m) { + return flatmap( + m.single_tensor_movements, + [](AbstractedSingleTensorMovement const &m) -> std::set { + return abstracted_single_tensor_movement_get_dst_layers(m); + }); } TensorSetMovement concretize_abstracted_tensor_set_movement( AbstractedTensorSetMovement const &abstracted, - std::map const - &pre_machine_stencils, + std::map const &pre_machine_stencils, std::map const &post_machine_stencils) { diff --git a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc index 9fbbdf0cfe..cf2702c461 100644 --- a/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc +++ b/lib/compiler/src/compiler/machine_mapping/abstracted_tensor_set_movement/get_abstracted_tensor_set_movement_across_split.cc @@ -10,17 +10,17 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" +#include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" #include "utils/containers/binary_cartesian_product.h" #include "utils/containers/flatmap.h" #include "utils/containers/get_only.h" #include "utils/containers/group_by.h" #include "utils/containers/map_from_pairs.h" #include "utils/containers/merge_maps_with.h" -#include "utils/containers/transform.h" #include "utils/containers/multiset_of.h" +#include "utils/containers/transform.h" #include "utils/containers/values.h" #include "utils/containers/vector_of.h" -#include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" namespace FlexFlow { @@ -43,8 +43,8 @@ AbstractedSingleTensorMovement get_abstracted_single_tensor_movement_along_edge( bidict coord_mapping = op_to_op_get_coord_mapping(mapping); - std::map - single_comms = map_from_pairs(transform( + std::map single_comms = + map_from_pairs(transform( unstructured_relation_from_bidict(coord_mapping), [&](std::pair const & src_dst) -> std::pair const &edges) - { - return merge_abstracted_single_tensor_movements(transform( - multiset_of(edges.unwrap_as_set()), - to_abstracted_single_tensor_movement)); + [&](nonempty_set const &edges) { + return merge_abstracted_single_tensor_movements( + transform(multiset_of(edges.unwrap_as_set()), + to_abstracted_single_tensor_movement)); }), }; } diff --git a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc index fede093227..9d483e4f30 100644 --- a/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/src/compiler/machine_mapping/allowed_machine_views.cc @@ -10,13 +10,13 @@ #include "utils/containers/filter.h" #include "utils/containers/get_all_permutations_with_repetition.h" #include "utils/containers/map_from_keys_and_values.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/product.h" #include "utils/containers/range.h" #include "utils/containers/repeat_element.h" +#include "utils/containers/set_of.h" #include "utils/containers/sorted.h" #include "utils/containers/transform.h" -#include "utils/containers/multiset_of.h" -#include "utils/containers/set_of.h" #include "utils/containers/zip.h" #include "utils/nonnegative_int/nonnegative_range.h" #include "utils/nonnegative_int/num_elements.h" @@ -64,9 +64,9 @@ static std::set positive_int{min_num_devices_with_full_stride_volume}); }; - auto get_candidate_strides = [&](std::vector const &tensor_dims, - positive_int total_devices) - -> std::multiset { + auto get_candidate_strides = + [&](std::vector const &tensor_dims, + positive_int total_devices) -> std::multiset { positive_int max_stride_upper_bound = get_max_stride_upper_bound(tensor_dims, total_devices); @@ -76,10 +76,9 @@ static std::set max_stride_upper_bound.nonnegative_int_from_positive_int() + 1_n), [](nonnegative_int stride) { return stride_t{positive_int{stride}}; }); - std::multiset> raw_stride_vectors = - cartesian_product( - repeat_element(/*num_times=*/num_elements(tensor_dims), - /*element=*/single_stride_range)); + std::multiset> raw_stride_vectors = cartesian_product( + repeat_element(/*num_times=*/num_elements(tensor_dims), + /*element=*/single_stride_range)); std::multiset strides = transform(raw_stride_vectors, [](auto const &stride_vec) { diff --git a/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc index 198fa29326..517256cbea 100644 --- a/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/apply_substitution_and_update_machine_mapping.cc @@ -39,9 +39,8 @@ SearchResult apply_substitution_and_update_machine_mapping( std::map post_node_data = get_sub_pcg_data(post_substitution_graph).node_data; - std::set - substitution_output_parallel_layers = - spcg_get_parallel_layers(substitution_output_result.first); + std::set substitution_output_parallel_layers = + spcg_get_parallel_layers(substitution_output_result.first); std::map machine_views = mapped_pcg.machine_mapping.machine_views; @@ -61,12 +60,11 @@ SearchResult apply_substitution_and_update_machine_mapping( ASSERT(is_subseteq_of(keys(post_node_data), keys(machine_views))); - std::map - post_node_machine_views = - filter(machine_views, - [&](std::pair const &p) { - return post_node_data.count(p.first); - }); + std::map post_node_machine_views = + filter(machine_views, + [&](std::pair const &p) { + return post_node_data.count(p.first); + }); ASSERT(keys(post_node_data) == keys(post_node_machine_views)); diff --git a/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc index 521217e599..f36d11418a 100644 --- a/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/get_optimal_machine_mapping.cc @@ -99,28 +99,26 @@ MachineMappingResult ASSERT(get_all_layers(sub_constraints) == get_all_leaf_paths(root)); - std::set unconstrained_boundary_layers = - set_minus(boundary_layers, set_of(get_constrained_layers(sub_constraints))); - - std::map> - allowed = generate_map( - unconstrained_boundary_layers, - [&](BinaryTreePath const &l) -> std::set { - UnmappedRuntimeOnlyOpCostEstimateKey leaf = - mm_problem_tree_get_subtree_at_path(root, l) - .value() - .get(); - return context.allowed_machine_views(leaf, resources); - }); - - std::set> - assignments = get_all_assignments(allowed); - - return transform( - assignments, - [](std::map const &m) { - return ParallelLayerGuidObliviousMachineMapping{m}; - }); + std::set unconstrained_boundary_layers = set_minus( + boundary_layers, set_of(get_constrained_layers(sub_constraints))); + + std::map> allowed = + generate_map(unconstrained_boundary_layers, + [&](BinaryTreePath const &l) -> std::set { + UnmappedRuntimeOnlyOpCostEstimateKey leaf = + mm_problem_tree_get_subtree_at_path(root, l) + .value() + .get(); + return context.allowed_machine_views(leaf, resources); + }); + + std::set> assignments = + get_all_assignments(allowed); + + return transform(assignments, + [](std::map const &m) { + return ParallelLayerGuidObliviousMachineMapping{m}; + }); }; auto eval_pre_boundary_mapping = diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc index 13bb389efb..a103dbe0c4 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping.cc @@ -7,8 +7,8 @@ #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/containers/are_disjoint.h" -#include "utils/containers/keys.h" #include "utils/containers/binary_merge_disjoint_maps.h" +#include "utils/containers/keys.h" namespace FlexFlow { @@ -16,11 +16,9 @@ MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapping(ParallelComputationGraph const &pcg, MachineMapping const &mapping) { - std::set pcg_layers = - pcg_get_parallel_layers(pcg); + std::set pcg_layers = pcg_get_parallel_layers(pcg); - std::set mapped_layers = - keys(mapping.machine_views); + std::set mapped_layers = keys(mapping.machine_views); ASSERT(mapped_layers == pcg_layers); @@ -29,8 +27,8 @@ MappedParallelComputationGraph ComputationGraphOpAttrs op_attrs = assert_unwrap( compgraph_op_attrs_from_pcg_op_attrs(pcg_get_op_attrs(pcg, l))); - std::map - inputs_dim_degrees = get_incoming_input_degrees(pcg, l); + std::map inputs_dim_degrees = + get_incoming_input_degrees(pcg, l); ASSERT(contains_key(mapping.machine_views, l)); MachineView machine_view = mapping.machine_views.at(l); diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc index 3b73b98a47..49ab1c9ad4 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_constraints.cc @@ -103,8 +103,7 @@ MachineMappingConstraints with_additional_constraints( std::optional require_only_root(MachineMappingConstraints const &constraints) { - ASSERT(keys(constraints.machine_views) == - std::set{binary_tree_root_path()}, + ASSERT(keys(constraints.machine_views) == std::set{binary_tree_root_path()}, fmt::format("require_only_root expected constraints to have only a " "single key (the root path), but received {}", constraints)); diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc index 28f049f08f..c1b38654ce 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_mutation_set.cc @@ -16,9 +16,8 @@ std::optional std::map machine_views; for (parallel_layer_guid_t layer : layers) { OperatorTaskSpace task = get_operator_task_space(pcg, layer); - std::set allowed_machine_views = - get_allowed_machine_views( - compute_slice_from_specification(resources), task); + std::set allowed_machine_views = get_allowed_machine_views( + compute_slice_from_specification(resources), task); if (allowed_machine_views.empty()) { return std::nullopt; } diff --git a/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc b/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc index f632c99190..5f2655706d 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_mapping_problem_tree/get_machine_mapping_problem_tree.cc @@ -18,13 +18,12 @@ bool is_valid_machine_mapping_problem_tree( AbstractedTensorSetMovement tensor_movement = series_split.tensor_set_movement; - auto contains_paths = - [](MachineMappingProblemTree const &t, - std::set const &paths) { - return all_of(paths, [&](BinaryTreePath const &p) { - return mm_problem_tree_get_subtree_at_path(t, p).has_value(); - }); - }; + auto contains_paths = [](MachineMappingProblemTree const &t, + std::set const &paths) { + return all_of(paths, [&](BinaryTreePath const &p) { + return mm_problem_tree_get_subtree_at_path(t, p).has_value(); + }); + }; return contains_paths(series_split.get_left_child(), get_src_layers(tensor_movement)) && diff --git a/lib/compiler/src/compiler/machine_mapping/machine_view.cc b/lib/compiler/src/compiler/machine_mapping/machine_view.cc index f87f8adf82..15277d132a 100644 --- a/lib/compiler/src/compiler/machine_mapping/machine_view.cc +++ b/lib/compiler/src/compiler/machine_mapping/machine_view.cc @@ -200,11 +200,11 @@ static OperatorAtomicTaskShardBinding mv_task_space_coord_for_machine_space_coord( machine_view, op_task_space, machine_space_coord); - std::map - mappings = get_operator_to_ptensor_mappings(op_attrs, inputs_dim_degrees); + std::map mappings = + get_operator_to_ptensor_mappings(op_attrs, inputs_dim_degrees); - std::map - ptensor_coords = generate_map( + std::map ptensor_coords = + generate_map( keys(inputs_dim_degrees), [&](TensorSlotName const &slot_name) -> ParallelTensorSpaceCoordinate { diff --git a/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc b/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc index ab54c034a1..bf1f89c9e7 100644 --- a/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc +++ b/lib/compiler/src/compiler/machine_mapping/memory_optimization/get_optimal_machine_mapping_with_memory.cc @@ -83,22 +83,19 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( [&](MachineMappingProblemTree const &root, std::set const &boundary_layers) -> std::set { - std::map> - allowed = generate_map( - boundary_layers, - [&](BinaryTreePath const &l) -> std::set { - UnmappedRuntimeOnlyOpCostEstimateKey leaf = - mm_problem_tree_get_subtree_at_path(root, l) - .value() - .get(); - return context.allowed_machine_views(leaf, resources); - }); - - return transform( - get_all_assignments(allowed), - [](std::map const &m) { - return ParallelLayerGuidObliviousMachineMapping{m}; + std::map> allowed = generate_map( + boundary_layers, [&](BinaryTreePath const &l) -> std::set { + UnmappedRuntimeOnlyOpCostEstimateKey leaf = + mm_problem_tree_get_subtree_at_path(root, l) + .value() + .get(); + return context.allowed_machine_views(leaf, resources); }); + + return transform(get_all_assignments(allowed), + [](std::map const &m) { + return ParallelLayerGuidObliviousMachineMapping{m}; + }); }; auto eval_pre_boundary_mapping = @@ -226,9 +223,8 @@ MachineMappingWithMemoryResult get_optimal_machine_mapping_with_memory( return parallel_combine(resource_split, left_result, right_result); }; - std::set parallel_results = - transform(get_machine_resource_splits(resources), - evaluate_resource_split); + std::set parallel_results = transform( + get_machine_resource_splits(resources), evaluate_resource_split); return minimize_runtime(series_result, get_mapping_with_minimal_runtime(parallel_results)); diff --git a/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc b/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc index 4f09c569d2..6114e8db50 100644 --- a/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc +++ b/lib/compiler/src/compiler/machine_mapping/memory_optimization/machine_mapping_with_memory_result.cc @@ -6,8 +6,8 @@ #include "utils/containers/set_union.h" #include "utils/containers/transform.h" #include "utils/full_binary_tree/binary_tree_path.h" -#include "utils/hash/tuple.h" #include "utils/hash/set.h" +#include "utils/hash/tuple.h" namespace FlexFlow { diff --git a/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc index 96ea3c6b30..6d272ae8fa 100644 --- a/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/memory_optimization/pareto_optimal_machine_mapping.cc @@ -4,9 +4,8 @@ namespace FlexFlow { -bool is_pareto_optimal_in( - ParetoOptimalMachineMapping const &m, - std::set const &others) { +bool is_pareto_optimal_in(ParetoOptimalMachineMapping const &m, + std::set const &others) { return is_pareto_optimal_in( m.cost, transform(others, [](ParetoOptimalMachineMapping const &m) { return m.cost; diff --git a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc index 1c61ba4fd9..19453a4e58 100644 --- a/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc +++ b/lib/compiler/src/compiler/machine_mapping/parallel_layer_guid_oblivious_machine_mapping.cc @@ -5,11 +5,11 @@ #include "op-attrs/get_operator_task_space.h" #include "op-attrs/parallel_tensor_shape.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/map_keys.h" #include "utils/containers/require_same.h" #include "utils/containers/try_at.h" #include "utils/full_binary_tree/binary_tree_path.h" -#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -47,12 +47,11 @@ std::map std::set leaf_paths = require_same( pcg_sp_tree_get_all_leaf_paths(decomposition), keys(mapping.raw_mapping)); - std::map - path_to_op_task_space_map = - map_values(pcg_sp_tree_get_path_to_leaf_map(decomposition), - [&](parallel_layer_guid_t l) -> OperatorTaskSpace { - return get_operator_task_space(pcg, l); - }); + std::map path_to_op_task_space_map = + map_values(pcg_sp_tree_get_path_to_leaf_map(decomposition), + [&](parallel_layer_guid_t l) -> OperatorTaskSpace { + return get_operator_task_space(pcg, l); + }); return generate_map( leaf_paths, [&](BinaryTreePath const &p) -> MachineSpaceStencil { @@ -68,8 +67,8 @@ std::map> MachineMappingProblemTree const &tree, ParallelLayerGuidObliviousMachineMapping const &mapping) { - std::map - tree_leaf_map = mm_problem_tree_get_path_to_leaf_map(tree); + std::map tree_leaf_map = + mm_problem_tree_get_path_to_leaf_map(tree); std::set mapping_paths = keys(mapping.raw_mapping); std::set tree_paths = keys(tree_leaf_map); @@ -88,11 +87,10 @@ std::map> ComputationGraphOpAttrs leaf_op_attrs = compgraph_op_attrs_from_pcg_op_attrs(leaf.op_attrs).value(); - std::map - leaf_input_degrees = - map_values(leaf.input_shapes, [](ParallelTensorShape const &s) { - return get_parallel_degrees(s); - }); + std::map leaf_input_degrees = + map_values(leaf.input_shapes, [](ParallelTensorShape const &s) { + return get_parallel_degrees(s); + }); return MachineSpaceStencil{ /*operator_task_space=*/get_operator_task_space(leaf_op_attrs, diff --git a/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc b/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc index 16f418cec2..357f1b3260 100644 --- a/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc +++ b/lib/compiler/src/compiler/machine_mapping/transitive_reduced_pcg.cc @@ -45,8 +45,8 @@ std::set binary_series_split_from_pcg_series_split(split); std::set> raw_edges = - set_of(get_transitive_reduced_kwarg_dataflow_edges_across_split(raw_tr_g, - raw_split)); + set_of(get_transitive_reduced_kwarg_dataflow_edges_across_split( + raw_tr_g, raw_split)); return transform(raw_edges, [](KwargDataflowEdge const &e) { return ParallelComputationGraphEdge{e}; @@ -63,8 +63,8 @@ std::set binary_series_split_from_pcg_series_split(split); std::set> raw_outputs = - set_of(get_transitive_reduced_kwarg_dataflow_outputs_across_split(raw_tr_g, - raw_split)); + set_of(get_transitive_reduced_kwarg_dataflow_outputs_across_split( + raw_tr_g, raw_split)); return transform(raw_outputs, [](KwargDataflowOutput const &o) { diff --git a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc index 351f097780..e659c42900 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/pcg_task_graph.cc @@ -11,10 +11,10 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" #include "utils/bidict/bidict.h" +#include "utils/containers/set_of.h" #include "utils/graph/instances/adjacency_digraph.h" #include #include -#include "utils/containers/set_of.h" namespace FlexFlow { diff --git a/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc b/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc index e708d4b0f6..488a0703ed 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc @@ -47,8 +47,7 @@ TaskGraphExecutionTrace simulate_task_graph_execution( }; auto dependencies_are_satisfied = [&](Node const &task) { - std::set incoming_dependencies = - get_predecessors(task_graph, task); + std::set incoming_dependencies = get_predecessors(task_graph, task); return is_subseteq_of(incoming_dependencies, execution_state.finished_tasks); }; @@ -80,9 +79,9 @@ TaskGraphExecutionTrace simulate_task_graph_execution( while (!is_processing_done()) { auto ready_tasks_copy = execution_state.ready_tasks; for (Node const &task : ready_tasks_copy) { - std::set raw_in_progress_tasks = transform( - set_of(execution_state.in_progress_tasks.contents()), - [](InProgressTask const &t) { return t.node; }); + std::set raw_in_progress_tasks = + transform(set_of(execution_state.in_progress_tasks.contents()), + [](InProgressTask const &t) { return t.node; }); if (constraint.is_satisfied( task, raw_in_progress_tasks, execution_state.finished_tasks)) { diff --git a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc index 77812b4174..6df8a2c85e 100644 --- a/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/src/compiler/task_graph_simulator/task_simulator.cc @@ -41,10 +41,9 @@ milliseconds_t task_simulator_estimate_forward_pass_time( return running_time.unwrap_milliseconds(); }; - auto is_allowed_to_run = - [&](Node const &task, - std::set const &in_progress_tasks, - std::set const &finished_tasks) -> bool { + auto is_allowed_to_run = [&](Node const &task, + std::set const &in_progress_tasks, + std::set const &finished_tasks) -> bool { PCGTask current_task = task_graph.node_to_task.at_l(task); if (current_task.is_tensor_movement()) { @@ -52,14 +51,14 @@ milliseconds_t task_simulator_estimate_forward_pass_time( } assert(current_task.is_operator()); - auto get_devices = - [&](Node const &n) -> std::set { + auto get_devices = [&](Node const &n) -> std::set { return task_graph.node_to_devices.at(n); }; std::set devices_occupied = set_union(transform(in_progress_tasks, get_devices)); - std::set required_devices = set_of(get_devices(task)); + std::set required_devices = + set_of(get_devices(task)); return set_intersection(devices_occupied, required_devices).empty(); }; diff --git a/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc b/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc index 8a81e97255..2e9f0ca6dc 100644 --- a/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc +++ b/lib/compiler/src/compiler/unity_algorithm/graph_optimize_state.cc @@ -6,13 +6,13 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.h" #include "utils/bidict/algorithms/bidict_from_map.h" +#include "utils/containers/multiset_of.h" +#include "utils/containers/transform.h" #include "utils/containers/zip_values_strict.h" #include "utils/containers/zip_values_strict_with.h" -#include "utils/hash/tuple.h" #include "utils/hash/map.h" #include "utils/hash/multiset.h" -#include "utils/containers/transform.h" -#include "utils/containers/multiset_of.h" +#include "utils/hash/tuple.h" namespace FlexFlow { @@ -31,9 +31,9 @@ static std::multiset std::tuple>, + std::tuple>, std::map> { ParallelLayerAttrs layer_attrs = get_parallel_layer_attrs(pcg, l); @@ -54,11 +54,10 @@ static std::multiset outputs = - map_values(get_outgoing_tensors(pcg, l), - [&](parallel_tensor_guid_t const &o) { - return get_parallel_tensor_attrs(pcg, o); - }); + std::map outputs = map_values( + get_outgoing_tensors(pcg, l), [&](parallel_tensor_guid_t const &o) { + return get_parallel_tensor_attrs(pcg, o); + }); return { layer_attrs, diff --git a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc index bfea9eb0ee..b191520be5 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/allowed_machine_views.cc @@ -1,8 +1,8 @@ #include "compiler/machine_mapping/allowed_machine_views.h" #include "utils/containers/extend.h" #include "utils/containers/range.h" -#include "utils/containers/transform.h" #include "utils/containers/set_of.h" +#include "utils/containers/transform.h" #include "utils/containers/zip.h" #include "utils/fmt/set.h" #include @@ -63,8 +63,7 @@ TEST_SUITE(FF_TEST_SUITE) { make_machine_view(0_n, 0_n, 2_p, intra), }; - std::set result = - get_allowed_machine_views(ms, task); + std::set result = get_allowed_machine_views(ms, task); CHECK(correct == result); } @@ -94,8 +93,7 @@ TEST_SUITE(FF_TEST_SUITE) { 0_n, 0_n, /*stride_1=*/2_p, intra, /*stride_2=*/1_p, inter), }; - std::set result = - get_allowed_machine_views(ms, task); + std::set result = get_allowed_machine_views(ms, task); CHECK(correct == result); } diff --git a/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc b/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc index 0e58b518f8..be07642a74 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/get_optimal_machine_mapping.cc @@ -344,7 +344,7 @@ TEST_SUITE(FF_TEST_SUITE) { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( std::map{{ + RuntimeOnlyOpCostMetrics>{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 4), @@ -401,7 +401,7 @@ TEST_SUITE(FF_TEST_SUITE) { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( std::map{{ + RuntimeOnlyOpCostMetrics>{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 4), @@ -481,7 +481,7 @@ TEST_SUITE(FF_TEST_SUITE) { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( std::map{{ + RuntimeOnlyOpCostMetrics>{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 4), @@ -530,7 +530,7 @@ TEST_SUITE(FF_TEST_SUITE) { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( std::map{{ + RuntimeOnlyOpCostMetrics>{{ mk_cost_entry(k1, mv_stride_1, 1), mk_cost_entry(k1, mv_stride_2, 3), mk_cost_entry(k2, mv_stride_1, 3), @@ -593,7 +593,7 @@ TEST_SUITE(FF_TEST_SUITE) { RuntimeOnlyCostEstimator runtime_only_cost_estimator = make_fake_runtime_only_cost_estimator( std::map{{ + RuntimeOnlyOpCostMetrics>{{ mk_cost_entry(k1, mv_stride_1, 3), mk_cost_entry(k1, mv_stride_2, 1), mk_cost_entry(k2, mv_stride_1, 4), diff --git a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc index 50ee6ee443..ecc96279d4 100644 --- a/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc +++ b/lib/compiler/test/src/compiler/machine_mapping/start_invariant_machine_view.cc @@ -127,10 +127,9 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_machine_space_offsets") { - std::set correct = { - MachineSpaceOffset{0, 0}, - MachineSpaceOffset{0, 2}, - MachineSpaceOffset{0, 4}}; + std::set correct = {MachineSpaceOffset{0, 0}, + MachineSpaceOffset{0, 2}, + MachineSpaceOffset{0, 4}}; std::set result = get_machine_space_offsets(task, simv); CHECK(correct == result); @@ -205,11 +204,10 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_machine_space_offsets") { - std::set correct = { - MachineSpaceOffset{0, 0}, - MachineSpaceOffset{0, 2}, - MachineSpaceOffset{1, 0}, - MachineSpaceOffset{1, 2}}; + std::set correct = {MachineSpaceOffset{0, 0}, + MachineSpaceOffset{0, 2}, + MachineSpaceOffset{1, 0}, + MachineSpaceOffset{1, 2}}; std::set result = get_machine_space_offsets(task, simv); CHECK(correct == result); diff --git a/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc b/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc index 3fd3bdcb56..eba839156d 100644 --- a/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc +++ b/lib/compiler/test/src/compiler/task_graph_simulator/simulate_task_graph_execution.cc @@ -25,10 +25,11 @@ TEST_SUITE(FF_TEST_SUITE) { auto cost_function = lookup_in_map( {{n.at(0), 1}, {n.at(1), 10}, {n.at(2), 100}, {n.at(3), 1000}}); - auto is_allowed_to_run = - [&](Node const &n, - std::set const &in_progress_tasks, - std::set const &finished_tasks) { return true; }; + auto is_allowed_to_run = [&](Node const &n, + std::set const &in_progress_tasks, + std::set const &finished_tasks) { + return true; + }; TaskExecutionConstraint constraint = TaskExecutionConstraint{is_allowed_to_run}; @@ -57,12 +58,11 @@ TEST_SUITE(FF_TEST_SUITE) { {{n.at(0), 10}, {n.at(1), 15}, {n.at(2), 20}, {n.at(3), 25}}); SUBCASE("no processing constraints") { - auto is_allowed_to_run = - [&](Node const &n, - std::set const &in_progress_tasks, - std::set const &finished_tasks) { - return true; - }; + auto is_allowed_to_run = [&](Node const &n, + std::set const &in_progress_tasks, + std::set const &finished_tasks) { + return true; + }; TaskExecutionConstraint constraint = TaskExecutionConstraint{is_allowed_to_run}; @@ -78,12 +78,11 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("one node at a time") { - auto is_allowed_to_run = - [&](Node const &n, - std::set const &in_progress_tasks, - std::set const &finished_tasks) { - return in_progress_tasks.size() == 0; - }; + auto is_allowed_to_run = [&](Node const &n, + std::set const &in_progress_tasks, + std::set const &finished_tasks) { + return in_progress_tasks.size() == 0; + }; TaskExecutionConstraint constraint = TaskExecutionConstraint{is_allowed_to_run}; @@ -121,12 +120,11 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), 35}}); SUBCASE("no processing constraints") { - auto is_allowed_to_run = - [&](Node const &n, - std::set const &in_progress_tasks, - std::set const &finished_tasks) { - return true; - }; + auto is_allowed_to_run = [&](Node const &n, + std::set const &in_progress_tasks, + std::set const &finished_tasks) { + return true; + }; TaskExecutionConstraint constraint = TaskExecutionConstraint{is_allowed_to_run}; @@ -144,12 +142,11 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("one node at a time") { - auto is_allowed_to_run = - [&](Node const &n, - std::set const &in_progress_tasks, - std::set const &finished_tasks) { - return in_progress_tasks.size() == 0; - }; + auto is_allowed_to_run = [&](Node const &n, + std::set const &in_progress_tasks, + std::set const &finished_tasks) { + return in_progress_tasks.size() == 0; + }; TaskExecutionConstraint constraint = TaskExecutionConstraint{is_allowed_to_run}; @@ -185,12 +182,11 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(4), 20}}); SUBCASE("at most two nodes at a time") { - auto is_allowed_to_run = - [&](Node const &n, - std::set const &in_progress_tasks, - std::set const &finished_tasks) { - return in_progress_tasks.size() < 2; - }; + auto is_allowed_to_run = [&](Node const &n, + std::set const &in_progress_tasks, + std::set const &finished_tasks) { + return in_progress_tasks.size() < 2; + }; TaskExecutionConstraint constraint = TaskExecutionConstraint{is_allowed_to_run}; diff --git a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc index 7720b38053..11901865fe 100644 --- a/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc +++ b/lib/compiler/test/src/compiler/task_graph_simulator/task_simulator.cc @@ -26,8 +26,8 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_source_nodes.h" #include "utils/nonnegative_int/nonnegative_int.h" #include -#include #include +#include #include namespace FlexFlow { diff --git a/lib/compiler/test/src/internal/cost_estimator_for_test.cc b/lib/compiler/test/src/internal/cost_estimator_for_test.cc index dd161c411d..bcbdcfc4bb 100644 --- a/lib/compiler/test/src/internal/cost_estimator_for_test.cc +++ b/lib/compiler/test/src/internal/cost_estimator_for_test.cc @@ -36,8 +36,7 @@ CostEstimator make_fake_cost_estimator( CostEstimator make_fake_cost_estimator( std::map const &op_cost_map, - std::map const - &comm_cost_map) { + std::map const &comm_cost_map) { return make_fake_cost_estimator( [op_cost_map](OpCostEstimateKey const &k) { ASSERT(contains_key(op_cost_map, k), k); diff --git a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc index 5a27e78d15..f5ced50d29 100644 --- a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc +++ b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.cc @@ -26,10 +26,9 @@ RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( } RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( - std::map const &op_cost_map, - std::map const - &comm_cost_map) { + std::map const + &op_cost_map, + std::map const &comm_cost_map) { return make_fake_runtime_only_cost_estimator( [op_cost_map](RuntimeOnlyOpCostEstimateKey const &k) { ASSERT(contains_key(op_cost_map, k), k); diff --git a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h index a09b75c967..f5d2a8f99a 100644 --- a/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h +++ b/lib/compiler/test/src/internal/runtime_only_cost_estimator_for_test.h @@ -12,8 +12,8 @@ RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( &get_communication_cost); RuntimeOnlyCostEstimator make_fake_runtime_only_cost_estimator( - std::map const &op_cost_map, + std::map const + &op_cost_map, std::map const &comm_cost_map); RuntimeOnlyCostEstimator make_fake_constant_runtime_only_cost_estimator( diff --git a/lib/kernels/src/kernels/accessor.cc b/lib/kernels/src/kernels/accessor.cc index 75f144a57a..0500c1695e 100644 --- a/lib/kernels/src/kernels/accessor.cc +++ b/lib/kernels/src/kernels/accessor.cc @@ -98,19 +98,23 @@ bool GenericTensorAccessorW::operator!=( return this->tie() != other.tie(); } -bool GenericTensorAccessorW::operator<(GenericTensorAccessorW const &other) const { +bool GenericTensorAccessorW::operator<( + GenericTensorAccessorW const &other) const { return this->tie() < other.tie(); } -bool GenericTensorAccessorW::operator<=(GenericTensorAccessorW const &other) const { +bool GenericTensorAccessorW::operator<=( + GenericTensorAccessorW const &other) const { return this->tie() <= other.tie(); } -bool GenericTensorAccessorW::operator>(GenericTensorAccessorW const &other) const { +bool GenericTensorAccessorW::operator>( + GenericTensorAccessorW const &other) const { return this->tie() > other.tie(); } -bool GenericTensorAccessorW::operator>=(GenericTensorAccessorW const &other) const { +bool GenericTensorAccessorW::operator>=( + GenericTensorAccessorW const &other) const { return this->tie() >= other.tie(); } @@ -166,19 +170,23 @@ bool GenericTensorAccessorR::operator!=( return this->tie() != other.tie(); } -bool GenericTensorAccessorR::operator<(GenericTensorAccessorR const &other) const { +bool GenericTensorAccessorR::operator<( + GenericTensorAccessorR const &other) const { return this->tie() < other.tie(); } -bool GenericTensorAccessorR::operator<=(GenericTensorAccessorR const &other) const { +bool GenericTensorAccessorR::operator<=( + GenericTensorAccessorR const &other) const { return this->tie() <= other.tie(); } -bool GenericTensorAccessorR::operator>(GenericTensorAccessorR const &other) const { +bool GenericTensorAccessorR::operator>( + GenericTensorAccessorR const &other) const { return this->tie() > other.tie(); } -bool GenericTensorAccessorR::operator>=(GenericTensorAccessorR const &other) const { +bool GenericTensorAccessorR::operator>=( + GenericTensorAccessorR const &other) const { return this->tie() >= other.tie(); } diff --git a/lib/local-execution/include/local-execution/computation_graph_instance.h b/lib/local-execution/include/local-execution/computation_graph_instance.h index 3ddc0ad406..997d37426e 100644 --- a/lib/local-execution/include/local-execution/computation_graph_instance.h +++ b/lib/local-execution/include/local-execution/computation_graph_instance.h @@ -14,8 +14,8 @@ #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" #include "task-spec/global_device_id_t.dtg.h" #include "utils/units/milliseconds_t.h" -#include #include +#include namespace FlexFlow { @@ -44,8 +44,7 @@ ComputationGraphInstance create_computation_graph_instance( ComputationGraph const &cg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::map const - &input_tensors, + std::map const &input_tensors, Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, diff --git a/lib/local-execution/include/local-execution/local_task_argument_accessor.h b/lib/local-execution/include/local-execution/local_task_argument_accessor.h index 414a6837ef..d0fc3cb38e 100644 --- a/lib/local-execution/include/local-execution/local_task_argument_accessor.h +++ b/lib/local-execution/include/local-execution/local_task_argument_accessor.h @@ -45,8 +45,7 @@ struct LocalTaskArgumentAccessor : public ITaskArgumentAccessor { private: Allocator allocator; - std::map - tensor_slots_backing; + std::map tensor_slots_backing; ProfilingSettings profiling_settings; device_handle_t ff_handle; diff --git a/lib/local-execution/include/local-execution/tensor_allocation.h b/lib/local-execution/include/local-execution/tensor_allocation.h index ad2b4b0de5..aa89909447 100644 --- a/lib/local-execution/include/local-execution/tensor_allocation.h +++ b/lib/local-execution/include/local-execution/tensor_allocation.h @@ -16,8 +16,7 @@ DynamicValueAttrs perform_tensor_allocation_for_value(DynamicValueAttrs const &, DynamicOpenDataflowGraph perform_tensor_allocation( DynamicOpenDataflowGraph const &, - std::map const - &preallocated, + std::map const &preallocated, Allocator &); } // namespace FlexFlow diff --git a/lib/local-execution/src/local-execution/computation_graph_instance.cc b/lib/local-execution/src/local-execution/computation_graph_instance.cc index 5dbb8768be..2882199cf7 100644 --- a/lib/local-execution/src/local-execution/computation_graph_instance.cc +++ b/lib/local-execution/src/local-execution/computation_graph_instance.cc @@ -14,8 +14,8 @@ #include "task-spec/dynamic_graph/update_insertion.h" #include "task-spec/per_device_op_state.h" #include "task-spec/task_argument_accessor/task_argument_accessor.h" -#include "utils/containers/transform.h" #include "utils/containers/map_from_pairs.h" +#include "utils/containers/transform.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" #include "utils/optional.h" #include @@ -62,8 +62,7 @@ ComputationGraphInstance create_computation_graph_instance( ComputationGraph const &cg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::map const - &input_tensors, + std::map const &input_tensors, Allocator &allocator, ProfilingSettings const &profiling_settings, device_handle_t const &device_handle, @@ -71,8 +70,7 @@ ComputationGraphInstance create_computation_graph_instance( DynamicOpenDataflowGraph dg = make_dynamic_open_dataflow_graph_from_cg(cg); dg = perform_pass_expansion(dg); - std::map inputs = - input_tensors; + std::map inputs = input_tensors; std::optional logit_grad_value; if (loss.has_value()) { auto [loss_attrs, label_tensor, logit_tensor] = assert_unwrap(loss); @@ -144,8 +142,8 @@ std::map> global_device_id_t device_idx) { std::vector execution_order = instance.get_execution_order(); - std::map> - result = execute_dynamic_node_invocation_set( + std::map> result = + execute_dynamic_node_invocation_set( /*invocations=*/execution_order, /*allocator=*/instance.get_allocator(), /*optimizer_attrs=*/instance.get_optimizer_attrs(), diff --git a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc index f238016d1f..42df950a5c 100644 --- a/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc +++ b/lib/local-execution/src/local-execution/cost_estimator/local_cost_estimator.cc @@ -16,9 +16,9 @@ #include "utils/containers/map_values.h" #include "utils/containers/maximum.h" #include "utils/containers/require_only_key.h" +#include "utils/containers/set_of.h" #include "utils/containers/sum.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/exception.h" #include "utils/optional.h" @@ -129,15 +129,15 @@ OpCostMetrics LocalCostEstimator::estimate_cost( // execute layer dynamic_layer_guid_t operator_layer_guid{get_layer_by_name(cg, "operator")}; - std::map> - fwd_timing = perform_forward_pass_for_computation_graph_instance( + std::map> fwd_timing = + perform_forward_pass_for_computation_graph_instance( instance, this->profiling_settings, this->device_handle, this->device_idx); milliseconds_t fwd = fwd_timing.at(operator_layer_guid).value(); - std::map> - bwd_timing = perform_backward_pass_for_computation_graph_instance( + std::map> bwd_timing = + perform_backward_pass_for_computation_graph_instance( instance, this->profiling_settings, this->device_handle, diff --git a/lib/local-execution/src/local-execution/task_execution.cc b/lib/local-execution/src/local-execution/task_execution.cc index 330712a0d3..8cca67da43 100644 --- a/lib/local-execution/src/local-execution/task_execution.cc +++ b/lib/local-execution/src/local-execution/task_execution.cc @@ -56,8 +56,8 @@ TaskArgumentAccessor make_task_argument_accessor_for_invocation( auto get_accessor = [](DynamicValueAttrs const &value) { return assert_unwrap(value.accessor); }; - std::map - tensor_slots_backing = binary_merge_disjoint_maps( + std::map tensor_slots_backing = + binary_merge_disjoint_maps( map_keys_and_values(invocation.inputs, make_param, get_accessor), map_keys_and_values(invocation.outputs, make_param, get_accessor)); diff --git a/lib/local-execution/src/local-execution/tensor_allocation.cc b/lib/local-execution/src/local-execution/tensor_allocation.cc index bf3e4de4f4..544b2f00ef 100644 --- a/lib/local-execution/src/local-execution/tensor_allocation.cc +++ b/lib/local-execution/src/local-execution/tensor_allocation.cc @@ -50,8 +50,7 @@ DynamicValueAttrs DynamicOpenDataflowGraph perform_tensor_allocation( DynamicOpenDataflowGraph const &g, - std::map const - &preallocated, + std::map const &preallocated, Allocator &allocator) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); @@ -59,8 +58,7 @@ DynamicOpenDataflowGraph perform_tensor_allocation( ASSERT(v.accessor == std::nullopt); } - std::set all_values = - set_of(get_dynamic_values(g)); + std::set all_values = set_of(get_dynamic_values(g)); bidict unallocated_to_allocated = generate_bidict( diff --git a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc index 44f4d2cfd1..31f7bc3e30 100644 --- a/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc +++ b/lib/local-execution/test/src/local-execution/local_task_argument_accessor.cc @@ -39,8 +39,8 @@ TEST_SUITE(FF_TEST_SUITE) { VARIADIC_TENSORS, }; - std::map - tensor_slots_backing = { + std::map tensor_slots_backing = + { { make_task_tensor_parameter_fwd(TensorSlotName::LHS_INPUT), DynamicTensorAccessor{input}, diff --git a/lib/op-attrs/include/op-attrs/get_operator_task_space.h b/lib/op-attrs/include/op-attrs/get_operator_task_space.h index 462cdddc4d..c2a3cb31c2 100644 --- a/lib/op-attrs/include/op-attrs/get_operator_task_space.h +++ b/lib/op-attrs/include/op-attrs/get_operator_task_space.h @@ -10,8 +10,7 @@ namespace FlexFlow { OperatorTaskSpace get_operator_task_space( ComputationGraphOpAttrs const &attrs, - std::map const - &inputs_degrees); + std::map const &inputs_degrees); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/initializer_attrs.h b/lib/op-attrs/include/op-attrs/initializer_attrs.h index 4fba656cec..639fe09e0c 100644 --- a/lib/op-attrs/include/op-attrs/initializer_attrs.h +++ b/lib/op-attrs/include/op-attrs/initializer_attrs.h @@ -8,11 +8,12 @@ namespace FlexFlow { InitializerAttrs make_zero_initializer(); InitializerAttrs make_kaiming_uniform( - TensorDims const &, - float a = 0.0, - KaimingInitializerMode mode = KaimingInitializerMode::FAN_IN, - KaimingInitializerNonlinearity nonlinearity = KaimingInitializerNonlinearity::LEAKY_RELU, - int seed = 0); + TensorDims const &, + float a = 0.0, + KaimingInitializerMode mode = KaimingInitializerMode::FAN_IN, + KaimingInitializerNonlinearity nonlinearity = + KaimingInitializerNonlinearity::LEAKY_RELU, + int seed = 0); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/ops/attention.h b/lib/op-attrs/include/op-attrs/ops/attention.h index 76f63f6780..0ef33dee5f 100644 --- a/lib/op-attrs/include/op-attrs/ops/attention.h +++ b/lib/op-attrs/include/op-attrs/ops/attention.h @@ -106,8 +106,7 @@ tl::expected ParallelTensorShape const &input_k, ParallelTensorShape const &input_v); -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(MultiHeadAttentionAttrs const &, ParallelTensorShape const &input_q, ParallelTensorShape const &input_k, diff --git a/lib/op-attrs/include/op-attrs/ops/batch_norm.h b/lib/op-attrs/include/op-attrs/ops/batch_norm.h index 9422c14c6c..50dcd4bb5c 100644 --- a/lib/op-attrs/include/op-attrs/ops/batch_norm.h +++ b/lib/op-attrs/include/op-attrs/ops/batch_norm.h @@ -36,8 +36,7 @@ tl::expected get_beta_weights_parallel_dim_degrees(BatchNormAttrs const &, ParallelTensorDimDegrees const &); -tl::expected, - std::string> +tl::expected, std::string> get_weight_parallel_dim_degrees( BatchNormAttrs const &attrs, ParallelTensorDimDegrees const &input_degrees); @@ -50,8 +49,7 @@ tl::expected tl::expected get_beta_weights_shape(BatchNormAttrs const &, ParallelTensorShape const &); -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(BatchNormAttrs const &attrs, ParallelTensorShape const &input_shape); diff --git a/lib/op-attrs/include/op-attrs/ops/layer_norm.h b/lib/op-attrs/include/op-attrs/ops/layer_norm.h index 7e1b06483d..1001a420f7 100644 --- a/lib/op-attrs/include/op-attrs/ops/layer_norm.h +++ b/lib/op-attrs/include/op-attrs/ops/layer_norm.h @@ -33,8 +33,7 @@ tl::expected tl::expected get_beta_weights_shape(LayerNormAttrs const &, ParallelTensorShape const &); -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(LayerNormAttrs const &attrs, ParallelTensorShape const &input_shape); diff --git a/lib/op-attrs/include/op-attrs/ops/linear.h b/lib/op-attrs/include/op-attrs/ops/linear.h index 817652abc6..281f4fc7cb 100644 --- a/lib/op-attrs/include/op-attrs/ops/linear.h +++ b/lib/op-attrs/include/op-attrs/ops/linear.h @@ -49,8 +49,7 @@ tl::expected get_output_shape(LinearAttrs const &attrs, ParallelTensorShape const &input); -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(LinearAttrs const &attrs, ParallelTensorShape const &input_shape); diff --git a/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h b/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h index e48798a860..9308a4fed6 100644 --- a/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h +++ b/lib/op-attrs/include/op-attrs/parallel_tensor_shape.h @@ -38,8 +38,7 @@ ParallelTensorShape TensorShape get_piece_shape(ParallelTensorShape const &); num_bytes_t get_piece_size_in_bytes(ParallelTensorShape const &); -std::set - replica_dims(ParallelTensorShape const &); +std::set replica_dims(ParallelTensorShape const &); positive_int get_num_replica_dims(ParallelTensorShape const &); positive_int get_num_replicas(ParallelTensorShape const &); diff --git a/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h b/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h index 02ec2d4611..ebeaa80edd 100644 --- a/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h +++ b/lib/op-attrs/include/op-attrs/replica_parallel_dim_set.h @@ -10,8 +10,7 @@ namespace FlexFlow { ReplicaParallelDimSet empty_replica_parallel_dim_set(); positive_int get_degree_of_replica_type(ReplicaParallelDimSet const &, ReplicaType); -std::set - get_replica_dims(ReplicaParallelDimSet const &); +std::set get_replica_dims(ReplicaParallelDimSet const &); } // namespace FlexFlow diff --git a/lib/op-attrs/include/op-attrs/shape_inference.h b/lib/op-attrs/include/op-attrs/shape_inference.h index 37c0c8536a..aca65c1e63 100644 --- a/lib/op-attrs/include/op-attrs/shape_inference.h +++ b/lib/op-attrs/include/op-attrs/shape_inference.h @@ -19,13 +19,11 @@ std::map get_weight_shapes( std::map get_output_shapes( PCGOperatorAttrs const &, - std::map const - &input_shapes); + std::map const &input_shapes); std::map get_weight_shapes( PCGOperatorAttrs const &, - std::map const - &input_shapes); + std::map const &input_shapes); } // namespace FlexFlow diff --git a/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc b/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc index c9f851369e..79d29a61f6 100644 --- a/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc +++ b/lib/op-attrs/src/op-attrs/ff_ordered/ff_ordered_from_map.cc @@ -5,7 +5,8 @@ namespace FlexFlow { using T = value_type<0>; -template FFOrdered ff_ordered_from_map(std::unordered_map const &); +template FFOrdered + ff_ordered_from_map(std::unordered_map const &); template FFOrdered ff_ordered_from_map(std::map const &); diff --git a/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc b/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc index f698dce0c2..91806b937c 100644 --- a/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc +++ b/lib/op-attrs/src/op-attrs/ff_ordered/map_from_ff_ordered.cc @@ -5,7 +5,6 @@ namespace FlexFlow { using T = value_type<0>; -template std::map - map_from_ff_ordered(FFOrdered const &); +template std::map map_from_ff_ordered(FFOrdered const &); } // namespace FlexFlow diff --git a/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc b/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc index 1df85a7134..a12fc32e96 100644 --- a/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc +++ b/lib/op-attrs/src/op-attrs/get_incoming_tensor_roles.cc @@ -10,17 +10,16 @@ namespace FlexFlow { -std::map - get_incoming_tensor_roles( - ComputationGraphOpAttrs const &comp_graph_op_attrs) { +std::map get_incoming_tensor_roles( + ComputationGraphOpAttrs const &comp_graph_op_attrs) { return get_incoming_tensor_roles( pcg_op_attrs_from_compgraph_op_attrs(comp_graph_op_attrs)); } std::map get_incoming_tensor_roles(PCGOperatorAttrs const &pcg_op_attrs) { - return pcg_op_attrs - .visit>(overload{ + return pcg_op_attrs.visit>( + overload{ [](BatchNormAttrs const &attrs) { return get_batch_norm_incoming_tensor_roles(attrs); }, diff --git a/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc b/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc index 0e180dc820..a51c0609e9 100644 --- a/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc +++ b/lib/op-attrs/src/op-attrs/get_operator_space_to_parallel_tensor_space_mappings.cc @@ -8,11 +8,11 @@ #include "op-attrs/ops/weight.h" #include "utils/containers/filtrans.h" #include "utils/containers/get_only.h" +#include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/require_only_key.h" #include "utils/containers/require_two_keys.h" #include "utils/containers/zip_values_strict.h" #include "utils/overload.h" -#include "utils/containers/merge_disjoint_maps.h" namespace FlexFlow { @@ -22,97 +22,98 @@ std::map std::map const &inputs_degrees) { return comp_graph_op_attrs.visit< - std::map>(overload{ - [&](ElementBinaryAttrs const &attrs) - -> std::map { - ASSERT(inputs_degrees.size() == 2); + std::map>( + overload{ + [&](ElementBinaryAttrs const &attrs) + -> std::map { + ASSERT(inputs_degrees.size() == 2); - ParallelTensorDimDegrees lhs_degrees = - inputs_degrees.at(TensorSlotName::LHS_INPUT); - ParallelTensorDimDegrees rhs_degrees = - inputs_degrees.at(TensorSlotName::RHS_INPUT); + ParallelTensorDimDegrees lhs_degrees = + inputs_degrees.at(TensorSlotName::LHS_INPUT); + ParallelTensorDimDegrees rhs_degrees = + inputs_degrees.at(TensorSlotName::RHS_INPUT); - return { - { - TensorSlotName::LHS_INPUT, - get_operator_to_lhs_input_mapping( - attrs, lhs_degrees, rhs_degrees), - }, - { - TensorSlotName::RHS_INPUT, - get_operator_to_rhs_input_mapping( - attrs, lhs_degrees, rhs_degrees), - }, - }; - }, - [&](ElementUnaryAttrs const &attrs) - -> std::map { - ParallelTensorDimDegrees input_degrees = - require_only_key(inputs_degrees, TensorSlotName::INPUT); + return { + { + TensorSlotName::LHS_INPUT, + get_operator_to_lhs_input_mapping( + attrs, lhs_degrees, rhs_degrees), + }, + { + TensorSlotName::RHS_INPUT, + get_operator_to_rhs_input_mapping( + attrs, lhs_degrees, rhs_degrees), + }, + }; + }, + [&](ElementUnaryAttrs const &attrs) + -> std::map { + ParallelTensorDimDegrees input_degrees = + require_only_key(inputs_degrees, TensorSlotName::INPUT); - return { - { - TensorSlotName::INPUT, - get_operator_to_input_mapping(attrs, input_degrees), - }, - }; - }, - [&](InputAttrs const &) { - ASSERT(inputs_degrees.size() == 0); + return { + { + TensorSlotName::INPUT, + get_operator_to_input_mapping(attrs, input_degrees), + }, + }; + }, + [&](InputAttrs const &) { + ASSERT(inputs_degrees.size() == 0); - return std::map{}; - }, - [&](LinearAttrs const &attrs) - -> std::map { - ParallelTensorDimDegrees input_degrees = - require_only_key(inputs_degrees, TensorSlotName::INPUT); + return std::map{}; + }, + [&](LinearAttrs const &attrs) + -> std::map { + ParallelTensorDimDegrees input_degrees = + require_only_key(inputs_degrees, TensorSlotName::INPUT); - std::map - result = { - {TensorSlotName::INPUT, - get_operator_to_input_mapping(attrs, input_degrees)}, - {TensorSlotName::WEIGHT, - get_operator_to_projection_mapping(attrs, input_degrees)}, - }; + std::map + result = { + {TensorSlotName::INPUT, + get_operator_to_input_mapping(attrs, input_degrees)}, + {TensorSlotName::WEIGHT, + get_operator_to_projection_mapping(attrs, input_degrees)}, + }; - if (attrs.use_bias) { - result.insert({TensorSlotName::BIAS, - get_operator_to_bias_mapping(attrs, input_degrees)}); - }; + if (attrs.use_bias) { + result.insert( + {TensorSlotName::BIAS, + get_operator_to_bias_mapping(attrs, input_degrees)}); + }; - return result; - }, - [&](TransposeAttrs const &attrs) - -> std::map { - ParallelTensorDimDegrees input_degrees = - require_only_key(inputs_degrees, TensorSlotName::INPUT); + return result; + }, + [&](TransposeAttrs const &attrs) + -> std::map { + ParallelTensorDimDegrees input_degrees = + require_only_key(inputs_degrees, TensorSlotName::INPUT); - return { - { - TensorSlotName::INPUT, - get_operator_to_input_mapping(attrs, input_degrees), - }, - }; - }, - [&](WeightAttrs const &) { - ASSERT(inputs_degrees.size() == 0); + return { + { + TensorSlotName::INPUT, + get_operator_to_input_mapping(attrs, input_degrees), + }, + }; + }, + [&](WeightAttrs const &) { + ASSERT(inputs_degrees.size() == 0); - return std::map{}; - }, - [](auto const &attrs) - -> std::map { - PANIC("Missing implmentation of get_operator_to_input_mappings", attrs); - }, - }); + return std::map{}; + }, + [](auto const &attrs) + -> std::map { + PANIC("Missing implmentation of get_operator_to_input_mappings", + attrs); + }, + }); } std::map @@ -170,92 +171,94 @@ std::map &inputs_degrees) { return comp_graph_op_attrs.visit< - std::map>(overload{ - [&](ElementBinaryAttrs const &attrs) - -> std::map { - auto [lhs_degrees, rhs_degrees] = - require_two_keys(inputs_degrees, - TensorSlotName::LHS_INPUT, - TensorSlotName::RHS_INPUT); + std::map>( + overload{ + [&](ElementBinaryAttrs const &attrs) + -> std::map { + auto [lhs_degrees, rhs_degrees] = + require_two_keys(inputs_degrees, + TensorSlotName::LHS_INPUT, + TensorSlotName::RHS_INPUT); - return { - { - TensorSlotName::OUTPUT, - get_operator_to_output_mapping(attrs, lhs_degrees, rhs_degrees), - }, - }; - }, - [&](ElementUnaryAttrs const &attrs) - -> std::map { - ParallelTensorDimDegrees input_degrees = - require_only_key(inputs_degrees, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + get_operator_to_output_mapping( + attrs, lhs_degrees, rhs_degrees), + }, + }; + }, + [&](ElementUnaryAttrs const &attrs) + -> std::map { + ParallelTensorDimDegrees input_degrees = + require_only_key(inputs_degrees, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_operator_to_output_mapping(attrs, input_degrees), - }, - }; - }, - [&](LinearAttrs const &attrs) - -> std::map { - ParallelTensorDimDegrees input_degrees = - require_only_key(inputs_degrees, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + get_operator_to_output_mapping(attrs, input_degrees), + }, + }; + }, + [&](LinearAttrs const &attrs) + -> std::map { + ParallelTensorDimDegrees input_degrees = + require_only_key(inputs_degrees, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_operator_to_output_mapping(attrs, input_degrees), - }, - }; - }, - [&](InputAttrs const &attrs) - -> std::map { - ASSERT(inputs_degrees.size() == 0); + return { + { + TensorSlotName::OUTPUT, + get_operator_to_output_mapping(attrs, input_degrees), + }, + }; + }, + [&](InputAttrs const &attrs) + -> std::map { + ASSERT(inputs_degrees.size() == 0); - return { - { - TensorSlotName::OUTPUT, - get_operator_to_output_mapping(attrs), - }, - }; - }, - [&](TransposeAttrs const &attrs) - -> std::map { - ParallelTensorDimDegrees input_degrees = - require_only_key(inputs_degrees, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + get_operator_to_output_mapping(attrs), + }, + }; + }, + [&](TransposeAttrs const &attrs) + -> std::map { + ParallelTensorDimDegrees input_degrees = + require_only_key(inputs_degrees, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_operator_to_output_mapping(attrs, input_degrees), - }, - }; - }, - [&](WeightAttrs const &attrs) - -> std::map { - ASSERT(inputs_degrees.size() == 0); + return { + { + TensorSlotName::OUTPUT, + get_operator_to_output_mapping(attrs, input_degrees), + }, + }; + }, + [&](WeightAttrs const &attrs) + -> std::map { + ASSERT(inputs_degrees.size() == 0); - return { - { - TensorSlotName::OUTPUT, - get_operator_to_output_mapping(attrs), - }, - }; - }, - [](auto const &attrs) - -> std::map { - PANIC("Missing implmentation of get_operator_to_input_mappings", attrs); - }, - }); + return { + { + TensorSlotName::OUTPUT, + get_operator_to_output_mapping(attrs), + }, + }; + }, + [](auto const &attrs) + -> std::map { + PANIC("Missing implmentation of get_operator_to_input_mappings", + attrs); + }, + }); } std::map diff --git a/lib/op-attrs/src/op-attrs/get_operator_task_space.cc b/lib/op-attrs/src/op-attrs/get_operator_task_space.cc index f6b0733328..ce81aea3b9 100644 --- a/lib/op-attrs/src/op-attrs/get_operator_task_space.cc +++ b/lib/op-attrs/src/op-attrs/get_operator_task_space.cc @@ -15,8 +15,7 @@ namespace FlexFlow { OperatorTaskSpace get_operator_task_space( ComputationGraphOpAttrs const &attrs, - std::map const - &inputs_degrees) { + std::map const &inputs_degrees) { return attrs.visit(overload{ [&](ElementUnaryAttrs const &attrs) { ParallelTensorDimDegrees input = diff --git a/lib/op-attrs/src/op-attrs/initializer_attrs.cc b/lib/op-attrs/src/op-attrs/initializer_attrs.cc index c0dc827ea3..8a41106782 100644 --- a/lib/op-attrs/src/op-attrs/initializer_attrs.cc +++ b/lib/op-attrs/src/op-attrs/initializer_attrs.cc @@ -46,11 +46,12 @@ static float // from pytorch: // see // https://github.com/pytorch/pytorch/blob/bd019c0bb485904a99fb38589444b1461ab1e486/torch/nn/init.py#L456-L518 -InitializerAttrs make_kaiming_uniform(TensorDims const &dims, - float a, - KaimingInitializerMode mode, - KaimingInitializerNonlinearity nonlinearity, - int seed) { +InitializerAttrs + make_kaiming_uniform(TensorDims const &dims, + float a, + KaimingInitializerMode mode, + KaimingInitializerNonlinearity nonlinearity, + int seed) { positive_int fan = calculate_fan_for_mode(dims, mode); float gain = gain_for_nonlinearity(nonlinearity, a); diff --git a/lib/op-attrs/src/op-attrs/operator_task_space.cc b/lib/op-attrs/src/op-attrs/operator_task_space.cc index ef7c8dde09..37793739a2 100644 --- a/lib/op-attrs/src/op-attrs/operator_task_space.cc +++ b/lib/op-attrs/src/op-attrs/operator_task_space.cc @@ -10,8 +10,8 @@ #include "utils/containers/maximum.h" #include "utils/containers/product.h" #include "utils/containers/range.h" -#include "utils/containers/transform.h" #include "utils/containers/set_of.h" +#include "utils/containers/transform.h" #include "utils/containers/vector_of.h" #include "utils/fmt/set.h" #include "utils/nonnegative_int/nonnegative_range.h" diff --git a/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc b/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc index 3d70c366a1..5c66947fca 100644 --- a/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc +++ b/lib/op-attrs/src/op-attrs/operator_task_space_to_operator_task_space_mapping.cc @@ -40,9 +40,10 @@ OperatorTaskSpace op_mapping_get_dst_space( bidict op_to_op_get_coord_mapping( OperatorTaskSpaceToOperatorTaskSpaceMapping const &mapping) { - return bidict_transform_values(bidict_transform_keys(mapping.raw_mapping.coord_mapping, - task_space_coordinate_from_dim_coord), - task_space_coordinate_from_dim_coord); + return bidict_transform_values( + bidict_transform_keys(mapping.raw_mapping.coord_mapping, + task_space_coordinate_from_dim_coord), + task_space_coordinate_from_dim_coord); } OperatorTaskSpaceToOperatorTaskSpaceMapping diff --git a/lib/op-attrs/src/op-attrs/ops/attention.cc b/lib/op-attrs/src/op-attrs/ops/attention.cc index aa86b9e673..f3b4c7293f 100644 --- a/lib/op-attrs/src/op-attrs/ops/attention.cc +++ b/lib/op-attrs/src/op-attrs/ops/attention.cc @@ -416,8 +416,7 @@ positive_int get_oSize(TensorShape const &) { NOT_IMPLEMENTED(); } -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(MultiHeadAttentionAttrs const &attrs, ParallelTensorShape const &input_q, ParallelTensorShape const &input_k, diff --git a/lib/op-attrs/src/op-attrs/ops/batch_norm.cc b/lib/op-attrs/src/op-attrs/ops/batch_norm.cc index d045b829d4..57024fdb64 100644 --- a/lib/op-attrs/src/op-attrs/ops/batch_norm.cc +++ b/lib/op-attrs/src/op-attrs/ops/batch_norm.cc @@ -210,8 +210,7 @@ tl::expected return get_gamma_weights_parallel_dim_degrees(attrs, input_degrees); } -tl::expected, - std::string> +tl::expected, std::string> get_weight_parallel_dim_degrees( BatchNormAttrs const &attrs, ParallelTensorDimDegrees const &input_degrees) { @@ -310,8 +309,7 @@ tl::expected return lift_to_parallel_with_degrees(unpar, degrees); } -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(BatchNormAttrs const &attrs, ParallelTensorShape const &input_shape) { diff --git a/lib/op-attrs/src/op-attrs/ops/layer_norm.cc b/lib/op-attrs/src/op-attrs/ops/layer_norm.cc index c732ace44e..cdc4f32a6f 100644 --- a/lib/op-attrs/src/op-attrs/ops/layer_norm.cc +++ b/lib/op-attrs/src/op-attrs/ops/layer_norm.cc @@ -220,8 +220,7 @@ tl::expected return get_gamma_weights_shape(attrs, input_shape); } -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(LayerNormAttrs const &attrs, ParallelTensorShape const &input_shape) { diff --git a/lib/op-attrs/src/op-attrs/ops/linear.cc b/lib/op-attrs/src/op-attrs/ops/linear.cc index 358929f1af..476b0dd384 100644 --- a/lib/op-attrs/src/op-attrs/ops/linear.cc +++ b/lib/op-attrs/src/op-attrs/ops/linear.cc @@ -205,8 +205,7 @@ ParallelTensorDimDegrees }; } -tl::expected, - std::string> +tl::expected, std::string> get_weight_shapes(LinearAttrs const &attrs, ParallelTensorShape const &input_shape) { diff --git a/lib/op-attrs/src/op-attrs/ops/transpose.cc b/lib/op-attrs/src/op-attrs/ops/transpose.cc index 4aa58aa2d3..3710c7e8c7 100644 --- a/lib/op-attrs/src/op-attrs/ops/transpose.cc +++ b/lib/op-attrs/src/op-attrs/ops/transpose.cc @@ -52,7 +52,8 @@ static ParallelTensorSpaceToParallelTensorSpaceMapping EqProjection inp_to_out = EqProjection{ bidict_transform_keys( - bidict_transform_values(attrs.permutation.as_bidict(), ff_dim_to_pt_dim), + bidict_transform_values(attrs.permutation.as_bidict(), + ff_dim_to_pt_dim), ff_dim_to_pt_dim), }; diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc index a334fac056..49e15d1fb7 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_dim_degrees.cc @@ -5,6 +5,7 @@ #include "op-attrs/parallel_tensor_dim_idx_t.dtg.h" #include "op-attrs/parallel_tensor_dim_idx_t.h" #include "op-attrs/parallel_tensor_space_coordinate.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/filtrans.h" #include "utils/containers/generate_map.h" @@ -12,14 +13,12 @@ #include "utils/containers/map_keys.h" #include "utils/containers/map_values.h" #include "utils/containers/range.h" +#include "utils/containers/set_of.h" #include "utils/containers/set_union.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/nonnegative_int/nonnegative_range.h" #include "utils/nonnegative_int/num_elements.h" #include "utils/orthotope/minimal_dim_domain.h" -#include "utils/containers/binary_merge_disjoint_maps.h" -#include "utils/containers/generate_map.h" namespace FlexFlow { @@ -88,13 +87,11 @@ positive_int get_degree_for_parallel_tensor_dim_idx( std::map get_parallel_tensor_degree_map(ParallelTensorDimDegrees const °rees) { - std::map - replica_dim_degrees = { - {parallel_tensor_dim_idx_t{ReplicaType::SUM}, - degrees.sum_degree.value}, - {parallel_tensor_dim_idx_t{ReplicaType::DISCARD_COPY}, - degrees.discard_copy_degree.value}, - }; + std::map replica_dim_degrees = { + {parallel_tensor_dim_idx_t{ReplicaType::SUM}, degrees.sum_degree.value}, + {parallel_tensor_dim_idx_t{ReplicaType::DISCARD_COPY}, + degrees.discard_copy_degree.value}, + }; std::map shard_dim_degrees = generate_map(get_idxs(degrees.shard_degrees), [&](ff_dim_t const &dim) { @@ -108,23 +105,22 @@ std::map })); } -std::set - get_parallel_tensor_space_coordinates( - ParallelTensorDimDegrees const °rees) { +std::set get_parallel_tensor_space_coordinates( + ParallelTensorDimDegrees const °rees) { std::map degree_map = get_parallel_tensor_degree_map(degrees); - std::map> + std::map> possible_per_dim_coords = map_values(degree_map, [](positive_int degree) { return set_of(nonnegative_range(degree)); }); return transform( get_all_assignments(possible_per_dim_coords), - [](std::map const - &m) { return parallel_tensor_space_coord_from_map(m); }); + [](std::map const &m) { + return parallel_tensor_space_coord_from_map(m); + }); } DimDomain diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc index 73fd7fcbec..5433f02a75 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_dims.cc @@ -24,8 +24,7 @@ FFOrdered ff_ordered_shard_degrees(ParallelTensorDims const &d) { [](ShardParallelDim const &d) { return d.degree; }); } -std::set - replica_dims(ParallelTensorDims const &d) { +std::set replica_dims(ParallelTensorDims const &d) { return get_replica_dims(d.replica_dims); } diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc index 86788671a9..64e5f421a6 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_shape.cc @@ -19,8 +19,7 @@ num_ptensor_shard_dims_t num_shard_dims(ParallelTensorShape const &s) { return num_shard_dims(s.dims); } -std::set - replica_dims(ParallelTensorShape const &s) { +std::set replica_dims(ParallelTensorShape const &s) { return replica_dims(s.dims); } diff --git a/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc b/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc index 101ac648c6..f18409b5da 100644 --- a/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc +++ b/lib/op-attrs/src/op-attrs/parallel_tensor_space_coordinate.cc @@ -3,8 +3,8 @@ #include "op-attrs/parallel_tensor_dim_idx_t.h" #include "utils/containers/contains_key.h" #include "utils/containers/filtermap_keys.h" -#include "utils/nonnegative_int/num_elements.h" #include "utils/containers/generate_map.h" +#include "utils/nonnegative_int/num_elements.h" namespace FlexFlow { @@ -22,11 +22,10 @@ num_ptensor_shard_dims_t }; } -std::set - get_dim_idxs_in_ptensor_space_coord( - ParallelTensorSpaceCoordinate const &coord) { +std::set get_dim_idxs_in_ptensor_space_coord( + ParallelTensorSpaceCoordinate const &coord) { - std::set result = + std::set result = dim_idxs_for_num_shard_dims(ptensor_coord_num_shard_dims(coord)); result.insert(sum_dim_idx()); result.insert(discard_copy_dim_idx()); diff --git a/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc b/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc index 36b24d8eaf..bc96838341 100644 --- a/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc +++ b/lib/op-attrs/src/op-attrs/replica_parallel_dim_set.cc @@ -20,8 +20,7 @@ positive_int get_degree_of_replica_type(ReplicaParallelDimSet const &s, } } -std::set - get_replica_dims(ReplicaParallelDimSet const &s) { +std::set get_replica_dims(ReplicaParallelDimSet const &s) { return std::set{ ReplicaParallelDim{s.sum_degree.value, ReplicaType::SUM}, ReplicaParallelDim{s.discard_copy_degree.value, diff --git a/lib/op-attrs/src/op-attrs/shape_inference.cc b/lib/op-attrs/src/op-attrs/shape_inference.cc index 0d1e0ee82f..49ea93b514 100644 --- a/lib/op-attrs/src/op-attrs/shape_inference.cc +++ b/lib/op-attrs/src/op-attrs/shape_inference.cc @@ -31,11 +31,10 @@ namespace FlexFlow { template -static std::tuple - require_3(std::map const &v, - TensorSlotName k1, - TensorSlotName k2, - TensorSlotName k3) { +static std::tuple require_3(std::map const &v, + TensorSlotName k1, + TensorSlotName k2, + TensorSlotName k3) { ASSERT(v.size() == 3); return {v.at(k1), v.at(k2), v.at(k3)}; @@ -61,359 +60,328 @@ static std::vector std::map get_output_shapes( ComputationGraphOpAttrs const &op_attrs, std::map const &input_shapes) { - return op_attrs.visit>( - overload{ - [&](BatchNormAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - - return { - { - TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, input)), - }, - }; - }, - [&](CastAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - - return { - { - TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, input)), - }, - }; - }, - [&](ConcatAttrs const &attrs) - -> std::map { - std::vector inputs = require_only_slots_sequence( - input_shapes, get_variadic_inputs_slot_name_sequence()); - - return { - { - TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, inputs)), - }, - }; - }, - [&](Conv2DAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs, input), - }, - }; - }, - [&](DropoutAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs, input), - }, - }; - }, - [&](ElementBinaryAttrs const &attrs) - -> std::map { - auto [lhs, rhs] = require_two_keys(input_shapes, - TensorSlotName::LHS_INPUT, - TensorSlotName::RHS_INPUT); - - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs, lhs, rhs), - }, - }; - }, - [&](ElementUnaryAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return op_attrs.visit>(overload{ + [&](BatchNormAttrs const &attrs) + -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs, input), - }, - }; - }, - [&](EmbeddingAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, input)), + }, + }; + }, + [&](CastAttrs const &attrs) -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, input)), - }, - }; - }, - [&](FlatAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, input)), + }, + }; + }, + [&](ConcatAttrs const &attrs) -> std::map { + std::vector inputs = require_only_slots_sequence( + input_shapes, get_variadic_inputs_slot_name_sequence()); + + return { + { + TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, inputs)), + }, + }; + }, + [&](Conv2DAttrs const &attrs) -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs, input), - }, - }; - }, - [&](GatherAttrs const &attrs) - -> std::map { - auto [input, index] = require_two_keys( - input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs, input), + }, + }; + }, + [&](DropoutAttrs const &attrs) -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs, input, index), - }, - }; - }, - [&](InputAttrs const &attrs) - -> std::map { - ASSERT(input_shapes.size() == 0); + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs, input), + }, + }; + }, + [&](ElementBinaryAttrs const &attrs) + -> std::map { + auto [lhs, rhs] = require_two_keys( + input_shapes, TensorSlotName::LHS_INPUT, TensorSlotName::RHS_INPUT); + + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs, lhs, rhs), + }, + }; + }, + [&](ElementUnaryAttrs const &attrs) + -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs), - }, - }; - }, - [&](LayerNormAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs, input), + }, + }; + }, + [&](EmbeddingAttrs const &attrs) + -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, input)), - }, - }; - }, - [&](LinearAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, input)), + }, + }; + }, + [&](FlatAttrs const &attrs) -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, input)), - }, - }; - }, - [&](MultiHeadAttentionAttrs const &attrs) - -> std::map { - auto [query, key, value] = require_3(input_shapes, - TensorSlotName::QUERY, - TensorSlotName::KEY, - TensorSlotName::VALUE); + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs, input), + }, + }; + }, + [&](GatherAttrs const &attrs) -> std::map { + auto [input, index] = require_two_keys( + input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); + + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs, input, index), + }, + }; + }, + [&](InputAttrs const &attrs) -> std::map { + ASSERT(input_shapes.size() == 0); + + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs), + }, + }; + }, + [&](LayerNormAttrs const &attrs) + -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - {TensorSlotName::OUTPUT, - throw_if_unexpected( - get_output_shape(attrs, query, key, value))}, - }; - }, - [&](Pool2DAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, input)), + }, + }; + }, + [&](LinearAttrs const &attrs) -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - {TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, input))}, - }; - }, - [&](SoftmaxAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return { + { + TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, input)), + }, + }; + }, + [&](MultiHeadAttentionAttrs const &attrs) + -> std::map { + auto [query, key, value] = require_3(input_shapes, + TensorSlotName::QUERY, + TensorSlotName::KEY, + TensorSlotName::VALUE); + + return { + {TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, query, key, value))}, + }; + }, + [&](Pool2DAttrs const &attrs) -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - {TensorSlotName::OUTPUT, - throw_if_unexpected(get_output_shape(attrs, input))}, - }; - }, - [&](TransposeAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); + return { + {TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, input))}, + }; + }, + [&](SoftmaxAttrs const &attrs) -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs, input), - }, - }; - }, - [&](WeightAttrs const &attrs) - -> std::map { - ASSERT(input_shapes.size() == 0); + return { + {TensorSlotName::OUTPUT, + throw_if_unexpected(get_output_shape(attrs, input))}, + }; + }, + [&](TransposeAttrs const &attrs) + -> std::map { + TensorShape input = + require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::OUTPUT, - get_output_shape(attrs), - }, - }; - }, - [&](auto const &attrs) - -> std::map { - NOT_IMPLEMENTED(); - }, - }); + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs, input), + }, + }; + }, + [&](WeightAttrs const &attrs) -> std::map { + ASSERT(input_shapes.size() == 0); + + return { + { + TensorSlotName::OUTPUT, + get_output_shape(attrs), + }, + }; + }, + [&](auto const &attrs) -> std::map { + NOT_IMPLEMENTED(); + }, + }); } std::map get_weight_shapes( ComputationGraphOpAttrs const &op_attrs, std::map const &input_shapes) { - return op_attrs.visit>( - overload{ - [&](BatchNormAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - - return throw_if_unexpected(get_weight_shapes(attrs, input)); - }, - [&](CastAttrs const &attrs) - -> std::map { + return op_attrs.visit>(overload{ + [&](BatchNormAttrs const &attrs) + -> std::map { + TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); - return {}; - }, - [&](ConcatAttrs const &attrs) - -> std::map { - require_only_slots_sequence( - input_shapes, get_variadic_inputs_slot_name_sequence()); - - return {}; - }, - [&](Conv2DAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - return get_weight_shapes(attrs, input); - }, - [&](DropoutAttrs const &attrs) - -> std::map { - require_only_key(input_shapes, TensorSlotName::INPUT); - return {}; - }, - [&](ElementBinaryAttrs const &attrs) - -> std::map { - require_two_keys(input_shapes, - TensorSlotName::LHS_INPUT, - TensorSlotName::RHS_INPUT); - return {}; - }, - [&](ElementUnaryAttrs const &attrs) - -> std::map { + return throw_if_unexpected(get_weight_shapes(attrs, input)); + }, + [&](CastAttrs const &attrs) -> std::map { + require_only_key(input_shapes, TensorSlotName::INPUT); + return {}; + }, + [&](ConcatAttrs const &attrs) -> std::map { + require_only_slots_sequence(input_shapes, + get_variadic_inputs_slot_name_sequence()); + + return {}; + }, + [&](Conv2DAttrs const &attrs) -> std::map { + TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); - return {}; - }, - [&](EmbeddingAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - return { - { - TensorSlotName::WEIGHT, - TensorShape{ - throw_if_unexpected(get_weights_shape(attrs, input)), - }, - }, - }; - }, - [&](FlatAttrs const &attrs) - -> std::map { + return get_weight_shapes(attrs, input); + }, + [&](DropoutAttrs const &attrs) -> std::map { + require_only_key(input_shapes, TensorSlotName::INPUT); + return {}; + }, + [&](ElementBinaryAttrs const &attrs) + -> std::map { + require_two_keys( + input_shapes, TensorSlotName::LHS_INPUT, TensorSlotName::RHS_INPUT); + return {}; + }, + [&](ElementUnaryAttrs const &attrs) + -> std::map { + require_only_key(input_shapes, TensorSlotName::INPUT); + return {}; + }, + [&](EmbeddingAttrs const &attrs) + -> std::map { + TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); - return {}; - }, - [&](GatherAttrs const &attrs) - -> std::map { - require_two_keys( - input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); - return {}; - }, - [&](InputAttrs const &attrs) - -> std::map { - ASSERT(input_shapes.size() == 0); - return {}; - }, - [&](LayerNormAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - return throw_if_unexpected(get_weight_shapes(attrs, input)); - }, - [&](LinearAttrs const &attrs) - -> std::map { - TensorShape input = - require_only_key(input_shapes, TensorSlotName::INPUT); - - return throw_if_unexpected(get_weight_shapes(attrs, input)); - }, - [&](MultiHeadAttentionAttrs const &attrs) - -> std::map { - auto [query, key, value] = require_3(input_shapes, - TensorSlotName::QUERY, - TensorSlotName::KEY, - TensorSlotName::VALUE); - - return throw_if_unexpected( - get_weight_shapes(attrs, query, key, value)); - }, - [&](Pool2DAttrs const &attrs) - -> std::map { + return { + { + TensorSlotName::WEIGHT, + TensorShape{ + throw_if_unexpected(get_weights_shape(attrs, input)), + }, + }, + }; + }, + [&](FlatAttrs const &attrs) -> std::map { + require_only_key(input_shapes, TensorSlotName::INPUT); + return {}; + }, + [&](GatherAttrs const &attrs) -> std::map { + require_two_keys( + input_shapes, TensorSlotName::INPUT, TensorSlotName::INDEX); + return {}; + }, + [&](InputAttrs const &attrs) -> std::map { + ASSERT(input_shapes.size() == 0); + return {}; + }, + [&](LayerNormAttrs const &attrs) + -> std::map { + TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); - return {}; - }, - [&](SoftmaxAttrs const &attrs) - -> std::map { + return throw_if_unexpected(get_weight_shapes(attrs, input)); + }, + [&](LinearAttrs const &attrs) -> std::map { + TensorShape input = require_only_key(input_shapes, TensorSlotName::INPUT); - return {}; - }, - [&](WeightAttrs const &attrs) - -> std::map { - ASSERT(input_shapes.size() == 0); - return {}; - }, - [&](auto const &attrs) - -> std::map { - NOT_IMPLEMENTED(); - }, - }); + return throw_if_unexpected(get_weight_shapes(attrs, input)); + }, + [&](MultiHeadAttentionAttrs const &attrs) + -> std::map { + auto [query, key, value] = require_3(input_shapes, + TensorSlotName::QUERY, + TensorSlotName::KEY, + TensorSlotName::VALUE); + + return throw_if_unexpected(get_weight_shapes(attrs, query, key, value)); + }, + [&](Pool2DAttrs const &attrs) -> std::map { + require_only_key(input_shapes, TensorSlotName::INPUT); + + return {}; + }, + [&](SoftmaxAttrs const &attrs) -> std::map { + require_only_key(input_shapes, TensorSlotName::INPUT); + + return {}; + }, + [&](WeightAttrs const &attrs) -> std::map { + ASSERT(input_shapes.size() == 0); + return {}; + }, + [&](auto const &attrs) -> std::map { + NOT_IMPLEMENTED(); + }, + }); } std::map get_output_shapes( PCGOperatorAttrs const &pcg_op_attrs, - std::map const - &input_shapes) { - return pcg_op_attrs - .visit>(overload{ + std::map const &input_shapes) { + return pcg_op_attrs.visit>( + overload{ [&](BatchNormAttrs const &attrs) -> std::map { ParallelTensorShape input = @@ -678,10 +646,9 @@ std::map get_output_shapes( std::map get_weight_shapes( PCGOperatorAttrs const &pcg_op_attrs, - std::map const - &input_shapes) { - return pcg_op_attrs - .visit>(overload{ + std::map const &input_shapes) { + return pcg_op_attrs.visit>( + overload{ [&](BatchNormAttrs const &attrs) -> std::map { ParallelTensorShape input = diff --git a/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc b/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc index bd75207834..2df1cf0325 100644 --- a/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc +++ b/lib/op-attrs/src/op-attrs/tensor_dim_permutation.cc @@ -16,8 +16,7 @@ namespace FlexFlow { -static void - check_are_contiguous_from_one(std::set const &idxs) { +static void check_are_contiguous_from_one(std::set const &idxs) { if (idxs.empty()) { return; } diff --git a/lib/op-attrs/src/op-attrs/tensor_dims.cc b/lib/op-attrs/src/op-attrs/tensor_dims.cc index 1693cc92d2..40e74a6b94 100644 --- a/lib/op-attrs/src/op-attrs/tensor_dims.cc +++ b/lib/op-attrs/src/op-attrs/tensor_dims.cc @@ -13,8 +13,8 @@ #include "utils/containers/contains.h" #include "utils/containers/product.h" #include "utils/containers/reversed.h" -#include "utils/containers/transform.h" #include "utils/containers/set_of.h" +#include "utils/containers/transform.h" #include "utils/containers/vector_of.h" #include "utils/containers/zip.h" #include "utils/integer_conversions.h" diff --git a/lib/op-attrs/test/src/op-attrs/operator_task_space.cc b/lib/op-attrs/test/src/op-attrs/operator_task_space.cc index d4bd194ce9..57ac7e49a0 100644 --- a/lib/op-attrs/test/src/op-attrs/operator_task_space.cc +++ b/lib/op-attrs/test/src/op-attrs/operator_task_space.cc @@ -12,8 +12,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::set correct = { TaskSpaceCoordinate{OrthotopeCoord{{}}}}; - std::set result = - get_task_space_coordinates(task); + std::set result = get_task_space_coordinates(task); CHECK(correct == result); } @@ -28,8 +27,7 @@ TEST_SUITE(FF_TEST_SUITE) { TaskSpaceCoordinate{OrthotopeCoord{{1_n, 0_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{1_n, 1_n}}}, }}; - std::set result = - get_task_space_coordinates(task); + std::set result = get_task_space_coordinates(task); CHECK(correct == result); } @@ -52,8 +50,7 @@ TEST_SUITE(FF_TEST_SUITE) { TaskSpaceCoordinate{OrthotopeCoord{{2_n, 1_n, 0_n}}}, TaskSpaceCoordinate{OrthotopeCoord{{2_n, 1_n, 1_n}}}, }}; - std::set result = - get_task_space_coordinates(task); + std::set result = get_task_space_coordinates(task); CHECK(correct == result); } } diff --git a/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc b/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc index a7c01f6431..8d6272ab95 100644 --- a/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc +++ b/lib/op-attrs/test/src/op-attrs/parallel_tensor_dim_degrees.cc @@ -1,6 +1,5 @@ #include "op-attrs/parallel_tensor_dim_degrees.h" #include "op-attrs/parallel_tensor_dim_idx_t.h" -#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/set.h" #include diff --git a/lib/op-attrs/test/src/op-attrs/tensor_dims.cc b/lib/op-attrs/test/src/op-attrs/tensor_dims.cc index d191e8e482..e76a1bff5e 100644 --- a/lib/op-attrs/test/src/op-attrs/tensor_dims.cc +++ b/lib/op-attrs/test/src/op-attrs/tensor_dims.cc @@ -121,8 +121,7 @@ TEST_SUITE(FF_TEST_SUITE) { FFOrdered{3_p, 1_p, 2_p}, }; - std::set result = - get_tensor_dims_coord_set(input); + std::set result = get_tensor_dims_coord_set(input); std::set correct = { TensorDimsCoord{FFOrdered{0_n, 0_n, 0_n}}, TensorDimsCoord{FFOrdered{0_n, 0_n, 1_n}}, @@ -138,8 +137,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("TensorDims is zero-dimensional") { TensorDims input = TensorDims{FFOrdered{}}; - std::set result = - get_tensor_dims_coord_set(input); + std::set result = get_tensor_dims_coord_set(input); std::set correct = { TensorDimsCoord{FFOrdered{}}, }; diff --git a/lib/pcg/include/pcg/computation_graph.h b/lib/pcg/include/pcg/computation_graph.h index 286ac6d027..354b89e3ac 100644 --- a/lib/pcg/include/pcg/computation_graph.h +++ b/lib/pcg/include/pcg/computation_graph.h @@ -21,8 +21,8 @@ LayerAddedResult add_layer( LayerAttrs const &attrs, std::map const &inputs, std::map const &weights, - std::optional> const - &outputs = std::nullopt); + std::optional> const &outputs = + std::nullopt); LayerAddedResult add_input_layer(ComputationGraph &computation_graph, TensorShape const &tensor_shape, @@ -60,9 +60,8 @@ std::set std::set get_subgraph_outgoing_edges(ComputationGraph const &, std::set const &); -std::set - get_subgraph_successors(ComputationGraph const &, - std::set const &); +std::set get_subgraph_successors(ComputationGraph const &, + std::set const &); LayerAttrs get_layer_attrs(ComputationGraph const &cg, layer_guid_t const &n); diff --git a/lib/pcg/include/pcg/computation_graph_builder.h b/lib/pcg/include/pcg/computation_graph_builder.h index f183aabbe1..4e7e3a5a98 100644 --- a/lib/pcg/include/pcg/computation_graph_builder.h +++ b/lib/pcg/include/pcg/computation_graph_builder.h @@ -266,8 +266,8 @@ struct ComputationGraphBuilder { LayerAttrs const &layer, std::map const &inputs, std::map const &weights, - std::optional> const - &outputs = std::nullopt); + std::optional> const &outputs = + std::nullopt); tensor_guid_t broadcast(tensor_guid_t const &, TensorDims const &, std::string const &); diff --git a/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h b/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h index c1fed1ffb2..d6bff416ce 100644 --- a/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h +++ b/lib/pcg/include/pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h @@ -7,9 +7,9 @@ #include "utils/bidict/algorithms/bidict_from_enumerating.h" #include "utils/containers/enumerate.h" #include "utils/containers/generate_map.h" +#include "utils/containers/set_of.h" #include "utils/containers/sorted.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" @@ -59,8 +59,7 @@ V1KwargDataflowGraph } template -std::pair, - std::map> +std::pair, std::map> from_v1_including_node_numbering(V1KwargDataflowGraph const &v1) { std::map node_map = generate_map(v1.nodes, [](nonnegative_int n) { diff --git a/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h b/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h index eee45c900d..d9ede4f0ce 100644 --- a/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h +++ b/lib/pcg/include/pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.h @@ -4,10 +4,10 @@ #include "pcg/file_format/v1/graphs/v1_kwarg_dataflow_graph.h" #include "pcg/file_format/v1/graphs/v1_labelled_kwarg_dataflow_graph.dtg.h" #include "utils/bidict/algorithms/bidict_from_enumerating.h" +#include "utils/containers/map_from_pairs.h" #include "utils/containers/map_keys.h" #include "utils/containers/map_values.h" #include "utils/containers/transform.h" -#include "utils/containers/map_from_pairs.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h" #include "utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph_view.h" @@ -26,8 +26,8 @@ std::pair, V1KwargDataflowGraph unlabelled = to_v1(g, nodes.reversed()); - std::map node_labels = map_values( - nodes.as_map(), [&](Node const &n) { return g.at(n); }); + std::map node_labels = + map_values(nodes.as_map(), [&](Node const &n) { return g.at(n); }); std::map, OutputLabel> output_labels = map_from_pairs( diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h index 89d09b9132..1dccacd198 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h @@ -42,8 +42,8 @@ bidict get_tensor_bindings_for_slot_name(MappedOperatorTaskGroup const &, TensorSlotName const &); -std::set get_slot_names_for_task_group(MappedOperatorTaskGroup const &); - +std::set + get_slot_names_for_task_group(MappedOperatorTaskGroup const &); nlohmann::json mapped_operator_task_group_as_dot_json(MappedOperatorTaskGroup const &); diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h index c33bc7e7a4..cefd171503 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h @@ -2,8 +2,8 @@ #define _FLEXFLOW_LIB_PCG_INCLUDE_PCG_MAPPED_PARALLEL_COMPUTATION_GRAPH_MAPPED_PARALLEL_COMPUTATION_GRAPH_H #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.h" -#include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.h" +#include "pcg/parallel_computation_graph/parallel_computation_graph.h" namespace FlexFlow { diff --git a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h index dcda3a977b..f421a6c056 100644 --- a/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h +++ b/lib/pcg/include/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h @@ -1,16 +1,15 @@ #ifndef _FLEXFLOW_LIB_PCG_INCLUDE_PCG_MAPPED_PARALLEL_COMPUTATION_GRAPH_MAPPED_PARALLEL_LAYER_INVOCATION_INFO_H #define _FLEXFLOW_LIB_PCG_INCLUDE_PCG_MAPPED_PARALLEL_COMPUTATION_GRAPH_MAPPED_PARALLEL_LAYER_INVOCATION_INFO_H +#include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.h" #include "pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.h" -#include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" namespace FlexFlow { MappedParallelLayerInvocationInfo - mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping( - ParallelLayerInvocationInfo const &, - MappedOperatorTaskGroup const &); + mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping( + ParallelLayerInvocationInfo const &, MappedOperatorTaskGroup const &); } // namespace FlexFlow diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h index d3cc9f0149..e2500ac764 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph.h @@ -9,10 +9,10 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.dtg.h" #include "pcg/parallel_computation_graph/parallel_layer_added_result.dtg.h" #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" +#include "pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.h" #include "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h" #include "pcg/parallel_computation_graph/parallel_tensor_use_t.dtg.h" #include -#include "pcg/parallel_computation_graph/parallel_layer_invocation_info.dtg.h" namespace FlexFlow { @@ -28,8 +28,8 @@ ParallelLayerAddedResult add_parallel_layer( ParallelLayerAttrs const &layer_attrs, std::map const &inputs, std::map const &weights, - std::optional> const - &outputs = std::nullopt); + std::optional> const &outputs = + std::nullopt); ParallelLayerAddedResult pcg_add_input_layer(ParallelComputationGraph &pcg, @@ -40,11 +40,11 @@ OperatorTaskSpace get_operator_task_space(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &layer); std::set - pcg_get_invocation_info_set(ParallelComputationGraph const &); + pcg_get_invocation_info_set(ParallelComputationGraph const &); ParallelLayerInvocationInfo - pcg_get_invocation_info_for_layer(ParallelComputationGraph const &, - parallel_layer_guid_t); + pcg_get_invocation_info_for_layer(ParallelComputationGraph const &, + parallel_layer_guid_t); std::set get_pcg_edges_from_layer_to_layer(ParallelComputationGraph const &pcg, @@ -99,9 +99,8 @@ std::map get_incoming_input_degrees(ParallelComputationGraph const &, parallel_layer_guid_t const &); -std::set - get_successors(ParallelComputationGraph const &, - parallel_layer_guid_t const &); +std::set get_successors(ParallelComputationGraph const &, + parallel_layer_guid_t const &); std::set get_subgraph_successors(ParallelComputationGraph const &, diff --git a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h index 23d1e0ff9c..28ef8aa215 100644 --- a/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h +++ b/lib/pcg/include/pcg/parallel_computation_graph/parallel_computation_graph_builder.h @@ -141,8 +141,7 @@ struct ParallelComputationGraphBuilder { std::map add_layer( ParallelLayerAttrs const &layer, std::map const &inputs, - std::map const - &weight_initializers); + std::map const &weight_initializers); parallel_tensor_guid_t add_weight(ParallelTensorShape const &weight_tensor_shape, diff --git a/lib/pcg/src/pcg/computation_graph.cc b/lib/pcg/src/pcg/computation_graph.cc index 1ea5420951..08ae7a95ff 100644 --- a/lib/pcg/src/pcg/computation_graph.cc +++ b/lib/pcg/src/pcg/computation_graph.cc @@ -2,6 +2,7 @@ #include "op-attrs/computation_graph_op_attrs.h" #include "op-attrs/get_incoming_tensor_roles.h" #include "op-attrs/shape_inference.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/filter_values.h" #include "utils/containers/filtrans.h" @@ -34,7 +35,6 @@ #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/node/algorithms.h" #include "utils/record_formatter.h" -#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -52,13 +52,13 @@ std::set get_layers(ComputationGraph const &cg) { [&](Node const &n) { return layer_guid_t{n}; }); } -LayerAddedResult add_layer( - ComputationGraph &computation_graph, - LayerAttrs const &layer_attrs, - std::map const &inputs, - std::map const &weights, - std::optional> const - &maybe_output_flags) { +LayerAddedResult + add_layer(ComputationGraph &computation_graph, + LayerAttrs const &layer_attrs, + std::map const &inputs, + std::map const &weights, + std::optional> const + &maybe_output_flags) { std::map input_shapes = map_values(inputs, [&](tensor_guid_t const &i) { @@ -73,13 +73,13 @@ LayerAddedResult add_layer( std::map expected_weight_shapes = get_weight_shapes(layer_attrs.op_attrs, input_shapes); - std::map> - raw_inputs = map_values( - inputs, [&](tensor_guid_t const &t) { return t.raw_graph_output; }); + std::map> raw_inputs = + map_values(inputs, + [&](tensor_guid_t const &t) { return t.raw_graph_output; }); - std::map> - raw_weights = map_values( - weights, [&](tensor_guid_t const &t) { return t.raw_graph_output; }); + std::map> raw_weights = + map_values(weights, + [&](tensor_guid_t const &t) { return t.raw_graph_output; }); std::map output_shapes = get_output_shapes(layer_attrs.op_attrs, input_shapes); @@ -230,9 +230,9 @@ std::map return all_tensor_attrs; } -std::set get_subgraph_incoming_edges( - ComputationGraph const &cg, - std::set const &subgraph_nodes) { +std::set + get_subgraph_incoming_edges(ComputationGraph const &cg, + std::set const &subgraph_nodes) { std::set raw_subgraph_nodes = transform( subgraph_nodes, [](layer_guid_t const &l) { return l.raw_node; }); @@ -246,9 +246,9 @@ std::set get_subgraph_incoming_edges( }); } -std::set get_subgraph_outgoing_edges( - ComputationGraph const &cg, - std::set const &subgraph_nodes) { +std::set + get_subgraph_outgoing_edges(ComputationGraph const &cg, + std::set const &subgraph_nodes) { std::set raw_subgraph_nodes = transform( subgraph_nodes, [](layer_guid_t const &l) { return l.raw_node; }); @@ -262,9 +262,9 @@ std::set get_subgraph_outgoing_edges( }); } -std::set get_subgraph_successors( - ComputationGraph const &cg, - std::set const &subgraph_nodes) { +std::set + get_subgraph_successors(ComputationGraph const &cg, + std::set const &subgraph_nodes) { std::set raw_subgraph_nodes = transform( subgraph_nodes, [](layer_guid_t const &l) { return l.raw_node; }); @@ -347,10 +347,11 @@ std::string as_dot(ComputationGraph const &cg) { return result; }; - std::function( - std::set const &)> - order_slots = [](std::set const &unordered) - -> nlohmann::json { return sorted(unordered); }; + std::function(std::set const &)> + order_slots = + [](std::set const &unordered) -> nlohmann::json { + return sorted(unordered); + }; return labelled_open_kwarg_dataflow_graph_view_as_dot( view_as_labelled_open_kwarg_dataflow_graph -#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -77,18 +77,17 @@ tensor_guid_t ComputationGraphBuilder::create_input( maybe_name, }; - return require_only_key( - this->add_layer(/*layer=*/layer_attrs, - /*inputs=*/{}, - /*weights=*/{}, - /*outputs=*/ - std::map{ - { - TensorSlotName::OUTPUT, - create_grad, - }, - }), - TensorSlotName::OUTPUT); + return require_only_key(this->add_layer(/*layer=*/layer_attrs, + /*inputs=*/{}, + /*weights=*/{}, + /*outputs=*/ + std::map{ + { + TensorSlotName::OUTPUT, + create_grad, + }, + }), + TensorSlotName::OUTPUT); } tensor_guid_t ComputationGraphBuilder::create_weight( @@ -107,10 +106,10 @@ tensor_guid_t ComputationGraphBuilder::create_weight( TensorSlotName::OUTPUT); } -static void check_incoming_tensor_roles( - LayerAttrs const &layer, - std::set const &input_slots, - std::set const &weight_slots) { +static void + check_incoming_tensor_roles(LayerAttrs const &layer, + std::set const &input_slots, + std::set const &weight_slots) { std::map correct = restrict_keys(get_incoming_tensor_roles(layer.op_attrs), set_union(input_slots, weight_slots)); @@ -127,14 +126,11 @@ static void check_incoming_tensor_roles( "check_incoming_tensor_roles found deviation in incoming tensors"); } -std::map - ComputationGraphBuilder::add_layer( - LayerAttrs const &layer, - std::map const &inputs, - std::map const - &weight_initializers, - std::optional> const - &outputs) { +std::map ComputationGraphBuilder::add_layer( + LayerAttrs const &layer, + std::map const &inputs, + std::map const &weight_initializers, + std::optional> const &outputs) { ASSERT(are_disjoint(keys(inputs), keys(weight_initializers))); check_incoming_tensor_roles(layer, keys(inputs), keys(weight_initializers)); @@ -144,13 +140,12 @@ std::map std::map weight_shapes = get_weight_shapes(layer.op_attrs, input_shapes); - std::map weights = - zip_values_strict_with( - weight_shapes, - weight_initializers, - [&](TensorShape const &shape, InitializerAttrs const &initializer) { - return this->create_weight(shape, initializer); - }); + std::map weights = zip_values_strict_with( + weight_shapes, + weight_initializers, + [&](TensorShape const &shape, InitializerAttrs const &initializer) { + return this->create_weight(shape, initializer); + }); LayerAddedResult added = ::FlexFlow::add_layer( this->computation_graph, layer, inputs, weights, outputs); @@ -853,8 +848,7 @@ tensor_guid_t ComputationGraphBuilder::concat( LayerAttrs layer = LayerAttrs{ComputationGraphOpAttrs{attrs}, name}; return require_only_key( - this->add_layer( - layer, map_from_pairs(zip(input_slot_names, inputs)), {}), + this->add_layer(layer, map_from_pairs(zip(input_slot_names, inputs)), {}), TensorSlotName::OUTPUT); } diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc index 0c9bdb3029..1f69e86573 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_operator_task_group.cc @@ -3,22 +3,22 @@ #include "op-attrs/operator_task_space.h" #include "op-attrs/parallel_tensor_space_coordinate.h" #include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.h" +#include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" #include "utils/bidict/algorithms/bidict_transform_values.h" +#include "utils/bidict/algorithms/right_entries.h" #include "utils/bidict/generate_bidict.h" #include "utils/containers/are_all_distinct.h" +#include "utils/containers/contains.h" +#include "utils/containers/keys.h" +#include "utils/containers/map_values.h" #include "utils/containers/require_all_same.h" +#include "utils/containers/require_all_same1.h" +#include "utils/containers/set_of.h" #include "utils/containers/sorted.h" #include "utils/containers/transform.h" #include "utils/containers/vector_of.h" #include "utils/hash/tuple.h" #include "utils/nonnegative_int/num_elements.h" -#include "utils/containers/require_all_same1.h" -#include "utils/containers/set_of.h" -#include "utils/containers/keys.h" -#include "utils/containers/contains.h" -#include "utils/bidict/algorithms/right_entries.h" -#include "utils/containers/map_values.h" -#include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" namespace FlexFlow { @@ -26,12 +26,11 @@ MappedOperatorTaskGroup::MappedOperatorTaskGroup( bidict const &shard_bindings) : shard_bindings(shard_bindings) { - std::vector> binding_slot_sets = - transform(vector_of(shard_bindings.right_values()), - [&](OperatorAtomicTaskShardBinding const &s) - -> std::set { - return keys(s.tensor_coords); - }); + std::vector> binding_slot_sets = transform( + vector_of(shard_bindings.right_values()), + [&](OperatorAtomicTaskShardBinding const &s) -> std::set { + return keys(s.tensor_coords); + }); std::set slot_names = require_all_same(binding_slot_sets).value(); @@ -100,24 +99,28 @@ bidict const & bidict get_tensor_bindings_for_slot_name(MappedOperatorTaskGroup const &task_group, TensorSlotName const &slot_name) { - std::set slot_names = get_slot_names_for_task_group(task_group); + std::set slot_names = + get_slot_names_for_task_group(task_group); ASSERT(contains(slot_names, slot_name)); std::map m = - map_values(task_group.get_shard_bindings().as_map(), - [&](OperatorAtomicTaskShardBinding const &b) -> ParallelTensorSpaceCoordinate { - return ptensor_space_coord_for_slot_name(b, slot_name); - }); + map_values(task_group.get_shard_bindings().as_map(), + [&](OperatorAtomicTaskShardBinding const &b) + -> ParallelTensorSpaceCoordinate { + return ptensor_space_coord_for_slot_name(b, slot_name); + }); return bidict_from_unstructured_relation(set_of(m)).reversed(); } -std::set get_slot_names_for_task_group(MappedOperatorTaskGroup const &g) { +std::set + get_slot_names_for_task_group(MappedOperatorTaskGroup const &g) { return require_all_same1( - transform(vector_of(right_entries(g.get_shard_bindings())), - [&](OperatorAtomicTaskShardBinding const &shard_bindings) -> std::set { - return keys(shard_bindings.tensor_coords); - })); + transform(vector_of(right_entries(g.get_shard_bindings())), + [&](OperatorAtomicTaskShardBinding const &shard_bindings) + -> std::set { + return keys(shard_bindings.tensor_coords); + })); } nlohmann::json diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 38d99f561b..687f65539c 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -1,21 +1,21 @@ #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" #include "op-attrs/pcg_operator_attrs.h" +#include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_attrs.h" +#include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" #include "utils/bidict/algorithms/bidict_from_map.h" #include "utils/bidict/algorithms/bidict_transform_keys.h" +#include "utils/containers/keys.h" +#include "utils/containers/require_all_of.h" +#include "utils/containers/set_of.h" +#include "utils/containers/set_union.h" #include "utils/containers/transform.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/find_isomorphism_between_kwarg_dataflow_graphs.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/materialize_labelled_kwarg_dataflow_graph_view.h" #include "utils/graph/labelled_kwarg_dataflow_graph/algorithms/rewrite_labelled_kwarg_dataflow_graph_node_labels.h" #include "utils/many_to_one/many_to_one_from_map.h" -#include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.h" -#include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" -#include "utils/containers/set_of.h" -#include "utils/containers/set_union.h" -#include "utils/containers/keys.h" -#include "utils/containers/require_all_of.h" namespace FlexFlow { @@ -25,14 +25,14 @@ std::set } std::set - mpcg_get_invocation_set(MappedParallelComputationGraph const &mpcg) -{ + mpcg_get_invocation_set(MappedParallelComputationGraph const &mpcg) { auto mk_mapped_invocation = [&](ParallelLayerInvocationInfo const &invocation) - -> MappedParallelLayerInvocationInfo - { - MappedOperatorTaskGroup mapping = mpcg_get_mapping_for_layer(mpcg, invocation.layer_info.guid); + -> MappedParallelLayerInvocationInfo { + MappedOperatorTaskGroup mapping = + mpcg_get_mapping_for_layer(mpcg, invocation.layer_info.guid); - return mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping(invocation, mapping); + return mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping( + invocation, mapping); }; ParallelComputationGraph pcg = pcg_from_mpcg(mpcg); @@ -134,21 +134,24 @@ MappedParallelComputationGraph mapped_pcg_from_pcg_and_mapped_op_task_groups( return mapped_op_task_groups.at(l); }; - auto slot_names_for_layer = [&](parallel_layer_guid_t l) -> std::set { - return set_union(keys(get_incoming_tensors(pcg, l)), keys(get_outgoing_tensors(pcg, l))); + auto slot_names_for_layer = + [&](parallel_layer_guid_t l) -> std::set { + return set_union(keys(get_incoming_tensors(pcg, l)), + keys(get_outgoing_tensors(pcg, l))); }; - auto slot_names_for_layer_mapping = [&](parallel_layer_guid_t l) -> std::set { + auto slot_names_for_layer_mapping = + [&](parallel_layer_guid_t l) -> std::set { return get_slot_names_for_task_group(mapping_for_layer(l)); }; - require_all_of( - pcg_get_parallel_layers(pcg), - [&](parallel_layer_guid_t l) -> void { - std::set for_layer = slot_names_for_layer(l); - std::set for_layer_mapping = slot_names_for_layer_mapping(l); - ASSERT(for_layer == for_layer_mapping); - }); + require_all_of(pcg_get_parallel_layers(pcg), + [&](parallel_layer_guid_t l) -> void { + std::set for_layer = slot_names_for_layer(l); + std::set for_layer_mapping = + slot_names_for_layer_mapping(l); + ASSERT(for_layer == for_layer_mapping); + }); auto mpcg_layer_attrs_from_pcg_layer_attrs = [&](Node const &node, ParallelLayerAttrs const &pcg_layer_attrs) @@ -236,8 +239,7 @@ std::string mapped_pcg_as_dot(MappedParallelComputationGraph const &mpcg) { return fmt::to_string(slot_name); }; - std::function( - std::set const &)> + std::function(std::set const &)> order_slots = [](std::set const &slot_names) -> std::vector { return sorted(slot_names); }; diff --git a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc index bda6afb60c..e4f93b3242 100644 --- a/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc +++ b/lib/pcg/src/pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.cc @@ -3,20 +3,19 @@ namespace FlexFlow { MappedParallelLayerInvocationInfo - mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping( - ParallelLayerInvocationInfo const &invocation_info, - MappedOperatorTaskGroup const &mapping) -{ + mapped_parallel_layer_invocation_info_from_pcg_invocation_and_mapping( + ParallelLayerInvocationInfo const &invocation_info, + MappedOperatorTaskGroup const &mapping) { return MappedParallelLayerInvocationInfo{ - /*incoming=*/invocation_info.incoming, - /*layer_info=*/MappedParallelLayerInfo{ - /*guid=*/invocation_info.layer_info.guid, - /*attrs=*/invocation_info.layer_info.attrs, - /*mapping=*/mapping, - }, - /*outgoing=*/invocation_info.outgoing, + /*incoming=*/invocation_info.incoming, + /*layer_info=*/ + MappedParallelLayerInfo{ + /*guid=*/invocation_info.layer_info.guid, + /*attrs=*/invocation_info.layer_info.attrs, + /*mapping=*/mapping, + }, + /*outgoing=*/invocation_info.outgoing, }; } - } // namespace FlexFlow diff --git a/lib/pcg/src/pcg/optimizer_attrs.cc b/lib/pcg/src/pcg/optimizer_attrs.cc index 05c5dd95ac..0961697d33 100644 --- a/lib/pcg/src/pcg/optimizer_attrs.cc +++ b/lib/pcg/src/pcg/optimizer_attrs.cc @@ -26,8 +26,7 @@ OptimizerAttrs std::set get_slot_names_for_optimizer(OptimizerAttrs const &attrs) { return attrs.visit>(overload{ - [](SGDOptimizerAttrs const &sgd_attrs) - -> std::set { + [](SGDOptimizerAttrs const &sgd_attrs) -> std::set { if (sgd_attrs.momentum > 0.0f) { return {OptimizerSlotName::SGD_V}; } else { diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index 1a27c88303..3af6840f13 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -10,14 +10,15 @@ #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.dtg.h" #include "pcg/parallel_computation_graph/parallel_computation_graph_edge.h" #include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/extend.h" #include "utils/containers/filter_values.h" #include "utils/containers/filtrans.h" #include "utils/containers/get_only.h" #include "utils/containers/repeat_element.h" -#include "utils/containers/transform.h" #include "utils/containers/set_of.h" +#include "utils/containers/transform.h" #include "utils/containers/zip_values_strict_with.h" #include "utils/containers/zip_with_strict.h" #include "utils/graph/digraph/algorithms/get_initial_nodes.h" @@ -37,7 +38,6 @@ #include "utils/graph/node/node.dtg.h" #include "utils/record_formatter.h" #include -#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -75,9 +75,8 @@ ParallelLayerAddedResult add_parallel_layer( return get_parallel_tensor_shape(pcg, i); }); - std::map - correct_weight_shapes = - get_weight_shapes(layer_attrs.op_attrs, input_shapes); + std::map correct_weight_shapes = + get_weight_shapes(layer_attrs.op_attrs, input_shapes); ASSERT(weight_shapes == correct_weight_shapes, "add_parallel_layer received incorrect weight shapes"); @@ -154,20 +153,17 @@ OperatorTaskSpace get_operator_task_space(ParallelComputationGraph const &pcg, std::map inputs = get_incoming_inputs(pcg, layer); - std::map input_degrees = - map_values(get_incoming_inputs(pcg, layer), - [&](parallel_tensor_guid_t input_guid) { - return get_parallel_degrees( - get_parallel_tensor_shape(pcg, input_guid)); - }); + std::map input_degrees = map_values( + get_incoming_inputs(pcg, layer), [&](parallel_tensor_guid_t input_guid) { + return get_parallel_degrees(get_parallel_tensor_shape(pcg, input_guid)); + }); return get_operator_task_space( compgraph_op_attrs_from_pcg_op_attrs(op_attrs).value(), input_degrees); } std::set - pcg_get_invocation_info_set(ParallelComputationGraph const &pcg) -{ + pcg_get_invocation_info_set(ParallelComputationGraph const &pcg) { return transform(set_of(pcg_get_parallel_layers(pcg)), [&](parallel_layer_guid_t l) -> ParallelLayerInvocationInfo { return pcg_get_invocation_info_for_layer(pcg, l); @@ -175,33 +171,34 @@ std::set } ParallelLayerInvocationInfo - pcg_get_invocation_info_for_layer(ParallelComputationGraph const &pcg, - parallel_layer_guid_t l) -{ + pcg_get_invocation_info_for_layer(ParallelComputationGraph const &pcg, + parallel_layer_guid_t l) { ParallelLayerAttrs l_attrs = get_parallel_layer_attrs(pcg, l); std::map incoming = - get_incoming_tensors(pcg, l); + get_incoming_tensors(pcg, l); std::map outgoing = - get_outgoing_tensors(pcg, l); + get_outgoing_tensors(pcg, l); - auto get_parallel_tensor_info = [&](parallel_tensor_guid_t t) -> ParallelTensorInfo { + auto get_parallel_tensor_info = + [&](parallel_tensor_guid_t t) -> ParallelTensorInfo { ParallelTensorAttrs t_attrs = get_parallel_tensor_attrs(pcg, t); return ParallelTensorInfo{ - /*guid=*/t, - /*attrs=*/t_attrs, + /*guid=*/t, + /*attrs=*/t_attrs, }; }; return ParallelLayerInvocationInfo{ - /*incoming=*/map_values(incoming, get_parallel_tensor_info), - /*layer_info=*/ParallelLayerInfo{ - /*guid=*/l, - /*attrs=*/l_attrs, - }, - /*outgoing=*/map_values(outgoing, get_parallel_tensor_info), + /*incoming=*/map_values(incoming, get_parallel_tensor_info), + /*layer_info=*/ + ParallelLayerInfo{ + /*guid=*/l, + /*attrs=*/l_attrs, + }, + /*outgoing=*/map_values(outgoing, get_parallel_tensor_info), }; } @@ -229,10 +226,9 @@ std::set get_outgoing_edges(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { std::set> raw_edges = - set_of( - get_outgoing_kwarg_dataflow_edges_for_node(pcg.raw_graph, - l.raw_graph_node) - .right_values()); + set_of(get_outgoing_kwarg_dataflow_edges_for_node(pcg.raw_graph, + l.raw_graph_node) + .right_values()); return transform(raw_edges, [](KwargDataflowEdge const &e) { return ParallelComputationGraphEdge{e}; }); @@ -241,9 +237,9 @@ std::set std::map get_incoming_edges(ParallelComputationGraph const &pcg, parallel_layer_guid_t const &l) { - std::map> - raw_edges = get_incoming_kwarg_dataflow_edges_for_node(pcg.raw_graph, - l.raw_graph_node); + std::map> raw_edges = + get_incoming_kwarg_dataflow_edges_for_node(pcg.raw_graph, + l.raw_graph_node); return map_values(raw_edges, [](KwargDataflowEdge const &e) { return ParallelComputationGraphEdge{e}; }); @@ -435,8 +431,7 @@ std::vector std::map get_parallel_layer_attrs_mapping(ParallelComputationGraph const &pcg) { - std::map - layer_attrs_mapping; + std::map layer_attrs_mapping; for (parallel_layer_guid_t const &layer_guid : pcg_get_parallel_layers(pcg)) { layer_attrs_mapping.insert( {layer_guid, get_parallel_layer_attrs(pcg, layer_guid)}); @@ -512,8 +507,7 @@ std::string pcg_as_dot(ParallelComputationGraph const &cg) { return fmt::to_string(slot_name); }; - std::function( - std::set const &)> + std::function(std::set const &)> order_slots = [](std::set const &slot_names) -> std::vector { return sorted(slot_names); }; diff --git a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc index 0943800fd4..c774c1a6b5 100644 --- a/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc +++ b/lib/pcg/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc @@ -24,6 +24,7 @@ #include "op-attrs/shape_inference.h" #include "pcg/parallel_computation_graph/generate_weight_transform.h" #include "pcg/parallel_computation_graph/parallel_computation_graph.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/count.h" #include "utils/containers/enumerate_vector.h" @@ -33,7 +34,6 @@ #include "utils/containers/transform.h" #include "utils/containers/zip_values_strict_with.h" #include "utils/containers/zip_with.h" -#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -308,15 +308,14 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::multihead_attention( ParallelLayerAttrs layer = ParallelLayerAttrs{PCGOperatorAttrs{attrs}, name}; - std::map initializers = - throw_if_unexpected( - get_initializers(attrs, - get_reduced_shape(this->get_shape(query)), - get_reduced_shape(this->get_shape(key)), - get_reduced_shape(this->get_shape(value)), - maybe_weights_initializer, - maybe_input_bias_initializer, - maybe_output_bias_initializer)); + std::map initializers = throw_if_unexpected( + get_initializers(attrs, + get_reduced_shape(this->get_shape(query)), + get_reduced_shape(this->get_shape(key)), + get_reduced_shape(this->get_shape(value)), + maybe_weights_initializer, + maybe_input_bias_initializer, + maybe_output_bias_initializer)); return require_only_key(this->add_layer(layer, { @@ -643,10 +642,10 @@ parallel_tensor_guid_t ParallelComputationGraphBuilder::add_weight( return current_weight_tensor; } -static void check_incoming_tensor_roles( - ParallelLayerAttrs const &layer, - std::set const &input_slots, - std::set const &weight_slots) { +static void + check_incoming_tensor_roles(ParallelLayerAttrs const &layer, + std::set const &input_slots, + std::set const &weight_slots) { std::map correct = get_incoming_tensor_roles(layer.op_attrs); std::map current = @@ -665,10 +664,8 @@ static void check_incoming_tensor_roles( std::map ParallelComputationGraphBuilder::add_layer( ParallelLayerAttrs const &layer, - std::map const - &inputs, - std::map const - &weight_initializers) { + std::map const &inputs, + std::map const &weight_initializers) { ASSERT(are_disjoint(keys(inputs), keys(weight_initializers))); check_incoming_tensor_roles(layer, keys(inputs), keys(weight_initializers)); diff --git a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc index 54c582b66e..be692a5f6a 100644 --- a/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc +++ b/lib/pcg/test/src/pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.cc @@ -79,71 +79,69 @@ TEST_SUITE(FF_TEST_SUITE) { MappedOperatorTaskGroup partition_mapping = MappedOperatorTaskGroup{ bidict{ { - machine_coord(0_n), - OperatorAtomicTaskShardBinding{ - { - {TensorSlotName::INPUT, ptensor_coord(0_n)}, - {TensorSlotName::OUTPUT, ptensor_coord(0_n)}, - }, - }, + machine_coord(0_n), + OperatorAtomicTaskShardBinding{ + { + {TensorSlotName::INPUT, ptensor_coord(0_n)}, + {TensorSlotName::OUTPUT, ptensor_coord(0_n)}, + }, + }, }, { - machine_coord(1_n), - OperatorAtomicTaskShardBinding{ - { - {TensorSlotName::INPUT, ptensor_coord(0_n)}, - {TensorSlotName::OUTPUT, ptensor_coord(1_n)}, - }, - }, + machine_coord(1_n), + OperatorAtomicTaskShardBinding{ + { + {TensorSlotName::INPUT, ptensor_coord(0_n)}, + {TensorSlotName::OUTPUT, ptensor_coord(1_n)}, + }, + }, }, }, }; - std::map - mapped_tasks = { - { - l_input1, - input_mapping, - }, - { - l_input2, - input_mapping, - }, - { - l_partition1, - partition_mapping, - }, - { - l_partition2, - partition_mapping, - }, - {l_add, - MappedOperatorTaskGroup{ - bidict{ - { - machine_coord(0_n), - OperatorAtomicTaskShardBinding{ - { + std::map mapped_tasks = { + { + l_input1, + input_mapping, + }, + { + l_input2, + input_mapping, + }, + { + l_partition1, + partition_mapping, + }, + { + l_partition2, + partition_mapping, + }, + {l_add, + MappedOperatorTaskGroup{ + bidict{ + { + machine_coord(0_n), + OperatorAtomicTaskShardBinding{ + { {TensorSlotName::LHS_INPUT, ptensor_coord(0_n)}, {TensorSlotName::RHS_INPUT, ptensor_coord(0_n)}, {TensorSlotName::OUTPUT, ptensor_coord(0_n)}, - }, - }, + }, }, - { - machine_coord(1_n), - OperatorAtomicTaskShardBinding{ - { + }, + { + machine_coord(1_n), + OperatorAtomicTaskShardBinding{ + { {TensorSlotName::LHS_INPUT, ptensor_coord(1_n)}, {TensorSlotName::RHS_INPUT, ptensor_coord(1_n)}, {TensorSlotName::OUTPUT, ptensor_coord(1_n)}, - }, - }, + }, }, }, - }}, - }; + }, + }}, + }; return mapped_pcg_from_pcg_and_mapped_op_task_groups(pcg, mapped_tasks); }; diff --git a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc index 22c294dbce..27184282cb 100644 --- a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc +++ b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph.cc @@ -562,7 +562,8 @@ TEST_SUITE(FF_TEST_SUITE) { DimDomain layer_2_task_space = layer_1_task_space; - auto make_coord = [](nonnegative_int x) -> DimCoord { + auto make_coord = + [](nonnegative_int x) -> DimCoord { return DimCoord{ std::map{ {operator_task_space_dim_idx_t{0_n}, x}, diff --git a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc index 662bac0a7c..ebb11b4c91 100644 --- a/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc +++ b/lib/pcg/test/src/pcg/parallel_computation_graph/parallel_computation_graph_builder.cc @@ -175,11 +175,10 @@ TEST_SUITE(FF_TEST_SUITE) { /*paddingH=*/paddingH, /*paddingW=*/paddingW); - std::map layers = - generate_map(pcg_get_parallel_layers(b.pcg), - [&](parallel_layer_guid_t const &l) { - return get_parallel_layer_attrs(b.pcg, l); - }); + std::map layers = generate_map( + pcg_get_parallel_layers(b.pcg), [&](parallel_layer_guid_t const &l) { + return get_parallel_layer_attrs(b.pcg, l); + }); CHECK_MESSAGE(layers.size() == 7, "Incorrect layers ", layers); auto num_attrs_of_type = [&](OperatorType op_type) -> nonnegative_int { diff --git a/lib/realm-execution/include/realm-execution/dependency_set.h b/lib/realm-execution/include/realm-execution/dependency_set.h index 1afee5b915..28b5261944 100644 --- a/lib/realm-execution/include/realm-execution/dependency_set.h +++ b/lib/realm-execution/include/realm-execution/dependency_set.h @@ -29,8 +29,7 @@ struct DependencySet { private: Realm::Event precondition; - std::map - atomic_dependencies; + std::map atomic_dependencies; }; } // namespace FlexFlow diff --git a/lib/realm-execution/include/realm-execution/distributed_ff_handle.h b/lib/realm-execution/include/realm-execution/distributed_ff_handle.h index fc11120c64..ade3119df3 100644 --- a/lib/realm-execution/include/realm-execution/distributed_ff_handle.h +++ b/lib/realm-execution/include/realm-execution/distributed_ff_handle.h @@ -17,15 +17,13 @@ struct DistributedFfHandle { DistributedFfHandle() = delete; explicit DistributedFfHandle( std::map> const - &handles); + DeviceSpecificPtr> const &handles); DeviceSpecificPtr const & at(Realm::Processor processor) const; private: - std::map> + std::map> handles; }; diff --git a/lib/realm-execution/include/realm-execution/instance_allocation.h b/lib/realm-execution/include/realm-execution/instance_allocation.h index d1e6c9e72c..c7f03d09ee 100644 --- a/lib/realm-execution/include/realm-execution/instance_allocation.h +++ b/lib/realm-execution/include/realm-execution/instance_allocation.h @@ -25,8 +25,7 @@ std::pair */ TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, - std::map const - &preallocated, + std::map const &preallocated, RealmContext &ctx); /** diff --git a/lib/realm-execution/include/realm-execution/pcg_instance.h b/lib/realm-execution/include/realm-execution/pcg_instance.h index 6b2ddb4fe4..76509af2fd 100644 --- a/lib/realm-execution/include/realm-execution/pcg_instance.h +++ b/lib/realm-execution/include/realm-execution/pcg_instance.h @@ -83,8 +83,7 @@ PCGInstance create_pcg_instance( MappedParallelComputationGraph const &mpcg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::map const - &input_tensors, + std::map const &input_tensors, ProfilingSettings const &profiling_settings, DistributedFfHandle const &ff_handle, DeviceType device_type); diff --git a/lib/realm-execution/include/realm-execution/realm_context.h b/lib/realm-execution/include/realm-execution/realm_context.h index f2be9727d1..71e7e48420 100644 --- a/lib/realm-execution/include/realm-execution/realm_context.h +++ b/lib/realm-execution/include/realm-execution/realm_context.h @@ -12,8 +12,8 @@ #include "realm-execution/tasks/task_id_t.dtg.h" #include "task-spec/global_device_id_t.dtg.h" #include "task-spec/local_device_id_t.dtg.h" -#include #include +#include namespace FlexFlow { diff --git a/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc b/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc index dd43e8e7ef..04bf34768f 100644 --- a/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc +++ b/lib/realm-execution/src/realm-execution/distributed_ff_handle.cc @@ -7,8 +7,7 @@ namespace FlexFlow { DistributedFfHandle::DistributedFfHandle( std::map> const - &handles) + DeviceSpecificPtr> const &handles) : handles(handles) {} DeviceSpecificPtr const & @@ -21,8 +20,7 @@ DistributedFfHandle size_t workSpaceSize, bool allowTensorOpMathConversion, Realm::Event precondition) { - std::map> + std::map> handles; // Allocate space for the result before launching any tasks diff --git a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc index 797951311b..27d7ea85e8 100644 --- a/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc +++ b/lib/realm-execution/src/realm-execution/distributed_per_device_op_state_initialization.cc @@ -10,8 +10,8 @@ #include "utils/containers/maybe_get_only.h" #include "utils/containers/values.h" #include "utils/optional.h" -#include #include +#include #include namespace FlexFlow { @@ -28,8 +28,7 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( // Initialize all operators and save the per-device op state ASSERT(no_nodes_are_initialized(dg)); - std::map *> + std::map *> device_state_map; for (DynamicNodeInvocation const &invocation : dg.invocations) { // Nodes mapped to multiple devices are always parallel operators and don't @@ -40,8 +39,8 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( continue; } - Realm::Processor target_proc = ctx.processor_from_global_device_id( - assert_unwrap(device_id)); + Realm::Processor target_proc = + ctx.processor_from_global_device_id(assert_unwrap(device_id)); TensorInstanceBacking tensor_backing = subset_tensor_instance_backing_for_invocation(tensor_instance_backing, @@ -73,8 +72,8 @@ PerDeviceOpStateBacking perform_distributed_per_device_op_state_initialization( ctx.get_outstanding_events().wait(); auto deref = [](DeviceSpecificPtr *const &p) { return *p; }; - std::map> - result = map_values(device_state_map, deref); + std::map> result = + map_values(device_state_map, deref); for (DeviceSpecificPtr *device_state_ptr : values(device_state_map)) { diff --git a/lib/realm-execution/src/realm-execution/instance_allocation.cc b/lib/realm-execution/src/realm-execution/instance_allocation.cc index a4782c85ac..b677f24a0b 100644 --- a/lib/realm-execution/src/realm-execution/instance_allocation.cc +++ b/lib/realm-execution/src/realm-execution/instance_allocation.cc @@ -36,8 +36,7 @@ std::pair TensorInstanceBacking perform_instance_allocation( DynamicOpenDataflowGraph const &g, - std::map const - &preallocated, + std::map const &preallocated, RealmContext &ctx) { ASSERT(no_tensors_are_allocated(g)); ASSERT(tensors_are_ready_for_allocation(g)); diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 2b60580b73..1814af9cb0 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -83,8 +83,7 @@ PCGInstance create_pcg_instance( MappedParallelComputationGraph const &mpcg, OptimizerAttrs const &optimizer_attrs, std::optional const &loss, - std::map const - &input_tensors, + std::map const &input_tensors, ProfilingSettings const &profiling_settings, DistributedFfHandle const &device_handle, DeviceType device_type) { @@ -93,8 +92,7 @@ PCGInstance create_pcg_instance( make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg, device_type); dg = perform_pass_expansion(dg); - std::map inputs = - input_tensors; + std::map inputs = input_tensors; std::optional logit_grad_value; if (loss.has_value()) { ParallelLossConfig loss_config = assert_unwrap(loss); @@ -245,12 +243,11 @@ static Realm::Event spawn_dynamic_node_invocation( Realm::Event result = precondition; for (auto const &[p, d] : assert_unwrap(output_grad.mapping).raw) { DynamicValueAttrs replica_key = output_grad; - replica_key.mapping = - ParallelTensorMapping{ - bidict{ + replica_key.mapping = ParallelTensorMapping{ + bidict{ {p, d}, - }, - }; + }, + }; replica_key.shard_coord = p; Realm::RegionInstance src_inst = diff --git a/lib/realm-execution/src/realm-execution/realm_context.cc b/lib/realm-execution/src/realm-execution/realm_context.cc index 56b1d11c65..1e90d4573f 100644 --- a/lib/realm-execution/src/realm-execution/realm_context.cc +++ b/lib/realm-execution/src/realm-execution/realm_context.cc @@ -15,8 +15,8 @@ #include "realm-execution/tasks/task_id_t.h" #include "task-spec/global_device_id_t.h" #include "utils/bidict/algorithms/bidict_from_enumerating.h" -#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" #include "utils/bidict/algorithms/bidict_transform_values.h" +#include "utils/bidict/algorithms/merge_disjoint_bidicts.h" #include "utils/containers/are_all_same.h" #include "utils/containers/contains_key.h" #include "utils/containers/group_by.h" @@ -56,13 +56,14 @@ bidict build_local_machine_topology( set_of(by_proc_kind.at_l(k).unwrap_as_unordered_set())) .reversed(); - bidict result = bidict_transform_values( - enumerated, [&](nonnegative_int idx) -> local_device_id_t { - return local_device_id_t{ - /*idx=*/device_in_node_idx_t{idx}, - /*device_type=*/device_type_from_processor_kind(k), - }; - }); + bidict result = + bidict_transform_values( + enumerated, [&](nonnegative_int idx) -> local_device_id_t { + return local_device_id_t{ + /*idx=*/device_in_node_idx_t{idx}, + /*device_type=*/device_type_from_processor_kind(k), + }; + }); return result; }; diff --git a/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc b/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc index 1d53824c29..d11e340d5c 100644 --- a/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc +++ b/lib/realm-execution/src/realm-execution/tasks/serializer/serializable_tensor_instance_backing.cc @@ -9,27 +9,27 @@ namespace FlexFlow { SerializableTensorInstanceBacking tensor_instance_backing_to_serializable( TensorInstanceBacking const &backing) { return SerializableTensorInstanceBacking{ - /*backing=*/map_keys_and_values( - backing.backing, - dynamic_value_attrs_to_serializable, - [](std::pair const &p) { - return std::pair{realm_instance_to_serializable(p.first), - realm_event_to_serializable(p.second)}; - }), + /*backing=*/map_keys_and_values( + backing.backing, + dynamic_value_attrs_to_serializable, + [](std::pair const &p) { + return std::pair{realm_instance_to_serializable(p.first), + realm_event_to_serializable(p.second)}; + }), }; } TensorInstanceBacking tensor_instance_backing_from_serializable( SerializableTensorInstanceBacking const &backing) { return TensorInstanceBacking{ - /*backing=*/map_keys_and_values( - backing.backing, - dynamic_value_attrs_from_serializable, - [](std::pair const - &p) { - return std::pair{realm_instance_from_serializable(p.first), - realm_event_from_serializable(p.second)}; - }), + /*backing=*/map_keys_and_values( + backing.backing, + dynamic_value_attrs_from_serializable, + [](std::pair const + &p) { + return std::pair{realm_instance_from_serializable(p.first), + realm_event_from_serializable(p.second)}; + }), }; } diff --git a/lib/realm-execution/test/src/realm-execution/test_e2e.cc b/lib/realm-execution/test/src/realm-execution/test_e2e.cc index d7f777e686..5f681a6a3d 100644 --- a/lib/realm-execution/test/src/realm-execution/test_e2e.cc +++ b/lib/realm-execution/test/src/realm-execution/test_e2e.cc @@ -233,8 +233,7 @@ TEST_SUITE(FF_TEST_SUITE) { GenericTensorAccessorW label_tensor = allocator.allocate_tensor(cfg.label_shape); - std::map - input_tensors; + std::map input_tensors; DistributedFfHandle device_handle = create_distributed_ff_handle(ctx, @@ -312,8 +311,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { GenericTensorAccessorW label_tensor = allocator.allocate_tensor(cfg.label_shape); - std::map - input_tensors; + std::map input_tensors; DistributedFfHandle device_handle = create_distributed_ff_handle( ctx, diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc index 708fb73d73..955e7d73d2 100644 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc @@ -252,8 +252,7 @@ TEST_SUITE(FF_TEST_SUITE) { MappedParallelComputationGraph mpcg = make_test_mpcg_for_device_type(DeviceType::CPU); - std::map - input_tensors; + std::map input_tensors; OptimizerAttrs optimizer_attrs = OptimizerAttrs{ SGDOptimizerAttrs{ @@ -317,8 +316,7 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { }, }; - std::map - input_tensors; + std::map input_tensors; DistributedFfHandle device_handle = create_distributed_ff_handle( ctx, diff --git a/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h b/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h index 3276e704c7..76921ec34a 100644 --- a/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h +++ b/lib/substitutions/include/substitutions/apply_substitution/perform_shape_inference.h @@ -33,8 +33,8 @@ LabelledOpenKwargDataflowGraphView const &g, - std::map, - ParallelTensorShape> const &input_shapes); + std::map, ParallelTensorShape> const + &input_shapes); } // namespace FlexFlow diff --git a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h index 178ba80bbf..70a6cfd18d 100644 --- a/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h +++ b/lib/substitutions/include/substitutions/sub_parallel_computation_graph.h @@ -40,12 +40,12 @@ std::map get_outgoing_tensors(SubParallelComputationGraph const &, parallel_layer_guid_t const &); -std::set get_subgraph_incoming_edges( - SubParallelComputationGraph const &, - std::set const &); -std::set get_subgraph_outgoing_edges( - SubParallelComputationGraph const &, - std::set const &); +std::set + get_subgraph_incoming_edges(SubParallelComputationGraph const &, + std::set const &); +std::set + get_subgraph_outgoing_edges(SubParallelComputationGraph const &, + std::set const &); std::set get_open_parallel_tensor_uses(SubParallelComputationGraph const &, diff --git a/lib/substitutions/include/substitutions/substitution_builder.h b/lib/substitutions/include/substitutions/substitution_builder.h index 4f180f6eb4..7d296311b0 100644 --- a/lib/substitutions/include/substitutions/substitution_builder.h +++ b/lib/substitutions/include/substitutions/substitution_builder.h @@ -20,16 +20,13 @@ struct SubstitutionBuilder { std::map add_pattern_node( OperatorAttributePattern const &node_pattern, std::map const &inputs, - std::map const - &output_patterns, + std::map const &output_patterns, std::optional const &name = std::nullopt); - std::map - add_output_graph_node( - OutputOperatorAttrsAssignment const &node_expr, - std::map const - &inputs, - std::set const &output_slots); + std::map add_output_graph_node( + OutputOperatorAttrsAssignment const &node_expr, + std::map const &inputs, + std::set const &output_slots); PatternNode pattern_node_named(std::string const &) const; PatternInput pattern_input_named(std::string const &) const; diff --git a/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h b/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h index 06a76f320f..ea883a6bed 100644 --- a/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h +++ b/lib/substitutions/include/substitutions/unlabelled/unlabelled_graph_pattern.h @@ -12,18 +12,14 @@ namespace FlexFlow { size_t num_nodes(UnlabelledGraphPattern const &); bool is_singleton_pattern(UnlabelledGraphPattern const &); -std::set - get_pattern_nodes(UnlabelledGraphPattern const &); -std::set - get_pattern_values(UnlabelledGraphPattern const &); +std::set get_pattern_nodes(UnlabelledGraphPattern const &); +std::set get_pattern_values(UnlabelledGraphPattern const &); std::vector get_topological_ordering(UnlabelledGraphPattern const &); -std::set - get_pattern_inputs(UnlabelledGraphPattern const &); +std::set get_pattern_inputs(UnlabelledGraphPattern const &); -std::set - get_pattern_edges(UnlabelledGraphPattern const &); +std::set get_pattern_edges(UnlabelledGraphPattern const &); std::map get_inputs_to_pattern_node(UnlabelledGraphPattern const &, diff --git a/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h b/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h index 6311a5a40e..f713e05633 100644 --- a/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h +++ b/lib/substitutions/include/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.h @@ -11,8 +11,7 @@ namespace FlexFlow { UnlabelledKwargDataflowGraphPatternMatch empty_unlabelled_pattern_match(); -std::set - matched_nodes(UnlabelledKwargDataflowGraphPatternMatch const &); +std::set matched_nodes(UnlabelledKwargDataflowGraphPatternMatch const &); std::optional merge_unlabelled_dataflow_graph_pattern_matches( UnlabelledKwargDataflowGraphPatternMatch const &subpattern_1, diff --git a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc index 6699870669..ffe323717e 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/apply_substitution.cc @@ -9,11 +9,11 @@ #include "substitutions/sub_parallel_computation_graph_data.dtg.h" #include "substitutions/sub_parallel_computation_graph_data.h" #include "substitutions/sub_parallel_computation_graph_edge.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/keys.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" #include "utils/containers/values.h" -#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -49,27 +49,24 @@ SubParallelComputationGraph apply_substitution_from_output_result( SubParallelComputationGraphData pre_data = get_sub_pcg_data(spcg); require_sub_parallel_computation_graph_data_is_valid(pre_data); - std::set pre_nodes = - keys(pre_data.node_data); + std::set pre_nodes = keys(pre_data.node_data); std::set matched_nodes = set_of(values(match.node_assignment)); std::set post_nodes_from_original_graph = set_minus(pre_nodes, matched_nodes); - std::map post_node_data = - [&] { - std::map - post_node_data_from_orig = restrict_keys( - pre_data.node_data, post_nodes_from_original_graph); - std::map - post_node_data_from_sub = output_graph_data.node_data; - - return binary_merge_disjoint_maps(post_node_data_from_orig, - post_node_data_from_sub); - }(); + std::map post_node_data = [&] { + std::map + post_node_data_from_orig = + restrict_keys(pre_data.node_data, post_nodes_from_original_graph); + std::map + post_node_data_from_sub = output_graph_data.node_data; - std::set post_inputs = - pre_data.inputs; + return binary_merge_disjoint_maps(post_node_data_from_orig, + post_node_data_from_sub); + }(); + + std::set post_inputs = pre_data.inputs; std::set post_edges = [&] { std::set post_edges_from_orig = @@ -86,11 +83,10 @@ SubParallelComputationGraph apply_substitution_from_output_result( } }); - std::set post_edges_from_sub = - filter(output_graph_data.edges, - [&](SubParallelComputationGraphEdge const &e) { - return e.raw_edge.is_internal_edge(); - }); + std::set post_edges_from_sub = filter( + output_graph_data.edges, [&](SubParallelComputationGraphEdge const &e) { + return e.raw_edge.is_internal_edge(); + }); bidict output_orig_pattern_mapping = get_output_mapping_for_pcg_pattern_match( @@ -109,10 +105,9 @@ SubParallelComputationGraph apply_substitution_from_output_result( input_parallel_tensor_guid_t output_graph_input = output_expr_to_result_sub_pcg_mapping.input_mapping.at_r( output_expr_input); - std::set uses = - get_open_parallel_tensor_uses( - substitution_output_graph, - open_parallel_tensor_guid_from_input(output_graph_input)); + std::set uses = get_open_parallel_tensor_uses( + substitution_output_graph, + open_parallel_tensor_guid_from_input(output_graph_input)); for (parallel_tensor_use_t const &use : uses) { SubParallelComputationGraphEdge new_edge = subpcg_edge_from_tensor_and_use(base_graph_tensor, use); @@ -148,8 +143,8 @@ SubParallelComputationGraph apply_substitution_from_output_result( }); }(); - std::map - post_value_data = [&] { + std::map post_value_data = + [&] { std::map post_value_data_from_orig = filter_keys( pre_data.value_data, [&](open_parallel_tensor_guid_t const &t) { @@ -169,7 +164,7 @@ SubParallelComputationGraph apply_substitution_from_output_result( std::map post_value_data_from_sub = output_graph_data.value_data; return binary_merge_disjoint_maps(post_value_data_from_orig, - post_value_data_from_sub); + post_value_data_from_sub); }(); SubParallelComputationGraphData post_data = SubParallelComputationGraphData{ diff --git a/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc b/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc index b308bc0bde..a5462ab5d0 100644 --- a/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc +++ b/lib/substitutions/src/substitutions/apply_substitution/evaluate_substitution_output.cc @@ -24,11 +24,10 @@ std::pair evaluate_substitution_output(SubParallelComputationGraph const &spcg, Substitution const &sub, PCGPatternMatch const &match) { - std::map node_match = - map_values(match.node_assignment.as_map(), - [&](parallel_layer_guid_t const &n) { - return get_operator_attrs(spcg, n); - }); + std::map node_match = map_values( + match.node_assignment.as_map(), [&](parallel_layer_guid_t const &n) { + return get_operator_attrs(spcg, n); + }); bidict new_node_id_permutation = generate_new_node_id_permutation(sub.output_graph_expr.raw_graph); @@ -72,9 +71,9 @@ std::pair bidict result_input_map = bidict_transform_keys( bidict_transform_values(new_input_id_permutation, - [](KwargDataflowGraphInput const &i) { - return OutputGraphExprInput{i}; - }), + [](KwargDataflowGraphInput const &i) { + return OutputGraphExprInput{i}; + }), [](KwargDataflowGraphInput const &i) { return input_parallel_tensor_guid_t{i}; }); @@ -86,16 +85,16 @@ std::pair [](Node const &n) { return OutputGraphExprNode{n}; }), [](NewNode const &n) { return parallel_layer_guid_t{n.raw_node}; }); - std::map, ParallelTensorShape> - input_shapes = map_values( - map_keys(match.input_assignment, - [&](PatternInput const &i) { - return result_input_map.at_r(sub.inputs_mapping.at_l(i)) - .raw_dataflow_graph_input; - }), - [&](open_parallel_tensor_guid_t const &v) { - return spcg.raw_graph.at(v.raw_open_dataflow_value).shape; - }); + std::map, ParallelTensorShape> input_shapes = + map_values(map_keys(match.input_assignment, + [&](PatternInput const &i) { + return result_input_map + .at_r(sub.inputs_mapping.at_l(i)) + .raw_dataflow_graph_input; + }), + [&](open_parallel_tensor_guid_t const &v) { + return spcg.raw_graph.at(v.raw_open_dataflow_value).shape; + }); LabelledOpenKwargDataflowGraphView const &g, - std::map, - ParallelTensorShape> const &input_shapes) { + std::map, ParallelTensorShape> const + &input_shapes) { - std::map, - ParallelTensorShape> + std::map, ParallelTensorShape> inferred = map_keys(input_shapes, [](KwargDataflowGraphInput const &i) @@ -53,8 +52,8 @@ LabelledOpenKwargDataflowGraphView - incoming_tensor_roles = get_incoming_tensor_roles(n_attrs.op_attrs); + std::map incoming_tensor_roles = + get_incoming_tensor_roles(n_attrs.op_attrs); ASSERT(is_subseteq_of(keys(incoming_shapes), keys(incoming_tensor_roles))); @@ -75,17 +74,16 @@ LabelledOpenKwargDataflowGraphView - inferred_weight_shapes = - get_weight_shapes(n_attrs.op_attrs, input_shapes); + std::map inferred_weight_shapes = + get_weight_shapes(n_attrs.op_attrs, input_shapes); ASSERT(weight_shapes == inferred_weight_shapes); std::map output_shapes = get_output_shapes(n_attrs.op_attrs, input_shapes); - std::map> - outputs = get_outgoing_kwarg_dataflow_outputs_for_node(g, n); + std::map> outputs = + get_outgoing_kwarg_dataflow_outputs_for_node(g, n); for (auto const &[output, shape] : values(zip_values_strict(outputs, output_shapes))) { diff --git a/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc b/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc index 529b6d908b..af36c0c87c 100644 --- a/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc +++ b/lib/substitutions/src/substitutions/output_graph/materialize_operator_from_attrs_map.cc @@ -6,8 +6,7 @@ namespace FlexFlow { struct Accessor { - Accessor( - std::map const &m) + Accessor(std::map const &m) : m(m) {} std::map const &m; @@ -29,8 +28,7 @@ struct Accessor { }; PCGOperatorAttrs materialize_operator_from_attrs_map( - std::map const - &attrs) { + std::map const &attrs) { OperatorType op_type = attrs.at(OperatorAttributeKey::OP_TYPE).get(); diff --git a/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc b/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc index 1a2a1034a9..75af6d3038 100644 --- a/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc +++ b/lib/substitutions/src/substitutions/output_graph/output_graph_expr.cc @@ -16,9 +16,9 @@ std::set get_nodes(OutputGraphExpr const &g) { std::map get_node_outputs(OutputGraphExpr const &g, OutputGraphExprNode const &n) { - std::map> - raw_outputs = get_outgoing_kwarg_dataflow_outputs_for_node( - g.raw_graph, n.raw_graph_node); + std::map> raw_outputs = + get_outgoing_kwarg_dataflow_outputs_for_node(g.raw_graph, + n.raw_graph_node); return map_values(raw_outputs, [](KwargDataflowOutput const &o) { diff --git a/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc b/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc index 9298c5b35c..f7db80619e 100644 --- a/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc +++ b/lib/substitutions/src/substitutions/output_graph/output_operator_attrs_assignment.cc @@ -2,9 +2,9 @@ #include "substitutions/operator_pattern/get_attribute_map.h" #include "substitutions/output_graph/materialize_operator_from_attrs_map.h" #include "substitutions/output_graph/output_operator_attribute_expr.h" +#include "utils/containers/binary_merge_maps_with_right_dominating.h" #include "utils/containers/map_values.h" #include "utils/exception.h" -#include "utils/containers/binary_merge_maps_with_right_dominating.h" namespace FlexFlow { @@ -16,9 +16,8 @@ PCGOperatorAttrs materialize_output_operator_from_attrs_assignment( OutputOperatorAttrsAssignment const &attrs_assignment, std::map const &node_match) { - std::map - template_attrs_map = [&]() - -> std::map { + std::map template_attrs_map = + [&]() -> std::map { if (attrs_assignment.template_operator.has_value()) { PatternNode template_node = attrs_assignment.template_operator.value(); PCGOperatorAttrs template_op_attrs = node_match.at(template_node); @@ -28,16 +27,16 @@ PCGOperatorAttrs materialize_output_operator_from_attrs_assignment( } }(); - std::map - assignments_attrs_map = map_values( - attrs_assignment.assignments, - [&](OutputOperatorAttributeExpr const &expr) { - return evaluate_output_operator_attribute_expr(expr, node_match); - }); + std::map assignments_attrs_map = + map_values(attrs_assignment.assignments, + [&](OutputOperatorAttributeExpr const &expr) { + return evaluate_output_operator_attribute_expr(expr, + node_match); + }); - std::map - joined_attrs_map = binary_merge_maps_with_right_dominating( - template_attrs_map, assignments_attrs_map); + std::map joined_attrs_map = + binary_merge_maps_with_right_dominating(template_attrs_map, + assignments_attrs_map); return materialize_operator_from_attrs_map(joined_attrs_map); } diff --git a/lib/substitutions/src/substitutions/pcg_pattern.cc b/lib/substitutions/src/substitutions/pcg_pattern.cc index c8b521bf1d..d86b10c588 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern.cc @@ -100,9 +100,9 @@ std::set get_inputs(PCGPattern const &p) { std::map get_pattern_node_outputs(PCGPattern const &pattern, PatternNode const &node) { - std::map> - raw_outputs = get_outgoing_kwarg_dataflow_outputs_for_node( - pattern.raw_graph, node.raw_node); + std::map> raw_outputs = + get_outgoing_kwarg_dataflow_outputs_for_node(pattern.raw_graph, + node.raw_node); return map_values(raw_outputs, [](KwargDataflowOutput const &o) { diff --git a/lib/substitutions/src/substitutions/pcg_pattern_match.cc b/lib/substitutions/src/substitutions/pcg_pattern_match.cc index 0c0d4b11bb..b03eb8340b 100644 --- a/lib/substitutions/src/substitutions/pcg_pattern_match.cc +++ b/lib/substitutions/src/substitutions/pcg_pattern_match.cc @@ -4,9 +4,9 @@ #include "substitutions/unlabelled/unlabelled_graph_pattern.h" #include "utils/bidict/algorithms/bidict_from_keys_and_values.h" #include "utils/bidict/algorithms/bidict_from_map.h" +#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/bidict/algorithms/binary_merge_disjoint_bidicts.h" #include "utils/bidict/algorithms/exhaustive_relational_join.h" -#include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/map_values.h" #include "utils/containers/values.h" @@ -57,8 +57,7 @@ void assert_pcg_pattern_match_is_valid_for_pattern_and_subpcg( PCGPatternMatch const &match, PCGPattern const &pattern, SubParallelComputationGraph const &spcg) { - std::set spcg_nodes = - spcg_get_parallel_layers(spcg); + std::set spcg_nodes = spcg_get_parallel_layers(spcg); std::set match_nodes = match.node_assignment.right_values(); ASSERT(is_subseteq_of(match_nodes, spcg_nodes)); @@ -75,8 +74,7 @@ void assert_pcg_pattern_match_is_valid_for_pattern_and_subpcg( ASSERT(match_pattern_nodes == pattern_nodes); std::set pattern_inputs = get_inputs(pattern); - std::set match_pattern_inputs = - keys(match.input_assignment); + std::set match_pattern_inputs = keys(match.input_assignment); ASSERT(pattern_inputs == match_pattern_inputs); } diff --git a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc index 427ac6747d..8a27abafe5 100644 --- a/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc +++ b/lib/substitutions/src/substitutions/sub_parallel_computation_graph.cc @@ -100,9 +100,9 @@ std::map }); } -std::set get_subgraph_outgoing_edges( - SubParallelComputationGraph const &spcg, - std::set const &layers) { +std::set + get_subgraph_outgoing_edges(SubParallelComputationGraph const &spcg, + std::set const &layers) { std::set> raw_edges = get_kwarg_dataflow_subgraph_outgoing_edges( spcg.raw_graph, transform(layers, [](parallel_layer_guid_t const &l) { @@ -120,9 +120,9 @@ std::set get_subgraph_incoming_edges( transform(subgraph, [](parallel_layer_guid_t const &l) { return l.raw_graph_node; }); - std::set> - raw_incoming_edges = get_open_kwarg_dataflow_subgraph_incoming_edges( - spcg.raw_graph, raw_subgraph); + std::set> raw_incoming_edges = + get_open_kwarg_dataflow_subgraph_incoming_edges(spcg.raw_graph, + raw_subgraph); return transform(raw_incoming_edges, [](OpenKwargDataflowEdge const &e) { diff --git a/lib/substitutions/src/substitutions/substitution.cc b/lib/substitutions/src/substitutions/substitution.cc index d19a2329e8..093574dc17 100644 --- a/lib/substitutions/src/substitutions/substitution.cc +++ b/lib/substitutions/src/substitutions/substitution.cc @@ -132,10 +132,8 @@ bool is_isomorphic_to(Substitution const &l, Substitution const &r) { bool is_valid_substitution(Substitution const &sub) { { - std::set pattern_inputs = - get_inputs(sub.pcg_pattern); - std::set mapped_inputs = - left_entries(sub.inputs_mapping); + std::set pattern_inputs = get_inputs(sub.pcg_pattern); + std::set mapped_inputs = left_entries(sub.inputs_mapping); if (pattern_inputs != mapped_inputs) { return false; diff --git a/lib/substitutions/src/substitutions/substitution_builder.cc b/lib/substitutions/src/substitutions/substitution_builder.cc index 0bfcac3c08..80c493f92c 100644 --- a/lib/substitutions/src/substitutions/substitution_builder.cc +++ b/lib/substitutions/src/substitutions/substitution_builder.cc @@ -54,13 +54,11 @@ std::pair SubstitutionBuilder::add_input( }; } -std::map - SubstitutionBuilder::add_pattern_node( - OperatorAttributePattern const &node_pattern, - std::map const &inputs, - std::map const - &output_patterns, - std::optional const &maybe_name) { +std::map SubstitutionBuilder::add_pattern_node( + OperatorAttributePattern const &node_pattern, + std::map const &inputs, + std::map const &output_patterns, + std::optional const &maybe_name) { KwargNodeAddedResult node_added = this->pattern_g.add_node( node_pattern, map_values(inputs, raw_open_dataflow_value_from_pattern_value), diff --git a/lib/substitutions/src/substitutions/unity_substitution_set.cc b/lib/substitutions/src/substitutions/unity_substitution_set.cc index 551d1101a6..e57aa4b610 100644 --- a/lib/substitutions/src/substitutions/unity_substitution_set.cc +++ b/lib/substitutions/src/substitutions/unity_substitution_set.cc @@ -475,25 +475,24 @@ Substitution create_partition_attention_combine(positive_int num_heads, OutputGraphExprValue o_replicate_weight_output = insert_replicate(b, degree, o_weights); - std::map o_attention_inputs = + std::map o_attention_inputs = { { - { - TensorSlotName::QUERY, - o_partition_query_input_output, - }, - { - TensorSlotName::KEY, - o_partition_key_input_output, - }, - { - TensorSlotName::VALUE, - o_partition_value_input_output, - }, - { - TensorSlotName::WEIGHT, - o_replicate_weight_output, - }, - }; + TensorSlotName::QUERY, + o_partition_query_input_output, + }, + { + TensorSlotName::KEY, + o_partition_key_input_output, + }, + { + TensorSlotName::VALUE, + o_partition_value_input_output, + }, + { + TensorSlotName::WEIGHT, + o_replicate_weight_output, + }, + }; OutputOperatorAttrsAssignment attention_expr = OutputOperatorAttrsAssignment{ b.pattern_node_named(attention_name), @@ -569,25 +568,24 @@ Substitution create_replicate_attention_reduce(positive_int num_heads, OutputGraphExprValue o_partition_weight_output = insert_partition(b, degree, ff_dim_t{1_n}, o_weights); - std::map o_attention_inputs = + std::map o_attention_inputs = { { - { - TensorSlotName::QUERY, - o_replicate_query_input_output, - }, - { - TensorSlotName::KEY, - o_replicate_key_input_output, - }, - { - TensorSlotName::VALUE, - o_replicate_value_input_output, - }, - { - TensorSlotName::WEIGHT, - o_partition_weight_output, - }, - }; + TensorSlotName::QUERY, + o_replicate_query_input_output, + }, + { + TensorSlotName::KEY, + o_replicate_key_input_output, + }, + { + TensorSlotName::VALUE, + o_replicate_value_input_output, + }, + { + TensorSlotName::WEIGHT, + o_partition_weight_output, + }, + }; OutputOperatorAttrsAssignment attention_expr = OutputOperatorAttrsAssignment{ b.pattern_node_named(attention_name), diff --git a/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc b/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc index 3d3c2f8211..9899bd92d6 100644 --- a/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc +++ b/lib/substitutions/src/substitutions/unlabelled/find_pattern_matches.cc @@ -35,8 +35,7 @@ static std::optional std::map pattern_outputs = get_outputs_from_pattern_node(pattern, pattern_node); - std::map> + std::map> graph_outputs = map_values( get_outgoing_kwarg_dataflow_outputs_for_node(graph, graph_node), [](KwargDataflowOutput const &o) { @@ -49,15 +48,13 @@ static std::optional std::map pattern_node_inputs = get_inputs_to_pattern_node(pattern, pattern_node); - std::set pattern_graph_inputs = - get_pattern_inputs(pattern); + std::set pattern_graph_inputs = get_pattern_inputs(pattern); ASSERT(set_of(values(pattern_node_inputs)) == transform(pattern_graph_inputs, [](PatternInput const &i) { return PatternValue{i}; })); - std::map> + std::map> graph_node_inputs = get_incoming_open_kwarg_dataflow_values_for_node(graph, graph_node); diff --git a/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc b/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc index 931ec43af3..cd5661d32d 100644 --- a/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc +++ b/lib/substitutions/src/substitutions/unlabelled/pattern_matching.cc @@ -139,8 +139,8 @@ bool pattern_matches_subgraph_under( } } - std::set> - concrete_edges = get_all_open_kwarg_dataflow_edges(subgraph); + std::set> concrete_edges = + get_all_open_kwarg_dataflow_edges(subgraph); std::set> concrete_edge_from_match = transform(get_pattern_edges(pattern), @@ -153,8 +153,8 @@ bool pattern_matches_subgraph_under( return false; } - std::set> - concrete_values = get_all_open_kwarg_dataflow_values(subgraph); + std::set> concrete_values = + get_all_open_kwarg_dataflow_values(subgraph); std::set> concrete_values_from_match = transform(get_pattern_values(pattern), @@ -184,8 +184,7 @@ bool unlabelled_pattern_does_match( MatchAdditionalCriterion const &additional_criterion) { std::set> - matched_by_pattern_inputs = - set_of(values(match.input_assignment)); + matched_by_pattern_inputs = set_of(values(match.input_assignment)); ASSERT(left_entries(match.node_assignment) == get_pattern_nodes(pattern)); ASSERT( diff --git a/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc b/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc index 1b6741ef9e..652d614875 100644 --- a/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc +++ b/lib/substitutions/src/substitutions/unlabelled/unlabelled_graph_pattern.cc @@ -23,27 +23,23 @@ bool is_singleton_pattern(UnlabelledGraphPattern const &pattern) { return num_nodes(pattern) == 1; } -std::set - get_pattern_nodes(UnlabelledGraphPattern const &p) { +std::set get_pattern_nodes(UnlabelledGraphPattern const &p) { return transform(get_nodes(p.raw_graph), [](Node const &n) { return PatternNode{n}; }); } -std::set - get_pattern_values(UnlabelledGraphPattern const &p) { +std::set get_pattern_values(UnlabelledGraphPattern const &p) { return transform(get_all_open_kwarg_dataflow_values(p.raw_graph), pattern_value_from_raw_open_kwarg_dataflow_value); } -std::set - get_pattern_inputs(UnlabelledGraphPattern const &p) { +std::set get_pattern_inputs(UnlabelledGraphPattern const &p) { return transform( get_all_kwarg_dataflow_graph_inputs(p.raw_graph), [](KwargDataflowGraphInput const &i) { return PatternInput{i}; }); } -std::set - get_pattern_edges(UnlabelledGraphPattern const &p) { +std::set get_pattern_edges(UnlabelledGraphPattern const &p) { return transform(get_all_open_kwarg_dataflow_edges(p.raw_graph), pattern_edge_from_raw_open_dataflow_edge); } diff --git a/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc b/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc index ce46faa2f0..73476e9c9e 100644 --- a/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc +++ b/lib/substitutions/src/substitutions/unlabelled/unlabelled_kwarg_dataflow_graph_pattern_match.cc @@ -33,23 +33,20 @@ std::optional std::map> merged_input_assignment = ({ - std::map> + std::map> lifted_input_assignment_1 = map_keys( subpattern_1.input_assignment, [&](PatternInput const &pi1) { return merged_graph_values_to_inputs_of_1.at_r(pi1); }); - std::map> + std::map> lifted_input_assignment_2 = map_keys( subpattern_2.input_assignment, [&](PatternInput const &pi2) { return merged_graph_values_to_inputs_of_2.at_r(pi2); }); std::optional< - std::map>> - merged = try_merge_nondisjoint_maps( - lifted_input_assignment_1, lifted_input_assignment_2); + std::map>> + merged = try_merge_nondisjoint_maps(lifted_input_assignment_1, + lifted_input_assignment_2); if (!merged.has_value()) { return std::nullopt; } diff --git a/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc b/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc index dbdb5cb5ed..f6b9e138f0 100644 --- a/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc +++ b/lib/substitutions/test/src/substitutions/apply_substitution/evaluate_substitution_output.cc @@ -320,8 +320,7 @@ TEST_SUITE(FF_TEST_SUITE) { result_i_activation, result_i_weights, }, - std::map{ + std::map{ { open_parallel_tensor_guid_from_input(result_i_activation), correct_result_i_activation_attrs, diff --git a/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc b/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc index 46efd88cc9..020f6e8210 100644 --- a/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc +++ b/lib/substitutions/test/src/substitutions/apply_substitution/perform_shape_inference.cc @@ -174,10 +174,9 @@ TEST_SUITE(FF_TEST_SUITE) { KwargDataflowOutput o2 = require_only_key(n2_added_result.outputs, TensorSlotName::OUTPUT); - std::map, ParallelTensorShape> - input_shapes = { - {i0, i0_shape}, - }; + std::map, ParallelTensorShape> input_shapes = { + {i0, i0_shape}, + }; LabelledOpenKwargDataflowGraphView result = set_of( - find_pattern_matches(pattern, sub_pcg_from_full_pcg(pcg))); + std::set result = + set_of(find_pattern_matches(pattern, sub_pcg_from_full_pcg(pcg))); PCGPatternMatch match1 = PCGPatternMatch{ bidict{ @@ -350,8 +350,8 @@ TEST_SUITE(FF_TEST_SUITE) { PCGPattern pattern = PCGPattern{g}; - std::set result = set_of( - find_pattern_matches(pattern, sub_pcg_from_full_pcg(pcg))); + std::set result = + set_of(find_pattern_matches(pattern, sub_pcg_from_full_pcg(pcg))); CHECK(result.size() == 3); } diff --git a/lib/substitutions/test/src/substitutions/unity_substitution_set.cc b/lib/substitutions/test/src/substitutions/unity_substitution_set.cc index 301c1363de..8e05e85f9e 100644 --- a/lib/substitutions/test/src/substitutions/unity_substitution_set.cc +++ b/lib/substitutions/test/src/substitutions/unity_substitution_set.cc @@ -43,8 +43,8 @@ parallel_tensor_guid_t add_single_output_layer( ParallelLayerAttrs const &layer_attrs, std::map const &inputs, std::map const &weights, - std::optional> const - &outputs = std::nullopt) { + std::optional> const &outputs = + std::nullopt) { return get_single_output( add_parallel_layer(pcg, layer_attrs, inputs, weights, outputs)); diff --git a/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc b/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc index 60f7d2929a..7a2625ec60 100644 --- a/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc +++ b/lib/substitutions/test/src/substitutions/unlabelled/find_pattern_matches.cc @@ -136,8 +136,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict>{}}; - std::map> + std::map> n1_incoming = { { TensorSlotName::INPUT, @@ -152,21 +151,17 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_incoming_edges") { SUBCASE("n0") { - std::map> + std::map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n0); - std::map> + std::map> correct = {}; CHECK(result == correct); } SUBCASE("n1") { - std::map> + std::map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n1); - std::map> + std::map> correct = n1_incoming; CHECK(result == correct); } @@ -175,8 +170,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_open_kwarg_dataflow_subgraph_inputs") { std::set> result = get_open_kwarg_dataflow_subgraph_inputs(graph, {n0, n1}); - std::set> correct = - {}; + std::set> correct = {}; CHECK(result == correct); } @@ -194,8 +188,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("inputs") { - std::set> result = - g.get_inputs(); + std::set> result = g.get_inputs(); std::set> correct = {}; CHECK(result == correct); } diff --git a/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc b/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc index 4866389a20..6f64603e92 100644 --- a/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc +++ b/lib/substitutions/test/src/substitutions/unlabelled/pattern_matching.cc @@ -171,8 +171,7 @@ TEST_SUITE(FF_TEST_SUITE) { bidict>{}}; - std::map> + std::map> n1_incoming = { { TensorSlotName::INPUT, @@ -187,21 +186,17 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_incoming_open_kwarg_dataflow_edges_for_node") { SUBCASE("n0") { - std::map> + std::map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n0); - std::map> + std::map> correct = {}; CHECK(result == correct); } SUBCASE("n1") { - std::map> + std::map> result = get_incoming_open_kwarg_dataflow_edges_for_node(graph, n1); - std::map> + std::map> correct = n1_incoming; CHECK(result == correct); } @@ -210,8 +205,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_open_kwarg_dataflow_subgraph_inputs") { std::set> result = get_open_kwarg_dataflow_subgraph_inputs(graph, {n0, n1}); - std::set> correct = - {}; + std::set> correct = {}; CHECK(result == correct); } @@ -228,8 +222,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("inputs") { - std::set> result = - g.get_inputs(); + std::set> result = g.get_inputs(); std::set> correct = {}; CHECK(result == correct); } diff --git a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h index fa8776f5ef..046ac1add7 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/copy_insertion.h @@ -4,9 +4,9 @@ #include "task-spec/dynamic_graph/dynamic_node_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_copy_info.dtg.h" #include "task-spec/dynamic_graph/internal_dynamic_slot_site.dtg.h" -#include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" namespace FlexFlow { @@ -15,17 +15,18 @@ bool value_is_mapped(DynamicValueAttrs const &); void require_node_is_ready_for_copy_insertion(DynamicNodeAttrs const &); void require_value_is_ready_for_copy_insertion(DynamicValueAttrs const &); -void require_invocation_is_ready_for_copy_insertion(DynamicNodeInvocation const &); -void require_graph_is_ready_for_copy_insertion(DynamicOpenDataflowGraph const &); +void require_invocation_is_ready_for_copy_insertion( + DynamicNodeInvocation const &); +void require_graph_is_ready_for_copy_insertion( + DynamicOpenDataflowGraph const &); void require_value_is_copy_inserted(DynamicValueAttrs const &); void require_invocation_is_fully_copy_inserted(DynamicNodeInvocation const &); void require_graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &); -std::map - get_mappings_for_invocation( - DynamicNodeInvocation const &, - std::map const &); +std::map get_mappings_for_invocation( + DynamicNodeInvocation const &, + std::map const &); DynamicNodeInvocation apply_mappings_for_invocation( dynamic_invocation_id_t const &, @@ -44,9 +45,11 @@ std::set copies_for_internal_value( DynamicValueAttrs const &value_attrs, InternalDynamicSlotSite const &src_site, ParallelTensorMapping const &src_site_mapping, - std::map const &sink_site_mappings); + std::map const + &sink_site_mappings); -std::set infer_all_copies_in_graph(DynamicOpenDataflowGraph const &); +std::set + infer_all_copies_in_graph(DynamicOpenDataflowGraph const &); std::map resolve_tensor_mappings(DynamicOpenDataflowGraph const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h index 3ba9d69c3a..2178b877c6 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_invocation.h @@ -1,7 +1,6 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_INVOCATION_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_NODE_INVOCATION_H -#include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "pcg/tensor_direction.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" @@ -9,15 +8,19 @@ namespace FlexFlow { -bool invocation_fully_satisfies(DynamicNodeInvocation const &, - std::function const &node_condition, - std::function const &value_condition, - std::function const &slot_condition); +bool invocation_fully_satisfies( + DynamicNodeInvocation const &, + std::function const &node_condition, + std::function const &value_condition, + std::function const &slot_condition); -void require_invocation_fully_satisfies(DynamicNodeInvocation const &, - std::function const &require_node_condition, - std::function const &require_value_condition, - std::function const &require_slot_condition); +void require_invocation_fully_satisfies( + DynamicNodeInvocation const &, + std::function const &require_node_condition, + std::function const + &require_value_condition, + std::function const + &require_slot_condition); std::map get_slot_map_for_direction(DynamicNodeInvocation const &, TensorDirection); @@ -26,13 +29,15 @@ TrainingOpType dynamic_node_invocation_get_op_type(DynamicNodeInvocation const &); std::set - get_incoming_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &, DynamicNodeInvocation const &); + get_incoming_dynamic_slot_sites_for_invocation( + dynamic_invocation_id_t const &, DynamicNodeInvocation const &); -std::set - get_output_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &, DynamicNodeInvocation const &); +std::set get_output_dynamic_slot_sites_for_invocation( + dynamic_invocation_id_t const &, DynamicNodeInvocation const &); std::set - get_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &, DynamicNodeInvocation const &); + get_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &, + DynamicNodeInvocation const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h index 2e41f223f6..3b28ea5149 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_node_mapping.h @@ -10,8 +10,8 @@ bidict dynamic_node_mapping_get_shard_bindings(DynamicNodeMapping const &); OperatorAtomicTaskShardBinding - dynamic_node_mapping_get_shard_binding_for_device(DynamicNodeMapping const &, - global_device_id_t const &); + dynamic_node_mapping_get_shard_binding_for_device( + DynamicNodeMapping const &, global_device_id_t const &); bidict dynamic_node_mapping_bindings_for_slot_name(DynamicNodeMapping const &, diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h index 82bfe59a15..29339beb3e 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_open_dataflow_graph.h @@ -2,12 +2,12 @@ #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_OPEN_DATAFLOW_GRAPH_H #include "task-spec/dynamic_graph/dynamic_graph_edge.dtg.h" +#include "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_slot_site.dtg.h" -#include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h" -#include "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_id_t.dtg.h" +#include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h" namespace FlexFlow { @@ -44,28 +44,30 @@ std::set get_dynamic_invocation_set(DynamicOpenDataflowGraph const &); std::set - dynamic_graph_get_internal_values(DynamicOpenDataflowGraph const &); + dynamic_graph_get_internal_values(DynamicOpenDataflowGraph const &); std::set - dynamic_graph_get_external_values(DynamicOpenDataflowGraph const &); - -dynamic_invocation_id_t dynamic_graph_get_id_for_invocation(DynamicOpenDataflowGraph const &, - DynamicNodeInvocation const &); -DynamicNodeInvocation dynamic_graph_get_invocation_for_id(DynamicOpenDataflowGraph const &, - dynamic_invocation_id_t const &); - -dynamic_value_id_t dynamic_graph_get_id_for_value(DynamicOpenDataflowGraph const &, - DynamicValueAttrs const &); -DynamicValueAttrs dynamic_graph_get_value_for_id(DynamicOpenDataflowGraph const &, - dynamic_value_id_t const &); + dynamic_graph_get_external_values(DynamicOpenDataflowGraph const &); + +dynamic_invocation_id_t + dynamic_graph_get_id_for_invocation(DynamicOpenDataflowGraph const &, + DynamicNodeInvocation const &); +DynamicNodeInvocation + dynamic_graph_get_invocation_for_id(DynamicOpenDataflowGraph const &, + dynamic_invocation_id_t const &); + +dynamic_value_id_t + dynamic_graph_get_id_for_value(DynamicOpenDataflowGraph const &, + DynamicValueAttrs const &); +DynamicValueAttrs + dynamic_graph_get_value_for_id(DynamicOpenDataflowGraph const &, + dynamic_value_id_t const &); std::set get_dynamic_graph_edges(DynamicOpenDataflowGraph const &); -std::set - get_dynamic_graph_edges_incoming_to_invocation( - DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); -std::set - get_dynamic_graph_edges_outgoing_from_invocation( - DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); +std::set get_dynamic_graph_edges_incoming_to_invocation( + DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); +std::set get_dynamic_graph_edges_outgoing_from_invocation( + DynamicOpenDataflowGraph const &, DynamicNodeInvocation const &); std::set get_internal_dynamic_slot_sites(DynamicOpenDataflowGraph const &); @@ -87,7 +89,9 @@ std::set dynamic_graph_find_sinks_of_value(DynamicOpenDataflowGraph const &, DynamicValueAttrs const &); -DynamicValueAttrs dynamic_value_attrs_for_slot_site(DynamicOpenDataflowGraph const &, DynamicSlotSite const &); +DynamicValueAttrs + dynamic_value_attrs_for_slot_site(DynamicOpenDataflowGraph const &, + DynamicSlotSite const &); std::optional find_output_value_attrs(DynamicOpenDataflowGraph const &, diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h index ac365b1c26..851facc2fd 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.h @@ -9,9 +9,9 @@ namespace FlexFlow { DynamicValueAttrs decide_dynamic_value_attrs_role(DynamicValueAttrs const &, DynamicTensorRole); -DynamicValueAttrs decide_dynamic_value_attrs_mapping( - DynamicValueAttrs const &, - ParallelTensorMapping const &); +DynamicValueAttrs + decide_dynamic_value_attrs_mapping(DynamicValueAttrs const &, + ParallelTensorMapping const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h index cae4d37229..8370bfe498 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h @@ -2,16 +2,16 @@ #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_DYNAMIC_OPEN_DATAFLOW_GRAPH_FROM_MPCG_H #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.dtg.h" -#include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_layer_invocation_info.dtg.h" +#include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" namespace FlexFlow { DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( - MappedParallelLayerInvocationInfo const &, - DeviceType device_type); + MappedParallelLayerInvocationInfo const &, DeviceType device_type); -DynamicNodeInvocation build_replicate_invocation(MappedParallelLayerInvocationInfo const &); +DynamicNodeInvocation + build_replicate_invocation(MappedParallelLayerInvocationInfo const &); DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( MappedParallelComputationGraph const &, DeviceType device_type); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h index 14bd7600b1..ab55751d77 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/parallel_tensor_mapping.h @@ -5,13 +5,17 @@ namespace FlexFlow { -global_device_id_t pt_mapping_get_device_for_coord(ParallelTensorMapping const &, - ParallelTensorSpaceCoordinate const &); -ParallelTensorSpaceCoordinate pt_mapping_get_coord_for_device(ParallelTensorMapping const &, - global_device_id_t const &); +global_device_id_t + pt_mapping_get_device_for_coord(ParallelTensorMapping const &, + ParallelTensorSpaceCoordinate const &); +ParallelTensorSpaceCoordinate + pt_mapping_get_coord_for_device(ParallelTensorMapping const &, + global_device_id_t const &); -std::set pt_mapping_get_coord_set(ParallelTensorMapping const &); -std::set pt_mapping_get_device_set(ParallelTensorMapping const &); +std::set + pt_mapping_get_coord_set(ParallelTensorMapping const &); +std::set + pt_mapping_get_device_set(ParallelTensorMapping const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h b/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h index 8aeefca69c..3adf66d117 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/pass_expansion.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_PASS_EXPANSION_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_PASS_EXPANSION_H +#include "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" -#include "task-spec/dynamic_graph/dynamic_invocation_id_t.dtg.h" namespace FlexFlow { @@ -16,23 +16,24 @@ void require_value_is_pass_expanded(DynamicValueAttrs const &); void require_value_is_not_pass_expanded(DynamicValueAttrs const &); void require_invocation_is_fully_pass_expanded(DynamicNodeInvocation const &); -void require_invocation_is_ready_for_pass_expansion(DynamicNodeInvocation const &); +void require_invocation_is_ready_for_pass_expansion( + DynamicNodeInvocation const &); void require_graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &); -void require_graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &); +void require_graph_is_ready_for_pass_expansion( + DynamicOpenDataflowGraph const &); std::set - determine_intermediate_values_needed_to_compute_gradients_of_value( - DynamicOpenDataflowGraph const &, - DynamicValueAttrs const &); + determine_intermediate_values_needed_to_compute_gradients_of_value( + DynamicOpenDataflowGraph const &, DynamicValueAttrs const &); std::set - determine_intermediate_values_needed_for_gradient_computation( - DynamicOpenDataflowGraph const &); + determine_intermediate_values_needed_for_gradient_computation( + DynamicOpenDataflowGraph const &); std::set - determine_invocations_needed_in_backward_pass_for_gradient_computation( - DynamicOpenDataflowGraph const &); + determine_invocations_needed_in_backward_pass_for_gradient_computation( + DynamicOpenDataflowGraph const &); DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &, FwbTensorType); DynamicValueAttrs pass_expand_value(DynamicValueAttrs const &, FwbTensorType); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h index d1c38a6694..f8716a754d 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h @@ -1,13 +1,14 @@ #ifndef _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_SERIALIZABLE_DYNAMIC_OPEN_DATAFLOW_GRAPH_H #define _FLEXFLOW_LIB_TASK_SPEC_INCLUDE_TASK_SPEC_DYNAMIC_GRAPH_SERIALIZABLE_DYNAMIC_OPEN_DATAFLOW_GRAPH_H -#include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" +#include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.dtg.h" namespace FlexFlow { SerializableDynamicOpenDataflowGraph - dynamic_open_dataflow_graph_to_serializable(DynamicOpenDataflowGraph const &); + dynamic_open_dataflow_graph_to_serializable( + DynamicOpenDataflowGraph const &); DynamicOpenDataflowGraph dynamic_open_dataflow_graph_from_serializable( SerializableDynamicOpenDataflowGraph const &); diff --git a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h index 68253369d7..338256ae1e 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/shard_expansion.h @@ -3,9 +3,9 @@ #include "task-spec/dynamic_graph/dynamic_node_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_invocation.dtg.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs_sharding_info.dtg.h" -#include "task-spec/dynamic_graph/dynamic_node_invocation_sharding_info.dtg.h" namespace FlexFlow { @@ -16,29 +16,29 @@ void require_graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &); void require_node_is_ready_for_shard_expansion(DynamicNodeAttrs const &); void require_value_is_ready_for_shard_expansion(DynamicValueAttrs const &); -void require_invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &); -void require_graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &); +void require_invocation_is_ready_for_shard_expansion( + DynamicNodeInvocation const &); +void require_graph_is_ready_for_shard_expansion( + DynamicOpenDataflowGraph const &); -[[nodiscard]] DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( - DynamicNodeAttrs const &, - MachineSpaceCoordinate const &); +[[nodiscard]] DynamicNodeAttrs + apply_dynamic_node_attrs_sharding_info(DynamicNodeAttrs const &, + MachineSpaceCoordinate const &); [[nodiscard]] DynamicValueAttrs apply_dynamic_value_attrs_sharding_info( - DynamicValueAttrs const &, - DynamicValueAttrsShardingInfo const &); + DynamicValueAttrs const &, DynamicValueAttrsShardingInfo const &); [[nodiscard]] DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( - DynamicNodeInvocation const &, - DynamicNodeInvocationShardingInfo const &); + DynamicNodeInvocation const &, DynamicNodeInvocationShardingInfo const &); [[nodiscard]] std::set - generate_shard_expansion_for_invocation(DynamicNodeInvocation const &); + generate_shard_expansion_for_invocation(DynamicNodeInvocation const &); [[nodiscard]] std::set - perform_shard_expansion_for_invocation(DynamicNodeInvocation const &); + perform_shard_expansion_for_invocation(DynamicNodeInvocation const &); [[nodiscard]] DynamicOpenDataflowGraph - perform_shard_expansion(DynamicOpenDataflowGraph const &); + perform_shard_expansion(DynamicOpenDataflowGraph const &); } // namespace FlexFlow diff --git a/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h b/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h index f38518e927..f274449b9e 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h +++ b/lib/task-spec/include/task-spec/dynamic_graph/update_insertion.h @@ -7,12 +7,14 @@ namespace FlexFlow { bool node_has_already_had_update_insertion_performed(DynamicNodeAttrs const &); -bool value_has_already_had_update_insertion_performed(DynamicValueAttrs const &); +bool value_has_already_had_update_insertion_performed( + DynamicValueAttrs const &); bool node_is_ready_for_update_insertion(DynamicNodeAttrs const &); bool value_is_ready_for_update_insertion(DynamicValueAttrs const &); -bool no_part_of_graph_has_had_update_insertion_performed(DynamicOpenDataflowGraph const &); +bool no_part_of_graph_has_had_update_insertion_performed( + DynamicOpenDataflowGraph const &); bool graph_is_ready_for_update_insertion(DynamicOpenDataflowGraph const &); std::set diff --git a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc index 544f249243..ef5695792f 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/copy_insertion.cc @@ -15,12 +15,17 @@ #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.h" #include "task-spec/dynamic_graph/parallel_tensor_mapping.dtg.h" +#include "task-spec/dynamic_graph/training_operation_attrs.h" +#include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" +#include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/contains_key.h" #include "utils/containers/count.h" #include "utils/containers/filter_values.h" #include "utils/containers/filtermap_keys.h" #include "utils/containers/filtrans.h" #include "utils/containers/flatmap.h" +#include "utils/containers/get_only.h" #include "utils/containers/map_values2.h" #include "utils/containers/merge_disjoint_maps.h" #include "utils/containers/set_difference.h" @@ -29,12 +34,7 @@ #include "utils/containers/values.h" #include "utils/containers/zip_values_strict_with.h" #include "utils/optional.h" -#include "task-spec/dynamic_graph/training_operation_attrs.h" -#include "utils/bidict/algorithms/bidict_from_unstructured_relation.h" -#include "utils/bidict/algorithms/unstructured_relation_from_bidict.h" #include "utils/overload.h" -#include "utils/containers/binary_merge_disjoint_maps.h" -#include "utils/containers/get_only.h" namespace FlexFlow { @@ -46,7 +46,8 @@ bool value_is_mapped(DynamicValueAttrs const &n) { return n.mapping.has_value(); } -void require_node_is_ready_for_copy_insertion(DynamicNodeAttrs const &node_attrs) { +void require_node_is_ready_for_copy_insertion( + DynamicNodeAttrs const &node_attrs) { ASSERT(node_attrs.op_attrs.has_value()); ASSERT(node_attrs.mapping.has_value()); } @@ -55,19 +56,21 @@ void require_value_is_ready_for_copy_insertion(DynamicValueAttrs const &v) { ASSERT(!v.mapping.has_value(), v); } -void require_invocation_is_ready_for_copy_insertion(DynamicNodeInvocation const &i) { - auto require_slot_is_ready_for_copy_insertion = [](DynamicTensorSlot const &) { - return; - }; +void require_invocation_is_ready_for_copy_insertion( + DynamicNodeInvocation const &i) { + auto require_slot_is_ready_for_copy_insertion = + [](DynamicTensorSlot const &) { return; }; - require_invocation_fully_satisfies(i, + require_invocation_fully_satisfies(i, require_node_is_ready_for_copy_insertion, require_value_is_ready_for_copy_insertion, require_slot_is_ready_for_copy_insertion); } -void require_graph_is_ready_for_copy_insertion(DynamicOpenDataflowGraph const &g) { - require_full_dynamic_graph_satisfies(g, require_invocation_is_ready_for_copy_insertion); +void require_graph_is_ready_for_copy_insertion( + DynamicOpenDataflowGraph const &g) { + require_full_dynamic_graph_satisfies( + g, require_invocation_is_ready_for_copy_insertion); } void require_value_is_copy_inserted(DynamicValueAttrs const &v) { @@ -75,25 +78,23 @@ void require_value_is_copy_inserted(DynamicValueAttrs const &v) { } void require_invocation_is_fully_copy_inserted(DynamicNodeInvocation const &i) { - auto require_node_is_copy_inserted = [](DynamicNodeAttrs const &) { - return; - }; + auto require_node_is_copy_inserted = [](DynamicNodeAttrs const &) { return; }; auto require_slot_is_copy_inserted = [](DynamicTensorSlot const &) { return; }; - require_invocation_fully_satisfies(i, + require_invocation_fully_satisfies(i, require_node_is_copy_inserted, require_value_is_copy_inserted, require_slot_is_copy_inserted); } void require_graph_is_fully_copy_inserted(DynamicOpenDataflowGraph const &g) { - require_full_dynamic_graph_satisfies(g, require_invocation_is_fully_copy_inserted); + require_full_dynamic_graph_satisfies( + g, require_invocation_is_fully_copy_inserted); } - std::map get_mappings_for_invocation_id( dynamic_invocation_id_t const &i, @@ -121,11 +122,11 @@ DynamicNodeInvocation apply_mappings_for_invocation( std::map i_mappings = get_mappings_for_invocation_id(id, all_mappings); - std::map - i_input_mappings = restrict_keys(i_mappings, keys(i.inputs)); + std::map i_input_mappings = + restrict_keys(i_mappings, keys(i.inputs)); - std::map - i_output_mappings = restrict_keys(i_mappings, keys(i.outputs)); + std::map i_output_mappings = + restrict_keys(i_mappings, keys(i.outputs)); auto apply_mapping = [&](DynamicValueAttrs const &v, @@ -147,26 +148,28 @@ DynamicNodeInvocation apply_mappings_for_invocation( return result; } -DynamicNodeInvocation make_copy_invocation(DynamicValueCopyInfo const ©_info) { +DynamicNodeInvocation + make_copy_invocation(DynamicValueCopyInfo const ©_info) { DynamicNodeInvocation result = DynamicNodeInvocation{ /*inputs=*/{ { DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }, - decide_dynamic_value_attrs_mapping(copy_info.value_attrs, copy_info.src_mapping), + decide_dynamic_value_attrs_mapping(copy_info.value_attrs, + copy_info.src_mapping), }, }, /*node_attrs=*/ DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs*/ TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, }, /*outputs=*/ { @@ -176,7 +179,8 @@ DynamicNodeInvocation make_copy_invocation(DynamicValueCopyInfo const ©_info /*slot_tensor_role=*/std::nullopt, /*task_shard=*/std::nullopt, }, - decide_dynamic_value_attrs_mapping(copy_info.value_attrs, copy_info.dst_mapping), + decide_dynamic_value_attrs_mapping(copy_info.value_attrs, + copy_info.dst_mapping), }, }, }; @@ -190,29 +194,30 @@ std::set copies_for_value( DynamicValueAttrs const &v, DynamicSlotSite const &src_site, std::set const &dst_sites, - std::map const &all_mappings) { + std::map const + &all_mappings) { require_value_is_ready_for_copy_insertion(v); - return src_site.visit>(overload { - [&](ExternalDynamicSlotSite const &) -> std::set { - return {}; - }, - [&](InternalDynamicSlotSite const &s) -> std::set { - ParallelTensorMapping src_mapping = all_mappings.at(s); - std::map sink_site_mappings = - restrict_keys(all_mappings, dst_sites); + return src_site.visit>(overload{ + [&](ExternalDynamicSlotSite const &) -> std::set { + return {}; + }, + [&](InternalDynamicSlotSite const &s) -> std::set { + ParallelTensorMapping src_mapping = all_mappings.at(s); + std::map + sink_site_mappings = restrict_keys(all_mappings, dst_sites); - return copies_for_internal_value(v, s, src_mapping, sink_site_mappings); - } - }); + return copies_for_internal_value(v, s, src_mapping, sink_site_mappings); + }}); } std::set copies_for_internal_value( DynamicValueAttrs const &v, InternalDynamicSlotSite const &src_site, ParallelTensorMapping const &src_mapping, - std::map const &sink_site_mappings) { + std::map const + &sink_site_mappings) { require_value_is_ready_for_copy_insertion(v); @@ -225,9 +230,9 @@ std::set copies_for_internal_value( auto make_copy_to = [&](ParallelTensorMapping const &sink_mapping) -> DynamicValueCopyInfo { return DynamicValueCopyInfo{ - /*value_attrs=*/v, - /*src_mapping=*/src_mapping, - /*sink_mapping=*/sink_mapping, + /*value_attrs=*/v, + /*src_mapping=*/src_mapping, + /*sink_mapping=*/sink_mapping, }; }; @@ -235,24 +240,26 @@ std::set copies_for_internal_value( } std::map - resolve_tensor_mappings(DynamicOpenDataflowGraph const &g) -{ + resolve_tensor_mappings(DynamicOpenDataflowGraph const &g) { require_graph_is_ready_for_copy_insertion(g); - std::map resolved_from_node_mappings = - resolve_partial_tensor_mappings_from_node_mappings(g); + std::map + resolved_from_node_mappings = + resolve_partial_tensor_mappings_from_node_mappings(g); - std::map resolved_from_adjacent_values = - resolve_missing_tensor_mappings_from_adjacent_values(g, resolved_from_node_mappings); + std::map + resolved_from_adjacent_values = + resolve_missing_tensor_mappings_from_adjacent_values( + g, resolved_from_node_mappings); std::map result = - binary_merge_disjoint_maps(resolved_from_node_mappings, resolved_from_adjacent_values); + binary_merge_disjoint_maps(resolved_from_node_mappings, + resolved_from_adjacent_values); { std::set all_internal_slot_sites = - get_internal_dynamic_slot_sites(g); - std::set resolved_slot_sites = - keys(result); + get_internal_dynamic_slot_sites(g); + std::set resolved_slot_sites = keys(result); ASSERT(resolved_slot_sites == all_internal_slot_sites); } @@ -264,19 +271,21 @@ std::map DynamicOpenDataflowGraph const &g) { require_graph_is_ready_for_copy_insertion(g); - - auto slots_to_map_for_replicate = [](dynamic_invocation_id_t const &invocation_id, - DynamicNodeInvocation const &invocation) - -> std::set - { + + auto slots_to_map_for_replicate = + [](dynamic_invocation_id_t const &invocation_id, + DynamicNodeInvocation const &invocation) + -> std::set { TrainingOpType op_type = dynamic_node_invocation_get_op_type(invocation); ASSERT(op_type == TrainingOpType{OperatorType::REPLICATE}); - std::set slot_sites = - (invocation.node_attrs.task_type == DynamicTaskType::BWD) - ? get_incoming_dynamic_slot_sites_for_invocation(invocation_id, invocation) - : get_output_dynamic_slot_sites_for_invocation(invocation_id, invocation); + std::set slot_sites = + (invocation.node_attrs.task_type == DynamicTaskType::BWD) + ? get_incoming_dynamic_slot_sites_for_invocation(invocation_id, + invocation) + : get_output_dynamic_slot_sites_for_invocation(invocation_id, + invocation); { InternalDynamicSlotSite slot_site = get_only(slot_sites); @@ -286,11 +295,12 @@ std::map return slot_sites; }; - auto get_mappings_for_invocation = [&](DynamicNodeInvocation const &invocation) + auto get_mappings_for_invocation = + [&](DynamicNodeInvocation const &invocation) -> std::map { - TrainingOpType op_type = dynamic_node_invocation_get_op_type(invocation); - dynamic_invocation_id_t invocation_id = dynamic_graph_get_id_for_invocation(g, invocation); + dynamic_invocation_id_t invocation_id = + dynamic_graph_get_id_for_invocation(g, invocation); std::set slot_sites_to_resolve = [&] { if (op_type == TrainingOpType{OperatorType::REPLICATE}) { @@ -305,7 +315,8 @@ std::map [&](InternalDynamicSlotSite const &s) -> ParallelTensorMapping { return ParallelTensorMapping{ dynamic_node_mapping_bindings_for_slot_name( - assert_unwrap(invocation.node_attrs.mapping), s.slot_name.slot_name), + assert_unwrap(invocation.node_attrs.mapping), + s.slot_name.slot_name), }; }); }; @@ -319,22 +330,21 @@ std::map std::map resolve_missing_tensor_mappings_from_adjacent_values( - DynamicOpenDataflowGraph const &g, - std::map const &resolved_mappings) -{ + DynamicOpenDataflowGraph const &g, + std::map const + &resolved_mappings) { require_graph_is_ready_for_copy_insertion(g); std::set all_internal_slot_sites = - get_internal_dynamic_slot_sites(g); + get_internal_dynamic_slot_sites(g); std::set missing_mappings = - set_minus(all_internal_slot_sites, keys(resolved_mappings)); + set_minus(all_internal_slot_sites, keys(resolved_mappings)); auto get_mapping_for_slot_site_from_adjacent_values = - [&](InternalDynamicSlotSite const &slot_site) - -> ParallelTensorMapping { - - DynamicNodeInvocation invocation = dynamic_graph_get_invocation_for_id(g, slot_site.invocation_id); + [&](InternalDynamicSlotSite const &slot_site) -> ParallelTensorMapping { + DynamicNodeInvocation invocation = + dynamic_graph_get_invocation_for_id(g, slot_site.invocation_id); TrainingOpType op_type = dynamic_node_invocation_get_op_type(invocation); std::optional task_type = invocation.node_attrs.task_type; @@ -345,19 +355,19 @@ std::map ASSERT(slot_site.direction == TensorDirection::OUTPUT); InternalDynamicSlotSite slot_site_sink = - get_only(dynamic_graph_find_sinks_of_slot_site(g, slot_site)); + get_only(dynamic_graph_find_sinks_of_slot_site(g, slot_site)); ASSERT(contains_key(resolved_mappings, slot_site_sink)); return resolved_mappings.at(slot_site_sink); - } else if ( - op_type == replicate_op_type - && (task_type == std::nullopt || task_type == DynamicTaskType::FWD) - ) { + } else if (op_type == replicate_op_type && + (task_type == std::nullopt || + task_type == DynamicTaskType::FWD)) { ASSERT(slot_site.direction == TensorDirection::INCOMING); InternalDynamicSlotSite slot_site_src = - dynamic_graph_find_source_of_slot_site(g, slot_site).require_internal(); + dynamic_graph_find_source_of_slot_site(g, slot_site) + .require_internal(); ASSERT(contains_key(resolved_mappings, slot_site_src)); @@ -371,22 +381,21 @@ std::map get_mapping_for_slot_site_from_adjacent_values); } -std::set infer_all_copies_in_graph(DynamicOpenDataflowGraph const &g) -{ +std::set + infer_all_copies_in_graph(DynamicOpenDataflowGraph const &g) { std::map - fully_resolved_tensor_mappings = - resolve_tensor_mappings(g); - - std::set all_copies = - flatmap(set_of(get_dynamic_values(g)), - [&](DynamicValueAttrs const &v) - -> std::set { + fully_resolved_tensor_mappings = resolve_tensor_mappings(g); - DynamicSlotSite src_site = dynamic_graph_find_source_of_value(g, v); - std::set sinks = dynamic_graph_find_sinks_of_value(g, v); + std::set all_copies = flatmap( + set_of(get_dynamic_values(g)), + [&](DynamicValueAttrs const &v) -> std::set { + DynamicSlotSite src_site = dynamic_graph_find_source_of_value(g, v); + std::set sinks = + dynamic_graph_find_sinks_of_value(g, v); - return copies_for_value(v, src_site, sinks, fully_resolved_tensor_mappings); - }); + return copies_for_value( + v, src_site, sinks, fully_resolved_tensor_mappings); + }); return all_copies; } @@ -397,20 +406,22 @@ DynamicOpenDataflowGraph require_graph_is_ready_for_copy_insertion(g); std::map - fully_resolved_tensor_mappings = - resolve_tensor_mappings(g); + fully_resolved_tensor_mappings = resolve_tensor_mappings(g); std::set all_copies = infer_all_copies_in_graph(g); - std::set all_copy_invocations = transform(all_copies, make_copy_invocation); + std::set all_copy_invocations = + transform(all_copies, make_copy_invocation); - std::set mapped_invocations = transform( - get_dynamic_invocation_set(g), - [&](DynamicNodeInvocation const &i) -> DynamicNodeInvocation { - dynamic_invocation_id_t id = dynamic_graph_get_id_for_invocation(g, i); + std::set mapped_invocations = + transform(get_dynamic_invocation_set(g), + [&](DynamicNodeInvocation const &i) -> DynamicNodeInvocation { + dynamic_invocation_id_t id = + dynamic_graph_get_id_for_invocation(g, i); - return apply_mappings_for_invocation(id, i, fully_resolved_tensor_mappings); - }); + return apply_mappings_for_invocation( + id, i, fully_resolved_tensor_mappings); + }); DynamicOpenDataflowGraph result = dynamic_open_dataflow_graph_from_invocation_set( diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc index d10224a947..fb6857f81e 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_invocation.cc @@ -1,31 +1,34 @@ #include "task-spec/dynamic_graph/dynamic_node_invocation.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/training_operation_attrs.h" +#include "utils/containers/all_of.h" #include "utils/containers/are_disjoint.h" +#include "utils/containers/keys.h" #include "utils/containers/set_union.h" -#include "utils/optional.h" #include "utils/containers/values.h" -#include "utils/containers/keys.h" -#include "utils/containers/all_of.h" +#include "utils/optional.h" namespace FlexFlow { -bool invocation_fully_satisfies(DynamicNodeInvocation const &i, - std::function const &node_condition, - std::function const &value_condition, - std::function const &slot_condition) -{ - return node_condition(i.node_attrs) - && all_of(values(i.inputs), value_condition) - && all_of(keys(i.inputs), slot_condition) - && all_of(values(i.outputs), value_condition) - && all_of(keys(i.outputs), slot_condition); +bool invocation_fully_satisfies( + DynamicNodeInvocation const &i, + std::function const &node_condition, + std::function const &value_condition, + std::function const &slot_condition) { + return node_condition(i.node_attrs) && + all_of(values(i.inputs), value_condition) && + all_of(keys(i.inputs), slot_condition) && + all_of(values(i.outputs), value_condition) && + all_of(keys(i.outputs), slot_condition); } -void require_invocation_fully_satisfies(DynamicNodeInvocation const &i, - std::function const &require_node_condition, - std::function const &require_value_condition, - std::function const &require_slot_condition) { +void require_invocation_fully_satisfies( + DynamicNodeInvocation const &i, + std::function const &require_node_condition, + std::function const + &require_value_condition, + std::function const + &require_slot_condition) { require_node_condition(i.node_attrs); for (DynamicTensorSlot const &k : keys(i.inputs)) { require_slot_condition(k); @@ -59,7 +62,8 @@ TrainingOpType } std::set - get_incoming_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { + get_incoming_dynamic_slot_sites_for_invocation( + dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { std::set incoming_slots = transform(set_of(i.inputs), @@ -75,8 +79,8 @@ std::set return incoming_slots; } -std::set - get_output_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { +std::set get_output_dynamic_slot_sites_for_invocation( + dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { std::set output_slots = transform(set_of(i.outputs), @@ -93,10 +97,13 @@ std::set } std::set - get_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &id, DynamicNodeInvocation const &i) { + get_dynamic_slot_sites_for_invocation(dynamic_invocation_id_t const &id, + DynamicNodeInvocation const &i) { - std::set incoming_slots = get_incoming_dynamic_slot_sites_for_invocation(id, i); - std::set output_slots = get_output_dynamic_slot_sites_for_invocation(id, i); + std::set incoming_slots = + get_incoming_dynamic_slot_sites_for_invocation(id, i); + std::set output_slots = + get_output_dynamic_slot_sites_for_invocation(id, i); ASSERT(are_disjoint(incoming_slots, output_slots)); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc index 610d6609c3..e66b61c6d4 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_node_mapping.cc @@ -1,27 +1,26 @@ #include "task-spec/dynamic_graph/dynamic_node_mapping.h" +#include "utils/bidict/algorithms/bidict_transform_keys.h" #include "utils/bidict/algorithms/bidict_transform_values.h" #include "utils/containers/transform.h" -#include "utils/bidict/algorithms/bidict_transform_keys.h" namespace FlexFlow { bidict - dynamic_node_mapping_get_shard_bindings(DynamicNodeMapping const &m) -{ + dynamic_node_mapping_get_shard_bindings(DynamicNodeMapping const &m) { return bidict_transform_keys( - m.op_task_group.get_shard_bindings(), - [&](MachineSpaceCoordinate const &mc) -> global_device_id_t { - return global_device_id_t{ - /*coord=*/mc, - /*device_type=*/m.device_type, - }; - }); + m.op_task_group.get_shard_bindings(), + [&](MachineSpaceCoordinate const &mc) -> global_device_id_t { + return global_device_id_t{ + /*coord=*/mc, + /*device_type=*/m.device_type, + }; + }); } OperatorAtomicTaskShardBinding - dynamic_node_mapping_get_shard_binding_for_device(DynamicNodeMapping const &mapping, - global_device_id_t const &device_id) -{ + dynamic_node_mapping_get_shard_binding_for_device( + DynamicNodeMapping const &mapping, + global_device_id_t const &device_id) { ASSERT(device_id.device_type == mapping.device_type); return mapping.op_task_group.get_shard_bindings().at_l(device_id.coord); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index af6c288288..c3c4f149b9 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -8,13 +8,17 @@ #include "task-spec/dynamic_graph/serializable_dynamic_node_attrs.h" #include "task-spec/dynamic_graph/serializable_dynamic_value_attrs.h" #include "utils/containers/all_of.h" +#include "utils/containers/at_idx.h" #include "utils/containers/concat_vectors.h" #include "utils/containers/contains_duplicates.h" #include "utils/containers/contains_value.h" #include "utils/containers/filter_values.h" #include "utils/containers/flatmap.h" #include "utils/containers/get_only.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/multiset_union.h" +#include "utils/containers/repeat.h" +#include "utils/containers/require_all_of.h" #include "utils/containers/transform.h" #include "utils/containers/zip_strict.h" #include "utils/containers/zip_values_strict.h" @@ -29,10 +33,6 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/kwarg_dataflow_graph_input.dtg.h" #include "utils/many_to_one/many_to_one.h" -#include "utils/containers/require_all_of.h" -#include "utils/containers/multiset_of.h" -#include "utils/containers/repeat.h" -#include "utils/containers/at_idx.h" namespace FlexFlow { @@ -105,8 +105,8 @@ bool no_part_of_dynamic_graph_satisfies( void require_full_dynamic_graph_satisfies( DynamicOpenDataflowGraph const &g, - std::function const &invocation_condition) -{ + std::function const + &invocation_condition) { require_all_of(g.invocations, invocation_condition); } @@ -120,21 +120,20 @@ std::multiset std::multiset get_dynamic_values(DynamicOpenDataflowGraph const &g) { - return flatmap(multiset_of(g.invocations), - [&](DynamicNodeInvocation const &i) - -> std::multiset { - return multiset_union(values(i.inputs), values(i.outputs)); - }); + return flatmap( + multiset_of(g.invocations), + [&](DynamicNodeInvocation const &i) -> std::multiset { + return multiset_union(values(i.inputs), values(i.outputs)); + }); } std::multiset get_dynamic_tensor_slots(DynamicOpenDataflowGraph const &g) { - return flatmap(multiset_of(g.invocations), - [&](DynamicNodeInvocation const &i) - -> std::multiset { - return multiset_of( - set_union(keys(i.inputs), keys(i.outputs))); - }); + return flatmap( + multiset_of(g.invocations), + [&](DynamicNodeInvocation const &i) -> std::multiset { + return multiset_of(set_union(keys(i.inputs), keys(i.outputs))); + }); } std::set @@ -143,30 +142,27 @@ std::set } std::set - dynamic_graph_get_internal_values(DynamicOpenDataflowGraph const &g) -{ + dynamic_graph_get_internal_values(DynamicOpenDataflowGraph const &g) { std::set internal_slot_sites = get_internal_dynamic_slot_sites(g); - std::set internal_values = - filtrans(internal_slot_sites, - [&](InternalDynamicSlotSite const &s) - -> std::optional { - if (s.direction == TensorDirection::OUTPUT) { - return dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{s}); - } else { - return std::nullopt; - } - }); + std::set internal_values = filtrans( + internal_slot_sites, + [&](InternalDynamicSlotSite const &s) + -> std::optional { + if (s.direction == TensorDirection::OUTPUT) { + return dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{s}); + } else { + return std::nullopt; + } + }); return internal_values; } std::set - dynamic_graph_get_external_values(DynamicOpenDataflowGraph const &g) -{ - std::set all_values = - set_of(get_dynamic_values(g)); + dynamic_graph_get_external_values(DynamicOpenDataflowGraph const &g) { + std::set all_values = set_of(get_dynamic_values(g)); std::set internal_values = dynamic_graph_get_internal_values(g); @@ -175,48 +171,47 @@ std::set } dynamic_invocation_id_t dynamic_graph_get_id_for_invocation( - DynamicOpenDataflowGraph const &g, - DynamicNodeInvocation const &invocation) -{ + DynamicOpenDataflowGraph const &g, + DynamicNodeInvocation const &invocation) { return dynamic_invocation_id_t{ - nonnegative_int{assert_unwrap(index_of(g.invocations, invocation))}, + nonnegative_int{assert_unwrap(index_of(g.invocations, invocation))}, }; } -DynamicNodeInvocation dynamic_graph_get_invocation_for_id(DynamicOpenDataflowGraph const &g, - dynamic_invocation_id_t const &id) -{ +DynamicNodeInvocation + dynamic_graph_get_invocation_for_id(DynamicOpenDataflowGraph const &g, + dynamic_invocation_id_t const &id) { return at_idx(g.invocations, id.idx); } -dynamic_value_id_t dynamic_graph_get_id_for_value(DynamicOpenDataflowGraph const &g, - DynamicValueAttrs const &value) -{ +dynamic_value_id_t + dynamic_graph_get_id_for_value(DynamicOpenDataflowGraph const &g, + DynamicValueAttrs const &value) { auto idx_in_set = [](std::set const &s, - DynamicValueAttrs const &v) - -> nonnegative_int - { + DynamicValueAttrs const &v) -> nonnegative_int { return nonnegative_int{assert_unwrap(index_of(s, v))}; }; { - std::set internal_values = dynamic_graph_get_internal_values(g); + std::set internal_values = + dynamic_graph_get_internal_values(g); if (contains(internal_values, value)) { return dynamic_value_id_t{ - dynamic_internal_value_id_t{ - idx_in_set(internal_values, value), - }, + dynamic_internal_value_id_t{ + idx_in_set(internal_values, value), + }, }; } } { - std::set external_values = dynamic_graph_get_external_values(g); + std::set external_values = + dynamic_graph_get_external_values(g); if (contains(external_values, value)) { return dynamic_value_id_t{ - dynamic_external_value_id_t{ - idx_in_set(external_values, value), - }, + dynamic_external_value_id_t{ + idx_in_set(external_values, value), + }, }; } } @@ -224,62 +219,57 @@ dynamic_value_id_t dynamic_graph_get_id_for_value(DynamicOpenDataflowGraph const PANIC("Could not find id for value {}", value); } -DynamicValueAttrs dynamic_graph_get_value_for_id(DynamicOpenDataflowGraph const &g, - dynamic_value_id_t const &id) -{ - return id.visit(overload { - [&](dynamic_internal_value_id_t const &internal_id) -> DynamicValueAttrs { - std::set internal_values = - dynamic_graph_get_internal_values(g); +DynamicValueAttrs + dynamic_graph_get_value_for_id(DynamicOpenDataflowGraph const &g, + dynamic_value_id_t const &id) { + return id.visit(overload{ + [&](dynamic_internal_value_id_t const &internal_id) -> DynamicValueAttrs { + std::set internal_values = + dynamic_graph_get_internal_values(g); - return at_idx(internal_values, internal_id.idx); - }, - [&](dynamic_external_value_id_t const &external_id) -> DynamicValueAttrs { - std::set external_values = - dynamic_graph_get_external_values(g); + return at_idx(internal_values, internal_id.idx); + }, + [&](dynamic_external_value_id_t const &external_id) -> DynamicValueAttrs { + std::set external_values = + dynamic_graph_get_external_values(g); - return at_idx(external_values, external_id.idx); - } - }); + return at_idx(external_values, external_id.idx); + }}); } std::set get_dynamic_graph_edges(DynamicOpenDataflowGraph const &g) { - return flatmap(get_dynamic_invocation_set(g), - [&](DynamicNodeInvocation const &i) - -> std::set { - return get_dynamic_graph_edges_incoming_to_invocation(g, i); - }); + return flatmap( + get_dynamic_invocation_set(g), + [&](DynamicNodeInvocation const &i) -> std::set { + return get_dynamic_graph_edges_incoming_to_invocation(g, i); + }); } -std::set - get_dynamic_graph_edges_incoming_to_invocation( - DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { +std::set get_dynamic_graph_edges_incoming_to_invocation( + DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { return transform( - set_of(i.inputs), - [&](std::pair const &p) - -> DynamicGraphEdge { - DynamicSlotSite src = - dynamic_graph_find_source_of_value(g, p.second); - - InternalDynamicSlotSite dst = InternalDynamicSlotSite{ - /*invocation_id=*/dynamic_graph_get_id_for_invocation(g, i), - /*direction=*/TensorDirection::INCOMING, - /*slot_name=*/p.first, - }; + set_of(i.inputs), + [&](std::pair const &p) + -> DynamicGraphEdge { + DynamicSlotSite src = dynamic_graph_find_source_of_value(g, p.second); - return dynamic_graph_edge_from_slot_sites(src, dst); - }); + InternalDynamicSlotSite dst = InternalDynamicSlotSite{ + /*invocation_id=*/dynamic_graph_get_id_for_invocation(g, i), + /*direction=*/TensorDirection::INCOMING, + /*slot_name=*/p.first, + }; + + return dynamic_graph_edge_from_slot_sites(src, dst); + }); } -std::set - get_dynamic_graph_edges_outgoing_from_invocation( - DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { +std::set get_dynamic_graph_edges_outgoing_from_invocation( + DynamicOpenDataflowGraph const &g, DynamicNodeInvocation const &i) { return flatmap( set_of(i.outputs), [&](std::pair const &p) -> std::set { - DynamicSlotSite src = DynamicSlotSite{ InternalDynamicSlotSite{ /*invocation_id=*/dynamic_graph_get_id_for_invocation(g, i), @@ -296,18 +286,21 @@ std::set }); } -DynamicValueAttrs dynamic_value_attrs_for_slot_site(DynamicOpenDataflowGraph const &g, - DynamicSlotSite const &slot) { +DynamicValueAttrs + dynamic_value_attrs_for_slot_site(DynamicOpenDataflowGraph const &g, + DynamicSlotSite const &slot) { return slot.visit(overload{ [&](ExternalDynamicSlotSite const &external_slot) -> DynamicValueAttrs { - dynamic_value_id_t value_id = dynamic_value_id_t{external_slot.value_id}; + dynamic_value_id_t value_id = + dynamic_value_id_t{external_slot.value_id}; return dynamic_graph_get_value_for_id(g, value_id); }, [&](InternalDynamicSlotSite const &internal_slot) -> DynamicValueAttrs { - DynamicNodeInvocation invocation = dynamic_graph_get_invocation_for_id(g, internal_slot.invocation_id); + DynamicNodeInvocation invocation = + dynamic_graph_get_invocation_for_id(g, internal_slot.invocation_id); switch (internal_slot.direction) { case TensorDirection::INCOMING: return invocation.inputs.at(internal_slot.slot_name); @@ -321,14 +314,13 @@ DynamicValueAttrs dynamic_value_attrs_for_slot_site(DynamicOpenDataflowGraph con std::set get_internal_dynamic_slot_sites(DynamicOpenDataflowGraph const &g) { - return flatmap(get_dynamic_invocation_set(g), - [&](DynamicNodeInvocation const &i) - -> std::set { - - dynamic_invocation_id_t id = dynamic_graph_get_id_for_invocation(g, i); + return flatmap( + get_dynamic_invocation_set(g), + [&](DynamicNodeInvocation const &i) -> std::set { + dynamic_invocation_id_t id = dynamic_graph_get_id_for_invocation(g, i); - return get_dynamic_slot_sites_for_invocation(id, i); - }); + return get_dynamic_slot_sites_for_invocation(id, i); + }); } std::set @@ -344,7 +336,8 @@ std::set external_values, [&](DynamicValueAttrs const &external_value) -> ExternalDynamicSlotSite { dynamic_external_value_id_t value_id = - dynamic_graph_get_id_for_value(g, external_value).require_external(); + dynamic_graph_get_id_for_value(g, external_value) + .require_external(); return ExternalDynamicSlotSite{value_id}; }); @@ -372,21 +365,22 @@ std::set return found; } -DynamicSlotSite - dynamic_graph_find_source_of_slot_site(DynamicOpenDataflowGraph const &g, - InternalDynamicSlotSite const &slot_site) -{ - DynamicValueAttrs value_attrs = dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{slot_site}); +DynamicSlotSite dynamic_graph_find_source_of_slot_site( + DynamicOpenDataflowGraph const &g, + InternalDynamicSlotSite const &slot_site) { + DynamicValueAttrs value_attrs = + dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{slot_site}); DynamicSlotSite src_site = dynamic_graph_find_source_of_value(g, value_attrs); return src_site; } -std::set - dynamic_graph_find_sinks_of_slot_site(DynamicOpenDataflowGraph const &g, - InternalDynamicSlotSite const &slot_site) -{ - DynamicValueAttrs value_attrs = dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{slot_site}); - std::set sink_sites = dynamic_graph_find_sinks_of_value(g, value_attrs); +std::set dynamic_graph_find_sinks_of_slot_site( + DynamicOpenDataflowGraph const &g, + InternalDynamicSlotSite const &slot_site) { + DynamicValueAttrs value_attrs = + dynamic_value_attrs_for_slot_site(g, DynamicSlotSite{slot_site}); + std::set sink_sites = + dynamic_graph_find_sinks_of_value(g, value_attrs); return sink_sites; } @@ -476,8 +470,7 @@ std::pair all_values = - set_of(get_dynamic_values(g)); + std::set all_values = set_of(get_dynamic_values(g)); ManyToOne value_to_producer; for (DynamicNodeInvocation const &invocation : diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc index 0d239a9069..901a543f2c 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_tensor_slot.cc @@ -18,5 +18,4 @@ DynamicTensorSlot slot_without_task_shard(DynamicTensorSlot const &s) { return result; } - } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc index 38243fd05d..a847c0ba75 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/dynamic_value_attrs.cc @@ -13,10 +13,9 @@ DynamicValueAttrs return result; } -DynamicValueAttrs decide_dynamic_value_attrs_mapping( - DynamicValueAttrs const &attrs, - ParallelTensorMapping const &mapping) -{ +DynamicValueAttrs + decide_dynamic_value_attrs_mapping(DynamicValueAttrs const &attrs, + ParallelTensorMapping const &mapping) { ASSERT(!attrs.mapping.has_value()); DynamicValueAttrs result = attrs; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc index bd10c4f674..e7c4b34460 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/loss_insertion.cc @@ -44,20 +44,20 @@ LossInsertionResult perform_loss_insertion( DynamicNodeInvocation loss_invocation{ /*inputs=*/{ { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/label_value.role, - /*task_shard=*/std::nullopt, - }, - label_value, + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/label_value.role, + /*task_shard=*/std::nullopt, + }, + label_value, }, { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::LOGIT, - /*slot_tensor_role=*/logit_value.role, - /*task_shard=*/std::nullopt, - }, - logit_value, + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::LOGIT, + /*slot_tensor_role=*/logit_value.role, + /*task_shard=*/std::nullopt, + }, + logit_value, }, }, /*node_attrs=*/ @@ -72,12 +72,12 @@ LossInsertionResult perform_loss_insertion( /*outputs=*/ { { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::LOGIT, - /*slot_tensor_role=*/logit_grad_value.role, - /*task_shard=*/std::nullopt, - }, - logit_grad_value, + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::LOGIT, + /*slot_tensor_role=*/logit_grad_value.role, + /*task_shard=*/std::nullopt, + }, + logit_grad_value, }, }, }; diff --git a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc index 604c77016b..215d7398c1 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/machine_slicing.cc @@ -3,10 +3,9 @@ namespace FlexFlow { -std::set - perform_machine_slicing_for_invocation( - DynamicNodeInvocation const &invocation, - global_device_id_t const &device_id) { +std::set perform_machine_slicing_for_invocation( + DynamicNodeInvocation const &invocation, + global_device_id_t const &device_id) { ASSERT(invocation.node_attrs.device_ids.has_value()); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc index d88f35c3cb..3ce603f5ff 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_cg.cc @@ -8,10 +8,10 @@ #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" #include "utils/containers/generate_map.h" -#include +#include "utils/containers/map_from_unordered.h" #include +#include #include -#include "utils/containers/map_from_unordered.h" namespace FlexFlow { @@ -31,51 +31,49 @@ DynamicOpenDataflowGraph /*per_device_op_state=*/std::nullopt, }; - std::map result_inputs = - transform( - get_incoming_tensors(cg, layer), - [&](TensorSlotName const &slot_name, tensor_guid_t const &tensor) { - TensorAttrs attrs = get_tensor_attrs(cg, tensor); - return std::pair{ - DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, - /*parallel_tensor_shape=*/lift_to_parallel(attrs.shape), - /*create_grad=*/(attrs.create_grad == CreateGrad::YES), - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }, - }; - }); + std::map result_inputs = transform( + get_incoming_tensors(cg, layer), + [&](TensorSlotName const &slot_name, tensor_guid_t const &tensor) { + TensorAttrs attrs = get_tensor_attrs(cg, tensor); + return std::pair{ + DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, + /*parallel_tensor_shape=*/lift_to_parallel(attrs.shape), + /*create_grad=*/(attrs.create_grad == CreateGrad::YES), + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }, + }; + }); - std::map result_outputs = - transform( - get_outgoing_tensors(cg, layer), - [&](TensorSlotName const &slot_name, tensor_guid_t const &tensor) { - TensorAttrs attrs = get_tensor_attrs(cg, tensor); - return std::pair{ - DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, - /*parallel_tensor_shape=*/lift_to_parallel(attrs.shape), - /*create_grad=*/(attrs.create_grad == CreateGrad::YES), - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }, - }; - }); + std::map result_outputs = transform( + get_outgoing_tensors(cg, layer), + [&](TensorSlotName const &slot_name, tensor_guid_t const &tensor) { + TensorAttrs attrs = get_tensor_attrs(cg, tensor); + return std::pair{ + DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{tensor}, + /*parallel_tensor_shape=*/lift_to_parallel(attrs.shape), + /*create_grad=*/(attrs.create_grad == CreateGrad::YES), + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }, + }; + }); result.invocations.emplace(result_inputs, result_attrs, result_outputs); } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 9e11337368..13e9196416 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -9,62 +9,61 @@ #include "task-spec/dynamic_graph/dynamic_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include "utils/bidict/algorithms/bidict_unordered_set_of.h" #include "utils/bidict/algorithms/merge_disjoint_bidicts.h" #include "utils/containers/get_only.h" #include "utils/containers/map_keys_and_values.h" #include "utils/containers/require_only_key.h" #include "utils/containers/transform_pairs.h" -#include #include +#include #include -#include "utils/bidict/algorithms/bidict_unordered_set_of.h" namespace FlexFlow { DynamicNodeInvocation make_dynamic_node_invocation_from_mapped( MappedParallelLayerInvocationInfo const &invocation_info, - DeviceType device_type) -{ + DeviceType device_type) { DynamicNodeAttrs result_attrs{ /*task_type=*/std::nullopt, /*device_ids=*/std::nullopt, - /*mapping=*/DynamicNodeMapping{ - /*op_task_group=*/invocation_info.layer_info.mapping, - /*device_type=*/device_type, + /*mapping=*/ + DynamicNodeMapping{ + /*op_task_group=*/invocation_info.layer_info.mapping, + /*device_type=*/device_type, }, - /*op_attrs=*/TrainingOperationAttrs{invocation_info.layer_info.attrs.op_attrs}, + /*op_attrs=*/ + TrainingOperationAttrs{invocation_info.layer_info.attrs.op_attrs}, /*pcg_layer_guid=*/dynamic_layer_guid_t{invocation_info.layer_info.guid}, /*per_device_op_state=*/std::nullopt, }; - auto lift_kv_pair = - [&](TensorSlotName slot_name, - ParallelTensorInfo const &tensor) - -> std::pair - { + auto lift_kv_pair = [&](TensorSlotName slot_name, + ParallelTensorInfo const &tensor) + -> std::pair { return { - DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor.guid}, - /*parallel_tensor_shape=*/tensor.attrs.shape, - /*create_grad=*/(tensor.attrs.create_grad == CreateGrad::YES), - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }, + DynamicTensorSlot{ + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{tensor.guid}, + /*parallel_tensor_shape=*/tensor.attrs.shape, + /*create_grad=*/(tensor.attrs.create_grad == CreateGrad::YES), + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }, }; }; std::map result_inputs = - transform(invocation_info.incoming, lift_kv_pair); + transform(invocation_info.incoming, lift_kv_pair); std::map result_outputs = - transform(invocation_info.outgoing, lift_kv_pair); + transform(invocation_info.outgoing, lift_kv_pair); DynamicNodeInvocation invocation = DynamicNodeInvocation{ /*inputs=*/result_inputs, @@ -79,12 +78,12 @@ DynamicOpenDataflowGraph make_dynamic_open_dataflow_graph_from_mapped_pcg( MappedParallelComputationGraph const &mpcg, DeviceType device_type) { return dynamic_open_dataflow_graph_from_invocation_set( - transform(mpcg_get_invocation_set(mpcg), - [&](MappedParallelLayerInvocationInfo const &mpcg_invocation) - -> DynamicNodeInvocation - { - return make_dynamic_node_invocation_from_mapped(mpcg_invocation, device_type); - })); + transform(mpcg_get_invocation_set(mpcg), + [&](MappedParallelLayerInvocationInfo const &mpcg_invocation) + -> DynamicNodeInvocation { + return make_dynamic_node_invocation_from_mapped( + mpcg_invocation, device_type); + })); } } // namespace FlexFlow diff --git a/lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc b/lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc index 8d8d41a59d..1b5f692bac 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/parallel_tensor_mapping.cc @@ -2,21 +2,25 @@ namespace FlexFlow { -global_device_id_t pt_mapping_get_device_for_coord(ParallelTensorMapping const &m, - ParallelTensorSpaceCoordinate const &coord) { +global_device_id_t pt_mapping_get_device_for_coord( + ParallelTensorMapping const &m, + ParallelTensorSpaceCoordinate const &coord) { return m.raw.at_l(coord); } -ParallelTensorSpaceCoordinate pt_mapping_get_coord_for_device(ParallelTensorMapping const &m, - global_device_id_t const &device) { +ParallelTensorSpaceCoordinate + pt_mapping_get_coord_for_device(ParallelTensorMapping const &m, + global_device_id_t const &device) { return m.raw.at_r(device); } -std::set pt_mapping_get_coord_set(ParallelTensorMapping const &m) { +std::set + pt_mapping_get_coord_set(ParallelTensorMapping const &m) { return m.raw.left_values(); } -std::set pt_mapping_get_device_set(ParallelTensorMapping const &m) { +std::set + pt_mapping_get_device_set(ParallelTensorMapping const &m) { return m.raw.right_values(); } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc index 2a188aeb09..a095880118 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/pass_expansion.cc @@ -1,15 +1,15 @@ #include "task-spec/dynamic_graph/pass_expansion.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/training_operation_attrs.h" #include "utils/containers/are_all_same.h" +#include "utils/containers/flatmap.h" #include "utils/containers/get_only.h" +#include "utils/containers/map_values.h" #include "utils/containers/merge_disjoint_maps.h" -#include "utils/containers/transform.h" #include "utils/containers/repeat_until_converged.h" -#include "utils/containers/flatmap.h" -#include "task-spec/dynamic_graph/dynamic_node_invocation.h" -#include "utils/containers/map_values.h" +#include "utils/containers/transform.h" namespace FlexFlow { @@ -41,145 +41,140 @@ void require_value_is_not_pass_expanded(DynamicValueAttrs const &v) { ASSERT(!v.role.has_value(), v); } -void require_invocation_is_fully_pass_expanded(DynamicNodeInvocation const &invocation) { +void require_invocation_is_fully_pass_expanded( + DynamicNodeInvocation const &invocation) { auto require_slot_is_pass_expanded = [&](DynamicTensorSlot const &s) { - if (dynamic_node_invocation_get_op_type(invocation) == TrainingOpType{TrainingOnlyOpType::COPY}) { + if (dynamic_node_invocation_get_op_type(invocation) == + TrainingOpType{TrainingOnlyOpType::COPY}) { return; } ASSERT(s.slot_tensor_role.has_value(), s); }; - require_invocation_fully_satisfies( - invocation, - require_node_might_be_pass_expanded, - require_value_is_pass_expanded, - require_slot_is_pass_expanded); + require_invocation_fully_satisfies(invocation, + require_node_might_be_pass_expanded, + require_value_is_pass_expanded, + require_slot_is_pass_expanded); } -void require_invocation_is_ready_for_pass_expansion(DynamicNodeInvocation const &invocation) { - require_invocation_fully_satisfies( - invocation, - require_node_might_not_be_pass_expanded, - require_value_is_not_pass_expanded, - require_slot_is_not_pass_expanded); +void require_invocation_is_ready_for_pass_expansion( + DynamicNodeInvocation const &invocation) { + require_invocation_fully_satisfies(invocation, + require_node_might_not_be_pass_expanded, + require_value_is_not_pass_expanded, + require_slot_is_not_pass_expanded); } - void require_graph_is_fully_pass_expanded(DynamicOpenDataflowGraph const &g) { require_full_dynamic_graph_satisfies( - g, - require_invocation_is_fully_pass_expanded); + g, require_invocation_is_fully_pass_expanded); } -void require_graph_is_ready_for_pass_expansion(DynamicOpenDataflowGraph const &g) { +void require_graph_is_ready_for_pass_expansion( + DynamicOpenDataflowGraph const &g) { require_full_dynamic_graph_satisfies( - g, - require_invocation_is_ready_for_pass_expansion); + g, require_invocation_is_ready_for_pass_expansion); } std::set - determine_intermediate_values_needed_to_compute_gradients_of_value( - DynamicOpenDataflowGraph const &g, - DynamicValueAttrs const &val) -{ + determine_intermediate_values_needed_to_compute_gradients_of_value( + DynamicOpenDataflowGraph const &g, DynamicValueAttrs const &val) { auto get_values_immediately_needed_to_compute_gradients_of_needed = - [&](std::set const &needed) - -> std::set - { - std::set additional - = flatmap(needed, - [&](DynamicValueAttrs const &v) -> std::set { - std::set sinks = - dynamic_graph_find_sinks_of_value(g, v); - - return flatmap( - sinks, - [&](InternalDynamicSlotSite const &sink) -> std::set { - DynamicNodeInvocation sink_invocation - = dynamic_graph_get_invocation_for_id(g, sink.invocation_id); + [&](std::set const &needed) + -> std::set { + std::set additional = flatmap( + needed, [&](DynamicValueAttrs const &v) -> std::set { + std::set sinks = + dynamic_graph_find_sinks_of_value(g, v); + + return flatmap(sinks, + [&](InternalDynamicSlotSite const &sink) + -> std::set { + DynamicNodeInvocation sink_invocation = + dynamic_graph_get_invocation_for_id( + g, sink.invocation_id); return set_of(values(sink_invocation.outputs)); }); - }); + }); - return set_union(needed, additional); - }; + return set_union(needed, additional); + }; std::set result = repeat_until_converged( - std::set{val}, - get_values_immediately_needed_to_compute_gradients_of_needed); + std::set{val}, + get_values_immediately_needed_to_compute_gradients_of_needed); ASSERT(contains(result, val)); return result; } std::set - determine_intermediate_values_needed_for_gradient_computation( - DynamicOpenDataflowGraph const &g) -{ + determine_intermediate_values_needed_for_gradient_computation( + DynamicOpenDataflowGraph const &g) { auto value_is_fundamentally_required = - [&](DynamicValueAttrs const &v) -> bool { - DynamicSlotSite source = dynamic_graph_find_source_of_value(g, v); + [&](DynamicValueAttrs const &v) -> bool { + DynamicSlotSite source = dynamic_graph_find_source_of_value(g, v); - if (source.is_external()) { - return true; - } + if (source.is_external()) { + return true; + } - InternalDynamicSlotSite internal_source = source.require_internal(); - ASSERT(internal_source.direction == TensorDirection::OUTPUT); + InternalDynamicSlotSite internal_source = source.require_internal(); + ASSERT(internal_source.direction == TensorDirection::OUTPUT); - DynamicNodeInvocation source_invocation = + DynamicNodeInvocation source_invocation = dynamic_graph_get_invocation_for_id(g, internal_source.invocation_id); - TrainingOpType op_type = dynamic_node_invocation_get_op_type(source_invocation); + TrainingOpType op_type = + dynamic_node_invocation_get_op_type(source_invocation); - TrainingOpType weight_op_type = TrainingOpType{OperatorType::WEIGHT}; - TrainingOpType input_op_type = TrainingOpType{OperatorType::INPUT}; + TrainingOpType weight_op_type = TrainingOpType{OperatorType::WEIGHT}; + TrainingOpType input_op_type = TrainingOpType{OperatorType::INPUT}; - if (op_type == weight_op_type) { - return true; - } else if (op_type == input_op_type) { - return assert_unwrap(v.create_grad); - } else { - return false; - } - }; + if (op_type == weight_op_type) { + return true; + } else if (op_type == input_op_type) { + return assert_unwrap(v.create_grad); + } else { + return false; + } + }; - std::set fundamentally_required_values - = filter(set_of(get_dynamic_values(g)), value_is_fundamentally_required); + std::set fundamentally_required_values = + filter(set_of(get_dynamic_values(g)), value_is_fundamentally_required); - std::set required_values = - flatmap(fundamentally_required_values, - [&](DynamicValueAttrs const &fundamentally_required_value) -> std::set { - return determine_intermediate_values_needed_to_compute_gradients_of_value( - g, fundamentally_required_value); - }); + std::set required_values = flatmap( + fundamentally_required_values, + [&](DynamicValueAttrs const &fundamentally_required_value) + -> std::set { + return determine_intermediate_values_needed_to_compute_gradients_of_value( + g, fundamentally_required_value); + }); return required_values; } std::set - determine_invocations_needed_in_backward_pass_for_gradient_computation( - DynamicOpenDataflowGraph const &g) -{ - std::set required_values = - determine_intermediate_values_needed_for_gradient_computation(g); - - auto get_sink_invocations_for_value = - [&](DynamicValueAttrs const &v) -> std::set { - return - transform( - dynamic_graph_find_sinks_of_value(g, v), - [&](InternalDynamicSlotSite const &sink_site) -> dynamic_invocation_id_t { - ASSERT(sink_site.direction == TensorDirection::INCOMING); - return sink_site.invocation_id; - }); - }; + determine_invocations_needed_in_backward_pass_for_gradient_computation( + DynamicOpenDataflowGraph const &g) { + std::set required_values = + determine_intermediate_values_needed_for_gradient_computation(g); + + auto get_sink_invocations_for_value = + [&](DynamicValueAttrs const &v) -> std::set { + return transform(dynamic_graph_find_sinks_of_value(g, v), + [&](InternalDynamicSlotSite const &sink_site) + -> dynamic_invocation_id_t { + ASSERT(sink_site.direction == TensorDirection::INCOMING); + return sink_site.invocation_id; + }); + }; return flatmap(required_values, get_sink_invocations_for_value); } - + DynamicTensorSlot pass_expand_slot(DynamicTensorSlot const &s, FwbTensorType tensor_type) { require_slot_is_not_pass_expanded(s); @@ -238,14 +233,12 @@ DynamicNodeInvocation perform_fwd_pass_expansion_for_invocation( }; }; - DynamicNodeInvocation result = [&]() - -> DynamicNodeInvocation - { + DynamicNodeInvocation result = [&]() -> DynamicNodeInvocation { if (op_attrs.is_copy()) { return DynamicNodeInvocation{ - /*inputs=*/map_values(invocation.inputs, to_fwd_value), - /*node_attrs=*/invocation.node_attrs, - /*outputs=*/map_values(invocation.outputs, to_fwd_value), + /*inputs=*/map_values(invocation.inputs, to_fwd_value), + /*node_attrs=*/invocation.node_attrs, + /*outputs=*/map_values(invocation.outputs, to_fwd_value), }; } else { return DynamicNodeInvocation{ @@ -290,24 +283,24 @@ DynamicNodeInvocation perform_bwd_pass_expansion_for_invocation( }; }; - DynamicNodeInvocation result = [&]() - -> DynamicNodeInvocation - { + DynamicNodeInvocation result = [&]() -> DynamicNodeInvocation { if (op_attrs.is_copy()) { return DynamicNodeInvocation{ - /*inputs=*/map_values(invocation.outputs, to_grad_value), - /*node_attrs=*/invocation.node_attrs, - /*outputs=*/map_values(invocation.inputs, to_grad_value), + /*inputs=*/map_values(invocation.outputs, to_grad_value), + /*node_attrs=*/invocation.node_attrs, + /*outputs=*/map_values(invocation.inputs, to_grad_value), }; - } else if (training_op_attrs_has_op_type(op_attrs, OperatorType::REPLICATE)) { + } else if (training_op_attrs_has_op_type(op_attrs, + OperatorType::REPLICATE)) { return DynamicNodeInvocation{ /*inputs=*/{ - transform(invocation.outputs, to_grad), + transform(invocation.outputs, to_grad), }, /*node_attrs=*/ pass_expand_node(invocation.node_attrs, DynamicTaskType::BWD), - /*outputs=*/{ - transform(invocation.inputs, to_grad), + /*outputs=*/ + { + transform(invocation.inputs, to_grad), }, }; } else { @@ -336,13 +329,13 @@ DynamicOpenDataflowGraph require_graph_is_ready_for_pass_expansion(g); - std::set needed_in_bwd_pass = - determine_invocations_needed_in_backward_pass_for_gradient_computation(g); + std::set needed_in_bwd_pass = + determine_invocations_needed_in_backward_pass_for_gradient_computation(g); DynamicOpenDataflowGraph result = flatmap_dynamic_invocation_set( - g, - [&](DynamicNodeInvocation const &invocation) { - dynamic_invocation_id_t invocation_id = dynamic_graph_get_id_for_invocation(g, invocation); + g, [&](DynamicNodeInvocation const &invocation) { + dynamic_invocation_id_t invocation_id = + dynamic_graph_get_id_for_invocation(g, invocation); if (contains(needed_in_bwd_pass, invocation_id)) { return std::set{ diff --git a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc index 5118637702..2ceb8c1214 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.cc @@ -4,19 +4,19 @@ namespace FlexFlow { SerializableDynamicOpenDataflowGraph - dynamic_open_dataflow_graph_to_serializable(DynamicOpenDataflowGraph const &g) -{ + dynamic_open_dataflow_graph_to_serializable( + DynamicOpenDataflowGraph const &g) { return SerializableDynamicOpenDataflowGraph{ - /*invocations=*/transform(g.invocations, dynamic_node_invocation_to_serializable), + /*invocations=*/transform(g.invocations, + dynamic_node_invocation_to_serializable), }; } DynamicOpenDataflowGraph dynamic_open_dataflow_graph_from_serializable( - SerializableDynamicOpenDataflowGraph const &serializable) -{ + SerializableDynamicOpenDataflowGraph const &serializable) { return DynamicOpenDataflowGraph{ - /*invocations=*/transform(serializable.invocations, - dynamic_node_invocation_from_serializable), + /*invocations=*/transform(serializable.invocations, + dynamic_node_invocation_from_serializable), }; } diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index 74034d869e..c566a30f8d 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -1,33 +1,33 @@ #include "task-spec/dynamic_graph/shard_expansion.h" +#include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.h" +#include "task-spec/dynamic_graph/dynamic_node_invocation.h" #include "task-spec/dynamic_graph/dynamic_node_mapping.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" -#include "utils/bidict/algorithms/bidict_filter_keys.h" +#include "task-spec/dynamic_graph/parallel_tensor_mapping.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" #include "task-spec/dynamic_graph/shard_expansion.h" -#include "utils/containers/get_only.h" -#include "utils/containers/map_values2.h" -#include "utils/containers/require_same.h" -#include "utils/containers/transform.h" -#include "utils/optional.h" -#include "utils/containers/binary_merge_disjoint_maps.h" -#include "task-spec/dynamic_graph/dynamic_node_invocation.h" -#include "utils/containers/map_from_unordered.h" #include "task-spec/dynamic_graph/training_operation_attrs.h" -#include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.h" +#include "utils/bidict/algorithms/bidict_filter_keys.h" #include "utils/bidict/algorithms/bidict_filter_values.h" -#include "task-spec/dynamic_graph/dynamic_tensor_role.h" -#include "utils/containers/merge_disjoint_maps.h" -#include "utils/containers/map_keys.h" -#include "utils/containers/require_only_key.h" -#include "utils/containers/set_of.h" -#include "task-spec/dynamic_graph/parallel_tensor_mapping.h" -#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" -#include "utils/binary_relation/binary_relation_transform_right2.h" #include "utils/binary_relation/binary_relation_from_map.h" -#include "utils/containers/are_disjoint.h" +#include "utils/binary_relation/binary_relation_transform_right2.h" #include "utils/binary_relation/filter_binary_relation.h" -#include "utils/containers/map_from_pairs.h" +#include "utils/containers/are_disjoint.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/flatmap.h" +#include "utils/containers/get_only.h" +#include "utils/containers/map_from_pairs.h" +#include "utils/containers/map_from_unordered.h" +#include "utils/containers/map_keys.h" +#include "utils/containers/map_values2.h" +#include "utils/containers/merge_disjoint_maps.h" +#include "utils/containers/require_only_key.h" +#include "utils/containers/require_same.h" +#include "utils/containers/set_of.h" +#include "utils/containers/transform.h" +#include "utils/optional.h" namespace FlexFlow { @@ -39,20 +39,21 @@ void require_value_is_shard_expanded(DynamicValueAttrs const &n) { ASSERT(n.shard_coord.has_value()); } -void require_invocation_is_fully_shard_expanded(DynamicNodeInvocation const &i) { +void require_invocation_is_fully_shard_expanded( + DynamicNodeInvocation const &i) { auto require_slot_is_shard_expanded = [](DynamicTensorSlot const &) { return; }; - return require_invocation_fully_satisfies( - i, - require_node_is_shard_expanded, - require_value_is_shard_expanded, - require_slot_is_shard_expanded); + return require_invocation_fully_satisfies(i, + require_node_is_shard_expanded, + require_value_is_shard_expanded, + require_slot_is_shard_expanded); } void require_graph_is_fully_shard_expanded(DynamicOpenDataflowGraph const &g) { - return require_full_dynamic_graph_satisfies(g, require_invocation_is_fully_shard_expanded); + return require_full_dynamic_graph_satisfies( + g, require_invocation_is_fully_shard_expanded); } void require_node_is_ready_for_shard_expansion(DynamicNodeAttrs const &n) { @@ -67,20 +68,22 @@ void require_value_is_ready_for_shard_expansion(DynamicValueAttrs const &n) { return; } -void require_invocation_is_ready_for_shard_expansion(DynamicNodeInvocation const &i) { - auto require_slot_is_ready_for_shard_expansion = [](DynamicTensorSlot const &) { - return; - }; +void require_invocation_is_ready_for_shard_expansion( + DynamicNodeInvocation const &i) { + auto require_slot_is_ready_for_shard_expansion = + [](DynamicTensorSlot const &) { return; }; return require_invocation_fully_satisfies( - i, - require_node_is_ready_for_shard_expansion, - require_value_is_ready_for_shard_expansion, - require_slot_is_ready_for_shard_expansion); + i, + require_node_is_ready_for_shard_expansion, + require_value_is_ready_for_shard_expansion, + require_slot_is_ready_for_shard_expansion); } -void require_graph_is_ready_for_shard_expansion(DynamicOpenDataflowGraph const &g) { - require_full_dynamic_graph_satisfies(g, require_invocation_is_ready_for_shard_expansion); +void require_graph_is_ready_for_shard_expansion( + DynamicOpenDataflowGraph const &g) { + require_full_dynamic_graph_satisfies( + g, require_invocation_is_ready_for_shard_expansion); } static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( @@ -89,13 +92,16 @@ static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( OperatorAtomicTaskShardBinding const &binding) { auto shard_expand_value_attrs = - [&](DynamicTensorSlot const &s, DynamicValueAttrs const &v) -> DynamicValueAttrsShardingInfo { + [&](DynamicTensorSlot const &s, + DynamicValueAttrs const &v) -> DynamicValueAttrsShardingInfo { ParallelTensorSpaceCoordinate parallel_tensor_coord = binding.tensor_coords.at(s.slot_name); return DynamicValueAttrsShardingInfo{ - /*shard_coord=*/parallel_tensor_coord, - /*mapping=*/pt_mapping_get_device_for_coord(assert_unwrap(v.mapping), parallel_tensor_coord), + /*shard_coord=*/parallel_tensor_coord, + /*mapping=*/ + pt_mapping_get_device_for_coord(assert_unwrap(v.mapping), + parallel_tensor_coord), }; }; @@ -107,18 +113,19 @@ static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( DynamicNodeInvocationShardingInfo result = DynamicNodeInvocationShardingInfo{ /*device_coord=*/nonempty_set{device_id}, - /*value_sharding=*/binary_relation_transform_right2( + /*value_sharding=*/ + binary_relation_transform_right2( binary_relation_from_map( - binary_merge_disjoint_maps(i.inputs, i.outputs)), + binary_merge_disjoint_maps(i.inputs, i.outputs)), shard_expand_value_attrs), }; { - std::set invocation_slots = set_union( - keys(i.inputs), keys(i.outputs)); + std::set invocation_slots = + set_union(keys(i.inputs), keys(i.outputs)); - std::set sharding_info_slots - = result.value_sharding.left_values(); + std::set sharding_info_slots = + result.value_sharding.left_values(); ASSERT(invocation_slots == sharding_info_slots); } @@ -131,9 +138,10 @@ static bidict bidict const &mapping, ParallelTensorSpaceCoordinate const ¶llel_tensor_coord) { - return bidict_filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { - return p == parallel_tensor_coord; - }); + return bidict_filter_keys(mapping, + [&](ParallelTensorSpaceCoordinate const &p) { + return p == parallel_tensor_coord; + }); } static DynamicNodeInvocation shard_invocation_for_binding( @@ -162,14 +170,15 @@ static DynamicNodeInvocation shard_invocation_for_binding( DynamicNodeAttrs expanded_node_attrs = [&]() { DynamicNodeAttrs result = i.node_attrs; - result.device_ids = nonempty_set{device_id};; + result.device_ids = nonempty_set{device_id}; + ; return result; }(); return DynamicNodeInvocation{ - /*inputs=*/map_values2(i.inputs, shard_expand_value_attrs), - /*node_attrs=*/expanded_node_attrs, - /*outputs=*/map_values2(i.outputs, shard_expand_value_attrs), + /*inputs=*/map_values2(i.inputs, shard_expand_value_attrs), + /*node_attrs=*/expanded_node_attrs, + /*outputs=*/map_values2(i.outputs, shard_expand_value_attrs), }; } @@ -181,14 +190,14 @@ static std::set ParallelTensorMapping input_mapping = assert_unwrap(input.mapping); ParallelTensorMapping output_mapping = assert_unwrap(output.mapping); - std::set coord_set = - require_same( - pt_mapping_get_coord_set(input_mapping), - pt_mapping_get_coord_set(output_mapping)); + std::set coord_set = + require_same(pt_mapping_get_coord_set(input_mapping), + pt_mapping_get_coord_set(output_mapping)); return transform( coord_set, - [&](ParallelTensorSpaceCoordinate const &p) -> DynamicNodeInvocationShardingInfo { + [&](ParallelTensorSpaceCoordinate const &p) + -> DynamicNodeInvocationShardingInfo { // The machine coord for a copy is inherently nebulous because it // doesn't strictly run in any single location. Further, Realm has the // flexibility to issue a copy operation from anywhere in the machine, @@ -196,14 +205,16 @@ static std::set // because we expect this to align with the most efficient way to issue // copies in Realm, although the current Realm backend uses a // centralized controller and thus issues copies all from a single node. - global_device_id_t device_id = pt_mapping_get_device_for_coord(input_mapping, p); - - return invocation_sharding_info_for_binding(i, - device_id, - OperatorAtomicTaskShardBinding{{ - {input_slot.slot_name, p}, - {output_slot.slot_name, p}, - }}); + global_device_id_t device_id = + pt_mapping_get_device_for_coord(input_mapping, p); + + return invocation_sharding_info_for_binding( + i, + device_id, + OperatorAtomicTaskShardBinding{{ + {input_slot.slot_name, p}, + {output_slot.slot_name, p}, + }}); }); } @@ -218,99 +229,104 @@ static std::set DynamicNodeMapping node_mapping = assert_unwrap(i.node_attrs.mapping); DynamicTensorSlot expected_input_slot = DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }; DynamicValueAttrs input = require_only_key(i.inputs, expected_input_slot); DynamicTensorSlot expected_output_slot = DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }; DynamicValueAttrs output = require_only_key(i.outputs, expected_output_slot); ParallelTensorMapping input_value_mapping = assert_unwrap(input.mapping); - std::set input_tensor_shards = pt_mapping_get_coord_set(input_value_mapping); + std::set input_tensor_shards = + pt_mapping_get_coord_set(input_value_mapping); ParallelTensorMapping output_value_mapping = assert_unwrap(output.mapping); - auto get_task_shard_device_ids_for_input_tensor_shard - = [&](ParallelTensorSpaceCoordinate const &input_tensor_shard) - -> nonempty_set - { - bidict dependent_on_input_tensor_shard - = bidict_filter_values( - dynamic_node_mapping_get_shard_bindings(node_mapping), - [&](OperatorAtomicTaskShardBinding const &b) -> bool { - return ptensor_space_coord_for_slot_name(b, TensorSlotName::INPUT) == input_tensor_shard; - }); + auto get_task_shard_device_ids_for_input_tensor_shard = + [&](ParallelTensorSpaceCoordinate const &input_tensor_shard) + -> nonempty_set { + bidict + dependent_on_input_tensor_shard = bidict_filter_values( + dynamic_node_mapping_get_shard_bindings(node_mapping), + [&](OperatorAtomicTaskShardBinding const &b) -> bool { + return ptensor_space_coord_for_slot_name( + b, TensorSlotName::INPUT) == input_tensor_shard; + }); return nonempty_set(dependent_on_input_tensor_shard.left_values()); }; - auto invocation_sharding_info_for_input_tensor_shard = [&](ParallelTensorSpaceCoordinate const &c) - -> DynamicNodeInvocationShardingInfo - { + auto invocation_sharding_info_for_input_tensor_shard = + [&](ParallelTensorSpaceCoordinate const &c) + -> DynamicNodeInvocationShardingInfo { nonempty_set task_shard_device_ids = - get_task_shard_device_ids_for_input_tensor_shard(c); - - std::map output_sharding_infos = - generate_map(task_shard_device_ids.unwrap_as_set(), - [&](global_device_id_t const &device_id) - -> DynamicValueAttrsShardingInfo - { - ParallelTensorSpaceCoordinate pc - = pt_mapping_get_coord_for_device(output_value_mapping, device_id); - - return DynamicValueAttrsShardingInfo{ - /*shard_coord=*/pc, - /*mapping=*/device_id, - }; - }); - - std::map keyed_output_sharding_infos = - map_keys(output_sharding_infos, - [&](global_device_id_t const &device_id) -> DynamicTensorSlot { - return DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/device_id.coord, - }; - }); + get_task_shard_device_ids_for_input_tensor_shard(c); + + std::map + output_sharding_infos = + generate_map(task_shard_device_ids.unwrap_as_set(), + [&](global_device_id_t const &device_id) + -> DynamicValueAttrsShardingInfo { + ParallelTensorSpaceCoordinate pc = + pt_mapping_get_coord_for_device( + output_value_mapping, device_id); + + return DynamicValueAttrsShardingInfo{ + /*shard_coord=*/pc, + /*mapping=*/device_id, + }; + }); + + std::map + keyed_output_sharding_infos = map_keys( + output_sharding_infos, + [&](global_device_id_t const &device_id) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/device_id.coord, + }; + }); DynamicTensorSlot input_slot = DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }; - DynamicValueAttrsShardingInfo input_sharding_info = DynamicValueAttrsShardingInfo{ - /*shard_coord=*/c, - /*mapping=*/pt_mapping_get_device_for_coord(input_value_mapping, c), - }; + DynamicValueAttrsShardingInfo input_sharding_info = + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/c, + /*mapping=*/pt_mapping_get_device_for_coord(input_value_mapping, c), + }; std::map sharding_infos = - binary_merge_disjoint_maps( - keyed_output_sharding_infos, - std::map{ - { - input_slot, - input_sharding_info, - }, - }); + binary_merge_disjoint_maps( + keyed_output_sharding_infos, + std::map{ + { + input_slot, + input_sharding_info, + }, + }); return DynamicNodeInvocationShardingInfo{ - /*device_ids=*/task_shard_device_ids, - /*value_sharding=*/binary_relation_from_map(sharding_infos), + /*device_ids=*/task_shard_device_ids, + /*value_sharding=*/binary_relation_from_map(sharding_infos), }; }; - return transform(input_tensor_shards, invocation_sharding_info_for_input_tensor_shard); + return transform(input_tensor_shards, + invocation_sharding_info_for_input_tensor_shard); } static std::set @@ -320,120 +336,126 @@ static std::set DynamicNodeMapping node_mapping = assert_unwrap(i.node_attrs.mapping); DynamicTensorSlot expected_output_grad_slot = DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }; - DynamicValueAttrs output_grad = require_only_key(i.inputs, expected_output_grad_slot); + DynamicValueAttrs output_grad = + require_only_key(i.inputs, expected_output_grad_slot); DynamicTensorSlot expected_input_grad_slot = DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }; - DynamicValueAttrs input_grad = require_only_key(i.outputs, expected_input_grad_slot); + DynamicValueAttrs input_grad = + require_only_key(i.outputs, expected_input_grad_slot); - ParallelTensorMapping output_grad_value_mapping = assert_unwrap(output_grad.mapping); - ParallelTensorMapping input_grad_value_mapping = assert_unwrap(input_grad.mapping); + ParallelTensorMapping output_grad_value_mapping = + assert_unwrap(output_grad.mapping); + ParallelTensorMapping input_grad_value_mapping = + assert_unwrap(input_grad.mapping); - std::set input_grad_tensor_shards - = pt_mapping_get_coord_set(input_grad_value_mapping); + std::set input_grad_tensor_shards = + pt_mapping_get_coord_set(input_grad_value_mapping); - auto get_task_shard_device_ids_for_input_grad_tensor_shard - = [&](ParallelTensorSpaceCoordinate const &input_grad_tensor_shard) - -> nonempty_set - { - bidict produce_input_grad_tensor_shard - = bidict_filter_values( - dynamic_node_mapping_get_shard_bindings(node_mapping), - [&](OperatorAtomicTaskShardBinding const &b) -> bool { - return ptensor_space_coord_for_slot_name(b, TensorSlotName::INPUT) == input_grad_tensor_shard; - }); + auto get_task_shard_device_ids_for_input_grad_tensor_shard = + [&](ParallelTensorSpaceCoordinate const &input_grad_tensor_shard) + -> nonempty_set { + bidict + produce_input_grad_tensor_shard = bidict_filter_values( + dynamic_node_mapping_get_shard_bindings(node_mapping), + [&](OperatorAtomicTaskShardBinding const &b) -> bool { + return ptensor_space_coord_for_slot_name( + b, TensorSlotName::INPUT) == input_grad_tensor_shard; + }); return nonempty_set(produce_input_grad_tensor_shard.left_values()); }; - auto invocation_sharding_info_for_input_grad_tensor_shard = [&](ParallelTensorSpaceCoordinate const &c) - -> DynamicNodeInvocationShardingInfo - { + auto invocation_sharding_info_for_input_grad_tensor_shard = + [&](ParallelTensorSpaceCoordinate const &c) + -> DynamicNodeInvocationShardingInfo { nonempty_set task_shard_device_ids = - get_task_shard_device_ids_for_input_grad_tensor_shard(c); - - std::map output_grad_sharding_infos = - generate_map(task_shard_device_ids.unwrap_as_set(), - [&](global_device_id_t const &device_id) - -> DynamicValueAttrsShardingInfo - { - ParallelTensorSpaceCoordinate pc - = pt_mapping_get_coord_for_device(output_grad_value_mapping, device_id); - - return DynamicValueAttrsShardingInfo{ - /*shard_coord=*/pc, - /*mapping=*/device_id, - }; - }); - - std::map keyed_output_grad_sharding_infos = - map_keys(output_grad_sharding_infos, - [&](global_device_id_t const &device_id) -> DynamicTensorSlot { - return DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/device_id.coord, - }; - }); + get_task_shard_device_ids_for_input_grad_tensor_shard(c); + + std::map + output_grad_sharding_infos = + generate_map(task_shard_device_ids.unwrap_as_set(), + [&](global_device_id_t const &device_id) + -> DynamicValueAttrsShardingInfo { + ParallelTensorSpaceCoordinate pc = + pt_mapping_get_coord_for_device( + output_grad_value_mapping, device_id); + + return DynamicValueAttrsShardingInfo{ + /*shard_coord=*/pc, + /*mapping=*/device_id, + }; + }); + + std::map + keyed_output_grad_sharding_infos = map_keys( + output_grad_sharding_infos, + [&](global_device_id_t const &device_id) -> DynamicTensorSlot { + return DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/device_id.coord, + }; + }); DynamicTensorSlot input_grad_slot = DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }; - DynamicValueAttrsShardingInfo input_grad_sharding_info = DynamicValueAttrsShardingInfo{ - /*shard_coord=*/c, - /*mapping=*/pt_mapping_get_device_for_coord(input_grad_value_mapping, c), - }; + DynamicValueAttrsShardingInfo input_grad_sharding_info = + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/c, + /*mapping=*/ + pt_mapping_get_device_for_coord(input_grad_value_mapping, c), + }; std::map sharding_infos = - binary_merge_disjoint_maps( - keyed_output_grad_sharding_infos, - std::map{ - { - input_grad_slot, - input_grad_sharding_info, - }, - }); + binary_merge_disjoint_maps( + keyed_output_grad_sharding_infos, + std::map{ + { + input_grad_slot, + input_grad_sharding_info, + }, + }); return DynamicNodeInvocationShardingInfo{ - /*device_ids=*/task_shard_device_ids, - /*value_sharding=*/binary_relation_from_map(sharding_infos), + /*device_ids=*/task_shard_device_ids, + /*value_sharding=*/binary_relation_from_map(sharding_infos), }; }; - return transform(input_grad_tensor_shards, invocation_sharding_info_for_input_grad_tensor_shard); + return transform(input_grad_tensor_shards, + invocation_sharding_info_for_input_grad_tensor_shard); } std::set perform_shard_expansion_for_invocation(DynamicNodeInvocation const &i) { - std::set - shard_expansion_info = generate_shard_expansion_for_invocation(i); + std::set shard_expansion_info = + generate_shard_expansion_for_invocation(i); return transform( - shard_expansion_info, - [&](DynamicNodeInvocationShardingInfo const &s) - -> DynamicNodeInvocation - { - return apply_dynamic_node_invocation_sharding_info(i, s); - }); + shard_expansion_info, + [&](DynamicNodeInvocationShardingInfo const &s) -> DynamicNodeInvocation { + return apply_dynamic_node_invocation_sharding_info(i, s); + }); } DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( - DynamicNodeAttrs const &node_attrs, - nonempty_set const &device_ids) -{ + DynamicNodeAttrs const &node_attrs, + nonempty_set const &device_ids) { DynamicNodeAttrs result = node_attrs; result.device_ids = device_ids; @@ -441,17 +463,16 @@ DynamicNodeAttrs apply_dynamic_node_attrs_sharding_info( } DynamicValueAttrs apply_dynamic_value_attrs_sharding_info( - DynamicValueAttrs const &value_attrs, - DynamicValueAttrsShardingInfo const &value_sharding_info) -{ + DynamicValueAttrs const &value_attrs, + DynamicValueAttrsShardingInfo const &value_sharding_info) { DynamicValueAttrs result = value_attrs; result.shard_coord = value_sharding_info.shard_coord; if (result.mapping.has_value()) { ParallelTensorMapping value_mapping = assert_unwrap(result.mapping); - global_device_id_t from_mapping = - pt_mapping_get_device_for_coord(value_mapping, value_sharding_info.shard_coord); + global_device_id_t from_mapping = pt_mapping_get_device_for_coord( + value_mapping, value_sharding_info.shard_coord); global_device_id_t from_sharding_info = value_sharding_info.mapping; ASSERT(from_mapping == from_sharding_info); @@ -461,88 +482,87 @@ DynamicValueAttrs apply_dynamic_value_attrs_sharding_info( } DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( - DynamicNodeInvocation const &invocation, - DynamicNodeInvocationShardingInfo const &invocation_sharding_info) -{ + DynamicNodeInvocation const &invocation, + DynamicNodeInvocationShardingInfo const &invocation_sharding_info) { require_invocation_is_ready_for_shard_expansion(invocation); { - std::set invocation_slots = set_union( - keys(invocation.inputs), keys(invocation.outputs)); + std::set invocation_slots = + set_union(keys(invocation.inputs), keys(invocation.outputs)); - std::set shard_info_slots_ignoring_task_shard = - transform( - invocation_sharding_info.value_sharding.left_values(), - slot_without_task_shard); + std::set shard_info_slots_ignoring_task_shard = + transform(invocation_sharding_info.value_sharding.left_values(), + slot_without_task_shard); ASSERT(invocation_slots == shard_info_slots_ignoring_task_shard, dynamic_node_invocation_to_serializable(invocation), invocation_sharding_info); } - std::set shard_labelled = - filtrans(invocation_sharding_info.value_sharding.left_values(), - [](DynamicTensorSlot const &s) -> std::optional - { - if (s.task_shard.has_value()) { - return slot_without_task_shard(s); - } else { - return std::nullopt; - } - }); + std::set shard_labelled = filtrans( + invocation_sharding_info.value_sharding.left_values(), + [](DynamicTensorSlot const &s) -> std::optional { + if (s.task_shard.has_value()) { + return slot_without_task_shard(s); + } else { + return std::nullopt; + } + }); { - std::set not_shard_labelled = - filter(invocation_sharding_info.value_sharding.left_values(), - [](DynamicTensorSlot const &s) -> bool - { - return !s.task_shard.has_value(); - }); + std::set not_shard_labelled = + filter(invocation_sharding_info.value_sharding.left_values(), + [](DynamicTensorSlot const &s) -> bool { + return !s.task_shard.has_value(); + }); ASSERT(are_disjoint(shard_labelled, not_shard_labelled)); } - auto shard_value = [&](DynamicTensorSlot const &slot, DynamicValueAttrs const &value_attrs) - -> std::map - { + auto shard_value = [&](DynamicTensorSlot const &slot, + DynamicValueAttrs const &value_attrs) + -> std::map { ASSERT(!slot.task_shard.has_value()); if (contains(shard_labelled, slot)) { - BinaryRelation - for_slot = filter_binary_relation(invocation_sharding_info.value_sharding, - [&](DynamicTensorSlot const &s, - DynamicValueAttrsShardingInfo const &) - -> bool - { - return slot_without_task_shard(s) == slot; - }); - - std::set> result = transform( - for_slot.unwrap_as_set(), - [&](std::pair const &p) - -> std::pair - { - return {p.first, apply_dynamic_value_attrs_sharding_info(value_attrs, p.second)}; - }); + BinaryRelation + for_slot = filter_binary_relation( + invocation_sharding_info.value_sharding, + [&](DynamicTensorSlot const &s, + DynamicValueAttrsShardingInfo const &) -> bool { + return slot_without_task_shard(s) == slot; + }); + + std::set> result = + transform(for_slot.unwrap_as_set(), + [&](std::pair const &p) + -> std::pair { + return {p.first, + apply_dynamic_value_attrs_sharding_info( + value_attrs, p.second)}; + }); return map_from_pairs(result); } else { - DynamicValueAttrsShardingInfo sharding_info - = get_only(invocation_sharding_info.value_sharding.at_l(slot)); + DynamicValueAttrsShardingInfo sharding_info = + get_only(invocation_sharding_info.value_sharding.at_l(slot)); return { - { - slot, - apply_dynamic_value_attrs_sharding_info(value_attrs, sharding_info), - }, + { + slot, + apply_dynamic_value_attrs_sharding_info(value_attrs, + sharding_info), + }, }; } }; DynamicNodeInvocation result = DynamicNodeInvocation{ - /*inputs=*/flatmap(invocation.inputs, shard_value), - /*node_attrs=*/apply_dynamic_node_attrs_sharding_info( - invocation.node_attrs, invocation_sharding_info.device_ids), - /*outputs=*/flatmap(invocation.outputs, shard_value), + /*inputs=*/flatmap(invocation.inputs, shard_value), + /*node_attrs=*/ + apply_dynamic_node_attrs_sharding_info( + invocation.node_attrs, invocation_sharding_info.device_ids), + /*outputs=*/flatmap(invocation.outputs, shard_value), }; require_invocation_is_fully_shard_expanded(result); @@ -551,14 +571,14 @@ DynamicNodeInvocation apply_dynamic_node_invocation_sharding_info( } std::set - generate_shard_expansion_for_invocation(DynamicNodeInvocation const &i) -{ + generate_shard_expansion_for_invocation(DynamicNodeInvocation const &i) { require_invocation_is_ready_for_shard_expansion(i); std::set result = [&]() { if (i.node_attrs.op_attrs.value().is_copy()) { return generate_shard_expansion_for_copy(i); - } else if (training_op_attrs_has_op_type(i.node_attrs.op_attrs.value(), OperatorType::REPLICATE)) { + } else if (training_op_attrs_has_op_type(i.node_attrs.op_attrs.value(), + OperatorType::REPLICATE)) { DynamicTaskType task_type = assert_unwrap(i.node_attrs.task_type); switch (task_type) { case DynamicTaskType::FWD: @@ -575,11 +595,14 @@ std::set target_devices_of_dynamic_node_mapping(mapping); return transform(shard_machine_coords, - [&](global_device_id_t const &device_id) -> DynamicNodeInvocationShardingInfo { + [&](global_device_id_t const &device_id) + -> DynamicNodeInvocationShardingInfo { OperatorAtomicTaskShardBinding slot_bindings = - dynamic_node_mapping_get_shard_binding_for_device(mapping, device_id); + dynamic_node_mapping_get_shard_binding_for_device( + mapping, device_id); - return invocation_sharding_info_for_binding(i, device_id, slot_bindings); + return invocation_sharding_info_for_binding( + i, device_id, slot_bindings); }); } }(); diff --git a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc index 6541fe6e3c..ba9ab4428d 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/update_insertion.cc @@ -83,10 +83,9 @@ static DynamicNodeInvocation get_update_invocation_for_invocation( }; } -std::set - perform_update_insertion_for_invocation( - DynamicNodeInvocation const &invocation, - OptimizerAttrs const &optimizer_attrs) { +std::set perform_update_insertion_for_invocation( + DynamicNodeInvocation const &invocation, + OptimizerAttrs const &optimizer_attrs) { if (invocation.node_attrs.task_type.value() == DynamicTaskType::FWD && invocation.node_attrs.op_attrs.value().is_pcg_op() && diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc index 1f97f064e8..055586664a 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/copy_insertion.cc @@ -1,27 +1,26 @@ #include "task-spec/dynamic_graph/copy_insertion.h" +#include "op-attrs/ops/element_unary.h" #include "op-attrs/tensor_slot_name.dtg.h" #include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_task_type.dtg.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" #include "task-spec/dynamic_graph/dynamic_value_attrs.dtg.h" -#include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/check_kv.h" -#include #include "task-spec/dynamic_graph/dynamic_value_attrs.h" +#include "task-spec/dynamic_graph/pass_expansion.h" #include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" -#include "op-attrs/ops/element_unary.h" -#include "utils/containers/require_only_key.h" #include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h" -#include "task-spec/dynamic_graph/pass_expansion.h" +#include "test/utils/doctest/check_kv.h" +#include "test/utils/doctest/fmt/set.h" +#include "utils/containers/require_only_key.h" +#include using namespace ::FlexFlow; static DynamicValueAttrs - mk_value_attrs(size_t src_layer_guid, - TensorSlotName src_slot, - std::optional const &mapping) -{ + mk_value_attrs(size_t src_layer_guid, + TensorSlotName src_slot, + std::optional const &mapping) { return DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{ parallel_tensor_guid_t{ @@ -44,9 +43,9 @@ static DynamicValueAttrs static DynamicTensorSlot mk_slot(TensorSlotName slot_name) { return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }; } @@ -54,18 +53,20 @@ static DynamicNodeAttrs mk_node_attrs(size_t layer_guid, PCGOperatorAttrs const &op_attrs, DynamicNodeMapping const &mapping) { return DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/std::nullopt, - /*mapping=*/mapping, - /*op_attrs=*/TrainingOperationAttrs{ - op_attrs, - }, - /*layer_guid=*/dynamic_layer_guid_t{ - parallel_layer_guid_t{ - Node{layer_guid}, + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/mapping, + /*op_attrs=*/ + TrainingOperationAttrs{ + op_attrs, + }, + /*layer_guid=*/ + dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{layer_guid}, + }, }, - }, - /*per_device_op_state=*/std::nullopt, + /*per_device_op_state=*/std::nullopt, }; } @@ -80,208 +81,213 @@ static global_device_id_t mk_device_id(MachineSpaceCoordinate const &mc) { return global_device_id_t{mc, DeviceType::GPU}; }; -static DynamicOpenDataflowGraph mk_single_input_node_graph(MachineSpaceCoordinate const &input_device) -{ +static DynamicOpenDataflowGraph + mk_single_input_node_graph(MachineSpaceCoordinate const &input_device) { TensorShape input_shape = TensorShape{ - TensorDims{ - FFOrdered{ - 8_p, - 5_p, + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, }, - }, - DataType::FLOAT, + DataType::FLOAT, }; PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ - InputAttrs{ - input_shape, - }, + InputAttrs{ + input_shape, + }, }; PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ - make_relu_attrs(), + make_relu_attrs(), }; - auto mk_node_mapping = [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { + auto mk_node_mapping = + [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { return DynamicNodeMapping{ - /*op_task_group=*/op_task_group, - /*device_type=*/DeviceType::GPU, + /*op_task_group=*/op_task_group, + /*device_type=*/DeviceType::GPU, }; }; auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/idx, - /*shared_components=*/FFOrdered{ - 0_n, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/ + FFOrdered{ + 0_n, + 0_n, + }, }; }; DynamicValueAttrs input_op_output = mk_value_attrs(123, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); - MappedOperatorTaskGroup input_node_mapping = - MappedOperatorTaskGroup{ + MappedOperatorTaskGroup input_node_mapping = MappedOperatorTaskGroup{ bidict{ - { - input_device, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), + { + input_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }, }, - }, }, - }; + }; DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ - /*inputs=*/{}, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/123, - /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, - /*mapping=*/mk_node_mapping(input_node_mapping)), - /*outputs=*/{ + /*inputs=*/{}, + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, + /*mapping=*/mk_node_mapping(input_node_mapping)), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - input_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, }, - }, }; - DynamicOpenDataflowGraph g - = dynamic_open_dataflow_graph_from_invocation_set( - {input_invocation}); + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set({input_invocation}); return g; } -static DynamicOpenDataflowGraph mk_single_input_into_relu_graph(MachineSpaceCoordinate const &input_device, - MachineSpaceCoordinate const &relu1_device) -{ +static DynamicOpenDataflowGraph mk_single_input_into_relu_graph( + MachineSpaceCoordinate const &input_device, + MachineSpaceCoordinate const &relu1_device) { TensorShape input_shape = TensorShape{ - TensorDims{ - FFOrdered{ - 8_p, - 5_p, + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, }, - }, - DataType::FLOAT, + DataType::FLOAT, }; PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ - InputAttrs{ - input_shape, - }, + InputAttrs{ + input_shape, + }, }; PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ - make_relu_attrs(), + make_relu_attrs(), }; - auto mk_node_mapping = [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { + auto mk_node_mapping = + [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { return DynamicNodeMapping{ - /*op_task_group=*/op_task_group, - /*device_type=*/DeviceType::GPU, + /*op_task_group=*/op_task_group, + /*device_type=*/DeviceType::GPU, }; }; auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/idx, - /*shared_components=*/FFOrdered{ - 0_n, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/ + FFOrdered{ + 0_n, + 0_n, + }, }; }; DynamicValueAttrs input_op_output = mk_value_attrs(123, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); - MappedOperatorTaskGroup input_node_mapping = - MappedOperatorTaskGroup{ + MappedOperatorTaskGroup input_node_mapping = MappedOperatorTaskGroup{ bidict{ - { - input_device, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), + { + input_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }, }, - }, }, - }; + }; DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ - /*inputs=*/{}, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/123, - /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, - /*mapping=*/mk_node_mapping(input_node_mapping)), - /*outputs=*/{ + /*inputs=*/{}, + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, + /*mapping=*/mk_node_mapping(input_node_mapping)), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - input_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, }, - }, }; DynamicValueAttrs relu1_op_output = mk_value_attrs(124, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); - MappedOperatorTaskGroup relu1_node_mapping = - MappedOperatorTaskGroup{ + MappedOperatorTaskGroup relu1_node_mapping = MappedOperatorTaskGroup{ bidict{ - { - relu1_device, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::INPUT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), + { + relu1_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }, }, - }, }, - }; + }; DynamicNodeInvocation relu1_invocation = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - input_op_output, + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + input_op_output, + }, }, - }, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/124, - /*op_attrs=*/relu_attrs, - /*mapping=*/mk_node_mapping(relu1_node_mapping)), - /*outputs=*/{ + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/124, + /*op_attrs=*/relu_attrs, + /*mapping=*/mk_node_mapping(relu1_node_mapping)), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - relu1_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + relu1_op_output, + }, }, - }, }; - DynamicOpenDataflowGraph g - = dynamic_open_dataflow_graph_from_invocation_set( - {input_invocation, relu1_invocation}); + DynamicOpenDataflowGraph g = dynamic_open_dataflow_graph_from_invocation_set( + {input_invocation, relu1_invocation}); return g; } @@ -294,259 +300,268 @@ struct ExampleGraphTestCase { dynamic_invocation_id_t relu2_op_id; }; -static ExampleGraphTestCase mk_example_replicate_graph( - MachineSpaceCoordinate const &input_device, - MachineSpaceCoordinate const &relu1_device, - MachineSpaceCoordinate const &replicate_device1, - MachineSpaceCoordinate const &replicate_device2, - MachineSpaceCoordinate const &relu2_device1, - MachineSpaceCoordinate const &relu2_device2) -{ +static ExampleGraphTestCase + mk_example_replicate_graph(MachineSpaceCoordinate const &input_device, + MachineSpaceCoordinate const &relu1_device, + MachineSpaceCoordinate const &replicate_device1, + MachineSpaceCoordinate const &replicate_device2, + MachineSpaceCoordinate const &relu2_device1, + MachineSpaceCoordinate const &relu2_device2) { TensorShape input_shape = TensorShape{ - TensorDims{ - FFOrdered{ - 8_p, - 5_p, + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, }, - }, - DataType::FLOAT, + DataType::FLOAT, }; PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ - InputAttrs{ - input_shape, - }, + InputAttrs{ + input_shape, + }, }; PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ - make_relu_attrs(), + make_relu_attrs(), }; - auto mk_node_mapping = [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { + auto mk_node_mapping = + [](MappedOperatorTaskGroup const &op_task_group) -> DynamicNodeMapping { return DynamicNodeMapping{ - /*op_task_group=*/op_task_group, - /*device_type=*/DeviceType::GPU, + /*op_task_group=*/op_task_group, + /*device_type=*/DeviceType::GPU, }; }; auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/idx, - /*shared_components=*/FFOrdered{ - 0_n, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/ + FFOrdered{ + 0_n, + 0_n, + }, }; }; DynamicValueAttrs input_op_output = mk_value_attrs(123, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); - MappedOperatorTaskGroup input_node_mapping = - MappedOperatorTaskGroup{ + MappedOperatorTaskGroup input_node_mapping = MappedOperatorTaskGroup{ bidict{ - { - input_device, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), + { + input_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }, }, - }, }, - }; + }; DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ - /*inputs=*/{}, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/123, - /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, - /*mapping=*/mk_node_mapping(input_node_mapping)), - /*outputs=*/{ + /*inputs=*/{}, + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}, + /*mapping=*/mk_node_mapping(input_node_mapping)), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - input_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, }, - }, }; DynamicValueAttrs relu1_op_output = mk_value_attrs(124, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); - MappedOperatorTaskGroup relu1_node_mapping = - MappedOperatorTaskGroup{ + MappedOperatorTaskGroup relu1_node_mapping = MappedOperatorTaskGroup{ bidict{ - { - relu1_device, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::INPUT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), + { + relu1_device, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }, }, - }, }, - }; + }; DynamicNodeInvocation relu1_invocation = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - input_op_output, + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + input_op_output, + }, }, - }, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/124, - /*op_attrs=*/relu_attrs, - /*mapping=*/mk_node_mapping(relu1_node_mapping)), - /*outputs=*/{ + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/124, + /*op_attrs=*/relu_attrs, + /*mapping=*/mk_node_mapping(relu1_node_mapping)), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - relu1_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + relu1_op_output, + }, }, - }, }; DynamicValueAttrs replicate_op_output = mk_value_attrs(125, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); - MappedOperatorTaskGroup replicate_node_mapping = - MappedOperatorTaskGroup{ + MappedOperatorTaskGroup replicate_node_mapping = MappedOperatorTaskGroup{ bidict{ - { - replicate_device1, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::INPUT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), + { + replicate_device1, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }, }, - }, - { - replicate_device2, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::INPUT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(1_n), + { + replicate_device2, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }, }, - }, }, - }, }, - }; + }; DynamicNodeInvocation replicate_invocation = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - relu1_op_output, - }, - }, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/125, - /*op_attrs=*/PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, - }, + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + relu1_op_output, + }, }, - /*mapping=*/mk_node_mapping(replicate_node_mapping)), - /*outputs=*/{ + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/125, + /*op_attrs=*/ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, + }, + /*mapping=*/mk_node_mapping(replicate_node_mapping)), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - replicate_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + replicate_op_output, + }, }, - }, }; DynamicValueAttrs relu2_op_output = mk_value_attrs(126, TensorSlotName::OUTPUT, /*mapping=*/std::nullopt); - MappedOperatorTaskGroup relu2_node_mapping = - MappedOperatorTaskGroup{ + MappedOperatorTaskGroup relu2_node_mapping = MappedOperatorTaskGroup{ bidict{ - { - relu2_device1, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::INPUT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), + { + relu2_device1, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }, }, - }, }, - }, - { - relu2_device2, - OperatorAtomicTaskShardBinding{ - std::map{ - { - TensorSlotName::INPUT, - mk_pt_coord(1_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(1_n), + { + relu2_device2, + OperatorAtomicTaskShardBinding{ + std::map{ + { + TensorSlotName::INPUT, + mk_pt_coord(1_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }, }, - }, }, - }, }, - }; + }; DynamicNodeInvocation relu2_invocation = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - replicate_op_output, + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + replicate_op_output, + }, }, - }, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/125, - /*op_attrs=*/relu_attrs, - /*mapping=*/mk_node_mapping(relu2_node_mapping)), - /*outputs=*/{ + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/125, + /*op_attrs=*/relu_attrs, + /*mapping=*/mk_node_mapping(relu2_node_mapping)), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - relu2_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + relu2_op_output, + }, }, - }, }; - DynamicOpenDataflowGraph g - = dynamic_open_dataflow_graph_from_invocation_set( - {input_invocation, relu1_invocation, replicate_invocation, relu2_invocation}); + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set({input_invocation, + relu1_invocation, + replicate_invocation, + relu2_invocation}); return ExampleGraphTestCase{ - /*g=*/g, - /*input_op_id=*/dynamic_graph_get_id_for_invocation(g, input_invocation), - /*relu1_op_id=*/dynamic_graph_get_id_for_invocation(g, relu1_invocation), - /*replicate_op_id=*/dynamic_graph_get_id_for_invocation(g, replicate_invocation), - /*relu2_op_id=*/dynamic_graph_get_id_for_invocation(g, relu2_invocation), + /*g=*/g, + /*input_op_id=*/dynamic_graph_get_id_for_invocation(g, input_invocation), + /*relu1_op_id=*/dynamic_graph_get_id_for_invocation(g, relu1_invocation), + /*replicate_op_id=*/ + dynamic_graph_get_id_for_invocation(g, replicate_invocation), + /*relu2_op_id=*/dynamic_graph_get_id_for_invocation(g, relu2_invocation), }; }; @@ -556,126 +571,128 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate mc2 = mk_machine_coord(1_n); MachineSpaceCoordinate mc3 = mk_machine_coord(2_n); - auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + auto mk_pt_coord = + [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/idx, - /*shared_components=*/FFOrdered{ - 0_n, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/ + FFOrdered{ + 0_n, + 0_n, + }, }; }; - SUBCASE("dynamic graph is not pass expanded") { auto mk_correct_mappings = - [&](ExampleGraphTestCase const &tc, - MachineSpaceCoordinate const &input_out_mc, - MachineSpaceCoordinate const &relu1_in_mc, - MachineSpaceCoordinate const &relu1_out_mc, - MachineSpaceCoordinate const &replicate_in_mc, - MachineSpaceCoordinate const &replicate_out_mc1, - MachineSpaceCoordinate const &replicate_out_mc2, - MachineSpaceCoordinate const &relu2_in_mc1, - MachineSpaceCoordinate const &relu2_in_mc2, - MachineSpaceCoordinate const &relu2_out_mc1, - MachineSpaceCoordinate const &relu2_out_mc2) - -> std::map - { - DynamicOpenDataflowGraph g = tc.g; - dynamic_invocation_id_t input_op_id = tc.input_op_id; - dynamic_invocation_id_t relu1_op_id = tc.relu1_op_id; - dynamic_invocation_id_t replicate_op_id = tc.replicate_op_id; - dynamic_invocation_id_t relu2_op_id = tc.relu2_op_id; - - auto mk_inp_slot = [](dynamic_invocation_id_t invocation_id) -> InternalDynamicSlotSite { - return InternalDynamicSlotSite{ + [&](ExampleGraphTestCase const &tc, + MachineSpaceCoordinate const &input_out_mc, + MachineSpaceCoordinate const &relu1_in_mc, + MachineSpaceCoordinate const &relu1_out_mc, + MachineSpaceCoordinate const &replicate_in_mc, + MachineSpaceCoordinate const &replicate_out_mc1, + MachineSpaceCoordinate const &replicate_out_mc2, + MachineSpaceCoordinate const &relu2_in_mc1, + MachineSpaceCoordinate const &relu2_in_mc2, + MachineSpaceCoordinate const &relu2_out_mc1, + MachineSpaceCoordinate const &relu2_out_mc2) + -> std::map { + DynamicOpenDataflowGraph g = tc.g; + dynamic_invocation_id_t input_op_id = tc.input_op_id; + dynamic_invocation_id_t relu1_op_id = tc.relu1_op_id; + dynamic_invocation_id_t replicate_op_id = tc.replicate_op_id; + dynamic_invocation_id_t relu2_op_id = tc.relu2_op_id; + + auto mk_inp_slot = [](dynamic_invocation_id_t invocation_id) + -> InternalDynamicSlotSite { + return InternalDynamicSlotSite{ /*invocation_id=*/invocation_id, /*direction=*/TensorDirection::INCOMING, /*slot_name=*/mk_slot(TensorSlotName::INPUT), - }; }; + }; - auto mk_out_slot = [](dynamic_invocation_id_t invocation_id) -> InternalDynamicSlotSite { - return InternalDynamicSlotSite{ + auto mk_out_slot = [](dynamic_invocation_id_t invocation_id) + -> InternalDynamicSlotSite { + return InternalDynamicSlotSite{ /*invocation_id=*/invocation_id, /*direction=*/TensorDirection::OUTPUT, /*slot_name=*/mk_slot(TensorSlotName::OUTPUT), - }; }; + }; - auto mk_single_shard_mapping = [&](MachineSpaceCoordinate const &mc) -> ParallelTensorMapping { - return ParallelTensorMapping{ + auto mk_single_shard_mapping = + [&](MachineSpaceCoordinate const &mc) -> ParallelTensorMapping { + return ParallelTensorMapping{ bidict{ - { - mk_pt_coord(0_n), - mk_device_id(mc), - }, + { + mk_pt_coord(0_n), + mk_device_id(mc), + }, }, - }; }; + }; - auto mk_two_shard_mapping = [&](MachineSpaceCoordinate const &mc1, - MachineSpaceCoordinate const &mc2) -> ParallelTensorMapping { - return ParallelTensorMapping{ + auto mk_two_shard_mapping = + [&](MachineSpaceCoordinate const &mc1, + MachineSpaceCoordinate const &mc2) -> ParallelTensorMapping { + return ParallelTensorMapping{ bidict{ - { - mk_pt_coord(0_n), - mk_device_id(mc1), - }, - { - mk_pt_coord(1_n), - mk_device_id(mc2), - }, + { + mk_pt_coord(0_n), + mk_device_id(mc1), + }, + { + mk_pt_coord(1_n), + mk_device_id(mc2), + }, }, - }; }; + }; - InternalDynamicSlotSite input_op_out = mk_out_slot(input_op_id); - InternalDynamicSlotSite relu1_op_in = mk_inp_slot(relu1_op_id); - InternalDynamicSlotSite relu1_op_out = mk_out_slot(relu1_op_id); - InternalDynamicSlotSite replicate_op_in = mk_inp_slot(replicate_op_id); - InternalDynamicSlotSite replicate_op_out = mk_out_slot(replicate_op_id); - InternalDynamicSlotSite relu2_op_in = mk_inp_slot(relu2_op_id); - InternalDynamicSlotSite relu2_op_out = mk_out_slot(relu2_op_id); + InternalDynamicSlotSite input_op_out = mk_out_slot(input_op_id); + InternalDynamicSlotSite relu1_op_in = mk_inp_slot(relu1_op_id); + InternalDynamicSlotSite relu1_op_out = mk_out_slot(relu1_op_id); + InternalDynamicSlotSite replicate_op_in = mk_inp_slot(replicate_op_id); + InternalDynamicSlotSite replicate_op_out = mk_out_slot(replicate_op_id); + InternalDynamicSlotSite relu2_op_in = mk_inp_slot(relu2_op_id); + InternalDynamicSlotSite relu2_op_out = mk_out_slot(relu2_op_id); - return { + return { { - input_op_out, - mk_single_shard_mapping(input_out_mc), + input_op_out, + mk_single_shard_mapping(input_out_mc), }, { - relu1_op_in, - mk_single_shard_mapping(relu1_in_mc), + relu1_op_in, + mk_single_shard_mapping(relu1_in_mc), }, { - relu1_op_out, - mk_single_shard_mapping(relu1_out_mc), + relu1_op_out, + mk_single_shard_mapping(relu1_out_mc), }, { - replicate_op_in, - mk_single_shard_mapping(replicate_in_mc), + replicate_op_in, + mk_single_shard_mapping(replicate_in_mc), }, { - replicate_op_out, - mk_two_shard_mapping(replicate_out_mc1, replicate_out_mc2), + replicate_op_out, + mk_two_shard_mapping(replicate_out_mc1, replicate_out_mc2), }, { - relu2_op_in, - mk_two_shard_mapping(relu2_in_mc1, relu2_in_mc2), + relu2_op_in, + mk_two_shard_mapping(relu2_in_mc1, relu2_in_mc2), }, { - relu2_op_out, - mk_two_shard_mapping(relu2_out_mc1, relu2_out_mc2), + relu2_op_out, + mk_two_shard_mapping(relu2_out_mc1, relu2_out_mc2), }, - }; }; - + }; SUBCASE("replicate input matches tensor source") { - ExampleGraphTestCase tc - = mk_example_replicate_graph( + ExampleGraphTestCase tc = mk_example_replicate_graph( /*input_device=*/mc1, /*relu1_device=*/mc1, /*replicate_device1=*/mc1, @@ -684,28 +701,27 @@ TEST_SUITE(FF_TEST_SUITE) { /*relu2_device2=*/mc2); std::map result = - resolve_tensor_mappings(tc.g); + resolve_tensor_mappings(tc.g); std::map correct = - mk_correct_mappings( - /*tc=*/tc, - /*input_out_mc=*/mc1, - /*relu1_in_mc=*/mc1, - /*relu1_out_mc=*/mc1, - /*replicate_in_mc=*/mc1, - /*relicate_out_mc1=*/mc1, - /*relicate_out_mc2=*/mc2, - /*relu2_in_mc1=*/mc1, - /*relu2_in_mc2=*/mc2, - /*relu2_out_mc1=*/mc1, - /*relu2_out_mc1=*/mc2); + mk_correct_mappings( + /*tc=*/tc, + /*input_out_mc=*/mc1, + /*relu1_in_mc=*/mc1, + /*relu1_out_mc=*/mc1, + /*replicate_in_mc=*/mc1, + /*relicate_out_mc1=*/mc1, + /*relicate_out_mc2=*/mc2, + /*relu2_in_mc1=*/mc1, + /*relu2_in_mc2=*/mc2, + /*relu2_out_mc1=*/mc1, + /*relu2_out_mc1=*/mc2); ASSERT(result == correct); } SUBCASE("input invocation's output follows the input's mapping") { - ExampleGraphTestCase tc - = mk_example_replicate_graph( + ExampleGraphTestCase tc = mk_example_replicate_graph( /*input_device=*/mc3, /*relu1_device=*/mc1, /*replicate_device1=*/mc1, @@ -714,28 +730,27 @@ TEST_SUITE(FF_TEST_SUITE) { /*relu2_device2=*/mc2); std::map result = - resolve_tensor_mappings(tc.g); + resolve_tensor_mappings(tc.g); std::map correct = - mk_correct_mappings( - /*tc=*/tc, - /*input_out_mc=*/mc3, - /*relu1_in_mc=*/mc1, - /*relu1_out_mc=*/mc1, - /*replicate_in_mc=*/mc1, - /*relicate_out_mc1=*/mc1, - /*relicate_out_mc2=*/mc2, - /*relu2_in_mc1=*/mc1, - /*relu2_in_mc2=*/mc2, - /*relu2_out_mc1=*/mc1, - /*relu2_out_mc1=*/mc2); + mk_correct_mappings( + /*tc=*/tc, + /*input_out_mc=*/mc3, + /*relu1_in_mc=*/mc1, + /*relu1_out_mc=*/mc1, + /*replicate_in_mc=*/mc1, + /*relicate_out_mc1=*/mc1, + /*relicate_out_mc2=*/mc2, + /*relu2_in_mc1=*/mc1, + /*relu2_in_mc2=*/mc2, + /*relu2_out_mc1=*/mc1, + /*relu2_out_mc1=*/mc2); ASSERT(result == correct); } SUBCASE("src and sink can differ due to different invocation mappings") { - ExampleGraphTestCase tc - = mk_example_replicate_graph( + ExampleGraphTestCase tc = mk_example_replicate_graph( /*input_device=*/mc3, /*relu1_device=*/mc1, /*replicate_device1=*/mc2, @@ -744,21 +759,21 @@ TEST_SUITE(FF_TEST_SUITE) { /*relu2_device2=*/mc3); std::map result = - resolve_tensor_mappings(tc.g); + resolve_tensor_mappings(tc.g); std::map correct = - mk_correct_mappings( - /*tc=*/tc, - /*input_out_mc=*/mc3, - /*relu1_in_mc=*/mc1, - /*relu1_out_mc=*/mc1, - /*replicate_in_mc=*/mc1, - /*relicate_out_mc1=*/mc2, - /*relicate_out_mc2=*/mc3, - /*relu2_in_mc1=*/mc1, - /*relu2_in_mc2=*/mc3, - /*relu2_out_mc1=*/mc1, - /*relu2_out_mc2=*/mc3); + mk_correct_mappings( + /*tc=*/tc, + /*input_out_mc=*/mc3, + /*relu1_in_mc=*/mc1, + /*relu1_out_mc=*/mc1, + /*replicate_in_mc=*/mc1, + /*relicate_out_mc1=*/mc2, + /*relicate_out_mc2=*/mc3, + /*relu2_in_mc1=*/mc1, + /*relu2_in_mc2=*/mc3, + /*relu2_out_mc1=*/mc1, + /*relu2_out_mc2=*/mc3); ASSERT(result == correct); } @@ -768,12 +783,12 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("copies_for_value") { DynamicValueAttrs value_attrs = DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{ - parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{1}, - TensorSlotName::OUTPUT, + parallel_tensor_guid_t{ + KwargDataflowOutput{ + Node{1}, + TensorSlotName::OUTPUT, + }, }, - }, }, /*parallel_tensor_shape=*/std::nullopt, /*create_grad=*/std::nullopt, @@ -783,62 +798,66 @@ TEST_SUITE(FF_TEST_SUITE) { /*role=*/std::nullopt, }; - auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + auto mk_pt_coord = + [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/0_n, - /*shared_components=*/FFOrdered{ - idx, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shared_components=*/ + FFOrdered{ + idx, + 0_n, + }, }; }; auto mk_device = [](nonnegative_int idx) -> global_device_id_t { return global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/2_n, - /*device_idx=*/idx, - }, - /*device_type=*/DeviceType::GPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/2_n, + /*device_idx=*/idx, + }, + /*device_type=*/DeviceType::GPU, }; }; auto mk_slot_site = [](nonnegative_int idx) -> InternalDynamicSlotSite { return InternalDynamicSlotSite{ - /*invocation_id=*/dynamic_invocation_id_t{idx}, - /*direction=*/TensorDirection::INCOMING, - /*slot_name=*/DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/DynamicTensorRole{FwbTensorType::FORWARD}, // could be any role - /*task_shard=*/std::nullopt, - }, + /*invocation_id=*/dynamic_invocation_id_t{idx}, + /*direction=*/TensorDirection::INCOMING, + /*slot_name=*/ + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/ + DynamicTensorRole{FwbTensorType::FORWARD}, // could be any role + /*task_shard=*/std::nullopt, + }, }; }; SUBCASE("if src site is external, no copies no matter what") { DynamicSlotSite src_site = DynamicSlotSite{ - ExternalDynamicSlotSite{ - dynamic_external_value_id_t{0_n}, - }, + ExternalDynamicSlotSite{ + dynamic_external_value_id_t{0_n}, + }, }; ParallelTensorMapping mapping1 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, }; ParallelTensorMapping mapping2 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(1_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(1_n)}, + }, }; ParallelTensorMapping mapping3 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(2_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(2_n)}, + }, }; InternalDynamicSlotSite dst_site1 = mk_slot_site(2_n); @@ -846,13 +865,12 @@ TEST_SUITE(FF_TEST_SUITE) { InternalDynamicSlotSite dst_site3 = mk_slot_site(4_n); std::map site_mappings = { - {dst_site1, mapping2}, - {dst_site2, mapping1}, - {dst_site3, mapping3}, + {dst_site1, mapping2}, + {dst_site2, mapping1}, + {dst_site3, mapping3}, }; - std::set result = - copies_for_value( + std::set result = copies_for_value( /*value_attrs=*/value_attrs, /*src_site=*/DynamicSlotSite{src_site}, /*dst_sites=*/{dst_site1, dst_site2, dst_site3}, @@ -864,31 +882,31 @@ TEST_SUITE(FF_TEST_SUITE) { }; InternalDynamicSlotSite src_site = InternalDynamicSlotSite{ - /*invocation_id=*/dynamic_invocation_id_t{0_n}, - /*direction=*/TensorDirection::OUTPUT, - /*slot_name=*/DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, + /*invocation_id=*/dynamic_invocation_id_t{0_n}, + /*direction=*/TensorDirection::OUTPUT, + /*slot_name=*/ + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, }; SUBCASE("if src mapping is same as dst mapping don't copy") { ParallelTensorMapping mapping1 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, }; InternalDynamicSlotSite dst_site = mk_slot_site(1_n); std::map site_mappings = { - {src_site, mapping1}, - {dst_site, mapping1}, + {src_site, mapping1}, + {dst_site, mapping1}, }; - std::set result = - copies_for_value( + std::set result = copies_for_value( /*value_attrs=*/value_attrs, /*src_site=*/DynamicSlotSite{src_site}, /*dst_sites=*/{dst_site}, @@ -901,37 +919,36 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("if src mapping does not overlap dst mapping issue copy") { ParallelTensorMapping mapping1 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, }; ParallelTensorMapping mapping2 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(5_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(5_n)}, + }, }; InternalDynamicSlotSite dst_site = mk_slot_site(1_n); std::map site_mappings = { - {src_site, mapping1}, - {dst_site, mapping2}, + {src_site, mapping1}, + {dst_site, mapping2}, }; - std::set result = - copies_for_value( + std::set result = copies_for_value( /*value_attrs=*/value_attrs, /*src_site=*/DynamicSlotSite{src_site}, /*dst_sites=*/{dst_site}, /*all_mappings=*/site_mappings); std::set correct = { - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping2, - }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, }; CHECK(result == correct); @@ -939,39 +956,38 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("if src mapping overlaps dst mapping issue full copy") { ParallelTensorMapping mapping1 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - {mk_pt_coord(1_n), mk_device(1_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(1_n)}, + }, }; ParallelTensorMapping mapping2 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - {mk_pt_coord(1_n), mk_device(2_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(2_n)}, + }, }; InternalDynamicSlotSite dst_site = mk_slot_site(1_n); std::map site_mappings = { - {src_site, mapping1}, - {dst_site, mapping2}, + {src_site, mapping1}, + {dst_site, mapping2}, }; - std::set result = - copies_for_value( + std::set result = copies_for_value( /*value_attrs=*/value_attrs, /*src_site=*/DynamicSlotSite{src_site}, /*dst_sites=*/{dst_site}, /*all_mappings=*/site_mappings); std::set correct = { - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping2, - }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, }; CHECK(result == correct); @@ -979,53 +995,52 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("if src mapping overlaps multiple dst mappings issue both copies") { ParallelTensorMapping mapping1 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - {mk_pt_coord(1_n), mk_device(1_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(1_n)}, + }, }; ParallelTensorMapping mapping2 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - {mk_pt_coord(1_n), mk_device(2_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(2_n)}, + }, }; ParallelTensorMapping mapping3 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - {mk_pt_coord(1_n), mk_device(3_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + {mk_pt_coord(1_n), mk_device(3_n)}, + }, }; InternalDynamicSlotSite dst_site1 = mk_slot_site(1_n); InternalDynamicSlotSite dst_site2 = mk_slot_site(2_n); std::map site_mappings = { - {src_site, mapping1}, - {dst_site1, mapping2}, - {dst_site2, mapping3}, + {src_site, mapping1}, + {dst_site1, mapping2}, + {dst_site2, mapping3}, }; - std::set result = - copies_for_value( + std::set result = copies_for_value( /*value_attrs=*/value_attrs, /*src_site=*/DynamicSlotSite{src_site}, /*dst_sites=*/{dst_site1, dst_site2}, /*sink_site_mappings=*/site_mappings); std::set correct = { - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping2, - }, - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping3, - }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping3, + }, }; CHECK(result == correct); @@ -1033,15 +1048,15 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("only copy once if multiple sinks use the same mapping") { ParallelTensorMapping mapping1 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, }; ParallelTensorMapping mapping2 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(5_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(5_n)}, + }, }; InternalDynamicSlotSite dst_site1 = mk_slot_site(1_n); @@ -1049,53 +1064,53 @@ TEST_SUITE(FF_TEST_SUITE) { InternalDynamicSlotSite dst_site3 = mk_slot_site(3_n); std::map site_mappings = { - {src_site, mapping1}, - {dst_site1, mapping2}, - {dst_site2, mapping2}, - {dst_site3, mapping2}, + {src_site, mapping1}, + {dst_site1, mapping2}, + {dst_site2, mapping2}, + {dst_site3, mapping2}, }; - std::set result = - copies_for_value( + std::set result = copies_for_value( /*value_attrs=*/value_attrs, /*src_site=*/DynamicSlotSite{src_site}, /*dst_sites=*/{dst_site1, dst_site2, dst_site3}, /*all_mappings=*/site_mappings); std::set correct = { - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping2, - }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, }; CHECK(result == correct); } - SUBCASE("if src mapping matches one dst mapping, still issue copies for the rest") { + SUBCASE("if src mapping matches one dst mapping, still issue copies for " + "the rest") { ParallelTensorMapping mapping1 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(0_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(0_n)}, + }, }; ParallelTensorMapping mapping2 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(5_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(5_n)}, + }, }; ParallelTensorMapping mapping3 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(6_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(6_n)}, + }, }; ParallelTensorMapping mapping4 = ParallelTensorMapping{ - bidict{ - {mk_pt_coord(0_n), mk_device(7_n)}, - }, + bidict{ + {mk_pt_coord(0_n), mk_device(7_n)}, + }, }; InternalDynamicSlotSite dst_site1 = mk_slot_site(1_n); @@ -1105,37 +1120,36 @@ TEST_SUITE(FF_TEST_SUITE) { InternalDynamicSlotSite dst_site5 = mk_slot_site(5_n); std::map site_mappings = { - {src_site, mapping1}, - {dst_site1, mapping2}, - {dst_site2, mapping1}, - {dst_site3, mapping2}, - {dst_site4, mapping4}, - {dst_site5, mapping3}, + {src_site, mapping1}, + {dst_site1, mapping2}, + {dst_site2, mapping1}, + {dst_site3, mapping2}, + {dst_site4, mapping4}, + {dst_site5, mapping3}, }; - std::set result = - copies_for_value( + std::set result = copies_for_value( /*value_attrs=*/value_attrs, /*src_site=*/DynamicSlotSite{src_site}, /*dst_sites=*/{dst_site1, dst_site2, dst_site3, dst_site4, dst_site5}, /*all_mappings=*/site_mappings); std::set correct = { - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping2, - }, - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping3, - }, - DynamicValueCopyInfo{ - /*value_attrs=*/value_attrs, - /*src_mapping=*/mapping1, - /*dst_mapping=*/mapping4, - }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping2, + }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping3, + }, + DynamicValueCopyInfo{ + /*value_attrs=*/value_attrs, + /*src_mapping=*/mapping1, + /*dst_mapping=*/mapping4, + }, }; CHECK(result == correct); @@ -1157,14 +1171,16 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - auto mk_device_id = [&](nonnegative_int device_idx) -> global_device_id_t { + auto mk_device_id = + [&](nonnegative_int device_idx) -> global_device_id_t { return global_device_id_t{ mk_machine_coord(device_idx), DeviceType::GPU, }; }; - auto mk_pcg_layer_guid = [](size_t pcg_layer_guid) -> dynamic_layer_guid_t { + auto mk_pcg_layer_guid = + [](size_t pcg_layer_guid) -> dynamic_layer_guid_t { return dynamic_layer_guid_t{ parallel_layer_guid_t{ Node{pcg_layer_guid}, @@ -1175,8 +1191,7 @@ TEST_SUITE(FF_TEST_SUITE) { auto mk_node_attrs = [](dynamic_layer_guid_t layer_guid, std::optional const &mapping, - TrainingOperationAttrs const &op_attrs) - -> DynamicNodeAttrs { + TrainingOperationAttrs const &op_attrs) -> DynamicNodeAttrs { return DynamicNodeAttrs{ /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, @@ -1191,7 +1206,8 @@ TEST_SUITE(FF_TEST_SUITE) { nonnegative_int output_shard_idx) -> OperatorAtomicTaskShardBinding { return OperatorAtomicTaskShardBinding{ - /*tensor_coords=*/std::map{ + /*tensor_coords=*/std::map{ { TensorSlotName::INPUT, mk_ptensor_coord(input_shard_idx), @@ -1205,9 +1221,9 @@ TEST_SUITE(FF_TEST_SUITE) { }; TrainingOperationAttrs relu_attrs = TrainingOperationAttrs{ - PCGOperatorAttrs{ - make_relu_attrs(), - }, + PCGOperatorAttrs{ + make_relu_attrs(), + }, }; DynamicValueAttrs v1 = mk_value_attrs( @@ -1266,8 +1282,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(1), mapping1, relu_attrs), /*outputs=*/ { { @@ -1285,8 +1300,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(2), mapping2, relu_attrs), /*outputs=*/ { { @@ -1354,8 +1368,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(1), mapping1, relu_attrs), /*outputs=*/ { { @@ -1394,8 +1407,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(2), mapping2, relu_attrs), /*outputs=*/ { { @@ -1456,8 +1468,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(1), mapping1, relu_attrs), /*outputs=*/ { { @@ -1475,8 +1486,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(2), mapping2, relu_attrs), /*outputs=*/ { { @@ -1533,8 +1543,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(1), mapping1, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(1), mapping1, relu_attrs), /*outputs=*/ { { @@ -1552,8 +1561,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, }, /*node_attrs=*/ - mk_node_attrs( - mk_pcg_layer_guid(2), mapping2, relu_attrs), + mk_node_attrs(mk_pcg_layer_guid(2), mapping2, relu_attrs), /*outputs=*/ { { @@ -1575,14 +1583,16 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("replicate operator") { - auto mk_pt_coord = [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { + auto mk_pt_coord = + [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/idx, - /*shared_components=*/FFOrdered{ - 0_n, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shared_components=*/ + FFOrdered{ + 0_n, + 0_n, + }, }; }; @@ -1590,8 +1600,7 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate mc2 = mk_machine_coord(1_n); MachineSpaceCoordinate mc3 = mk_machine_coord(2_n); - ExampleGraphTestCase tc - = mk_example_replicate_graph( + ExampleGraphTestCase tc = mk_example_replicate_graph( /*input_device=*/mc3, /*relu1_device=*/mc1, /*replicate_device1=*/mc2, @@ -1605,171 +1614,176 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicOpenDataflowGraph result = perform_copy_insertion(tc.g); DynamicOpenDataflowGraph correct = [&] { - auto map_input_value = [](DynamicNodeInvocation const &invocation, - ParallelTensorMapping const &mapping) - -> DynamicNodeInvocation - { + auto map_input_value = + [](DynamicNodeInvocation const &invocation, + ParallelTensorMapping const &mapping) -> DynamicNodeInvocation { DynamicNodeInvocation result = invocation; DynamicTensorSlot input_slot = mk_slot(TensorSlotName::INPUT); result.inputs = { - { - input_slot, - decide_dynamic_value_attrs_mapping( - require_only_key(invocation.inputs, input_slot), - mapping), - }, + { + input_slot, + decide_dynamic_value_attrs_mapping( + require_only_key(invocation.inputs, input_slot), mapping), + }, }; return result; }; - auto map_output_value = [](DynamicNodeInvocation const &invocation, - ParallelTensorMapping const &mapping) - -> DynamicNodeInvocation - { + auto map_output_value = + [](DynamicNodeInvocation const &invocation, + ParallelTensorMapping const &mapping) -> DynamicNodeInvocation { DynamicNodeInvocation result = invocation; DynamicTensorSlot output_slot = mk_slot(TensorSlotName::OUTPUT); result.outputs = { - { - output_slot, - decide_dynamic_value_attrs_mapping( - require_only_key(invocation.outputs, output_slot), - mapping), - }, + { + output_slot, + decide_dynamic_value_attrs_mapping( + require_only_key(invocation.outputs, output_slot), + mapping), + }, }; return result; }; - auto map_input_and_output_values = [&](DynamicNodeInvocation const &invocation, - ParallelTensorMapping const &input_mapping, - ParallelTensorMapping const &output_mapping) - -> DynamicNodeInvocation - { - return map_input_value( - map_output_value( - invocation, - output_mapping), - input_mapping); + auto map_input_and_output_values = + [&](DynamicNodeInvocation const &invocation, + ParallelTensorMapping const &input_mapping, + ParallelTensorMapping const &output_mapping) + -> DynamicNodeInvocation { + return map_input_value(map_output_value(invocation, output_mapping), + input_mapping); }; - auto mk_single_shard_mapping = [&](MachineSpaceCoordinate const &mc) -> ParallelTensorMapping { + auto mk_single_shard_mapping = + [&](MachineSpaceCoordinate const &mc) -> ParallelTensorMapping { return ParallelTensorMapping{ - bidict{ - { - mk_pt_coord(0_n), - mk_device_id(mc), + bidict{ + { + mk_pt_coord(0_n), + mk_device_id(mc), + }, }, - }, }; }; - auto mk_two_shard_mapping = [&](MachineSpaceCoordinate const &mc1, - MachineSpaceCoordinate const &mc2) -> ParallelTensorMapping { + auto mk_two_shard_mapping = + [&](MachineSpaceCoordinate const &mc1, + MachineSpaceCoordinate const &mc2) -> ParallelTensorMapping { return ParallelTensorMapping{ - bidict{ - { - mk_pt_coord(0_n), - mk_device_id(mc1), - }, - { - mk_pt_coord(1_n), - mk_device_id(mc2), + bidict{ + { + mk_pt_coord(0_n), + mk_device_id(mc1), + }, + { + mk_pt_coord(1_n), + mk_device_id(mc2), + }, }, - }, }; }; - DynamicNodeInvocation input_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.input_op_id); - DynamicNodeInvocation relu1_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.relu1_op_id); - DynamicNodeInvocation replicate_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.replicate_op_id); - DynamicNodeInvocation relu2_invocation = dynamic_graph_get_invocation_for_id(tc.g, tc.relu2_op_id); + DynamicNodeInvocation input_invocation = + dynamic_graph_get_invocation_for_id(tc.g, tc.input_op_id); + DynamicNodeInvocation relu1_invocation = + dynamic_graph_get_invocation_for_id(tc.g, tc.relu1_op_id); + DynamicNodeInvocation replicate_invocation = + dynamic_graph_get_invocation_for_id(tc.g, tc.replicate_op_id); + DynamicNodeInvocation relu2_invocation = + dynamic_graph_get_invocation_for_id(tc.g, tc.relu2_op_id); DynamicNodeInvocation value_mapped_input_invocation = - map_output_value( - input_invocation, - mk_single_shard_mapping(mc3)); - + map_output_value(input_invocation, mk_single_shard_mapping(mc3)); DynamicNodeInvocation value_mapped_relu1_invocation = - map_input_and_output_values( - relu1_invocation, - mk_single_shard_mapping(mc1), - mk_single_shard_mapping(mc1)); + map_input_and_output_values(relu1_invocation, + mk_single_shard_mapping(mc1), + mk_single_shard_mapping(mc1)); DynamicNodeInvocation value_mapped_replicate_invocation = - map_input_and_output_values( - replicate_invocation, - mk_single_shard_mapping(mc1), - mk_two_shard_mapping(mc2, mc3)); + map_input_and_output_values(replicate_invocation, + mk_single_shard_mapping(mc1), + mk_two_shard_mapping(mc2, mc3)); DynamicNodeInvocation value_mapped_relu2_invocation = - map_input_and_output_values( - relu2_invocation, - mk_two_shard_mapping(mc1, mc3), - mk_two_shard_mapping(mc1, mc3)); + map_input_and_output_values(relu2_invocation, + mk_two_shard_mapping(mc1, mc3), + mk_two_shard_mapping(mc1, mc3)); DynamicNodeInvocation input_to_relu1_copy = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - require_only_key(value_mapped_input_invocation.outputs, mk_slot(TensorSlotName::OUTPUT)), + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + require_only_key(value_mapped_input_invocation.outputs, + mk_slot(TensorSlotName::OUTPUT)), + }, }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/ + dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - require_only_key(value_mapped_relu1_invocation.inputs, mk_slot(TensorSlotName::INPUT)), + { + mk_slot(TensorSlotName::OUTPUT), + require_only_key(value_mapped_relu1_invocation.inputs, + mk_slot(TensorSlotName::INPUT)), + }, }, - }, }; DynamicNodeInvocation replicate_to_relu2_copy = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - require_only_key(value_mapped_replicate_invocation.outputs, mk_slot(TensorSlotName::OUTPUT)), + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + require_only_key(value_mapped_replicate_invocation.outputs, + mk_slot(TensorSlotName::OUTPUT)), + }, }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, - /*layer_guid=*/dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/TrainingOperationAttrs{CopyAttrs{}}, + /*layer_guid=*/ + dynamic_layer_guid_t{dynamic_copy_layer_guid_t{}}, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - require_only_key(value_mapped_relu2_invocation.inputs, mk_slot(TensorSlotName::INPUT)), + { + mk_slot(TensorSlotName::OUTPUT), + require_only_key(value_mapped_relu2_invocation.inputs, + mk_slot(TensorSlotName::INPUT)), + }, }, - }, }; return dynamic_open_dataflow_graph_from_invocation_set({ - value_mapped_input_invocation, - input_to_relu1_copy, - value_mapped_relu1_invocation, - value_mapped_replicate_invocation, - replicate_to_relu2_copy, - value_mapped_relu2_invocation, + value_mapped_input_invocation, + input_to_relu1_copy, + value_mapped_relu1_invocation, + value_mapped_replicate_invocation, + replicate_to_relu2_copy, + value_mapped_relu2_invocation, }); }(); - nlohmann::json result_json = dynamic_open_dataflow_graph_to_serializable(result); - nlohmann::json correct_json = dynamic_open_dataflow_graph_to_serializable(correct); + nlohmann::json result_json = + dynamic_open_dataflow_graph_to_serializable(result); + nlohmann::json correct_json = + dynamic_open_dataflow_graph_to_serializable(correct); - CHECK_MESSAGE( - result == correct, - check_kv("result\n", result_json.dump()), - check_kv("correct\n", correct_json.dump())); + CHECK_MESSAGE(result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); } SUBCASE("copy insertion commutes with pass expansion") { @@ -1781,82 +1795,98 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicOpenDataflowGraph g = mk_single_input_node_graph(mc1); DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = - perform_copy_insertion(perform_pass_expansion(g)); + perform_copy_insertion(perform_pass_expansion(g)); DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = - perform_pass_expansion(perform_copy_insertion(g)); + perform_pass_expansion(perform_copy_insertion(g)); - nlohmann::json pass_expansion_before_copy_insertion_json - = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); - nlohmann::json copy_insertion_before_pass_expansion_json - = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + nlohmann::json pass_expansion_before_copy_insertion_json = + dynamic_open_dataflow_graph_to_serializable( + pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json = + dynamic_open_dataflow_graph_to_serializable( + copy_insertion_before_pass_expansion); CHECK_MESSAGE( - pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + pass_expansion_before_copy_insertion == + copy_insertion_before_pass_expansion, check_kv("pass_expansion_before_copy_insertion_json\n", pass_expansion_before_copy_insertion_json.dump()), check_kv("copy_insertion_before_pass_expansion_json\n", copy_insertion_before_pass_expansion_json.dump()), check_kv("pass_expansion_before_copy_insertion\n", - dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + dynamic_open_dataflow_graph_as_dot( + pass_expansion_before_copy_insertion)), check_kv("copy_insertion_before_pass_expansion\n", - dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + dynamic_open_dataflow_graph_as_dot( + copy_insertion_before_pass_expansion))); } SUBCASE("graph is single input node followed by relu") { SUBCASE("operations are mapped to the same device") { - DynamicOpenDataflowGraph g = mk_single_input_into_relu_graph(mc1, mc1); + DynamicOpenDataflowGraph g = + mk_single_input_into_relu_graph(mc1, mc1); DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = - perform_copy_insertion(perform_pass_expansion(g)); + perform_copy_insertion(perform_pass_expansion(g)); DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = - perform_pass_expansion(perform_copy_insertion(g)); + perform_pass_expansion(perform_copy_insertion(g)); - nlohmann::json pass_expansion_before_copy_insertion_json - = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); - nlohmann::json copy_insertion_before_pass_expansion_json - = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + nlohmann::json pass_expansion_before_copy_insertion_json = + dynamic_open_dataflow_graph_to_serializable( + pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json = + dynamic_open_dataflow_graph_to_serializable( + copy_insertion_before_pass_expansion); CHECK_MESSAGE( - pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + pass_expansion_before_copy_insertion == + copy_insertion_before_pass_expansion, check_kv("pass_expansion_before_copy_insertion_json\n", pass_expansion_before_copy_insertion_json.dump()), check_kv("copy_insertion_before_pass_expansion_json\n", copy_insertion_before_pass_expansion_json.dump()), check_kv("pass_expansion_before_copy_insertion\n", - dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + dynamic_open_dataflow_graph_as_dot( + pass_expansion_before_copy_insertion)), check_kv("copy_insertion_before_pass_expansion\n", - dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + dynamic_open_dataflow_graph_as_dot( + copy_insertion_before_pass_expansion))); } SUBCASE("operations are mapped to different devices") { - DynamicOpenDataflowGraph g = mk_single_input_into_relu_graph(mc1, mc2); + DynamicOpenDataflowGraph g = + mk_single_input_into_relu_graph(mc1, mc2); DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = - perform_copy_insertion(perform_pass_expansion(g)); + perform_copy_insertion(perform_pass_expansion(g)); DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = - perform_pass_expansion(perform_copy_insertion(g)); + perform_pass_expansion(perform_copy_insertion(g)); - nlohmann::json pass_expansion_before_copy_insertion_json - = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); - nlohmann::json copy_insertion_before_pass_expansion_json - = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + nlohmann::json pass_expansion_before_copy_insertion_json = + dynamic_open_dataflow_graph_to_serializable( + pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json = + dynamic_open_dataflow_graph_to_serializable( + copy_insertion_before_pass_expansion); CHECK_MESSAGE( - pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + pass_expansion_before_copy_insertion == + copy_insertion_before_pass_expansion, check_kv("pass_expansion_before_copy_insertion_json\n", pass_expansion_before_copy_insertion_json.dump()), check_kv("copy_insertion_before_pass_expansion_json\n", copy_insertion_before_pass_expansion_json.dump()), check_kv("pass_expansion_before_copy_insertion\n", - dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + dynamic_open_dataflow_graph_as_dot( + pass_expansion_before_copy_insertion)), check_kv("copy_insertion_before_pass_expansion\n", - dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + dynamic_open_dataflow_graph_as_dot( + copy_insertion_before_pass_expansion))); } } SUBCASE("multinode graph including replicate") { - ExampleGraphTestCase tc - = mk_example_replicate_graph( + ExampleGraphTestCase tc = mk_example_replicate_graph( /*input_device=*/mc3, /*relu1_device=*/mc1, /*replicate_device1=*/mc2, @@ -1865,25 +1895,30 @@ TEST_SUITE(FF_TEST_SUITE) { /*relu2_device2=*/mc3); DynamicOpenDataflowGraph pass_expansion_before_copy_insertion = - perform_copy_insertion(perform_pass_expansion(tc.g)); + perform_copy_insertion(perform_pass_expansion(tc.g)); DynamicOpenDataflowGraph copy_insertion_before_pass_expansion = - perform_pass_expansion(perform_copy_insertion(tc.g)); + perform_pass_expansion(perform_copy_insertion(tc.g)); - nlohmann::json pass_expansion_before_copy_insertion_json - = dynamic_open_dataflow_graph_to_serializable(pass_expansion_before_copy_insertion); - nlohmann::json copy_insertion_before_pass_expansion_json - = dynamic_open_dataflow_graph_to_serializable(copy_insertion_before_pass_expansion); + nlohmann::json pass_expansion_before_copy_insertion_json = + dynamic_open_dataflow_graph_to_serializable( + pass_expansion_before_copy_insertion); + nlohmann::json copy_insertion_before_pass_expansion_json = + dynamic_open_dataflow_graph_to_serializable( + copy_insertion_before_pass_expansion); CHECK_MESSAGE( - pass_expansion_before_copy_insertion == copy_insertion_before_pass_expansion, + pass_expansion_before_copy_insertion == + copy_insertion_before_pass_expansion, check_kv("pass_expansion_before_copy_insertion_json\n", pass_expansion_before_copy_insertion_json.dump()), check_kv("copy_insertion_before_pass_expansion_json\n", copy_insertion_before_pass_expansion_json.dump()), check_kv("pass_expansion_before_copy_insertion\n", - dynamic_open_dataflow_graph_as_dot(pass_expansion_before_copy_insertion)), + dynamic_open_dataflow_graph_as_dot( + pass_expansion_before_copy_insertion)), check_kv("copy_insertion_before_pass_expansion\n", - dynamic_open_dataflow_graph_as_dot(copy_insertion_before_pass_expansion))); + dynamic_open_dataflow_graph_as_dot( + copy_insertion_before_pass_expansion))); } } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc index ce0e710b7f..b5cd5d744a 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/dynamic_open_dataflow_graph.cc @@ -232,24 +232,24 @@ TEST_SUITE(FF_TEST_SUITE) { DynamicNodeInvocation invocation_1 = DynamicNodeInvocation{ /*inputs=*/std::map{ { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - value_1, + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + value_1, }, }, /*node_attrs=*/node_attrs, /*outputs=*/ std::map{ { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - value_2, + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + value_2, }, }, }; @@ -308,8 +308,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicOpenDataflowGraph g = - dynamic_open_dataflow_graph_from_invocation_set( - invocation_set); + dynamic_open_dataflow_graph_from_invocation_set(invocation_set); dynamic_invocation_id_t invocation_1_id = dynamic_graph_get_id_for_invocation(g, invocation_1); @@ -322,11 +321,10 @@ TEST_SUITE(FF_TEST_SUITE) { std::set result = get_dynamic_slot_sites(g); - auto mk_internal_slot_site = [](dynamic_invocation_id_t const &invocation_id, - TensorDirection direction, - TensorSlotName slot_name) - -> DynamicSlotSite - { + auto mk_internal_slot_site = + [](dynamic_invocation_id_t const &invocation_id, + TensorDirection direction, + TensorSlotName slot_name) -> DynamicSlotSite { return DynamicSlotSite{ InternalDynamicSlotSite{ /*invocation_id=*/invocation_id, diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc index 61fe183ac0..ccd8d2abfe 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.cc @@ -1,13 +1,13 @@ -#include #include "task-spec/dynamic_graph/make_dynamic_open_dataflow_graph_from_mapped_pcg.h" -#include "utils/containers/require_only_key.h" +#include "op-attrs/initializer_attrs.h" #include "op-attrs/ops/element_unary.h" #include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" -#include "op-attrs/initializer_attrs.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" #include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h" #include "test/utils/doctest/check_kv.h" -#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" +#include "utils/containers/require_only_key.h" +#include using namespace ::FlexFlow; @@ -17,47 +17,50 @@ TEST_SUITE(FF_TEST_SUITE) { MachineSpaceCoordinate gpu1 = MachineSpaceCoordinate{0_n, 1_n}; SUBCASE("Replicate") { - ParallelTensorSpaceCoordinate tensor_coord0 = ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/0_n, - /*shard_component=*/FFOrdered{0_n}, - }; + ParallelTensorSpaceCoordinate tensor_coord0 = + ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{0_n}, + }; - ParallelTensorSpaceCoordinate tensor_coord1 = ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/1_n, - /*shard_component=*/FFOrdered{0_n}, - }; + ParallelTensorSpaceCoordinate tensor_coord1 = + ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/1_n, + /*shard_component=*/FFOrdered{0_n}, + }; MappedOperatorTaskGroup mapped_op_task_group = MappedOperatorTaskGroup{ - { - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, }, - }, }; ParallelTensorShape input_shape = ParallelTensorShape{ - /*dims=*/ParallelTensorDims{ - /*shard_dims=*/FFOrdered{ - ShardParallelDim{8_p, 2_p}, - ShardParallelDim{5_p, 1_p}, - }, - /*replica_dims=*/ReplicaParallelDimSet{ - SumDegree{1_p}, - DiscardCopyDegree{1_p}, + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{8_p, 2_p}, + ShardParallelDim{5_p, 1_p}, + }, + /*replica_dims=*/ + ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{1_p}, + }, }, - }, - /*data_type=*/DataType::FLOAT, + /*data_type=*/DataType::FLOAT, }; ParallelTensorShape output_shape = [&] { @@ -67,114 +70,125 @@ TEST_SUITE(FF_TEST_SUITE) { }(); PCGOperatorAttrs op_attrs = PCGOperatorAttrs{ - ReplicateAttrs{ - 2_p, - }, + ReplicateAttrs{ + 2_p, + }, }; parallel_layer_guid_t layer_guid = parallel_layer_guid_t{Node{0}}; parallel_tensor_guid_t input_tensor_guid = parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{5}, - TensorSlotName::OUTPUT, - }, + KwargDataflowOutput{ + Node{5}, + TensorSlotName::OUTPUT, + }, }; parallel_tensor_guid_t output_tensor_guid = parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{0}, - TensorSlotName::OUTPUT, - }, + KwargDataflowOutput{ + Node{0}, + TensorSlotName::OUTPUT, + }, }; - MappedParallelLayerInvocationInfo input = MappedParallelLayerInvocationInfo{ - /*incoming=*/{ - { - TensorSlotName::INPUT, - ParallelTensorInfo{ - /*guid=*/input_tensor_guid, - /*attrs=*/ParallelTensorAttrs{ - /*shape=*/input_shape, - /*create_grad=*/CreateGrad::YES, + MappedParallelLayerInvocationInfo input = + MappedParallelLayerInvocationInfo{ + /*incoming=*/{ + { + TensorSlotName::INPUT, + ParallelTensorInfo{ + /*guid=*/input_tensor_guid, + /*attrs=*/ + ParallelTensorAttrs{ + /*shape=*/input_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, }, - }, - }, - }, - /*layer_info=*/MappedParallelLayerInfo{ - /*guid=*/layer_guid, - /*attrs=*/ParallelLayerAttrs{ - /*op_attrs=*/op_attrs, - /*name=*/std::nullopt, - }, - /*mapping=*/mapped_op_task_group, - }, - /*outgoing=*/{ - { - TensorSlotName::OUTPUT, - ParallelTensorInfo{ - /*guid=*/output_tensor_guid, - /*attrs=*/ParallelTensorAttrs{ - /*shape=*/output_shape, - /*create_grad=*/CreateGrad::YES, + /*layer_info=*/ + MappedParallelLayerInfo{ + /*guid=*/layer_guid, + /*attrs=*/ + ParallelLayerAttrs{ + /*op_attrs=*/op_attrs, + /*name=*/std::nullopt, + }, + /*mapping=*/mapped_op_task_group, }, - }, - }, - }, - }; + /*outgoing=*/ + { + { + TensorSlotName::OUTPUT, + ParallelTensorInfo{ + /*guid=*/output_tensor_guid, + /*attrs=*/ + ParallelTensorAttrs{ + /*shape=*/output_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + }, + }; - DynamicNodeInvocation result = make_dynamic_node_invocation_from_mapped(input, DeviceType::GPU); + DynamicNodeInvocation result = + make_dynamic_node_invocation_from_mapped(input, DeviceType::GPU); DynamicNodeInvocation correct = DynamicNodeInvocation{ - /*inputs=*/{ - { - DynamicTensorSlot{ - TensorSlotName::INPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, - /*parallel_tensor_shape=*/input_shape, - /*create_grad=*/true, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }, + /*inputs=*/{ + { + DynamicTensorSlot{ + TensorSlotName::INPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{input_tensor_guid}, + /*parallel_tensor_shape=*/input_shape, + /*create_grad=*/true, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }, + }, }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/DynamicNodeMapping{ - /*op_task_group=*/mapped_op_task_group, - /*device_type=*/DeviceType::GPU, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/ + DynamicNodeMapping{ + /*op_task_group=*/mapped_op_task_group, + /*device_type=*/DeviceType::GPU, + }, + /*op_attrs=*/TrainingOperationAttrs{op_attrs}, + /*layer_guid=*/dynamic_layer_guid_t{layer_guid}, + /*per_device_op_state=*/std::nullopt, }, - /*op_attrs=*/TrainingOperationAttrs{op_attrs}, - /*layer_guid=*/dynamic_layer_guid_t{layer_guid}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ + /*outputs=*/ { - DynamicTensorSlot{ - TensorSlotName::OUTPUT, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, - /*parallel_tensor_shape=*/output_shape, - /*create_grad=*/true, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, - }, - }, - } - }; + { + DynamicTensorSlot{ + TensorSlotName::OUTPUT, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrs{ + /*tensor_guid=*/dynamic_tensor_guid_t{output_tensor_guid}, + /*parallel_tensor_shape=*/output_shape, + /*create_grad=*/true, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, + }, + }, + }}; - nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); - nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + nlohmann::json result_json = + dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = + dynamic_node_invocation_to_serializable(correct); CHECK_MESSAGE(result == correct, check_kv("result\n", result_json.dump()), @@ -183,25 +197,23 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("standard op") { PCGOperatorAttrs op_attrs = PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/7_p, - /*use_bias=*/true, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, - }, + LinearAttrs{ + /*out_channels=*/7_p, + /*use_bias=*/true, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }; parallel_layer_guid_t layer_guid = parallel_layer_guid_t{Node{0}}; - auto mk_tensor_guid = [](size_t node_id) - -> parallel_tensor_guid_t - { + auto mk_tensor_guid = [](size_t node_id) -> parallel_tensor_guid_t { return parallel_tensor_guid_t{ - KwargDataflowOutput{ - Node{node_id}, - TensorSlotName::OUTPUT, - }, + KwargDataflowOutput{ + Node{node_id}, + TensorSlotName::OUTPUT, + }, }; }; @@ -211,230 +223,244 @@ TEST_SUITE(FF_TEST_SUITE) { parallel_tensor_guid_t output_tensor_guid = mk_tensor_guid(0); ParallelTensorShape input_shape = ParallelTensorShape{ - /*dims=*/ParallelTensorDims{ - /*shard_dims=*/FFOrdered{ - ShardParallelDim{8_p, 2_p}, - ShardParallelDim{5_p, 1_p}, - }, - /*replica_dims=*/ReplicaParallelDimSet{ - SumDegree{1_p}, - DiscardCopyDegree{1_p}, + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{8_p, 2_p}, + ShardParallelDim{5_p, 1_p}, + }, + /*replica_dims=*/ + ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{1_p}, + }, }, - }, - /*data_type=*/DataType::FLOAT, + /*data_type=*/DataType::FLOAT, }; ParallelTensorShape weight_shape = ParallelTensorShape{ - /*dims=*/ParallelTensorDims{ - /*shard_dims=*/FFOrdered{ - ShardParallelDim{5_p, 1_p}, - ShardParallelDim{7_p, 1_p}, - }, - /*replica_dims=*/ReplicaParallelDimSet{ - SumDegree{1_p}, - DiscardCopyDegree{2_p}, + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{5_p, 1_p}, + ShardParallelDim{7_p, 1_p}, + }, + /*replica_dims=*/ + ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{2_p}, + }, }, - }, - /*data_type=*/DataType::FLOAT, + /*data_type=*/DataType::FLOAT, }; ParallelTensorShape bias_shape = ParallelTensorShape{ - /*dims=*/ParallelTensorDims{ - /*shard_dims=*/FFOrdered{ - ShardParallelDim{7_p, 1_p}, - }, - /*replica_dims=*/ReplicaParallelDimSet{ - SumDegree{1_p}, - DiscardCopyDegree{2_p}, + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{7_p, 1_p}, + }, + /*replica_dims=*/ + ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{2_p}, + }, }, - }, - /*data_type=*/DataType::FLOAT, + /*data_type=*/DataType::FLOAT, }; ParallelTensorShape output_shape = ParallelTensorShape{ - /*dims=*/ParallelTensorDims{ - /*shard_dims=*/FFOrdered{ - ShardParallelDim{8_p, 2_p}, - ShardParallelDim{7_p, 1_p}, - }, - /*replica_dims=*/ReplicaParallelDimSet{ - SumDegree{1_p}, - DiscardCopyDegree{1_p}, + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{8_p, 2_p}, + ShardParallelDim{7_p, 1_p}, + }, + /*replica_dims=*/ + ReplicaParallelDimSet{ + SumDegree{1_p}, + DiscardCopyDegree{1_p}, + }, }, - }, - /*data_type=*/DataType::FLOAT, + /*data_type=*/DataType::FLOAT, }; - auto mk_2d_pt_coord = [](nonnegative_int replica_coord, nonnegative_int shard_coord) - -> ParallelTensorSpaceCoordinate - { + auto mk_2d_pt_coord = + [](nonnegative_int replica_coord, + nonnegative_int shard_coord) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_compnent=*/replica_coord, - /*shard_components=*/FFOrdered{ - shard_coord, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_compnent=*/replica_coord, + /*shard_components=*/ + FFOrdered{ + shard_coord, + 0_n, + }, }; }; - auto mk_1d_pt_coord = [](nonnegative_int replica_coord, nonnegative_int shard_coord) - -> ParallelTensorSpaceCoordinate - { + auto mk_1d_pt_coord = + [](nonnegative_int replica_coord, + nonnegative_int shard_coord) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_compnent=*/replica_coord, - /*shard_components=*/FFOrdered{ - shard_coord, - }, + /*sum_component=*/0_n, + /*discard_copy_compnent=*/replica_coord, + /*shard_components=*/ + FFOrdered{ + shard_coord, + }, }; }; MappedOperatorTaskGroup mapped_op_task_group = MappedOperatorTaskGroup{ - { - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, mk_2d_pt_coord(0_n, 0_n)}, - {TensorSlotName::WEIGHT, mk_2d_pt_coord(0_n, 0_n)}, - {TensorSlotName::BIAS, mk_1d_pt_coord(0_n, 0_n)}, - {TensorSlotName::OUTPUT, mk_2d_pt_coord(0_n, 0_n)}, - }}, - }, { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, mk_2d_pt_coord(0_n, 1_n)}, - {TensorSlotName::WEIGHT, mk_2d_pt_coord(1_n, 0_n)}, - {TensorSlotName::BIAS, mk_1d_pt_coord(1_n, 0_n)}, - {TensorSlotName::OUTPUT, mk_2d_pt_coord(0_n, 1_n)}, - }}, + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_2d_pt_coord(0_n, 0_n)}, + {TensorSlotName::WEIGHT, mk_2d_pt_coord(0_n, 0_n)}, + {TensorSlotName::BIAS, mk_1d_pt_coord(0_n, 0_n)}, + {TensorSlotName::OUTPUT, mk_2d_pt_coord(0_n, 0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_2d_pt_coord(0_n, 1_n)}, + {TensorSlotName::WEIGHT, mk_2d_pt_coord(1_n, 0_n)}, + {TensorSlotName::BIAS, mk_1d_pt_coord(1_n, 0_n)}, + {TensorSlotName::OUTPUT, mk_2d_pt_coord(0_n, 1_n)}, + }}, + }, }, - }, }; - MappedParallelLayerInvocationInfo input = MappedParallelLayerInvocationInfo{ - /*incoming=*/{ - { - TensorSlotName::INPUT, - ParallelTensorInfo{ - /*guid=*/input_tensor_guid, - /*attrs=*/ParallelTensorAttrs{ - /*shape=*/input_shape, - /*create_grad=*/CreateGrad::YES, + MappedParallelLayerInvocationInfo input = + MappedParallelLayerInvocationInfo{ + /*incoming=*/{ + { + TensorSlotName::INPUT, + ParallelTensorInfo{ + /*guid=*/input_tensor_guid, + /*attrs=*/ + ParallelTensorAttrs{ + /*shape=*/input_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + { + TensorSlotName::WEIGHT, + ParallelTensorInfo{ + /*guid=*/weight_tensor_guid, + /*attrs=*/ + ParallelTensorAttrs{ + /*shape=*/weight_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, + { + TensorSlotName::BIAS, + ParallelTensorInfo{ + /*guid=*/bias_tensor_guid, + /*attrs=*/ + ParallelTensorAttrs{ + /*shape=*/bias_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, }, - }, - }, - { - TensorSlotName::WEIGHT, - ParallelTensorInfo{ - /*guid=*/weight_tensor_guid, - /*attrs=*/ParallelTensorAttrs{ - /*shape=*/weight_shape, - /*create_grad=*/CreateGrad::YES, - }, - }, - }, - { - TensorSlotName::BIAS, - ParallelTensorInfo{ - /*guid=*/bias_tensor_guid, - /*attrs=*/ParallelTensorAttrs{ - /*shape=*/bias_shape, - /*create_grad=*/CreateGrad::YES, + /*layer_info=*/ + MappedParallelLayerInfo{ + /*guid=*/layer_guid, + /*attrs=*/ + ParallelLayerAttrs{ + /*op_attrs=*/op_attrs, + /*name=*/std::nullopt, + }, + /*mapping=*/mapped_op_task_group, }, - }, - }, - }, - /*layer_info=*/MappedParallelLayerInfo{ - /*guid=*/layer_guid, - /*attrs=*/ParallelLayerAttrs{ - /*op_attrs=*/op_attrs, - /*name=*/std::nullopt, - }, - /*mapping=*/mapped_op_task_group, - }, - /*outgoing=*/{ - { - TensorSlotName::OUTPUT, - ParallelTensorInfo{ - /*guid=*/output_tensor_guid, - /*attrs=*/ParallelTensorAttrs{ - /*shape=*/output_shape, - /*create_grad=*/CreateGrad::YES, + /*outgoing=*/ + { + { + TensorSlotName::OUTPUT, + ParallelTensorInfo{ + /*guid=*/output_tensor_guid, + /*attrs=*/ + ParallelTensorAttrs{ + /*shape=*/output_shape, + /*create_grad=*/CreateGrad::YES, + }, + }, + }, }, - }, - }, - }, - }; + }; - DynamicNodeInvocation result = make_dynamic_node_invocation_from_mapped(input, DeviceType::GPU); + DynamicNodeInvocation result = + make_dynamic_node_invocation_from_mapped(input, DeviceType::GPU); - DynamicNodeInvocation correct = [&]() - -> DynamicNodeInvocation - { - auto mk_slot = [](TensorSlotName slot_name) - -> DynamicTensorSlot - { + DynamicNodeInvocation correct = [&]() -> DynamicNodeInvocation { + auto mk_slot = [](TensorSlotName slot_name) -> DynamicTensorSlot { return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }; }; - auto mk_value = [](parallel_tensor_guid_t const &tensor_guid, - ParallelTensorShape const &shape) - -> DynamicValueAttrs - { + auto mk_value = + [](parallel_tensor_guid_t const &tensor_guid, + ParallelTensorShape const &shape) -> DynamicValueAttrs { return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor_guid}, - /*parallel_tensor_shape=*/shape, - /*create_grad=*/true, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, + /*tensor_guid=*/dynamic_tensor_guid_t{tensor_guid}, + /*parallel_tensor_shape=*/shape, + /*create_grad=*/true, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, }; }; return DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(input_tensor_guid, input_shape), + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_tensor_guid, input_shape), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(weight_tensor_guid, weight_shape), + }, + { + mk_slot(TensorSlotName::BIAS), + mk_value(bias_tensor_guid, bias_shape), + }, }, - { - mk_slot(TensorSlotName::WEIGHT), - mk_value(weight_tensor_guid, weight_shape), + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/ + DynamicNodeMapping{ + /*op_task_group=*/mapped_op_task_group, + /*device_type=*/DeviceType::GPU, + }, + /*op_attrs=*/TrainingOperationAttrs{op_attrs}, + /*layer_guid=*/dynamic_layer_guid_t{layer_guid}, + /*per_device_op_state=*/std::nullopt, }, + /*outputs=*/ { - mk_slot(TensorSlotName::BIAS), - mk_value(bias_tensor_guid, bias_shape), - }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/DynamicNodeMapping{ - /*op_task_group=*/mapped_op_task_group, - /*device_type=*/DeviceType::GPU, - }, - /*op_attrs=*/TrainingOperationAttrs{op_attrs}, - /*layer_guid=*/dynamic_layer_guid_t{layer_guid}, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(output_tensor_guid, output_shape), - }, - } - }; + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(output_tensor_guid, output_shape), + }, + }}; }(); - nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); - nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + nlohmann::json result_json = + dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = + dynamic_node_invocation_to_serializable(correct); CHECK_MESSAGE(result == correct, check_kv("result\n", result_json.dump()), @@ -448,7 +474,8 @@ TEST_SUITE(FF_TEST_SUITE) { positive_int hidden_dim = 32_p; positive_int output_dim = 1_p; - auto make_layer_attrs = [](PCGOperatorAttrs const &op_attrs) -> ParallelLayerAttrs { + auto make_layer_attrs = + [](PCGOperatorAttrs const &op_attrs) -> ParallelLayerAttrs { return ParallelLayerAttrs{ /*op_attrs=*/op_attrs, /*name=*/std::nullopt, @@ -473,62 +500,60 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelComputationGraph pcg = empty_parallel_computation_graph(); PCGOperatorAttrs input_op_attrs = PCGOperatorAttrs{ - InputAttrs{ - /*tensor_shape=*/input_tensor_shape - }, + InputAttrs{/*tensor_shape=*/input_tensor_shape}, }; PCGOperatorAttrs partition_input_op_attrs = PCGOperatorAttrs{ - RepartitionAttrs{ - /*repartition_dim=*/ff_dim_t{0_n}, - /*repartition_degree=*/2_p, - }, + RepartitionAttrs{ + /*repartition_dim=*/ff_dim_t{0_n}, + /*repartition_degree=*/2_p, + }, }; PCGOperatorAttrs weight_1_op_attrs = PCGOperatorAttrs{ - WeightAttrs{ - /*tensor_shape=*/weight_1_tensor_shape, - /*initializer=*/make_kaiming_uniform(weight_1_tensor_shape.dims), - }, + WeightAttrs{ + /*tensor_shape=*/weight_1_tensor_shape, + /*initializer=*/make_kaiming_uniform(weight_1_tensor_shape.dims), + }, }; PCGOperatorAttrs replicate_weight_1_op_attrs = PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, - }, + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }; PCGOperatorAttrs weight_2_op_attrs = PCGOperatorAttrs{ - WeightAttrs{ - /*tensor_shape=*/weight_2_tensor_shape, - /*initializer=*/make_kaiming_uniform(weight_1_tensor_shape.dims), - }, + WeightAttrs{ + /*tensor_shape=*/weight_2_tensor_shape, + /*initializer=*/make_kaiming_uniform(weight_1_tensor_shape.dims), + }, }; PCGOperatorAttrs replicate_weight_2_op_attrs = PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, - }, + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }; PCGOperatorAttrs linear_1_op_attrs = PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/hidden_dim, - /*use_bias=*/false, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, - }, + LinearAttrs{ + /*out_channels=*/hidden_dim, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }; PCGOperatorAttrs linear_2_op_attrs = PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/output_dim, - /*use_bias=*/false, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, - }, + LinearAttrs{ + /*out_channels=*/output_dim, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }; ParallelLayerAddedResult input_layer = @@ -537,17 +562,18 @@ TEST_SUITE(FF_TEST_SUITE) { require_only_key(input_layer.outputs, TensorSlotName::OUTPUT); ParallelLayerAddedResult partition_input_layer = - add_parallel_layer(pcg, + add_parallel_layer(pcg, make_layer_attrs(partition_input_op_attrs), - /*inputs=*/{ - {TensorSlotName::INPUT, t_input}, + /*inputs=*/ + { + {TensorSlotName::INPUT, t_input}, }, /*weights=*/{}); parallel_tensor_guid_t t_partitioned_input = require_only_key(partition_input_layer.outputs, TensorSlotName::OUTPUT); ParallelLayerAddedResult weight_1_layer = - add_parallel_layer(pcg, + add_parallel_layer(pcg, make_layer_attrs(weight_1_op_attrs), /*inputs=*/{}, /*weights=*/{}); @@ -555,17 +581,18 @@ TEST_SUITE(FF_TEST_SUITE) { require_only_key(weight_1_layer.outputs, TensorSlotName::OUTPUT); ParallelLayerAddedResult replicate_weight_1_layer = - add_parallel_layer(pcg, + add_parallel_layer(pcg, make_layer_attrs(replicate_weight_1_op_attrs), - /*inputs=*/{ - {TensorSlotName::INPUT, t_weight_1}, + /*inputs=*/ + { + {TensorSlotName::INPUT, t_weight_1}, }, /*weights=*/{}); - parallel_tensor_guid_t t_replicated_weight_1 = - require_only_key(replicate_weight_1_layer.outputs, TensorSlotName::OUTPUT); + parallel_tensor_guid_t t_replicated_weight_1 = require_only_key( + replicate_weight_1_layer.outputs, TensorSlotName::OUTPUT); ParallelLayerAddedResult weight_2_layer = - add_parallel_layer(pcg, + add_parallel_layer(pcg, make_layer_attrs(weight_2_op_attrs), /*inputs=*/{}, /*weights=*/{}); @@ -573,493 +600,513 @@ TEST_SUITE(FF_TEST_SUITE) { require_only_key(weight_2_layer.outputs, TensorSlotName::OUTPUT); ParallelLayerAddedResult replicate_weight_2_layer = - add_parallel_layer(pcg, + add_parallel_layer(pcg, make_layer_attrs(replicate_weight_2_op_attrs), - /*inputs=*/{ - {TensorSlotName::INPUT, t_weight_2}, + /*inputs=*/ + { + {TensorSlotName::INPUT, t_weight_2}, }, /*weights=*/{}); - parallel_tensor_guid_t t_replicated_weight_2 = - require_only_key(replicate_weight_2_layer.outputs, TensorSlotName::OUTPUT); + parallel_tensor_guid_t t_replicated_weight_2 = require_only_key( + replicate_weight_2_layer.outputs, TensorSlotName::OUTPUT); ParallelLayerAddedResult linear_1_layer = - add_parallel_layer(pcg, + add_parallel_layer(pcg, make_layer_attrs(linear_1_op_attrs), - /*inputs=*/{ - {TensorSlotName::INPUT, t_partitioned_input}, + /*inputs=*/ + { + {TensorSlotName::INPUT, t_partitioned_input}, }, - /*weights=*/{ - {TensorSlotName::WEIGHT, t_replicated_weight_1}, + /*weights=*/ + { + {TensorSlotName::WEIGHT, t_replicated_weight_1}, }); parallel_tensor_guid_t t_hidden_activation = require_only_key(linear_1_layer.outputs, TensorSlotName::OUTPUT); ParallelLayerAddedResult linear_2_layer = - add_parallel_layer(pcg, + add_parallel_layer(pcg, make_layer_attrs(linear_2_op_attrs), - /*inputs=*/{ - {TensorSlotName::INPUT, t_hidden_activation}, + /*inputs=*/ + { + {TensorSlotName::INPUT, t_hidden_activation}, }, - /*weights=*/{ - {TensorSlotName::WEIGHT, t_replicated_weight_2}, + /*weights=*/ + { + {TensorSlotName::WEIGHT, t_replicated_weight_2}, }); parallel_tensor_guid_t t_output = require_only_key(linear_2_layer.outputs, TensorSlotName::OUTPUT); - - MachineSpaceCoordinate gpu0 = MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, + /*node_idx=*/0_n, + /*device_idx=*/0_n, }; MachineSpaceCoordinate gpu1 = MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/1_n, + /*node_idx=*/0_n, + /*device_idx=*/1_n, }; - auto mk_shard_coord = [](nonnegative_int shard_idx) - -> ParallelTensorSpaceCoordinate - { + auto mk_shard_coord = + [](nonnegative_int shard_idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/0_n, - /*shard_component=*/FFOrdered{shard_idx, 0_n}, + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{shard_idx, 0_n}, }; }; - auto mk_replica_coord = [](nonnegative_int replica_idx) - -> ParallelTensorSpaceCoordinate - { + auto mk_replica_coord = + [](nonnegative_int replica_idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/replica_idx, - /*shard_component=*/FFOrdered{0_n, 0_n}, + /*sum_component=*/0_n, + /*discard_copy_component=*/replica_idx, + /*shard_component=*/FFOrdered{0_n, 0_n}, }; }; - ParallelTensorSpaceCoordinate nonparallel_coord = ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/0_n, - /*shard_component=*/FFOrdered{0_n, 0_n}, - }; + ParallelTensorSpaceCoordinate nonparallel_coord = + ParallelTensorSpaceCoordinate{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{0_n, 0_n}, + }; MappedOperatorTaskGroup input_op_mapping = MappedOperatorTaskGroup{ - { { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, nonparallel_coord}, - }}, + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, nonparallel_coord}, + }}, + }, }, - }, }; MappedOperatorTaskGroup weight_1_op_mapping = MappedOperatorTaskGroup{ - { { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, nonparallel_coord}, - }}, + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, nonparallel_coord}, + }}, + }, }, - }, }; MappedOperatorTaskGroup weight_2_op_mapping = MappedOperatorTaskGroup{ - { { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, nonparallel_coord}, - }}, + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, nonparallel_coord}, + }}, + }, }, - }, }; - MappedOperatorTaskGroup partition_input_op_mapping = MappedOperatorTaskGroup{ - { - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, nonparallel_coord}, - {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, - }}, - }, - { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, nonparallel_coord}, - {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, - }}, - }, - }, - }; + MappedOperatorTaskGroup partition_input_op_mapping = + MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, + }}, + }, + }, + }; - MappedOperatorTaskGroup replicate_weight_1_op_mapping = MappedOperatorTaskGroup{ - { - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, nonparallel_coord}, - {TensorSlotName::OUTPUT, mk_replica_coord(0_n)}, - }}, - }, - { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, nonparallel_coord}, - {TensorSlotName::OUTPUT, mk_replica_coord(1_n)}, - }}, - }, - }, - }; + MappedOperatorTaskGroup replicate_weight_1_op_mapping = + MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(1_n)}, + }}, + }, + }, + }; - MappedOperatorTaskGroup replicate_weight_2_op_mapping = MappedOperatorTaskGroup{ - { - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, nonparallel_coord}, - {TensorSlotName::OUTPUT, mk_replica_coord(0_n)}, - }}, - }, - { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, nonparallel_coord}, - {TensorSlotName::OUTPUT, mk_replica_coord(1_n)}, - }}, - }, - }, - }; + MappedOperatorTaskGroup replicate_weight_2_op_mapping = + MappedOperatorTaskGroup{ + { + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, nonparallel_coord}, + {TensorSlotName::OUTPUT, mk_replica_coord(1_n)}, + }}, + }, + }, + }; MappedOperatorTaskGroup linear_1_op_mapping = MappedOperatorTaskGroup{ - { - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, mk_shard_coord(0_n)}, - {TensorSlotName::WEIGHT, mk_replica_coord(0_n)}, - {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, - }}, - }, { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, mk_shard_coord(1_n)}, - {TensorSlotName::WEIGHT, mk_replica_coord(1_n)}, - {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, - }}, + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(0_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(0_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(1_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(1_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, + }}, + }, }, - }, }; MappedOperatorTaskGroup linear_2_op_mapping = MappedOperatorTaskGroup{ - { - { - gpu0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, mk_shard_coord(0_n)}, - {TensorSlotName::WEIGHT, mk_replica_coord(0_n)}, - {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, - }}, - }, { - gpu1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, mk_shard_coord(1_n)}, - {TensorSlotName::WEIGHT, mk_replica_coord(1_n)}, - {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, - }}, + { + gpu0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(0_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(0_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(0_n)}, + }}, + }, + { + gpu1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, mk_shard_coord(1_n)}, + {TensorSlotName::WEIGHT, mk_replica_coord(1_n)}, + {TensorSlotName::OUTPUT, mk_shard_coord(1_n)}, + }}, + }, }, - }, }; - MappedParallelComputationGraph mpcg = mapped_pcg_from_pcg_and_mapped_op_task_groups( - /*pcg=*/pcg, - /*mapped_op_task_groups=*/{ - { - input_layer.parallel_layer, - input_op_mapping, - }, - { - weight_1_layer.parallel_layer, - weight_1_op_mapping, - }, - { - weight_2_layer.parallel_layer, - weight_2_op_mapping, - }, - { - partition_input_layer.parallel_layer, - partition_input_op_mapping, - }, - { - replicate_weight_1_layer.parallel_layer, - replicate_weight_1_op_mapping, - }, - { - replicate_weight_2_layer.parallel_layer, - replicate_weight_2_op_mapping, - }, - { - linear_1_layer.parallel_layer, - linear_1_op_mapping, - }, - { - linear_2_layer.parallel_layer, - linear_2_op_mapping, - }, - }); - - - DynamicOpenDataflowGraph result = make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg, DeviceType::GPU); - - DynamicOpenDataflowGraph correct = [&]() - -> DynamicOpenDataflowGraph - { - auto mk_pt_shape = [](positive_int sum_degree, - positive_int discard_copy_degree, - positive_int shard_dim_size_0, - positive_int shard_degree_0, - positive_int shard_dim_size_1, - positive_int shard_degree_1) - -> ParallelTensorShape - { + MappedParallelComputationGraph mpcg = + mapped_pcg_from_pcg_and_mapped_op_task_groups( + /*pcg=*/pcg, + /*mapped_op_task_groups=*/{ + { + input_layer.parallel_layer, + input_op_mapping, + }, + { + weight_1_layer.parallel_layer, + weight_1_op_mapping, + }, + { + weight_2_layer.parallel_layer, + weight_2_op_mapping, + }, + { + partition_input_layer.parallel_layer, + partition_input_op_mapping, + }, + { + replicate_weight_1_layer.parallel_layer, + replicate_weight_1_op_mapping, + }, + { + replicate_weight_2_layer.parallel_layer, + replicate_weight_2_op_mapping, + }, + { + linear_1_layer.parallel_layer, + linear_1_op_mapping, + }, + { + linear_2_layer.parallel_layer, + linear_2_op_mapping, + }, + }); + + DynamicOpenDataflowGraph result = + make_dynamic_open_dataflow_graph_from_mapped_pcg(mpcg, DeviceType::GPU); + + DynamicOpenDataflowGraph correct = [&]() -> DynamicOpenDataflowGraph { + auto mk_pt_shape = + [](positive_int sum_degree, + positive_int discard_copy_degree, + positive_int shard_dim_size_0, + positive_int shard_degree_0, + positive_int shard_dim_size_1, + positive_int shard_degree_1) -> ParallelTensorShape { return ParallelTensorShape{ - /*dims=*/ParallelTensorDims{ - /*shard_dims=*/FFOrdered{ - ShardParallelDim{shard_dim_size_0, shard_degree_0}, - ShardParallelDim{shard_dim_size_1, shard_degree_1}, + /*dims=*/ParallelTensorDims{ + /*shard_dims=*/FFOrdered{ + ShardParallelDim{shard_dim_size_0, shard_degree_0}, + ShardParallelDim{shard_dim_size_1, shard_degree_1}, + }, + /*replica_dims=*/ + ReplicaParallelDimSet{ + /*sum_degree=*/SumDegree{sum_degree}, + /*discard_copy_degree=*/ + DiscardCopyDegree{discard_copy_degree}, + }, }, - /*replica_dims=*/ReplicaParallelDimSet{ - /*sum_degree=*/SumDegree{sum_degree}, - /*discard_copy_degree=*/DiscardCopyDegree{discard_copy_degree}, - }, - }, - /*data_type=*/DataType::FLOAT, + /*data_type=*/DataType::FLOAT, }; }; - ParallelTensorShape input_pt_shape - = mk_pt_shape(1_p, 1_p, batch_size, 1_p, data_dim, 1_p); - ParallelTensorShape weight_1_pt_shape - = mk_pt_shape(1_p, 1_p, hidden_dim, 1_p, data_dim, 1_p); - ParallelTensorShape weight_2_pt_shape - = mk_pt_shape(1_p, 1_p, output_dim, 1_p, hidden_dim, 1_p); - ParallelTensorShape partitioned_input_pt_shape - = mk_pt_shape(1_p, 1_p, batch_size, 2_p, data_dim, 1_p); - ParallelTensorShape replicated_weight_1_pt_shape - = mk_pt_shape(1_p, 2_p, hidden_dim, 1_p, data_dim, 1_p); - ParallelTensorShape replicated_weight_2_pt_shape - = mk_pt_shape(1_p, 2_p, output_dim, 1_p, hidden_dim, 1_p); - ParallelTensorShape hidden_activation_pt_shape - = mk_pt_shape(1_p, 1_p, batch_size, 2_p, hidden_dim, 1_p); - ParallelTensorShape output_pt_shape - = mk_pt_shape(1_p, 1_p, batch_size, 2_p, output_dim, 1_p); - - auto mk_node_attrs = [&](parallel_layer_guid_t const &layer_guid, - PCGOperatorAttrs const &op_attrs, - MappedOperatorTaskGroup const &mapped_op_task_group) - -> DynamicNodeAttrs - { + ParallelTensorShape input_pt_shape = + mk_pt_shape(1_p, 1_p, batch_size, 1_p, data_dim, 1_p); + ParallelTensorShape weight_1_pt_shape = + mk_pt_shape(1_p, 1_p, hidden_dim, 1_p, data_dim, 1_p); + ParallelTensorShape weight_2_pt_shape = + mk_pt_shape(1_p, 1_p, output_dim, 1_p, hidden_dim, 1_p); + ParallelTensorShape partitioned_input_pt_shape = + mk_pt_shape(1_p, 1_p, batch_size, 2_p, data_dim, 1_p); + ParallelTensorShape replicated_weight_1_pt_shape = + mk_pt_shape(1_p, 2_p, hidden_dim, 1_p, data_dim, 1_p); + ParallelTensorShape replicated_weight_2_pt_shape = + mk_pt_shape(1_p, 2_p, output_dim, 1_p, hidden_dim, 1_p); + ParallelTensorShape hidden_activation_pt_shape = + mk_pt_shape(1_p, 1_p, batch_size, 2_p, hidden_dim, 1_p); + ParallelTensorShape output_pt_shape = + mk_pt_shape(1_p, 1_p, batch_size, 2_p, output_dim, 1_p); + + auto mk_node_attrs = + [&](parallel_layer_guid_t const &layer_guid, + PCGOperatorAttrs const &op_attrs, + MappedOperatorTaskGroup const &mapped_op_task_group) + -> DynamicNodeAttrs { return DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_coord=*/std::nullopt, - /*mapping=*/DynamicNodeMapping{ - mapped_op_task_group, - DeviceType::GPU, - }, - /*op_attrs=*/TrainingOperationAttrs{op_attrs}, - /*pcg_layer_guid=*/dynamic_layer_guid_t{layer_guid}, - /*per_device_op_state=*/std::nullopt, + /*task_type=*/std::nullopt, + /*device_coord=*/std::nullopt, + /*mapping=*/ + DynamicNodeMapping{ + mapped_op_task_group, + DeviceType::GPU, + }, + /*op_attrs=*/TrainingOperationAttrs{op_attrs}, + /*pcg_layer_guid=*/dynamic_layer_guid_t{layer_guid}, + /*per_device_op_state=*/std::nullopt, }; }; - auto mk_slot = [&](TensorSlotName slot_name) - -> DynamicTensorSlot - { + auto mk_slot = [&](TensorSlotName slot_name) -> DynamicTensorSlot { return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }; }; auto mk_value = [&](parallel_tensor_guid_t const &tensor_guid, ParallelTensorShape const &shape, - bool create_grad = true) - -> DynamicValueAttrs - { + bool create_grad = true) -> DynamicValueAttrs { return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{tensor_guid}, - /*parallel_tensor_shape=*/shape, - /*create_grad=*/create_grad, - /*shard_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, + /*tensor_guid=*/dynamic_tensor_guid_t{tensor_guid}, + /*parallel_tensor_shape=*/shape, + /*create_grad=*/create_grad, + /*shard_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, }; }; DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{}, - /*node=*/mk_node_attrs(input_layer.parallel_layer, - input_op_attrs, - input_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_input, input_pt_shape, /*create_grad=*/false), + /*inputs=*/std::map{}, + /*node=*/ + mk_node_attrs( + input_layer.parallel_layer, input_op_attrs, input_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_input, input_pt_shape, /*create_grad=*/false), + }, }, - }, }; DynamicNodeInvocation weight_1_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{}, - /*node=*/mk_node_attrs(weight_1_layer.parallel_layer, - weight_1_op_attrs, - weight_1_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_weight_1, weight_1_pt_shape), + /*inputs=*/std::map{}, + /*node=*/ + mk_node_attrs(weight_1_layer.parallel_layer, + weight_1_op_attrs, + weight_1_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_weight_1, weight_1_pt_shape), + }, }, - }, }; DynamicNodeInvocation weight_2_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{}, - /*node=*/mk_node_attrs(weight_2_layer.parallel_layer, - weight_2_op_attrs, - weight_2_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_weight_2, weight_2_pt_shape), + /*inputs=*/std::map{}, + /*node=*/ + mk_node_attrs(weight_2_layer.parallel_layer, + weight_2_op_attrs, + weight_2_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_weight_2, weight_2_pt_shape), + }, }, - }, }; DynamicNodeInvocation partition_input_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(t_input, input_pt_shape, /*create_grad=*/false), + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_input, input_pt_shape, /*create_grad=*/false), + }, }, - }, - /*node=*/mk_node_attrs(partition_input_layer.parallel_layer, - partition_input_op_attrs, - partition_input_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_partitioned_input, partitioned_input_pt_shape), + /*node=*/ + mk_node_attrs(partition_input_layer.parallel_layer, + partition_input_op_attrs, + partition_input_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_partitioned_input, partitioned_input_pt_shape), + }, }, - }, }; - DynamicNodeInvocation replicate_weight_1_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(t_weight_1, weight_1_pt_shape), - }, - }, - /*node=*/mk_node_attrs(replicate_weight_1_layer.parallel_layer, - replicate_weight_1_op_attrs, - replicate_weight_1_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_replicated_weight_1, replicated_weight_1_pt_shape), - }, - }, - }; + DynamicNodeInvocation replicate_weight_1_invocation = + DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_weight_1, weight_1_pt_shape), + }, + }, + /*node=*/ + mk_node_attrs(replicate_weight_1_layer.parallel_layer, + replicate_weight_1_op_attrs, + replicate_weight_1_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_replicated_weight_1, + replicated_weight_1_pt_shape), + }, + }, + }; - DynamicNodeInvocation replicate_weight_2_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(t_weight_2, weight_2_pt_shape), - }, - }, - /*node=*/mk_node_attrs(replicate_weight_2_layer.parallel_layer, - replicate_weight_2_op_attrs, - replicate_weight_2_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_replicated_weight_2, replicated_weight_2_pt_shape), - }, - }, - }; + DynamicNodeInvocation replicate_weight_2_invocation = + DynamicNodeInvocation{ + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_weight_2, weight_2_pt_shape), + }, + }, + /*node=*/ + mk_node_attrs(replicate_weight_2_layer.parallel_layer, + replicate_weight_2_op_attrs, + replicate_weight_2_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_replicated_weight_2, + replicated_weight_2_pt_shape), + }, + }, + }; DynamicNodeInvocation linear_1_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(t_partitioned_input, partitioned_input_pt_shape), - }, - { - mk_slot(TensorSlotName::WEIGHT), - mk_value(t_replicated_weight_1, replicated_weight_1_pt_shape), + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_partitioned_input, partitioned_input_pt_shape), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(t_replicated_weight_1, replicated_weight_1_pt_shape), + }, }, - }, - /*node=*/mk_node_attrs(linear_1_layer.parallel_layer, - linear_1_op_attrs, - linear_1_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_hidden_activation, hidden_activation_pt_shape), + /*node=*/ + mk_node_attrs(linear_1_layer.parallel_layer, + linear_1_op_attrs, + linear_1_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_hidden_activation, hidden_activation_pt_shape), + }, }, - }, }; DynamicNodeInvocation linear_2_invocation = DynamicNodeInvocation{ - /*inputs=*/std::map{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(t_hidden_activation, hidden_activation_pt_shape), - }, - { - mk_slot(TensorSlotName::WEIGHT), - mk_value(t_replicated_weight_2, replicated_weight_2_pt_shape), + /*inputs=*/std::map{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(t_hidden_activation, hidden_activation_pt_shape), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(t_replicated_weight_2, replicated_weight_2_pt_shape), + }, }, - }, - /*node=*/mk_node_attrs(linear_2_layer.parallel_layer, - linear_2_op_attrs, - linear_2_op_mapping), - /*outputs=*/std::map{ - { - mk_slot(TensorSlotName::OUTPUT), - mk_value(t_output, output_pt_shape), + /*node=*/ + mk_node_attrs(linear_2_layer.parallel_layer, + linear_2_op_attrs, + linear_2_op_mapping), + /*outputs=*/ + std::map{ + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(t_output, output_pt_shape), + }, }, - }, }; return dynamic_open_dataflow_graph_from_invocation_set( - /*invocations=*/{ - input_invocation, - weight_1_invocation, - weight_2_invocation, - partition_input_invocation, - replicate_weight_1_invocation, - replicate_weight_2_invocation, - linear_1_invocation, - linear_2_invocation, - }); + /*invocations=*/{ + input_invocation, + weight_1_invocation, + weight_2_invocation, + partition_input_invocation, + replicate_weight_1_invocation, + replicate_weight_2_invocation, + linear_1_invocation, + linear_2_invocation, + }); }(); - nlohmann::json result_json - = dynamic_open_dataflow_graph_to_serializable(result); - nlohmann::json correct_json - = dynamic_open_dataflow_graph_to_serializable(correct); + nlohmann::json result_json = + dynamic_open_dataflow_graph_to_serializable(result); + nlohmann::json correct_json = + dynamic_open_dataflow_graph_to_serializable(correct); - CHECK_MESSAGE( - result == correct, - check_kv("result\n", result_json.dump()), - check_kv("correct\n", correct_json.dump())); + CHECK_MESSAGE(result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); } } diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc index d9e8ffa781..d3a915fb5c 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/pass_expansion.cc @@ -1,48 +1,48 @@ #include "task-spec/dynamic_graph/pass_expansion.h" +#include "op-attrs/initializer_attrs.h" #include "op-attrs/ops/element_unary.h" #include "task-spec/dynamic_graph/dynamic_open_dataflow_graph.h" #include "task-spec/dynamic_graph/dynamic_tensor_role.h" -#include #include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" -#include "op-attrs/initializer_attrs.h" -#include "test/utils/doctest/check_kv.h" #include "task-spec/dynamic_graph/serializable_dynamic_open_dataflow_graph.h" +#include "test/utils/doctest/check_kv.h" +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("determine_intermediate_values_needed_for_gradient_computation") { - auto mk_slot = [](TensorSlotName slot_name) -> DynamicTensorSlot { + auto mk_slot = [](TensorSlotName slot_name) -> DynamicTensorSlot { return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }; }; auto mk_node_attrs = [](size_t layer_guid, PCGOperatorAttrs const &op_attrs) { return DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/TrainingOperationAttrs{ - op_attrs, - }, - /*layer_guid=*/dynamic_layer_guid_t{ - parallel_layer_guid_t{ - Node{layer_guid}, + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/ + TrainingOperationAttrs{ + op_attrs, }, - }, - /*per_device_op_state=*/std::nullopt, + /*layer_guid=*/ + dynamic_layer_guid_t{ + parallel_layer_guid_t{ + Node{layer_guid}, + }, + }, + /*per_device_op_state=*/std::nullopt, }; }; auto mk_value_attrs = [](size_t src_layer_guid, TensorSlotName src_slot, - bool create_grad) - -> DynamicValueAttrs - { + bool create_grad) -> DynamicValueAttrs { return DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{ parallel_tensor_guid_t{ @@ -71,71 +71,74 @@ TEST_SUITE(FF_TEST_SUITE) { auto mk_test_graph = [&](bool input_create_grad) -> TestGraph { TensorShape input_shape = TensorShape{ - TensorDims{ - FFOrdered{ - 8_p, - 5_p, + TensorDims{ + FFOrdered{ + 8_p, + 5_p, + }, }, - }, - DataType::FLOAT, + DataType::FLOAT, }; - DynamicValueAttrs input_op_output = - mk_value_attrs(123, TensorSlotName::OUTPUT, /*create_grad=*/input_create_grad); + DynamicValueAttrs input_op_output = mk_value_attrs( + 123, TensorSlotName::OUTPUT, /*create_grad=*/input_create_grad); DynamicValueAttrs relu1_op_output = mk_value_attrs(124, TensorSlotName::OUTPUT, /*create_grad=*/true); - PCGOperatorAttrs input_attrs = PCGOperatorAttrs{ - InputAttrs{ - input_shape, - }, + InputAttrs{ + input_shape, + }, }; PCGOperatorAttrs relu_attrs = PCGOperatorAttrs{ - make_relu_attrs(), + make_relu_attrs(), }; DynamicNodeInvocation input_invocation = DynamicNodeInvocation{ - /*inputs=*/{}, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/123, - /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}), - /*outputs=*/{ + /*inputs=*/{}, + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/123, + /*op_attrs=*/PCGOperatorAttrs{InputAttrs{input_shape}}), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - input_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + input_op_output, + }, }, - }, }; DynamicNodeInvocation relu_invocation = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - input_op_output, + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + input_op_output, + }, }, - }, - /*node_attrs=*/mk_node_attrs( - /*layer_guid=*/124, - /*op_attrs=*/relu_attrs), - /*outputs=*/{ + /*node_attrs=*/ + mk_node_attrs( + /*layer_guid=*/124, + /*op_attrs=*/relu_attrs), + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - relu1_op_output, + { + mk_slot(TensorSlotName::OUTPUT), + relu1_op_output, + }, }, - }, }; - DynamicOpenDataflowGraph g - = dynamic_open_dataflow_graph_from_invocation_set( - {input_invocation, relu_invocation}); + DynamicOpenDataflowGraph g = + dynamic_open_dataflow_graph_from_invocation_set( + {input_invocation, relu_invocation}); return TestGraph{ - /*g=*/g, - /*input_op_output=*/input_op_output, - /*relu1_op_output=*/relu1_op_output, + /*g=*/g, + /*input_op_output=*/input_op_output, + /*relu1_op_output=*/relu1_op_output, }; }; @@ -143,7 +146,7 @@ TEST_SUITE(FF_TEST_SUITE) { TestGraph tg = mk_test_graph(/*input_create_grad=*/false); std::set result = - determine_intermediate_values_needed_for_gradient_computation(tg.g); + determine_intermediate_values_needed_for_gradient_computation(tg.g); std::set correct = {}; @@ -154,11 +157,11 @@ TEST_SUITE(FF_TEST_SUITE) { TestGraph tg = mk_test_graph(/*input_create_grad=*/true); std::set result = - determine_intermediate_values_needed_for_gradient_computation(tg.g); + determine_intermediate_values_needed_for_gradient_computation(tg.g); std::set correct = { - tg.input_op_output, - tg.relu1_op_output, + tg.input_op_output, + tg.relu1_op_output, }; ASSERT(result == correct); @@ -247,24 +250,24 @@ TEST_SUITE(FF_TEST_SUITE) { perform_fwd_pass_expansion_for_invocation(invocation); DynamicNodeInvocation correct = DynamicNodeInvocation{ - /*inputs=*/{ - {mk_slot(TensorSlotName::INPUT, fwd_role), v1_fwd}, - {mk_slot(TensorSlotName::WEIGHT, fwd_role), v2_fwd}, - {mk_slot(TensorSlotName::BIAS, fwd_role), v1_fwd}, - }, - /*node_attrs=*/ - DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, - /*device_coord=*/std::nullopt, - /*mapping=*/std::nullopt, - /*op_attrs=*/op_attrs, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/ - { - {mk_slot(TensorSlotName::OUTPUT, fwd_role), v3_fwd}, - }, + /*inputs=*/{ + {mk_slot(TensorSlotName::INPUT, fwd_role), v1_fwd}, + {mk_slot(TensorSlotName::WEIGHT, fwd_role), v2_fwd}, + {mk_slot(TensorSlotName::BIAS, fwd_role), v1_fwd}, + }, + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/DynamicTaskType::FWD, + /*device_coord=*/std::nullopt, + /*mapping=*/std::nullopt, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ + { + {mk_slot(TensorSlotName::OUTPUT, fwd_role), v3_fwd}, + }, }; ASSERT(result == correct); @@ -303,7 +306,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*node_attrs=*/ DynamicNodeAttrs{ - /*task_type=*/std::nullopt, + /*task_type=*/std::nullopt, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*op_attrs=*/op_attrs, @@ -316,7 +319,8 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); + ASSERT(dynamic_node_invocation_to_serializable(result) == + dynamic_node_invocation_to_serializable(correct)); } } @@ -432,7 +436,8 @@ TEST_SUITE(FF_TEST_SUITE) { }; }(); - ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); + ASSERT(dynamic_node_invocation_to_serializable(result) == + dynamic_node_invocation_to_serializable(correct)); } SUBCASE("replicate operator") { @@ -493,7 +498,8 @@ TEST_SUITE(FF_TEST_SUITE) { }; }(); - ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); + ASSERT(dynamic_node_invocation_to_serializable(result) == + dynamic_node_invocation_to_serializable(correct)); } SUBCASE("copy operator") { @@ -534,7 +540,7 @@ TEST_SUITE(FF_TEST_SUITE) { }, /*node_attrs=*/ DynamicNodeAttrs{ - /*pass_type=*/std::nullopt, + /*pass_type=*/std::nullopt, /*device_coord=*/std::nullopt, /*mapping=*/std::nullopt, /*op_attrs=*/op_attrs, @@ -548,7 +554,8 @@ TEST_SUITE(FF_TEST_SUITE) { }; }(); - ASSERT(dynamic_node_invocation_to_serializable(result) == dynamic_node_invocation_to_serializable(correct)); + ASSERT(dynamic_node_invocation_to_serializable(result) == + dynamic_node_invocation_to_serializable(correct)); } } @@ -569,8 +576,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; auto mk_value_attrs = - [](size_t node_id, - std::optional const &tensor_type) + [](size_t node_id, std::optional const &tensor_type) -> DynamicValueAttrs { return DynamicValueAttrs{ /*tensor_guid=*/dynamic_tensor_guid_t{parallel_tensor_guid_t{ @@ -622,41 +628,46 @@ TEST_SUITE(FF_TEST_SUITE) { TrainingOperationAttrs weight_op_attrs = TrainingOperationAttrs{ PCGOperatorAttrs{ - WeightAttrs{ - /*tensor_shape=*/weight_shape, - /*initializer=*/make_zero_initializer(), - }, + WeightAttrs{ + /*tensor_shape=*/weight_shape, + /*initializer=*/make_zero_initializer(), + }, }, }; TrainingOperationAttrs linear_op_attrs = TrainingOperationAttrs{ PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/6_p, - /*use_bias=*/false, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, - }, + LinearAttrs{ + /*out_channels=*/6_p, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }, }; DynamicOpenDataflowGraph input = [&]() -> DynamicOpenDataflowGraph { - DynamicNodeAttrs input_node = mk_node_attrs(10, input_op_attrs, std::nullopt); - DynamicNodeAttrs weight_node = mk_node_attrs(11, weight_op_attrs, std::nullopt); - DynamicNodeAttrs relu_node = mk_node_attrs(12, relu_op_attrs, std::nullopt); - DynamicNodeAttrs linear_node = mk_node_attrs(13, linear_op_attrs, std::nullopt); + DynamicNodeAttrs input_node = + mk_node_attrs(10, input_op_attrs, std::nullopt); + DynamicNodeAttrs weight_node = + mk_node_attrs(11, weight_op_attrs, std::nullopt); + DynamicNodeAttrs relu_node = + mk_node_attrs(12, relu_op_attrs, std::nullopt); + DynamicNodeAttrs linear_node = + mk_node_attrs(13, linear_op_attrs, std::nullopt); DynamicValueAttrs input_tensor = mk_value_attrs(0, std::nullopt); DynamicValueAttrs weight_tensor = mk_value_attrs(1, std::nullopt); DynamicValueAttrs relu_output = mk_value_attrs(2, std::nullopt); DynamicValueAttrs linear_output = mk_value_attrs(3, std::nullopt); - auto mk_dynamic_slot = [](TensorSlotName const &slot_name) -> DynamicTensorSlot { + auto mk_dynamic_slot = + [](TensorSlotName const &slot_name) -> DynamicTensorSlot { return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/std::nullopt, + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/std::nullopt, }; }; @@ -667,8 +678,8 @@ TEST_SUITE(FF_TEST_SUITE) { /*outputs=*/ std::map{ { - mk_dynamic_slot(TensorSlotName::OUTPUT), - input_tensor, + mk_dynamic_slot(TensorSlotName::OUTPUT), + input_tensor, }, }, }, @@ -678,44 +689,44 @@ TEST_SUITE(FF_TEST_SUITE) { /*outputs=*/ std::map{ { - mk_dynamic_slot(TensorSlotName::OUTPUT), - weight_tensor, + mk_dynamic_slot(TensorSlotName::OUTPUT), + weight_tensor, }, }, }, DynamicNodeInvocation{ /*inputs=*/std::map{ { - mk_dynamic_slot(TensorSlotName::INPUT), - input_tensor, + mk_dynamic_slot(TensorSlotName::INPUT), + input_tensor, }, }, /*node_attrs=*/relu_node, /*outputs=*/ std::map{ { - mk_dynamic_slot(TensorSlotName::OUTPUT), - relu_output, + mk_dynamic_slot(TensorSlotName::OUTPUT), + relu_output, }, }, }, DynamicNodeInvocation{ /*inputs=*/std::map{ { - mk_dynamic_slot(TensorSlotName::INPUT), - relu_output, + mk_dynamic_slot(TensorSlotName::INPUT), + relu_output, }, { - mk_dynamic_slot(TensorSlotName::WEIGHT), - weight_tensor, + mk_dynamic_slot(TensorSlotName::WEIGHT), + weight_tensor, }, }, /*node_attrs=*/linear_node, /*outputs=*/ std::map{ { - mk_dynamic_slot(TensorSlotName::OUTPUT), - linear_output, + mk_dynamic_slot(TensorSlotName::OUTPUT), + linear_output, }, }, }, @@ -753,7 +764,7 @@ TEST_SUITE(FF_TEST_SUITE) { mk_value_attrs(2, mk_dynamic_tensor_role_bwd()); DynamicValueAttrs linear_output_tensor_activation = mk_value_attrs(3, mk_dynamic_tensor_role_fwd()); - DynamicValueAttrs linear_output_tensor_gradient= + DynamicValueAttrs linear_output_tensor_gradient = mk_value_attrs(3, mk_dynamic_tensor_role_bwd()); auto mk_fwd_slot = [&](TensorSlotName slot_name) -> DynamicTensorSlot { @@ -868,14 +879,16 @@ TEST_SUITE(FF_TEST_SUITE) { return dynamic_open_dataflow_graph_from_invocation_set(invocation_set); }(); - CHECK(get_dynamic_invocation_set(result).size() == correct.invocations.size()); + CHECK(get_dynamic_invocation_set(result).size() == + correct.invocations.size()); - nlohmann::json result_json - = dynamic_open_dataflow_graph_to_serializable(result); - nlohmann::json correct_json - = dynamic_open_dataflow_graph_to_serializable(correct); + nlohmann::json result_json = + dynamic_open_dataflow_graph_to_serializable(result); + nlohmann::json correct_json = + dynamic_open_dataflow_graph_to_serializable(correct); - CHECK_MESSAGE(get_dynamic_invocation_set(result) == get_dynamic_invocation_set(correct), + CHECK_MESSAGE(get_dynamic_invocation_set(result) == + get_dynamic_invocation_set(correct), check_kv("result", result_json.dump()), check_kv("correct", correct_json.dump())); diff --git a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc index 89339751ef..9054aaef28 100644 --- a/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/test/src/task-spec/dynamic_graph/shard_expansion.cc @@ -1,20 +1,20 @@ #include "task-spec/dynamic_graph/shard_expansion.h" +#include "op-attrs/ops/element_unary.h" #include "pcg/mapped_parallel_computation_graph/mapped_operator_task_group.h" #include "task-spec/dynamic_graph/copy_attrs.dtg.h" #include "task-spec/dynamic_graph/dynamic_copy_layer_guid_t.dtg.h" #include "task-spec/dynamic_graph/dynamic_node_mapping.h" +#include "task-spec/dynamic_graph/dynamic_tensor_role.h" +#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" #include "task-spec/dynamic_graph/training_operation_attrs.dtg.h" +#include "test/utils/doctest/check_kv.h" #include "test/utils/doctest/fmt/set.h" -#include -#include "task-spec/dynamic_graph/dynamic_tensor_role.h" -#include "op-attrs/ops/element_unary.h" #include "utils/bidict/algorithms/bidict_filter_keys.h" #include "utils/bidict/algorithms/bidict_filter_values.h" -#include "utils/containers/map_from_pairs.h" -#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/binary_relation/binary_relation_from_map.h" -#include "task-spec/dynamic_graph/serializable_dynamic_node_invocation.h" -#include "test/utils/doctest/check_kv.h" +#include "utils/containers/binary_merge_disjoint_maps.h" +#include "utils/containers/map_from_pairs.h" +#include using namespace ::FlexFlow; @@ -48,8 +48,9 @@ static ParallelTensorSpaceCoordinate mk_pt_coord(nonnegative_int idx1, }; }; -DynamicTensorSlot mk_slot(TensorSlotName const &slot_name, - std::optional const &task_shard = std::nullopt) { +DynamicTensorSlot mk_slot( + TensorSlotName const &slot_name, + std::optional const &task_shard = std::nullopt) { return DynamicTensorSlot{ /*slot_name=*/slot_name, /*slot_tensor_role=*/std::nullopt, @@ -60,11 +61,13 @@ DynamicTensorSlot mk_slot(TensorSlotName const &slot_name, DynamicValueAttrs mk_value(size_t src_node_id, TensorSlotName src_slot_name, - bidict const &tensor_binding, + bidict const + &tensor_binding, std::optional const &shard_coord, std::optional const &role = std::nullopt) { - bidict mapping = tensor_binding; + bidict mapping = + tensor_binding; if (shard_coord.has_value()) { mapping = bidict_filter_keys(mapping, [&](ParallelTensorSpaceCoordinate const &p) { @@ -91,66 +94,65 @@ DynamicValueAttrs TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("apply_dynamic_node_invocation_sharding_info") { global_device_id_t device_0 = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/0_n, - /*device_idx=*/0_n, - }, - /*device_type=*/DeviceType::GPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/0_n, + /*device_idx=*/0_n, + }, + /*device_type=*/DeviceType::GPU, }; global_device_id_t device_1 = global_device_id_t{ - /*coord=*/MachineSpaceCoordinate{ - /*node_idx=*/2_n, - /*device_idx=*/1_n, - }, - /*device_type=*/DeviceType::GPU, + /*coord=*/MachineSpaceCoordinate{ + /*node_idx=*/2_n, + /*device_idx=*/1_n, + }, + /*device_type=*/DeviceType::GPU, }; - auto mk_slot = [](TensorSlotName slot_name, - std::optional const &task_shard = std::nullopt) - -> DynamicTensorSlot - { + auto mk_slot = [](TensorSlotName slot_name, + std::optional const &task_shard = + std::nullopt) -> DynamicTensorSlot { return DynamicTensorSlot{ - /*slot_name=*/slot_name, - /*slot_tensor_role=*/std::nullopt, - /*task_shard=*/task_shard, + /*slot_name=*/slot_name, + /*slot_tensor_role=*/std::nullopt, + /*task_shard=*/task_shard, }; }; - auto mk_value = [](size_t src_node_id, - TensorSlotName src_slot_name, - std::optional const &shard_coord = std::nullopt) - -> DynamicValueAttrs - { + auto mk_value = + [](size_t src_node_id, + TensorSlotName src_slot_name, + std::optional const &shard_coord = + std::nullopt) -> DynamicValueAttrs { return DynamicValueAttrs{ - /*tensor_guid=*/dynamic_tensor_guid_t{ - parallel_tensor_guid_t{ - KwargDataflowOutput{ - /*node=*/Node{src_node_id}, - /*slot_name=*/src_slot_name, - }, + /*tensor_guid=*/dynamic_tensor_guid_t{ + parallel_tensor_guid_t{ + KwargDataflowOutput{ + /*node=*/Node{src_node_id}, + /*slot_name=*/src_slot_name, + }, + }, }, - }, - /*parallel_tensor_shape=*/std::nullopt, - /*create_grad=*/std::nullopt, - /*shard_coord=*/shard_coord, - /*mapping=*/std::nullopt, - /*accessor=*/std::nullopt, - /*role=*/std::nullopt, + /*parallel_tensor_shape=*/std::nullopt, + /*create_grad=*/std::nullopt, + /*shard_coord=*/shard_coord, + /*mapping=*/std::nullopt, + /*accessor=*/std::nullopt, + /*role=*/std::nullopt, }; }; SUBCASE("sharding info creates additional arguments ie replicate") { - auto mk_pt_coord = [](nonnegative_int idx) - -> ParallelTensorSpaceCoordinate - { + auto mk_pt_coord = + [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/idx, - /*shard_components=*/FFOrdered{ - 0_n, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shard_components=*/ + FFOrdered{ + 0_n, + 0_n, + }, }; }; @@ -158,322 +160,352 @@ TEST_SUITE(FF_TEST_SUITE) { size_t replicate_layer_node_id = 13; DynamicNodeMapping node_mapping = DynamicNodeMapping{ - /*op_task_group=*/MappedOperatorTaskGroup{{ - { - device_0.coord, - OperatorAtomicTaskShardBinding{{ + /*op_task_group=*/MappedOperatorTaskGroup{{ { - TensorSlotName::INPUT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), - }, - }}, - }, - { - device_1.coord, - OperatorAtomicTaskShardBinding{{ - { - TensorSlotName::INPUT, - mk_pt_coord(0_n), + device_0.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }}, }, { - TensorSlotName::OUTPUT, - mk_pt_coord(1_n), + device_1.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }}, }, - }}, - }, - }}, - /*device_type=*/DeviceType::GPU, + }}, + /*device_type=*/DeviceType::GPU, }; TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ - PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }, - }, }; dynamic_layer_guid_t layer_guid = dynamic_layer_guid_t{ parallel_layer_guid_t{ - Node{replicate_layer_node_id}, - }, + Node{replicate_layer_node_id}, + }, }; DynamicNodeInvocation invocation = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(input_src_node_id, TensorSlotName::OUTPUT), + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, TensorSlotName::OUTPUT), + }, }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/std::nullopt, - /*mapping=*/node_mapping, - /*op_attrs=*/op_attrs, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, + }, + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - mk_value(replicate_layer_node_id, TensorSlotName::OUTPUT), + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(replicate_layer_node_id, TensorSlotName::OUTPUT), + }, }, - }, }; - DynamicNodeInvocationShardingInfo invocation_sharding_info = - DynamicNodeInvocationShardingInfo{ - /*device_ids=*/nonempty_set{ - device_0, - device_1, - }, - /*value_sharding=*/{ - { - mk_slot(TensorSlotName::INPUT), - DynamicValueAttrsShardingInfo{ - /*shard_coord=*/mk_pt_coord(0_n), - /*mapping=*/device_0, + DynamicNodeInvocationShardingInfo invocation_sharding_info = + DynamicNodeInvocationShardingInfo{ + /*device_ids=*/nonempty_set{ + device_0, + device_1, }, - }, - { - mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_0.coord), - DynamicValueAttrsShardingInfo{ - /*shard_coord=*/mk_pt_coord(0_n), - /*mapping=*/device_0, - }, - }, - { - mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_1.coord), - DynamicValueAttrsShardingInfo{ - /*shard_coord=*/mk_pt_coord(1_n), - /*mapping=*/device_1, + /*value_sharding=*/ + { + { + mk_slot(TensorSlotName::INPUT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(0_n), + /*mapping=*/device_0, + }, + }, + { + mk_slot(TensorSlotName::OUTPUT, + /*task_shard=*/device_0.coord), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(0_n), + /*mapping=*/device_0, + }, + }, + { + mk_slot(TensorSlotName::OUTPUT, + /*task_shard=*/device_1.coord), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(1_n), + /*mapping=*/device_1, + }, + }, }, - }, - }, - }; + }; - DynamicNodeInvocation result = - apply_dynamic_node_invocation_sharding_info(invocation, invocation_sharding_info); + DynamicNodeInvocation result = + apply_dynamic_node_invocation_sharding_info(invocation, + invocation_sharding_info); DynamicNodeInvocation correct = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(input_src_node_id, TensorSlotName::OUTPUT, /*shrad_coord=*/mk_pt_coord(0_n)), - }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/nonempty_set{ - device_0, - device_1, + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, + TensorSlotName::OUTPUT, + /*shrad_coord=*/mk_pt_coord(0_n)), + }, }, - /*mapping=*/node_mapping, - /*op_attrs=*/op_attrs, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ - { - mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_0.coord), - mk_value(replicate_layer_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(0_n)), + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/ + nonempty_set{ + device_0, + device_1, + }, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, }, + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT, /*task_shard=*/device_1.coord), - mk_value(replicate_layer_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(1_n)), + { + mk_slot(TensorSlotName::OUTPUT, + /*task_shard=*/device_0.coord), + mk_value(replicate_layer_node_id, + TensorSlotName::OUTPUT, + /*shard_coord=*/mk_pt_coord(0_n)), + }, + { + mk_slot(TensorSlotName::OUTPUT, + /*task_shard=*/device_1.coord), + mk_value(replicate_layer_node_id, + TensorSlotName::OUTPUT, + /*shard_coord=*/mk_pt_coord(1_n)), + }, }, - }, }; - nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); - nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + nlohmann::json result_json = + dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = + dynamic_node_invocation_to_serializable(correct); - CHECK_MESSAGE( - result == correct, - check_kv("result\n", result_json.dump()), - check_kv("correct\n", correct_json.dump()) - ); - } + CHECK_MESSAGE(result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); + } - SUBCASE("sharding info does not create additional arguments ie standard operator") { + SUBCASE("sharding info does not create additional arguments ie standard " + "operator") { size_t input_src_node_id = 234; size_t weight_src_node_id = 345; size_t linear_layer_node_id = 13; - auto mk_pt_coord = [](nonnegative_int idx) - -> ParallelTensorSpaceCoordinate - { + auto mk_pt_coord = + [](nonnegative_int idx) -> ParallelTensorSpaceCoordinate { return ParallelTensorSpaceCoordinate{ - /*sum_component=*/0_n, - /*discard_copy_component=*/idx, - /*shard_components=*/FFOrdered{ - 0_n, - 0_n, - }, + /*sum_component=*/0_n, + /*discard_copy_component=*/idx, + /*shard_components=*/ + FFOrdered{ + 0_n, + 0_n, + }, }; }; // note that the node mapping does not have to be accurate/real here. - // apply_dynamic_node_invocation_sharding_info should just function + // apply_dynamic_node_invocation_sharding_info should just function // based on what it is given. DynamicNodeMapping node_mapping = DynamicNodeMapping{ - /*op_task_group=*/MappedOperatorTaskGroup{{ - { - device_0.coord, - OperatorAtomicTaskShardBinding{{ - { - TensorSlotName::INPUT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::WEIGHT, - mk_pt_coord(0_n), - }, - { - TensorSlotName::OUTPUT, - mk_pt_coord(0_n), - }, - }}, - }, - { - device_1.coord, - OperatorAtomicTaskShardBinding{{ - { - TensorSlotName::INPUT, - mk_pt_coord(1_n), - }, + /*op_task_group=*/MappedOperatorTaskGroup{{ { - TensorSlotName::WEIGHT, - mk_pt_coord(2_n), + device_0.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::WEIGHT, + mk_pt_coord(0_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(0_n), + }, + }}, }, { - TensorSlotName::OUTPUT, - mk_pt_coord(1_n), + device_1.coord, + OperatorAtomicTaskShardBinding{{ + { + TensorSlotName::INPUT, + mk_pt_coord(1_n), + }, + { + TensorSlotName::WEIGHT, + mk_pt_coord(2_n), + }, + { + TensorSlotName::OUTPUT, + mk_pt_coord(1_n), + }, + }}, }, - }}, - }, - }}, - /*device_type=*/DeviceType::GPU, + }}, + /*device_type=*/DeviceType::GPU, }; TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ - PCGOperatorAttrs{ - LinearAttrs{ - /*out_channels=*/8_p, - /*use_bias=*/false, - /*data_type=*/DataType::FLOAT, - /*activation=*/std::nullopt, - /*regularizer=*/std::nullopt, + PCGOperatorAttrs{ + LinearAttrs{ + /*out_channels=*/8_p, + /*use_bias=*/false, + /*data_type=*/DataType::FLOAT, + /*activation=*/std::nullopt, + /*regularizer=*/std::nullopt, + }, }, - }, }; dynamic_layer_guid_t layer_guid = dynamic_layer_guid_t{ parallel_layer_guid_t{ - Node{linear_layer_node_id}, - }, + Node{linear_layer_node_id}, + }, }; DynamicNodeInvocation invocation = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(input_src_node_id, TensorSlotName::OUTPUT), + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, TensorSlotName::OUTPUT), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(weight_src_node_id, TensorSlotName::OUTPUT), + }, }, - { - mk_slot(TensorSlotName::WEIGHT), - mk_value(weight_src_node_id, TensorSlotName::OUTPUT), + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/std::nullopt, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/std::nullopt, - /*mapping=*/node_mapping, - /*op_attrs=*/op_attrs, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - mk_value(linear_layer_node_id, TensorSlotName::OUTPUT), + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(linear_layer_node_id, TensorSlotName::OUTPUT), + }, }, - }, }; - DynamicNodeInvocationShardingInfo invocation_sharding_info = - DynamicNodeInvocationShardingInfo{ - /*device_ids=*/nonempty_set{ - device_1, - }, - /*value_sharding=*/{ - { - mk_slot(TensorSlotName::INPUT), - DynamicValueAttrsShardingInfo{ - /*shard_coord=*/mk_pt_coord(1_n), - /*mapping=*/device_1, - }, - }, - { - mk_slot(TensorSlotName::WEIGHT), - DynamicValueAttrsShardingInfo{ - /*shard_coord=*/mk_pt_coord(2_n), - /*mapping=*/device_1, + DynamicNodeInvocationShardingInfo invocation_sharding_info = + DynamicNodeInvocationShardingInfo{ + /*device_ids=*/nonempty_set{ + device_1, }, - }, - { - mk_slot(TensorSlotName::OUTPUT), - DynamicValueAttrsShardingInfo{ - /*shard_coord=*/mk_pt_coord(1_n), - /*mapping=*/device_1, + /*value_sharding=*/ + { + { + mk_slot(TensorSlotName::INPUT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(1_n), + /*mapping=*/device_1, + }, + }, + { + mk_slot(TensorSlotName::WEIGHT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(2_n), + /*mapping=*/device_1, + }, + }, + { + mk_slot(TensorSlotName::OUTPUT), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/mk_pt_coord(1_n), + /*mapping=*/device_1, + }, + }, }, - }, - }, - }; + }; - DynamicNodeInvocation result = - apply_dynamic_node_invocation_sharding_info(invocation, invocation_sharding_info); + DynamicNodeInvocation result = + apply_dynamic_node_invocation_sharding_info(invocation, + invocation_sharding_info); DynamicNodeInvocation correct = DynamicNodeInvocation{ - /*inputs=*/{ - { - mk_slot(TensorSlotName::INPUT), - mk_value(input_src_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(1_n)), + /*inputs=*/{ + { + mk_slot(TensorSlotName::INPUT), + mk_value(input_src_node_id, + TensorSlotName::OUTPUT, + /*shard_coord=*/mk_pt_coord(1_n)), + }, + { + mk_slot(TensorSlotName::WEIGHT), + mk_value(weight_src_node_id, + TensorSlotName::OUTPUT, + /*shard_coord=*/mk_pt_coord(2_n)), + }, }, - { - mk_slot(TensorSlotName::WEIGHT), - mk_value(weight_src_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(2_n)), + /*node_attrs=*/ + DynamicNodeAttrs{ + /*task_type=*/std::nullopt, + /*device_ids=*/nonempty_set{device_1}, + /*mapping=*/node_mapping, + /*op_attrs=*/op_attrs, + /*layer_guid=*/layer_guid, + /*per_device_op_state=*/std::nullopt, }, - }, - /*node_attrs=*/DynamicNodeAttrs{ - /*task_type=*/std::nullopt, - /*device_ids=*/nonempty_set{device_1}, - /*mapping=*/node_mapping, - /*op_attrs=*/op_attrs, - /*layer_guid=*/layer_guid, - /*per_device_op_state=*/std::nullopt, - }, - /*outputs=*/{ + /*outputs=*/ { - mk_slot(TensorSlotName::OUTPUT), - mk_value(linear_layer_node_id, TensorSlotName::OUTPUT, /*shard_coord=*/mk_pt_coord(1_n)), + { + mk_slot(TensorSlotName::OUTPUT), + mk_value(linear_layer_node_id, + TensorSlotName::OUTPUT, + /*shard_coord=*/mk_pt_coord(1_n)), + }, }, - }, }; - nlohmann::json result_json = dynamic_node_invocation_to_serializable(result); - nlohmann::json correct_json = dynamic_node_invocation_to_serializable(correct); + nlohmann::json result_json = + dynamic_node_invocation_to_serializable(result); + nlohmann::json correct_json = + dynamic_node_invocation_to_serializable(correct); - CHECK_MESSAGE( - result == correct, - check_kv("result\n", result_json.dump()), - check_kv("correct\n", correct_json.dump()) - ); + CHECK_MESSAGE(result == correct, + check_kv("result\n", result_json.dump()), + check_kv("correct\n", correct_json.dump())); } } @@ -484,29 +516,29 @@ TEST_SUITE(FF_TEST_SUITE) { TensorSlotName use_slot_name, DynamicNodeMapping const &node_mapping, std::optional const &shard_coord, - std::optional const &role = std::nullopt) - -> DynamicValueAttrs { - - bidict - tensor_binding = dynamic_node_mapping_bindings_for_slot_name(node_mapping, - use_slot_name); - return mk_value(src_node_id, src_slot_name, tensor_binding, shard_coord, role); + std::optional const &role = + std::nullopt) -> DynamicValueAttrs { + bidict tensor_binding = + dynamic_node_mapping_bindings_for_slot_name(node_mapping, + use_slot_name); + return mk_value( + src_node_id, src_slot_name, tensor_binding, shard_coord, role); }; - auto mk_sharding_info = [&](TensorSlotName slot_name, - ParallelTensorSpaceCoordinate const &shard_coord, - DynamicNodeMapping const &node_mapping) - -> std::pair - { - bidict - tensor_binding = dynamic_node_mapping_bindings_for_slot_name(node_mapping, slot_name); + auto mk_sharding_info = + [&](TensorSlotName slot_name, + ParallelTensorSpaceCoordinate const &shard_coord, + DynamicNodeMapping const &node_mapping) + -> std::pair { + bidict tensor_binding = + dynamic_node_mapping_bindings_for_slot_name(node_mapping, slot_name); return std::pair{ - mk_slot(slot_name), - DynamicValueAttrsShardingInfo{ - /*shard_coord=*/shard_coord, - /*mapping=*/tensor_binding.at_l(shard_coord), - }, + mk_slot(slot_name), + DynamicValueAttrsShardingInfo{ + /*shard_coord=*/shard_coord, + /*mapping=*/tensor_binding.at_l(shard_coord), + }, }; }; @@ -551,11 +583,11 @@ TEST_SUITE(FF_TEST_SUITE) { std::optional const &shard_coord) -> DynamicValueAttrs { if (shard_coord.has_value()) { - tensor_binding = - bidict_filter_keys(tensor_binding, - [&](ParallelTensorSpaceCoordinate const &p) -> bool { - return p == shard_coord.value(); - }); + tensor_binding = bidict_filter_keys( + tensor_binding, + [&](ParallelTensorSpaceCoordinate const &p) -> bool { + return p == shard_coord.value(); + }); } return DynamicValueAttrs{ @@ -594,9 +626,9 @@ TEST_SUITE(FF_TEST_SUITE) { mk_pt_coord(0_n, 0_n, 0_n, 0_n); TrainingOperationAttrs op_attrs = TrainingOperationAttrs{ - PCGOperatorAttrs{ - make_relu_attrs(), - }, + PCGOperatorAttrs{ + make_relu_attrs(), + }, }; DynamicNodeMapping node_mapping = DynamicNodeMapping{ @@ -682,13 +714,20 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelTensorSpaceCoordinate const &output_2_shard_coord) -> DynamicNodeInvocationShardingInfo { return DynamicNodeInvocationShardingInfo{ - /*device_coord=*/nonempty_set{device_coord}, - /*value_sharding=*/{ - mk_sharding_info(TensorSlotName::INPUT, input_shard_coord, node_mapping), - mk_sharding_info(TensorSlotName::WEIGHT, weight_shard_coord, node_mapping), - mk_sharding_info(TensorSlotName::OUTPUT_1, output_1_shard_coord, node_mapping), - mk_sharding_info(TensorSlotName::OUTPUT_2, output_2_shard_coord, node_mapping), - }, + /*device_coord=*/nonempty_set{device_coord}, + /*value_sharding=*/ + { + mk_sharding_info( + TensorSlotName::INPUT, input_shard_coord, node_mapping), + mk_sharding_info( + TensorSlotName::WEIGHT, weight_shard_coord, node_mapping), + mk_sharding_info(TensorSlotName::OUTPUT_1, + output_1_shard_coord, + node_mapping), + mk_sharding_info(TensorSlotName::OUTPUT_2, + output_2_shard_coord, + node_mapping), + }, }; }; @@ -726,7 +765,8 @@ TEST_SUITE(FF_TEST_SUITE) { /*inputs=*/{ { mk_slot(TensorSlotName::INPUT), - mk_value( 0, TensorSlotName::OUTPUT, src_binding, std::nullopt), + mk_value( + 0, TensorSlotName::OUTPUT, src_binding, std::nullopt), }, }, /*node_attrs=*/ @@ -742,7 +782,8 @@ TEST_SUITE(FF_TEST_SUITE) { { { mk_slot(TensorSlotName::OUTPUT), - mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), + mk_value( + 20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), }, }, }; @@ -754,25 +795,25 @@ TEST_SUITE(FF_TEST_SUITE) { [&](global_device_id_t const &device_coord, ParallelTensorSpaceCoordinate const &tensor_shard_coord) -> DynamicNodeInvocationShardingInfo { - return DynamicNodeInvocationShardingInfo{ - /*device_coord=*/nonempty_set{device_coord}, - /*value_sharding=*/BinaryRelation{ - { - mk_slot(TensorSlotName::INPUT), - DynamicValueAttrsShardingInfo{ - tensor_shard_coord, - src_binding.at_l(tensor_shard_coord), - }, - }, - { - mk_slot(TensorSlotName::OUTPUT), - DynamicValueAttrsShardingInfo{ - tensor_shard_coord, - dst_binding.at_l(tensor_shard_coord), - }, + /*device_coord=*/nonempty_set{device_coord}, + /*value_sharding=*/ + BinaryRelation{ + { + mk_slot(TensorSlotName::INPUT), + DynamicValueAttrsShardingInfo{ + tensor_shard_coord, + src_binding.at_l(tensor_shard_coord), + }, + }, + { + mk_slot(TensorSlotName::OUTPUT), + DynamicValueAttrsShardingInfo{ + tensor_shard_coord, + dst_binding.at_l(tensor_shard_coord), + }, + }, }, - }, }; }; @@ -809,27 +850,27 @@ TEST_SUITE(FF_TEST_SUITE) { }; DynamicNodeMapping node_mapping = DynamicNodeMapping{ - /*op_task_group=*/MappedOperatorTaskGroup{ - bidict{ - { - dev1.coord, - mk_shard_binding(pt1, pt1), - }, - { - dev2.coord, - mk_shard_binding(pt1, pt2), - }, - { - dev3.coord, - mk_shard_binding(pt2, pt3), - }, - { - dev4.coord, - mk_shard_binding(pt2, pt4), - }, - }, - }, - /*device_type=*/DeviceType::GPU, + /*op_task_group=*/MappedOperatorTaskGroup{ + bidict{ + { + dev1.coord, + mk_shard_binding(pt1, pt1), + }, + { + dev2.coord, + mk_shard_binding(pt1, pt2), + }, + { + dev3.coord, + mk_shard_binding(pt2, pt3), + }, + { + dev4.coord, + mk_shard_binding(pt2, pt4), + }, + }, + }, + /*device_type=*/DeviceType::GPU, }; SUBCASE("fwd") { @@ -849,37 +890,41 @@ TEST_SUITE(FF_TEST_SUITE) { /*inputs=*/{ { DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, - mk_value(0, TensorSlotName::OUTPUT, src_binding, std::nullopt), + mk_value( + 0, TensorSlotName::OUTPUT, src_binding, std::nullopt), }, }, /*node_attrs=*/ DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::FWD, + /*task_type=*/DynamicTaskType::FWD, /*device_ids=*/std::nullopt, /*mapping=*/node_mapping, - /*op_attrs=*/TrainingOperationAttrs{ - PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, + /*op_attrs=*/ + TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }, - }, }, - /*layer_guid=*/dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, + /*layer_guid=*/ + dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, /*per_device_op_state=*/std::nullopt, }, /*outputs=*/ { { DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, }, - mk_value(20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), + mk_value( + 20, TensorSlotName::OUTPUT, dst_binding, std::nullopt), }, }, }; @@ -887,20 +932,18 @@ TEST_SUITE(FF_TEST_SUITE) { std::set result = generate_shard_expansion_for_invocation(input); - auto mk_output_binding = [&](global_device_id_t const &device) - -> std::pair - { + -> std::pair { return { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/device.coord, - }, - DynamicValueAttrsShardingInfo{ - dst_binding.at_r(device), - device, - }, + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/device.coord, + }, + DynamicValueAttrsShardingInfo{ + dst_binding.at_r(device), + device, + }, }; }; @@ -909,32 +952,31 @@ TEST_SUITE(FF_TEST_SUITE) { ParallelTensorSpaceCoordinate const &input_shard_coord, std::set const &output_task_shards) -> DynamicNodeInvocationShardingInfo { - return DynamicNodeInvocationShardingInfo{ - /*device_ids=*/device_ids, - /*value_sharding=*/ - binary_relation_from_map( - binary_merge_disjoint_maps( + /*device_ids=*/device_ids, + /*value_sharding=*/ + binary_relation_from_map(binary_merge_disjoint_maps( std::map{ - { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrsShardingInfo{ - input_shard_coord, - src_binding.at_l(input_shard_coord), + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_fwd(), + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrsShardingInfo{ + input_shard_coord, + src_binding.at_l(input_shard_coord), + }, }, - }, }, - map_from_pairs(transform(output_task_shards, mk_output_binding)))), + map_from_pairs( + transform(output_task_shards, mk_output_binding)))), }; }; std::set correct = { - mk_invocation_shard(nonempty_set{dev1, dev2}, pt1, {dev1, dev2}), - mk_invocation_shard(nonempty_set{dev3, dev4}, pt2, {dev3, dev4}), + mk_invocation_shard(nonempty_set{dev1, dev2}, pt1, {dev1, dev2}), + mk_invocation_shard(nonempty_set{dev3, dev4}, pt2, {dev3, dev4}), }; CHECK(result.size() == correct.size()); @@ -942,53 +984,63 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("bwd") { - bidict output_grad_binding{ - {pt1, dev1}, - {pt2, dev2}, - {pt3, dev3}, - {pt4, dev4}, - }; - - bidict input_grad_binding{ - {pt1, dev1}, - {pt2, dev2}, - }; + bidict + output_grad_binding{ + {pt1, dev1}, + {pt2, dev2}, + {pt3, dev3}, + {pt4, dev4}, + }; + + bidict + input_grad_binding{ + {pt1, dev1}, + {pt2, dev2}, + }; DynamicNodeInvocation input = DynamicNodeInvocation{ /*inputs=*/{ { DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }, - mk_value(0, TensorSlotName::OUTPUT, output_grad_binding, std::nullopt), + mk_value(0, + TensorSlotName::OUTPUT, + output_grad_binding, + std::nullopt), }, }, /*node_attrs=*/ DynamicNodeAttrs{ - /*task_type=*/DynamicTaskType::BWD, + /*task_type=*/DynamicTaskType::BWD, /*device_ids=*/std::nullopt, /*mapping=*/node_mapping, - /*op_attrs=*/TrainingOperationAttrs{ - PCGOperatorAttrs{ - ReplicateAttrs{ - /*replicate_degree=*/2_p, + /*op_attrs=*/ + TrainingOperationAttrs{ + PCGOperatorAttrs{ + ReplicateAttrs{ + /*replicate_degree=*/2_p, + }, }, - }, }, - /*layer_guid=*/dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, + /*layer_guid=*/ + dynamic_layer_guid_t{parallel_layer_guid_t{Node{20}}}, /*per_device_op_state=*/std::nullopt, }, /*outputs=*/ { { DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, }, - mk_value(20, TensorSlotName::INPUT, input_grad_binding, std::nullopt), + mk_value(20, + TensorSlotName::INPUT, + input_grad_binding, + std::nullopt), }, }, }; @@ -997,18 +1049,17 @@ TEST_SUITE(FF_TEST_SUITE) { generate_shard_expansion_for_invocation(input); auto mk_output_grad_binding = [&](global_device_id_t const &device) - -> std::pair - { + -> std::pair { return { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::OUTPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/device.coord, - }, - DynamicValueAttrsShardingInfo{ - output_grad_binding.at_r(device), - device, - }, + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::OUTPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/device.coord, + }, + DynamicValueAttrsShardingInfo{ + output_grad_binding.at_r(device), + device, + }, }; }; @@ -1017,32 +1068,31 @@ TEST_SUITE(FF_TEST_SUITE) { std::set const &output_grad_task_shards, ParallelTensorSpaceCoordinate const &input_grad_shard_coord) -> DynamicNodeInvocationShardingInfo { - return DynamicNodeInvocationShardingInfo{ - /*device_ids=*/device_ids, - /*value_sharding=*/ - binary_relation_from_map( - binary_merge_disjoint_maps( + /*device_ids=*/device_ids, + /*value_sharding=*/ + binary_relation_from_map(binary_merge_disjoint_maps( std::map{ - { - DynamicTensorSlot{ - /*slot_name=*/TensorSlotName::INPUT, - /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), - /*task_shard=*/std::nullopt, - }, - DynamicValueAttrsShardingInfo{ - input_grad_shard_coord, - input_grad_binding.at_l(input_grad_shard_coord), + { + DynamicTensorSlot{ + /*slot_name=*/TensorSlotName::INPUT, + /*slot_tensor_role=*/mk_dynamic_tensor_role_bwd(), + /*task_shard=*/std::nullopt, + }, + DynamicValueAttrsShardingInfo{ + input_grad_shard_coord, + input_grad_binding.at_l(input_grad_shard_coord), + }, }, - }, }, - map_from_pairs(transform(output_grad_task_shards, mk_output_grad_binding)))), + map_from_pairs(transform(output_grad_task_shards, + mk_output_grad_binding)))), }; }; std::set correct = { - mk_invocation_shard(nonempty_set{dev1, dev2}, {dev1, dev2}, pt1), - mk_invocation_shard(nonempty_set{dev3, dev4}, {dev3, dev4}, pt2), + mk_invocation_shard(nonempty_set{dev1, dev2}, {dev1, dev2}, pt1), + mk_invocation_shard(nonempty_set{dev3, dev4}, {dev3, dev4}, pt2), }; CHECK(result.size() == correct.size()); diff --git a/lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h b/lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h index ad43cd52b6..c03270c8d7 100644 --- a/lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h +++ b/lib/utils/include/utils/archetypes/jsonable_ordered_value_type.h @@ -56,7 +56,8 @@ std::string format_as(jsonable_ordered_value_type const &) { } template -std::ostream &operator<<(std::ostream &s, jsonable_ordered_value_type const &x) { +std::ostream &operator<<(std::ostream &s, + jsonable_ordered_value_type const &x) { PANIC(); } @@ -70,7 +71,8 @@ struct adl_serializer<::FlexFlow::jsonable_ordered_value_type> { PANIC(); } - static void to_json(json &, ::FlexFlow::jsonable_ordered_value_type const &) { + static void to_json(json &, + ::FlexFlow::jsonable_ordered_value_type const &) { PANIC(); } }; @@ -81,7 +83,8 @@ namespace std { template struct hash<::FlexFlow::jsonable_ordered_value_type> { - size_t operator()(::FlexFlow::jsonable_ordered_value_type const &) const { + size_t + operator()(::FlexFlow::jsonable_ordered_value_type const &) const { PANIC(); }; }; diff --git a/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h b/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h index 3125c95eb1..4794f6b393 100644 --- a/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_from_unstructured_relation.h @@ -2,11 +2,11 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_ALGORITHMS_BIDICT_FROM_UNSTRUCTURED_RELATION_H #include "utils/bidict/bidict.h" -#include "utils/containers/transform.h" -#include "utils/containers/get_element_counts.h" -#include #include "utils/containers/filter_values.h" +#include "utils/containers/get_element_counts.h" #include "utils/containers/multiset_of.h" +#include "utils/containers/transform.h" +#include namespace FlexFlow { @@ -15,43 +15,33 @@ bidict bidict_from_unstructured_relation( std::set> const &relation) { { - std::multiset l_values = transform( - multiset_of(relation), - [](std::pair const &p) -> L { - return p.first; - } - ); + std::multiset l_values = + transform(multiset_of(relation), + [](std::pair const &p) -> L { return p.first; }); std::map l_value_counts = get_element_counts(l_values); std::map duplicated_element_counts = - filter_values(l_value_counts, - [](positive_int num_occurences) -> bool { - return num_occurences > 1; - }); + filter_values(l_value_counts, [](positive_int num_occurences) -> bool { + return num_occurences > 1; + }); - ASSERT(duplicated_element_counts.empty(), - duplicated_element_counts); + ASSERT(duplicated_element_counts.empty(), duplicated_element_counts); } { - std::multiset r_values = transform( - multiset_of(relation), - [](std::pair const &p) -> R { - return p.second; - } - ); + std::multiset r_values = + transform(multiset_of(relation), + [](std::pair const &p) -> R { return p.second; }); std::map r_value_counts = get_element_counts(r_values); std::map duplicated_element_counts = - filter_values(r_value_counts, - [](positive_int num_occurences) -> bool { - return num_occurences > 1; - }); + filter_values(r_value_counts, [](positive_int num_occurences) -> bool { + return num_occurences > 1; + }); - ASSERT(duplicated_element_counts.empty(), - duplicated_element_counts); + ASSERT(duplicated_element_counts.empty(), duplicated_element_counts); } bidict result; diff --git a/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h b/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h index 25731f3d89..ff4e53080f 100644 --- a/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_transform_keys_and_values.h @@ -11,7 +11,8 @@ template , typename V2 = std::invoke_result_t> -bidict bidict_transform_keys_and_values(bidict const &m, KF &&kf, VF &&vf) { +bidict + bidict_transform_keys_and_values(bidict const &m, KF &&kf, VF &&vf) { bidict result; for (auto const &kv : m) { result.equate_strict(kf(kv.first), vf(kv.second)); diff --git a/lib/utils/include/utils/bidict/algorithms/bidict_unordered_set_of.h b/lib/utils/include/utils/bidict/algorithms/bidict_unordered_set_of.h index 251573d441..a59da5c906 100644 --- a/lib/utils/include/utils/bidict/algorithms/bidict_unordered_set_of.h +++ b/lib/utils/include/utils/bidict/algorithms/bidict_unordered_set_of.h @@ -7,7 +7,8 @@ namespace FlexFlow { template -std::unordered_set> bidict_unordered_set_of(bidict const &c) { +std::unordered_set> + bidict_unordered_set_of(bidict const &c) { std::unordered_set> result; for (auto const &lr : c) { diff --git a/lib/utils/include/utils/bidict/bidict.h b/lib/utils/include/utils/bidict/bidict.h index 79c88d6bd9..4e1c430fa0 100644 --- a/lib/utils/include/utils/bidict/bidict.h +++ b/lib/utils/include/utils/bidict/bidict.h @@ -1,22 +1,22 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_BIDICT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BIDICT_BIDICT_H +#include "utils/check_fmtable.h" +#include "utils/containers/contains_key.h" +#include "utils/containers/keys.h" #include "utils/containers/map_from_keys_and_values.h" +#include "utils/containers/require_same.h" +#include "utils/containers/set_of.h" +#include "utils/containers/unordered_map_from_map.h" +#include "utils/containers/values.h" +#include "utils/fmt/map.h" +#include "utils/hash/map.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" #include #include #include #include -#include "utils/containers/require_same.h" -#include "utils/containers/values.h" -#include "utils/containers/contains_key.h" -#include "utils/containers/unordered_map_from_map.h" -#include "utils/check_fmtable.h" -#include "utils/containers/keys.h" -#include "utils/containers/set_of.h" -#include "utils/hash/map.h" -#include "utils/fmt/map.h" namespace FlexFlow { @@ -95,17 +95,13 @@ struct bidict { } bool operator==(bidict const &other) const { - return require_same( - (this->fwd_map == other.fwd_map), - (this->bwd_map == other.bwd_map) - ); + return require_same((this->fwd_map == other.fwd_map), + (this->bwd_map == other.bwd_map)); } bool operator!=(bidict const &other) const { - return require_same( - (this->fwd_map != other.fwd_map), - (this->bwd_map != other.bwd_map) - ); + return require_same((this->fwd_map != other.fwd_map), + (this->bwd_map != other.bwd_map)); } R const &at_l(L const &l) const { @@ -221,7 +217,7 @@ struct bidict { return this->fwd_map; } - operator std::unordered_map () const { + operator std::unordered_map() const { return unordered_map_from_map(this->fwd_map); } @@ -241,8 +237,7 @@ struct bidict { return this->bwd_map; } - bidict(std::map const &fwd_map, - std::map const &bwd_map) + bidict(std::map const &fwd_map, std::map const &bwd_map) : fwd_map(fwd_map), bwd_map(bwd_map) {} bool operator<(bidict const &other) const { @@ -260,6 +255,7 @@ struct bidict { bool operator>=(bidict const &other) const { return this->fwd_map >= other.fwd_map; } + private: void check_invariants() const { std::set fwd_l_vals = keys(this->fwd_map); diff --git a/lib/utils/include/utils/binary_relation/binary_relation.h b/lib/utils/include/utils/binary_relation/binary_relation.h index 0d09045ec2..4619eba195 100644 --- a/lib/utils/include/utils/binary_relation/binary_relation.h +++ b/lib/utils/include/utils/binary_relation/binary_relation.h @@ -1,20 +1,20 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_H -#include -#include "utils/json/check_is_json_deserializable.h" +#include "utils/containers/filtrans.h" +#include "utils/containers/multiset_of.h" +#include "utils/containers/transform.h" +#include "utils/fmt/pair.h" +#include "utils/fmt/set.h" #include "utils/hash-utils.h" -#include "utils/hash/set.h" #include "utils/hash/pair.h" +#include "utils/hash/set.h" #include "utils/hash/tuple.h" -#include "utils/fmt/set.h" -#include "utils/fmt/pair.h" +#include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" #include #include -#include "utils/containers/multiset_of.h" -#include "utils/containers/filtrans.h" -#include "utils/containers/transform.h" +#include namespace FlexFlow { @@ -22,8 +22,7 @@ template struct BinaryRelation { BinaryRelation() : raw{} {} - BinaryRelation(std::set> const &raw) - : raw(raw) {} + BinaryRelation(std::set> const &raw) : raw(raw) {} BinaryRelation(std::initializer_list> init) : BinaryRelation(init.begin(), init.end()) {} @@ -68,59 +67,45 @@ struct BinaryRelation { } std::set at_l(L const &l) const { - return filtrans( - this->raw, - [&](std::pair const &p) -> std::optional { - if (p.first == l) { - return p.second; - } else { - return std::nullopt; - } - }); + return filtrans(this->raw, + [&](std::pair const &p) -> std::optional { + if (p.first == l) { + return p.second; + } else { + return std::nullopt; + } + }); } std::set at_r(R const &r) const { - return filtrans( - this->raw, - [&](std::pair const &p) -> std::optional { - if (p.second == r) { - return p.first; - } else { - return std::nullopt; - } - }); + return filtrans(this->raw, + [&](std::pair const &p) -> std::optional { + if (p.second == r) { + return p.first; + } else { + return std::nullopt; + } + }); } std::multiset left_value_occurences() const { - return transform( - multiset_of(this->raw), - [&](std::pair const &p) -> L { - return p.first; - }); + return transform(multiset_of(this->raw), + [&](std::pair const &p) -> L { return p.first; }); } std::multiset right_value_occurences() const { - return transform( - multiset_of(this->raw), - [&](std::pair const &p) -> R { - return p.second; - }); + return transform(multiset_of(this->raw), + [&](std::pair const &p) -> R { return p.second; }); } std::set left_values() const { - return transform( - this->raw, - [&](std::pair const &p) -> L { - return p.first; - }); + return transform(this->raw, + [&](std::pair const &p) -> L { return p.first; }); } std::set right_values() const { - return transform( - this->raw, - [&](std::pair const &p) -> R { - return p.second; - }); + return transform(this->raw, + [&](std::pair const &p) -> R { return p.second; }); } std::size_t size() const { @@ -134,12 +119,12 @@ struct BinaryRelation { std::set> const &unwrap_as_set() const { return this->raw; } + private: std::set> raw; private: - std::tuple - tie() const { + std::tuple tie() const { return std::tie(this->raw); } @@ -147,8 +132,7 @@ struct BinaryRelation { }; template -std::set> - format_as(BinaryRelation const &m) { +std::set> format_as(BinaryRelation const &m) { return m.unwrap_as_set(); } diff --git a/lib/utils/include/utils/binary_relation/binary_relation_from_map.h b/lib/utils/include/utils/binary_relation/binary_relation_from_map.h index 7ddb857722..151490fde3 100644 --- a/lib/utils/include/utils/binary_relation/binary_relation_from_map.h +++ b/lib/utils/include/utils/binary_relation/binary_relation_from_map.h @@ -1,16 +1,16 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_FROM_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_FROM_MAP_H -#include #include "utils/binary_relation/binary_relation.h" #include "utils/containers/set_of.h" +#include namespace FlexFlow { template BinaryRelation binary_relation_from_map(std::map const &m) { return BinaryRelation{ - set_of(m), + set_of(m), }; } diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_left.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_left.h index 60669bbfc1..e4bc65aefb 100644 --- a/lib/utils/include/utils/binary_relation/binary_relation_transform_left.h +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_left.h @@ -1,19 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT_H -#include #include "utils/binary_relation/binary_relation.h" +#include namespace FlexFlow { -template < - typename L, - typename R, - typename F, - typename L2 = std::invoke_result_t -> -BinaryRelation binary_relation_transform_left(BinaryRelation const &rel, - F &&f) { +template > +BinaryRelation + binary_relation_transform_left(BinaryRelation const &rel, F &&f) { BinaryRelation result; for (std::pair const &p : rel.unwrap_as_set()) { diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h index a317992edf..8f9dcd4100 100644 --- a/lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_left2.h @@ -1,19 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT2_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_LEFT2_H -#include #include "utils/binary_relation/binary_relation.h" +#include namespace FlexFlow { -template < - typename L, - typename R, - typename F, - typename L2 = std::invoke_result_t -> -BinaryRelation binary_relation_transform_left2(BinaryRelation const &rel, - F &&f) { +template > +BinaryRelation + binary_relation_transform_left2(BinaryRelation const &rel, F &&f) { BinaryRelation result; for (std::pair const &p : rel.unwrap_as_set()) { diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_right.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_right.h index a2e36a7441..e13f09f9dd 100644 --- a/lib/utils/include/utils/binary_relation/binary_relation_transform_right.h +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_right.h @@ -1,19 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT_H -#include #include "utils/binary_relation/binary_relation.h" +#include namespace FlexFlow { -template < - typename L, - typename R, - typename F, - typename R2 = std::invoke_result_t -> -BinaryRelation binary_relation_transform_right(BinaryRelation const &rel, - F &&f) { +template > +BinaryRelation + binary_relation_transform_right(BinaryRelation const &rel, F &&f) { BinaryRelation result; for (std::pair const &p : rel.unwrap_as_set()) { diff --git a/lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h b/lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h index 2ca789a8e0..6e73975d6a 100644 --- a/lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h +++ b/lib/utils/include/utils/binary_relation/binary_relation_transform_right2.h @@ -1,19 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT2_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_BINARY_RELATION_TRANSFORM_RIGHT2_H -#include #include "utils/binary_relation/binary_relation.h" +#include namespace FlexFlow { -template < - typename L, - typename R, - typename F, - typename R2 = std::invoke_result_t -> -BinaryRelation binary_relation_transform_right2(BinaryRelation const &rel, - F &&f) { +template > +BinaryRelation + binary_relation_transform_right2(BinaryRelation const &rel, F &&f) { BinaryRelation result; for (std::pair const &p : rel.unwrap_as_set()) { diff --git a/lib/utils/include/utils/binary_relation/filter_binary_relation.h b/lib/utils/include/utils/binary_relation/filter_binary_relation.h index 2f8026a3ed..065af2ba05 100644 --- a/lib/utils/include/utils/binary_relation/filter_binary_relation.h +++ b/lib/utils/include/utils/binary_relation/filter_binary_relation.h @@ -7,12 +7,13 @@ namespace FlexFlow { template -BinaryRelation filter_binary_relation(BinaryRelation const &rel, F &&f) { +BinaryRelation filter_binary_relation(BinaryRelation const &rel, + F &&f) { return BinaryRelation{ - filter(rel.unwrap_as_set(), - [&](std::pair const &p) -> bool { - return f(p.first, p.second); - }), + filter(rel.unwrap_as_set(), + [&](std::pair const &p) -> bool { + return f(p.first, p.second); + }), }; } diff --git a/lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h b/lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h index 152e0db15b..7ce9f00f75 100644 --- a/lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h +++ b/lib/utils/include/utils/binary_relation/require_binary_relation_is_left_unique.h @@ -7,7 +7,8 @@ namespace FlexFlow { template -OneToMany require_binary_relation_is_left_unique(BinaryRelation const &rel) { +OneToMany + require_binary_relation_is_left_unique(BinaryRelation const &rel) { OneToMany result; for (std::pair const &p : rel.unwrap_as_set()) { diff --git a/lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h b/lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h index 5d765bbb64..9c05e17de8 100644 --- a/lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h +++ b/lib/utils/include/utils/binary_relation/require_binary_relation_is_right_unique.h @@ -7,7 +7,8 @@ namespace FlexFlow { template -ManyToOne require_binary_relation_is_right_unique(BinaryRelation const &rel) { +ManyToOne + require_binary_relation_is_right_unique(BinaryRelation const &rel) { ManyToOne result; for (std::pair const &p : rel.unwrap_as_set()) { diff --git a/lib/utils/include/utils/binary_relation/transform_binary_relation.h b/lib/utils/include/utils/binary_relation/transform_binary_relation.h index 538c477ea1..cf69a97d0a 100644 --- a/lib/utils/include/utils/binary_relation/transform_binary_relation.h +++ b/lib/utils/include/utils/binary_relation/transform_binary_relation.h @@ -1,20 +1,18 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_TRANSFORM_BINARY_RELATION_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_BINARY_RELATION_TRANSFORM_BINARY_RELATION_H -#include #include "utils/binary_relation/binary_relation.h" +#include namespace FlexFlow { -template < - typename L, - typename R, - typename F, - typename L2 = std::invoke_result_t::first_type, - typename R2 = std::invoke_result_t::second_type -> -BinaryRelation binary_relation_transform_left(BinaryRelation const &rel, - F &&f) { +template ::first_type, + typename R2 = std::invoke_result_t::second_type> +BinaryRelation + binary_relation_transform_left(BinaryRelation const &rel, F &&f) { BinaryRelation result; for (std::pair const &p : rel.unwrap_as_set()) { diff --git a/lib/utils/include/utils/containers/are_disjoint.h b/lib/utils/include/utils/containers/are_disjoint.h index 3be7712dd2..2fd2813572 100644 --- a/lib/utils/include/utils/containers/are_disjoint.h +++ b/lib/utils/include/utils/containers/are_disjoint.h @@ -12,8 +12,7 @@ bool are_disjoint(std::unordered_set const &l, } template -bool are_disjoint(std::set const &l, - std::set const &r) { +bool are_disjoint(std::set const &l, std::set const &r) { return set_intersection(l, r).empty(); } diff --git a/lib/utils/include/utils/containers/binary_cartesian_product.h b/lib/utils/include/utils/containers/binary_cartesian_product.h index 471f99cc6e..3e92031620 100644 --- a/lib/utils/include/utils/containers/binary_cartesian_product.h +++ b/lib/utils/include/utils/containers/binary_cartesian_product.h @@ -6,9 +6,8 @@ namespace FlexFlow { template -std::set> - binary_cartesian_product(std::set const &lhs, - std::set const &rhs) { +std::set> binary_cartesian_product(std::set const &lhs, + std::set const &rhs) { std::set> result; for (A const &a : lhs) { diff --git a/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h b/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h index ff3e798897..c0854017db 100644 --- a/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h +++ b/lib/utils/include/utils/containers/binary_merge_disjoint_maps.h @@ -2,16 +2,15 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_DISJOINT_MAPS_H #include "utils/containers/binary_merge_maps_with.h" -#include #include "utils/containers/keys.h" #include "utils/containers/set_intersection.h" +#include namespace FlexFlow { template -std::map - binary_merge_disjoint_maps(std::map const &lhs, - std::map const &rhs) { +std::map binary_merge_disjoint_maps(std::map const &lhs, + std::map const &rhs) { std::set lhs_keys = keys(lhs); std::set rhs_keys = keys(rhs); diff --git a/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h b/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h index 761465532e..b51ad0787a 100644 --- a/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h +++ b/lib/utils/include/utils/containers/binary_merge_disjoint_unordered_maps.h @@ -1,17 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_DISJOINT_UNORDERED_MAPS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_DISJOINT_UNORDERED_MAPS_H -#include #include "utils/containers/binary_merge_unordered_maps_with.h" -#include "utils/containers/unordered_keys.h" #include "utils/containers/set_intersection.h" +#include "utils/containers/unordered_keys.h" +#include namespace FlexFlow { template std::unordered_map binary_merge_disjoint_unordered_maps(std::unordered_map const &lhs, - std::unordered_map const &rhs) { + std::unordered_map const &rhs) { std::unordered_set lhs_keys = unordered_keys(lhs); std::unordered_set rhs_keys = unordered_keys(rhs); diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with.h b/lib/utils/include/utils/containers/binary_merge_maps_with.h index eb404f02bf..5291459927 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with.h @@ -12,10 +12,9 @@ namespace FlexFlow { template -std::map - binary_merge_maps_with(std::map const &lhs, - std::map const &rhs, - F &&f) { +std::map binary_merge_maps_with(std::map const &lhs, + std::map const &rhs, + F &&f) { std::set l_keys = keys(lhs); std::set r_keys = keys(rhs); diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h b/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h index 25be62d9c3..5f1bfb482e 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with_left_dominating.h @@ -6,8 +6,9 @@ namespace FlexFlow { template -std::map binary_merge_maps_with_left_dominating( - std::map const &lhs, std::map const &rhs) { +std::map + binary_merge_maps_with_left_dominating(std::map const &lhs, + std::map const &rhs) { std::map result; merge_in_map(rhs, result); merge_in_map(lhs, result); diff --git a/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h b/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h index e4bfdd6d29..40e777a1d8 100644 --- a/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h +++ b/lib/utils/include/utils/containers/binary_merge_maps_with_right_dominating.h @@ -6,8 +6,9 @@ namespace FlexFlow { template -std::map binary_merge_maps_with_right_dominating( - std::map const &lhs, std::map const &rhs) { +std::map + binary_merge_maps_with_right_dominating(std::map const &lhs, + std::map const &rhs) { std::map result; merge_in_map(lhs, result); merge_in_map(rhs, result); diff --git a/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h index ef4ccebea3..4ba42d37db 100644 --- a/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h +++ b/lib/utils/include/utils/containers/binary_merge_unordered_maps_with.h @@ -2,11 +2,11 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_BINARY_MERGE_UNORDERED_MAPS_WITH_H #include "utils/containers/generate_unordered_map.h" -#include "utils/containers/set_intersection.h" -#include "utils/containers/unordered_keys.h" #include "utils/containers/merge_unordered_maps_with_right_dominating.h" #include "utils/containers/restrict_keys.h" +#include "utils/containers/set_intersection.h" #include "utils/containers/set_minus.h" +#include "utils/containers/unordered_keys.h" #include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/filter_values.h b/lib/utils/include/utils/containers/filter_values.h index 6d4be1254f..5b61ecad92 100644 --- a/lib/utils/include/utils/containers/filter_values.h +++ b/lib/utils/include/utils/containers/filter_values.h @@ -6,8 +6,7 @@ namespace FlexFlow { template -std::map filter_values(std::map const &m, - F const &f) { +std::map filter_values(std::map const &m, F const &f) { std::map result; for (auto const &kv : m) { if (f(kv.second)) { diff --git a/lib/utils/include/utils/containers/filtrans.h b/lib/utils/include/utils/containers/filtrans.h index 76e3bfaa84..6402dce195 100644 --- a/lib/utils/include/utils/containers/filtrans.h +++ b/lib/utils/include/utils/containers/filtrans.h @@ -87,7 +87,8 @@ std::multiset filtrans(std::multiset const &s, F &&f) { template >> -std::unordered_multiset filtrans(std::unordered_multiset const &s, F &&f) { +std::unordered_multiset filtrans(std::unordered_multiset const &s, + F &&f) { std::unordered_multiset result; for (In const &i : s) { diff --git a/lib/utils/include/utils/containers/find.h b/lib/utils/include/utils/containers/find.h index 68226479df..373b3295fd 100644 --- a/lib/utils/include/utils/containers/find.h +++ b/lib/utils/include/utils/containers/find.h @@ -13,8 +13,7 @@ typename Container::const_iterator } template -typename std::set::const_iterator - find(std::set const &c, V const &e) { +typename std::set::const_iterator find(std::set const &c, V const &e) { return c.find(e); } diff --git a/lib/utils/include/utils/containers/flatmap.h b/lib/utils/include/utils/containers/flatmap.h index bd8836d596..ad709d9cb1 100644 --- a/lib/utils/include/utils/containers/flatmap.h +++ b/lib/utils/include/utils/containers/flatmap.h @@ -1,13 +1,13 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_FLATMAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_FLATMAP_H +#include "utils/containers/binary_merge_disjoint_maps.h" +#include "utils/containers/binary_merge_disjoint_unordered_maps.h" #include "utils/containers/extend.h" #include "utils/containers/get_element_type.h" +#include #include #include -#include -#include "utils/containers/binary_merge_disjoint_maps.h" -#include "utils/containers/binary_merge_disjoint_unordered_maps.h" namespace FlexFlow { @@ -90,8 +90,7 @@ template < typename F, typename OutK = typename std::invoke_result_t::key_type, typename OutV = typename std::invoke_result_t::mapped_type> -std::map flatmap(std::map const &m, - F &&f) { +std::map flatmap(std::map const &m, F &&f) { std::map result; for (auto const &[k, v] : m) { diff --git a/lib/utils/include/utils/containers/generate_map.h b/lib/utils/include/utils/containers/generate_map.h index 08bfc86350..3788512fd0 100644 --- a/lib/utils/include/utils/containers/generate_map.h +++ b/lib/utils/include/utils/containers/generate_map.h @@ -14,7 +14,8 @@ template , typename V = std::invoke_result_t> std::map generate_map(C const &c, F &&f) { - static_assert(is_lt_comparable_v, "Key type should be ordered (but is not)"); + static_assert(is_lt_comparable_v, + "Key type should be ordered (but is not)"); auto transformed = vector_transform(vector_of(c), [&](K const &k) -> std::pair { diff --git a/lib/utils/include/utils/containers/get_all_assignments.h b/lib/utils/include/utils/containers/get_all_assignments.h index d1667700aa..d4c8f64475 100644 --- a/lib/utils/include/utils/containers/get_all_assignments.h +++ b/lib/utils/include/utils/containers/get_all_assignments.h @@ -3,21 +3,18 @@ #include "utils/containers/cartesian_product.h" #include "utils/containers/keys.h" -#include "utils/containers/transform.h" #include "utils/containers/map_from_pairs.h" #include "utils/containers/set_of.h" +#include "utils/containers/transform.h" +#include "utils/containers/unordered_keys.h" +#include "utils/containers/unordered_map_from_pairs.h" +#include "utils/containers/unordered_set_of.h" #include "utils/containers/vector_of.h" #include "utils/containers/zip.h" #include "utils/hash/unordered_map.h" #include #include #include -#include "utils/containers/keys.h" -#include "utils/containers/map_from_pairs.h" -#include "utils/containers/set_of.h" -#include "utils/containers/unordered_keys.h" -#include "utils/containers/unordered_set_of.h" -#include "utils/containers/unordered_map_from_pairs.h" namespace FlexFlow { @@ -50,8 +47,8 @@ std::unordered_set> get_all_assignments( * assignment is returned */ template -std::set> get_all_assignments( - std::map> const &options_per_key) { +std::set> + get_all_assignments(std::map> const &options_per_key) { if (options_per_key.empty()) { return {{}}; } @@ -60,16 +57,15 @@ std::set> get_all_assignments( std::vector> ordered_value_option_sets = transform( ordered_keys, [&](K const &k) { return options_per_key.at(k); }); - std::set> result = transform( - set_of(cartesian_product(ordered_value_option_sets)), - [&](std::vector const &chosen_values) { - return map_from_pairs(zip(ordered_keys, chosen_values)); - }); + std::set> result = + transform(set_of(cartesian_product(ordered_value_option_sets)), + [&](std::vector const &chosen_values) { + return map_from_pairs(zip(ordered_keys, chosen_values)); + }); return result; } - } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/get_element_counts.h b/lib/utils/include/utils/containers/get_element_counts.h index 121e5399d5..ddc0a4dc92 100644 --- a/lib/utils/include/utils/containers/get_element_counts.h +++ b/lib/utils/include/utils/containers/get_element_counts.h @@ -2,11 +2,11 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GET_ELEMENT_COUNTS_H #include "utils/containers/contains_key.h" -#include +#include "utils/positive_int/positive_int.h" #include -#include #include -#include "utils/positive_int/positive_int.h" +#include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/get_only.h b/lib/utils/include/utils/containers/get_only.h index 79b799edde..124a7760d6 100644 --- a/lib/utils/include/utils/containers/get_only.h +++ b/lib/utils/include/utils/containers/get_only.h @@ -2,8 +2,8 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_GET_ONLY_H #include "utils/containers/maybe_get_only.h" -#include #include "utils/optional.h" +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/group_by.h b/lib/utils/include/utils/containers/group_by.h index 057467bdfd..4aaf27e741 100644 --- a/lib/utils/include/utils/containers/group_by.h +++ b/lib/utils/include/utils/containers/group_by.h @@ -29,8 +29,7 @@ OneToMany group_by(std::set const &vs, F &&f) { } template > -std::map> group_by(std::vector const &vs, - F &&f) { +std::map> group_by(std::vector const &vs, F &&f) { std::map> result; for (V const &v : vs) { result[f(v)].push_back(v); diff --git a/lib/utils/include/utils/containers/invert_map.h b/lib/utils/include/utils/containers/invert_map.h index c1ada072dd..cd4bd0db28 100644 --- a/lib/utils/include/utils/containers/invert_map.h +++ b/lib/utils/include/utils/containers/invert_map.h @@ -1,17 +1,16 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_INVERT_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_INVERT_MAP_H -#include #include #include +#include #include #include namespace FlexFlow { template -std::map> - invert_map(std::map const &m) { +std::map> invert_map(std::map const &m) { std::map> result; for (auto const &[key, value] : m) { result[value].insert(key); diff --git a/lib/utils/include/utils/containers/invert_unordered_map.h b/lib/utils/include/utils/containers/invert_unordered_map.h index 9efc092227..816293a768 100644 --- a/lib/utils/include/utils/containers/invert_unordered_map.h +++ b/lib/utils/include/utils/containers/invert_unordered_map.h @@ -16,7 +16,6 @@ std::unordered_map> return result; } - } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/is_submapeq_of.h b/lib/utils/include/utils/containers/is_submapeq_of.h index a9d52c871b..479f84abfe 100644 --- a/lib/utils/include/utils/containers/is_submapeq_of.h +++ b/lib/utils/include/utils/containers/is_submapeq_of.h @@ -8,8 +8,7 @@ namespace FlexFlow { template -bool is_submapeq_of(std::map const &sub, - std::map const &m) { +bool is_submapeq_of(std::map const &sub, std::map const &m) { return restrict_keys(m, keys(sub)) == sub; } diff --git a/lib/utils/include/utils/containers/is_superseteq_of.h b/lib/utils/include/utils/containers/is_superseteq_of.h index 7f1580e874..6a01411a31 100644 --- a/lib/utils/include/utils/containers/is_superseteq_of.h +++ b/lib/utils/include/utils/containers/is_superseteq_of.h @@ -7,8 +7,7 @@ namespace FlexFlow { template -bool is_superseteq_of(std::set const &super, - std::set const &sub) { +bool is_superseteq_of(std::set const &super, std::set const &sub) { return is_subseteq_of(sub, super); } diff --git a/lib/utils/include/utils/containers/lift_optional_through_map.h b/lib/utils/include/utils/containers/lift_optional_through_map.h index 03aee73200..9358b62dcf 100644 --- a/lib/utils/include/utils/containers/lift_optional_through_map.h +++ b/lib/utils/include/utils/containers/lift_optional_through_map.h @@ -5,14 +5,14 @@ #include "utils/containers/map_values.h" #include "utils/containers/values.h" #include -#include #include +#include namespace FlexFlow { template -static std::optional> lift_optional_through_map( - std::map> const &m) { +static std::optional> + lift_optional_through_map(std::map> const &m) { ASSERT(!m.empty()); std::multiset> m_values = values(m); diff --git a/lib/utils/include/utils/containers/lookup_in_map.h b/lib/utils/include/utils/containers/lookup_in_map.h index 81161df9e6..990068916c 100644 --- a/lib/utils/include/utils/containers/lookup_in_map.h +++ b/lib/utils/include/utils/containers/lookup_in_map.h @@ -1,12 +1,12 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_LOOKUP_IN_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_LOOKUP_IN_MAP_H -#include "utils/fmt/map.h" #include "utils/containers/contains_key.h" +#include "utils/fmt/map.h" #include -#include -#include #include +#include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/map_from_keys_and_values.h b/lib/utils/include/utils/containers/map_from_keys_and_values.h index 7a9c8f450d..544b1d448d 100644 --- a/lib/utils/include/utils/containers/map_from_keys_and_values.h +++ b/lib/utils/include/utils/containers/map_from_keys_and_values.h @@ -3,15 +3,14 @@ #include "utils/containers/zip.h" #include -#include #include +#include namespace FlexFlow { template -std::map - map_from_keys_and_values(std::vector const &keys, - std::vector const &values) { +std::map map_from_keys_and_values(std::vector const &keys, + std::vector const &values) { ASSERT(keys.size() == values.size()); std::map result; diff --git a/lib/utils/include/utils/containers/map_keys.h b/lib/utils/include/utils/containers/map_keys.h index ff41248a30..d83e988b1b 100644 --- a/lib/utils/include/utils/containers/map_keys.h +++ b/lib/utils/include/utils/containers/map_keys.h @@ -2,13 +2,13 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_KEYS_H #include "utils/containers/keys.h" -#include "utils/containers/unordered_keys.h" #include "utils/containers/transform.h" +#include "utils/containers/unordered_keys.h" #include "utils/containers/unordered_multiset_of.h" #include "utils/exception.h" +#include #include #include -#include namespace FlexFlow { @@ -20,8 +20,7 @@ template > -std::unordered_map map_keys(std::unordered_map const &m, - F &&f) { +std::unordered_map map_keys(std::unordered_map const &m, F &&f) { std::unordered_map result; for (auto const &kv : m) { diff --git a/lib/utils/include/utils/containers/map_keys2.h b/lib/utils/include/utils/containers/map_keys2.h index 90320f848a..4b85a37739 100644 --- a/lib/utils/include/utils/containers/map_keys2.h +++ b/lib/utils/include/utils/containers/map_keys2.h @@ -11,8 +11,7 @@ template > -std::map map_keys2(std::map const &m, - F const &f) { +std::map map_keys2(std::map const &m, F const &f) { std::map result; for (auto const &kv : m) { diff --git a/lib/utils/include/utils/containers/map_keys_and_values.h b/lib/utils/include/utils/containers/map_keys_and_values.h index 1873421e17..eb9b4f88aa 100644 --- a/lib/utils/include/utils/containers/map_keys_and_values.h +++ b/lib/utils/include/utils/containers/map_keys_and_values.h @@ -2,8 +2,8 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_KEYS_AND_VALUES_H #include -#include #include +#include namespace FlexFlow { @@ -33,8 +33,8 @@ template , typename V2 = std::invoke_result_t> -std::map map_keys_and_values( - std::map const &m, FK const &fk, FV const &fv) { +std::map + map_keys_and_values(std::map const &m, FK const &fk, FV const &fv) { std::map result; for (auto const &kv : m) { diff --git a/lib/utils/include/utils/containers/map_keys_with_value_merging.h b/lib/utils/include/utils/containers/map_keys_with_value_merging.h index 34efff1de5..8a07febcca 100644 --- a/lib/utils/include/utils/containers/map_keys_with_value_merging.h +++ b/lib/utils/include/utils/containers/map_keys_with_value_merging.h @@ -2,8 +2,8 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_KEYS_WITH_VALUE_MERGING_H #include "utils/containers/contains_key.h" -#include #include +#include namespace FlexFlow { @@ -38,8 +38,9 @@ template > -std::map map_keys_with_value_merging( - std::map const &m, F &&key_func, MergeF &&merge_values) { +std::map map_keys_with_value_merging(std::map const &m, + F &&key_func, + MergeF &&merge_values) { std::map result; diff --git a/lib/utils/include/utils/containers/map_values.h b/lib/utils/include/utils/containers/map_values.h index 575fff977e..5196f7449f 100644 --- a/lib/utils/include/utils/containers/map_values.h +++ b/lib/utils/include/utils/containers/map_values.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_VALUES_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_VALUES_H +#include #include #include -#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/map_values2.h b/lib/utils/include/utils/containers/map_values2.h index dd943b02bb..3d571181cb 100644 --- a/lib/utils/include/utils/containers/map_values2.h +++ b/lib/utils/include/utils/containers/map_values2.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_VALUES2_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MAP_VALUES2_H +#include #include #include -#include namespace FlexFlow { @@ -24,8 +24,7 @@ template > -std::map map_values2(std::map const &m, - F &&f) { +std::map map_values2(std::map const &m, F &&f) { std::map result; for (std::pair const &kv : m) { result.insert(std::pair{kv.first, f(kv.first, kv.second)}); @@ -33,7 +32,6 @@ std::map map_values2(std::map const &m, return result; } - } // namespace FlexFlow #endif diff --git a/lib/utils/include/utils/containers/merge_disjoint_maps.h b/lib/utils/include/utils/containers/merge_disjoint_maps.h index b541fdbd53..c15cd6828e 100644 --- a/lib/utils/include/utils/containers/merge_disjoint_maps.h +++ b/lib/utils/include/utils/containers/merge_disjoint_maps.h @@ -13,8 +13,7 @@ std::map merge_disjoint_maps(C const &c) { std::map empty = {}; return foldl(c, /*init=*/empty, - [](std::map const &lhs, - std::map const &rhs) { + [](std::map const &lhs, std::map const &rhs) { return binary_merge_disjoint_maps(lhs, rhs); }); } diff --git a/lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h b/lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h index 1bd7fb019b..7f0b220bb7 100644 --- a/lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h +++ b/lib/utils/include/utils/containers/merge_disjoint_unordered_maps.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_DISJOINT_UNORDERED_MAPS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_MERGE_DISJOINT_UNORDERED_MAPS_H -#include "utils/containers/foldl.h" #include "utils/containers/binary_merge_disjoint_unordered_maps.h" +#include "utils/containers/foldl.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/merge_in_map.h b/lib/utils/include/utils/containers/merge_in_map.h index e41c1a6826..13a89a0ee7 100644 --- a/lib/utils/include/utils/containers/merge_in_map.h +++ b/lib/utils/include/utils/containers/merge_in_map.h @@ -6,8 +6,7 @@ namespace FlexFlow { template -void merge_in_map(std::map const &m, - std::map &result) { +void merge_in_map(std::map const &m, std::map &result) { for (auto const &[k, v] : m) { auto it = result.find(k); if (it != result.end()) { diff --git a/lib/utils/include/utils/containers/merge_maps_with.h b/lib/utils/include/utils/containers/merge_maps_with.h index 31658056ce..2dac8f103d 100644 --- a/lib/utils/include/utils/containers/merge_maps_with.h +++ b/lib/utils/include/utils/containers/merge_maps_with.h @@ -9,13 +9,11 @@ namespace FlexFlow { template -std::map - merge_maps_with(std::vector> const &to_merge, - F &&f) { +std::map merge_maps_with(std::vector> const &to_merge, + F &&f) { return foldl(to_merge, std::map{}, - [&](std::map const &accum, - std::map const &m) { + [&](std::map const &accum, std::map const &m) { return binary_merge_maps_with(accum, m, f); }); } diff --git a/lib/utils/include/utils/containers/merge_unordered_maps_with.h b/lib/utils/include/utils/containers/merge_unordered_maps_with.h index fee7fa2fa4..5a22a9031f 100644 --- a/lib/utils/include/utils/containers/merge_unordered_maps_with.h +++ b/lib/utils/include/utils/containers/merge_unordered_maps_with.h @@ -9,9 +9,8 @@ namespace FlexFlow { template -std::unordered_map - merge_unordered_maps_with(std::vector> const &to_merge, - F &&f) { +std::unordered_map merge_unordered_maps_with( + std::vector> const &to_merge, F &&f) { return foldl(to_merge, std::unordered_map{}, [&](std::unordered_map const &accum, diff --git a/lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h b/lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h index 1323378019..9af1f16cfd 100644 --- a/lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h +++ b/lib/utils/include/utils/containers/merge_unordered_maps_with_right_dominating.h @@ -8,7 +8,8 @@ namespace FlexFlow { template -std::unordered_map merge_unordered_maps_with_right_dominating(C const &c) { +std::unordered_map + merge_unordered_maps_with_right_dominating(C const &c) { std::unordered_map result; for (std::unordered_map const &m : c) { diff --git a/lib/utils/include/utils/containers/minimum.h b/lib/utils/include/utils/containers/minimum.h index bd17b50e74..5536a99d11 100644 --- a/lib/utils/include/utils/containers/minimum.h +++ b/lib/utils/include/utils/containers/minimum.h @@ -9,9 +9,7 @@ namespace FlexFlow { template typename C::value_type minimum(C const &c) { ASSERT( - c.size() > 0, - "minimum expected non-empty container but received {}", c - ); + c.size() > 0, "minimum expected non-empty container but received {}", c); return *std::min_element(c.begin(), c.end()); } diff --git a/lib/utils/include/utils/containers/require_only_key.h b/lib/utils/include/utils/containers/require_only_key.h index ef142921ec..f7c86b644e 100644 --- a/lib/utils/include/utils/containers/require_only_key.h +++ b/lib/utils/include/utils/containers/require_only_key.h @@ -3,8 +3,8 @@ #include "utils/containers/contains_key.h" #include -#include #include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/require_same.h b/lib/utils/include/utils/containers/require_same.h index 2f6251064c..a76fefdf75 100644 --- a/lib/utils/include/utils/containers/require_same.h +++ b/lib/utils/include/utils/containers/require_same.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REQUIRE_SAME_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_REQUIRE_SAME_H -#include #include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/require_two_keys.h b/lib/utils/include/utils/containers/require_two_keys.h index 044e6ff21a..6351103e0b 100644 --- a/lib/utils/include/utils/containers/require_two_keys.h +++ b/lib/utils/include/utils/containers/require_two_keys.h @@ -7,9 +7,8 @@ namespace FlexFlow { template -std::pair require_two_keys(std::map const &m, - K const &k1, - K const &k2) { +std::pair + require_two_keys(std::map const &m, K const &k1, K const &k2) { ASSERT(k1 != k2); ASSERT(m.size() == 2); diff --git a/lib/utils/include/utils/containers/restrict_keys.h b/lib/utils/include/utils/containers/restrict_keys.h index 353b2ec237..547586cea9 100644 --- a/lib/utils/include/utils/containers/restrict_keys.h +++ b/lib/utils/include/utils/containers/restrict_keys.h @@ -2,10 +2,10 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_RESTRICT_KEYS_H #include "utils/containers/contains.h" -#include -#include #include #include +#include +#include namespace FlexFlow { @@ -22,8 +22,7 @@ std::unordered_map restrict_keys(std::unordered_map const &m, } template -std::map restrict_keys(std::map const &m, - std::set const &mask) { +std::map restrict_keys(std::map const &m, std::set const &mask) { std::map result; for (auto const &kv : m) { if (contains(mask, kv.first)) { diff --git a/lib/utils/include/utils/containers/set_difference.h b/lib/utils/include/utils/containers/set_difference.h index b4250b8c32..174949e6ad 100644 --- a/lib/utils/include/utils/containers/set_difference.h +++ b/lib/utils/include/utils/containers/set_difference.h @@ -8,8 +8,7 @@ namespace FlexFlow { template -std::set set_difference(std::set const &l, - std::set const &r) { +std::set set_difference(std::set const &l, std::set const &r) { return filter(l, [&](T const &element) { return !contains(r, element); }); } diff --git a/lib/utils/include/utils/containers/set_of.h b/lib/utils/include/utils/containers/set_of.h index 14cb0f8ee0..d262057007 100644 --- a/lib/utils/include/utils/containers/set_of.h +++ b/lib/utils/include/utils/containers/set_of.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_SET_OF_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_SET_OF_H -#include #include +#include namespace FlexFlow { @@ -16,8 +16,7 @@ std::set set_of(C const &c) { } template -std::set> - set_of(std::map const &m) { +std::set> set_of(std::map const &m) { std::set> result; for (auto const &[k, v] : m) { result.insert({k, v}); diff --git a/lib/utils/include/utils/containers/transform.h b/lib/utils/include/utils/containers/transform.h index bb34b2b5a5..7c917e0565 100644 --- a/lib/utils/include/utils/containers/transform.h +++ b/lib/utils/include/utils/containers/transform.h @@ -6,11 +6,11 @@ #include #include #include -#include -#include #include -#include +#include #include +#include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h b/lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h index dbf380eb41..54a4028039 100644 --- a/lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h +++ b/lib/utils/include/utils/containers/try_merge_nondisjoint_maps.h @@ -2,8 +2,8 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_TRY_MERGE_NONDISJOINT_MAPS_H #include "utils/containers/contains_key.h" -#include #include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/unordered_items.h b/lib/utils/include/utils/containers/unordered_items.h index 1bd8da1498..be547e742e 100644 --- a/lib/utils/include/utils/containers/unordered_items.h +++ b/lib/utils/include/utils/containers/unordered_items.h @@ -1,8 +1,8 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_ITEMS_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_UNORDERED_ITEMS_H -#include #include "utils/hash/pair.h" +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h b/lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h index ff916f3704..13cf04a8c9 100644 --- a/lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h +++ b/lib/utils/include/utils/containers/unordered_map_from_keys_and_values.h @@ -11,7 +11,7 @@ namespace FlexFlow { template std::unordered_map unordered_map_from_keys_and_values(std::vector const &keys, - std::vector const &values) { + std::vector const &values) { ASSERT(keys.size() == values.size()); std::unordered_map result; diff --git a/lib/utils/include/utils/containers/vector_from_idx_map.h b/lib/utils/include/utils/containers/vector_from_idx_map.h index fae100dea9..49e856b10c 100644 --- a/lib/utils/include/utils/containers/vector_from_idx_map.h +++ b/lib/utils/include/utils/containers/vector_from_idx_map.h @@ -3,8 +3,8 @@ #include "utils/containers/contains_key.h" #include "utils/nonnegative_int/nonnegative_int.h" -#include #include +#include #include namespace FlexFlow { diff --git a/lib/utils/include/utils/containers/without_nullopts.h b/lib/utils/include/utils/containers/without_nullopts.h index 3c6a2e74d8..3efa8aab57 100644 --- a/lib/utils/include/utils/containers/without_nullopts.h +++ b/lib/utils/include/utils/containers/without_nullopts.h @@ -19,8 +19,7 @@ std::vector without_nullopts(std::vector> const &v) { } template -std::set - without_nullopts(std::set> const &s) { +std::set without_nullopts(std::set> const &s) { std::set result; for (std::optional const &t : s) { if (t.has_value()) { diff --git a/lib/utils/include/utils/containers/zip_values_strict.h b/lib/utils/include/utils/containers/zip_values_strict.h index b891490e39..b05db04cf9 100644 --- a/lib/utils/include/utils/containers/zip_values_strict.h +++ b/lib/utils/include/utils/containers/zip_values_strict.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_CONTAINERS_ZIP_VALUES_STRICT_H +#include "utils/containers/generate_map.h" #include "utils/containers/generate_unordered_map.h" -#include "utils/containers/unordered_keys.h" +#include "utils/containers/keys.h" #include "utils/containers/require_same.h" +#include "utils/containers/unordered_keys.h" #include -#include #include -#include "utils/containers/keys.h" -#include "utils/containers/generate_map.h" +#include namespace FlexFlow { @@ -19,18 +19,18 @@ std::unordered_map> ASSERT(unordered_keys(m1) == unordered_keys(m2)); - return generate_unordered_map(require_same(unordered_keys(m1), unordered_keys(m2)), [&](K const &k) { - return std::pair{ - m1.at(k), - m2.at(k), - }; - }); + return generate_unordered_map( + require_same(unordered_keys(m1), unordered_keys(m2)), [&](K const &k) { + return std::pair{ + m1.at(k), + m2.at(k), + }; + }); } template -std::map> - zip_values_strict(std::map const &m1, - std::map const &m2) { +std::map> zip_values_strict(std::map const &m1, + std::map const &m2) { ASSERT(keys(m1) == keys(m2)); diff --git a/lib/utils/include/utils/containers/zip_values_strict_with.h b/lib/utils/include/utils/containers/zip_values_strict_with.h index 1cd0165e1b..fec142504f 100644 --- a/lib/utils/include/utils/containers/zip_values_strict_with.h +++ b/lib/utils/include/utils/containers/zip_values_strict_with.h @@ -14,10 +14,9 @@ template > -std::map - zip_values_strict_with(std::map const &m1, - std::map const &m2, - F &&f) { +std::map zip_values_strict_with(std::map const &m1, + std::map const &m2, + F &&f) { ASSERT(keys(m1) == keys(m2)); diff --git a/lib/utils/include/utils/deduplicated_priority_queue.h b/lib/utils/include/utils/deduplicated_priority_queue.h index b4a59f69a2..37dbb04223 100644 --- a/lib/utils/include/utils/deduplicated_priority_queue.h +++ b/lib/utils/include/utils/deduplicated_priority_queue.h @@ -4,7 +4,6 @@ #include "utils/containers/contains.h" #include #include -#include #include namespace FlexFlow { diff --git a/lib/utils/include/utils/disjoint_set.h b/lib/utils/include/utils/disjoint_set.h index d8eea1d7a8..2a11fb1d49 100644 --- a/lib/utils/include/utils/disjoint_set.h +++ b/lib/utils/include/utils/disjoint_set.h @@ -5,7 +5,6 @@ #include #include #include -#include namespace FlexFlow { diff --git a/lib/utils/include/utils/dot/dot_file.h b/lib/utils/include/utils/dot/dot_file.h index e427beda22..bc8893363b 100644 --- a/lib/utils/include/utils/dot/dot_file.h +++ b/lib/utils/include/utils/dot/dot_file.h @@ -9,10 +9,9 @@ #include #include #include +#include #include #include -#include -#include #include namespace FlexFlow { diff --git a/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h b/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h index 9b6d981d30..da9f90eafd 100644 --- a/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h +++ b/lib/utils/include/utils/full_binary_tree/find_paths_to_leaf.h @@ -15,31 +15,29 @@ std::set find_paths_to_leaf( Tree const &tree, FullBinaryTreeImplementation const &impl, Leaf const &needle) { - auto visitor = FullBinaryTreeVisitor, - Tree, - Parent, - Leaf>{ - [&](Parent const &parent) -> std::set { - return set_union( - transform( - find_paths_to_leaf(impl.get_left_child(parent), impl, needle), - [](BinaryTreePath const &path) { - return nest_inside_left_child(path); - }), - transform( - find_paths_to_leaf(impl.get_right_child(parent), impl, needle), - [](BinaryTreePath const &path) { - return nest_inside_right_child(path); - })); - }, - [&](Leaf const &leaf) -> std::set { - if (leaf == needle) { - return {binary_tree_root_path()}; - } else { - return {}; - } - }, - }; + auto visitor = + FullBinaryTreeVisitor, Tree, Parent, Leaf>{ + [&](Parent const &parent) -> std::set { + return set_union( + transform(find_paths_to_leaf( + impl.get_left_child(parent), impl, needle), + [](BinaryTreePath const &path) { + return nest_inside_left_child(path); + }), + transform(find_paths_to_leaf( + impl.get_right_child(parent), impl, needle), + [](BinaryTreePath const &path) { + return nest_inside_right_child(path); + })); + }, + [&](Leaf const &leaf) -> std::set { + if (leaf == needle) { + return {binary_tree_root_path()}; + } else { + return {}; + } + }, + }; return visit(tree, impl, visitor); } diff --git a/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h b/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h index f3796d3f8e..053f800666 100644 --- a/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h +++ b/lib/utils/include/utils/full_binary_tree/get_all_leaf_paths.h @@ -15,25 +15,24 @@ template std::set get_all_leaf_paths( Tree const &tree, FullBinaryTreeImplementation const &impl) { - auto visitor = FullBinaryTreeVisitor, - Tree, - Parent, - Leaf>{ - [&](Parent const &parent) -> std::set { - return set_union( - transform(get_all_leaf_paths(impl.get_left_child(parent), impl), - [](BinaryTreePath const &path) { - return nest_inside_left_child(path); - }), - transform(get_all_leaf_paths(impl.get_right_child(parent), impl), - [](BinaryTreePath const &path) { - return nest_inside_right_child(path); - })); - }, - [&](Leaf const &leaf) -> std::set { - return {binary_tree_root_path()}; - }, - }; + auto visitor = + FullBinaryTreeVisitor, Tree, Parent, Leaf>{ + [&](Parent const &parent) -> std::set { + return set_union( + transform(get_all_leaf_paths(impl.get_left_child(parent), impl), + [](BinaryTreePath const &path) { + return nest_inside_left_child(path); + }), + transform( + get_all_leaf_paths(impl.get_right_child(parent), impl), + [](BinaryTreePath const &path) { + return nest_inside_right_child(path); + })); + }, + [&](Leaf const &leaf) -> std::set { + return {binary_tree_root_path()}; + }, + }; return visit(tree, impl, visitor); } diff --git a/lib/utils/include/utils/full_binary_tree/get_leaves.h b/lib/utils/include/utils/full_binary_tree/get_leaves.h index 929b5327f4..2051e4f4ce 100644 --- a/lib/utils/include/utils/full_binary_tree/get_leaves.h +++ b/lib/utils/include/utils/full_binary_tree/get_leaves.h @@ -13,17 +13,13 @@ std::multiset get_leaves(Tree const &tree, FullBinaryTreeImplementation const &impl) { - auto visitor = - FullBinaryTreeVisitor, Tree, Parent, Leaf>{ - [&](Parent const &parent) -> std::multiset { - return multiset_union( - get_leaves(impl.get_left_child(parent), impl), - get_leaves(impl.get_right_child(parent), impl)); - }, - [](Leaf const &leaf) -> std::multiset { - return {leaf}; - }, - }; + auto visitor = FullBinaryTreeVisitor, Tree, Parent, Leaf>{ + [&](Parent const &parent) -> std::multiset { + return multiset_union(get_leaves(impl.get_left_child(parent), impl), + get_leaves(impl.get_right_child(parent), impl)); + }, + [](Leaf const &leaf) -> std::multiset { return {leaf}; }, + }; return visit(tree, impl, visitor); } diff --git a/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h b/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h index 15df552bb8..c55d3e00b3 100644 --- a/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h +++ b/lib/utils/include/utils/full_binary_tree/get_path_to_leaf_map.h @@ -17,27 +17,29 @@ std::map get_path_to_leaf_map( Tree const &tree, FullBinaryTreeImplementation const &impl) { - auto visitor = FullBinaryTreeVisitor, - Tree, - Parent, - Leaf>{ - [&](Parent const &parent) -> std::map { - std::map left_map = map_keys( - get_path_to_leaf_map(impl.get_left_child(parent), impl), - [](BinaryTreePath const &p) { return nest_inside_left_child(p); }); - - std::map right_map = map_keys( - get_path_to_leaf_map(impl.get_right_child(parent), impl), - [](BinaryTreePath const &p) { return nest_inside_right_child(p); }); - - return binary_merge_disjoint_maps(left_map, right_map); - }, - [](Leaf const &leaf) -> std::map { - return std::map{ - {binary_tree_root_path(), leaf}, - }; - }, - }; + auto visitor = + FullBinaryTreeVisitor, Tree, Parent, Leaf>{ + [&](Parent const &parent) -> std::map { + std::map left_map = map_keys( + get_path_to_leaf_map(impl.get_left_child(parent), impl), + [](BinaryTreePath const &p) { + return nest_inside_left_child(p); + }); + + std::map right_map = map_keys( + get_path_to_leaf_map(impl.get_right_child(parent), impl), + [](BinaryTreePath const &p) { + return nest_inside_right_child(p); + }); + + return binary_merge_disjoint_maps(left_map, right_map); + }, + [](Leaf const &leaf) -> std::map { + return std::map{ + {binary_tree_root_path(), leaf}, + }; + }, + }; return visit(tree, impl, visitor); } diff --git a/lib/utils/include/utils/graph/algorithms.h b/lib/utils/include/utils/graph/algorithms.h index c54a8418cd..2cc4504e54 100644 --- a/lib/utils/include/utils/graph/algorithms.h +++ b/lib/utils/include/utils/graph/algorithms.h @@ -14,8 +14,7 @@ std::vector add_nodes(Graph &, int); std::vector add_nodes(UndirectedGraph &, int); std::vector add_nodes(DiGraph &, int); -std::set query_nodes(GraphView const &, - std::set const &); +std::set query_nodes(GraphView const &, std::set const &); void remove_node(DiGraph &, Node const &); void remove_node(UndirectedGraph &, Node const &); @@ -38,38 +37,32 @@ bool contains_edge(DiGraphView const &, DirectedEdge const &); bool contains_edge(UndirectedGraphView const &, UndirectedEdge const &); void remove_edges(DiGraph &, std::set const &); -void remove_edges(UndirectedGraph &, - std::set const &); +void remove_edges(UndirectedGraph &, std::set const &); std::set get_edges(UndirectedGraphView const &); std::set get_node_edges(UndirectedGraphView const &, - Node const &); + Node const &); std::set get_node_edges(UndirectedGraphView const &, - Node const &); -std::set - get_node_edges(UndirectedGraphView const &, - std::set const &); + Node const &); +std::set get_node_edges(UndirectedGraphView const &, + std::set const &); -std::set get_neighbors(UndirectedGraphView const &, - Node const &); +std::set get_neighbors(UndirectedGraphView const &, Node const &); std::set get_neighbors(DiGraphView const &, Node const &); -std::vector - get_dfs_ordering(DiGraphView const &, - std::set const &starting_points); +std::vector get_dfs_ordering(DiGraphView const &, + std::set const &starting_points); std::vector get_unchecked_dfs_ordering(DiGraphView const &, std::set const &starting_points); -std::vector - get_bfs_ordering(DiGraphView const &, - std::set const &starting_points); +std::vector get_bfs_ordering(DiGraphView const &, + std::set const &starting_points); std::vector get_unchecked_topological_ordering(DiGraphView const &); -std::set - get_transitive_reduction_delta(DiGraphView const &); +std::set get_transitive_reduction_delta(DiGraphView const &); UndirectedGraphView get_subgraph(UndirectedGraphView const &, std::set const &); diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms.h index 58de430af0..ece50ee24a 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms.h @@ -13,8 +13,7 @@ std::vector get_dataflow_inputs(DataflowGraphView const &, Node const &); std::vector get_outputs(DataflowGraphView const &, Node const &); -std::set - get_all_dataflow_outputs(DataflowGraphView const &); +std::set get_all_dataflow_outputs(DataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h index effce9f3cd..c353dddbd5 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/find_isomorphisms.h @@ -6,8 +6,8 @@ namespace FlexFlow { -std::set - find_isomorphisms(DataflowGraphView const &, DataflowGraphView const &); +std::set find_isomorphisms(DataflowGraphView const &, + DataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h index aa4408bc4f..8481d65822 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_incoming_edges.h @@ -7,9 +7,8 @@ namespace FlexFlow { std::vector get_incoming_edges(DataflowGraphView const &, Node const &); -std::set - get_incoming_edges(DataflowGraphView const &, - std::set const &); +std::set get_incoming_edges(DataflowGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h index d0f03751f6..d595ff2305 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.h @@ -6,10 +6,9 @@ namespace FlexFlow { std::set get_outgoing_edges(DataflowGraphView const &, - Node const &); -std::set - get_outgoing_edges(DataflowGraphView const &, - std::set const &); + Node const &); +std::set get_outgoing_edges(DataflowGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h index 0ec7d26796..97efea3424 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.h @@ -5,9 +5,8 @@ namespace FlexFlow { -std::set - get_subgraph_incoming_edges(DataflowGraphView const &, - std::set const &); +std::set get_subgraph_incoming_edges(DataflowGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h index 6a4898c341..f06d1d2294 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h +++ b/lib/utils/include/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.h @@ -5,9 +5,8 @@ namespace FlexFlow { -std::set - get_subgraph_outgoing_edges(DataflowGraphView const &, - std::set const &); +std::set get_subgraph_outgoing_edges(DataflowGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h index 5f3890e42d..2f4995d79f 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h +++ b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph.h @@ -19,8 +19,7 @@ struct DataflowGraph : virtual public DataflowGraphView { std::set query_nodes(NodeQuery const &) const; std::set query_edges(DataflowEdgeQuery const &) const; - std::set - query_outputs(DataflowOutputQuery const &) const; + std::set query_outputs(DataflowOutputQuery const &) const; template static typename std::enable_if::value, diff --git a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h index dce09008d9..2b63dcbc12 100644 --- a/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/dataflow_graph/dataflow_graph_view.h @@ -14,8 +14,7 @@ struct DataflowGraphView : virtual public DiGraphView { std::set query_nodes(NodeQuery const &) const; std::set query_edges(DataflowEdgeQuery const &) const; - std::set - query_outputs(DataflowOutputQuery const &) const; + std::set query_outputs(DataflowOutputQuery const &) const; template static typename std::enable_if::value, diff --git a/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h b/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h index 9fe8dc6fb7..a538276962 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/contract_node.h @@ -12,8 +12,7 @@ struct ContractNodeView : public IDiGraphView { Node const &into) : g(g), from(removed), to(into) {} - std::set - query_edges(DirectedEdgeQuery const &) const override; + std::set query_edges(DirectedEdgeQuery const &) const override; std::set query_nodes(NodeQuery const &) const override; ContractNodeView *clone() const override; diff --git a/lib/utils/include/utils/graph/digraph/algorithms/flipped.h b/lib/utils/include/utils/graph/digraph/algorithms/flipped.h index a19dac4908..18dbe0f773 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/flipped.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/flipped.h @@ -10,8 +10,7 @@ struct FlippedView : public IDiGraphView { FlippedView() = delete; explicit FlippedView(DiGraphView const &); - std::set - query_edges(DirectedEdgeQuery const &) const override; + std::set query_edges(DirectedEdgeQuery const &) const override; std::set query_nodes(NodeQuery const &) const override; FlippedView *clone() const override; diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h b/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h index 6831f4d321..085f7bf551 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_descendants.h @@ -12,8 +12,7 @@ namespace FlexFlow { * @note `starting_node` is not considered to be its own descendant, and is thus * not included in the returned set. **/ -std::set get_descendants(DiGraphView const &g, - Node const &starting_node); +std::set get_descendants(DiGraphView const &g, Node const &starting_node); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h index 6f5bc6fcef..bc2b502381 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators.h @@ -21,8 +21,7 @@ std::set get_dominators(DiGraphView const &, Node const &); * that all edges belonging to the set of nodes now pass through a single * unified node). */ -std::set get_dominators(DiGraphView const &, - std::set const &); +std::set get_dominators(DiGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h index b6845f3041..6b124b8a87 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_dominators_map.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::map> - get_dominators_map(DiGraphView const &); +std::map> get_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h b/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h index 06ab2953da..5f717130ad 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.h @@ -4,10 +4,8 @@ #include "utils/graph/digraph/digraph_view.h" namespace FlexFlow { -std::set - get_edges_from_subgraph_to_subgraph(DiGraphView const &, - std::set const &, - std::set const &); +std::set get_edges_from_subgraph_to_subgraph( + DiGraphView const &, std::set const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h index 60fcbecc36..32a26c6e42 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_imm_dominators_map.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::map> - get_imm_dominators_map(DiGraphView const &); +std::map> get_imm_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h index 53367b1743..6253b7e078 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_incoming_edges.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::set get_incoming_edges(DiGraphView const &, - Node const &); +std::set get_incoming_edges(DiGraphView const &, Node const &); std::map> get_incoming_edges(DiGraphView const &, std::set const &); diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h b/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h index b3880d9322..bd2e25e418 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_outgoing_edges.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::set get_outgoing_edges(DiGraphView const &, - Node const &); +std::set get_outgoing_edges(DiGraphView const &, Node const &); std::map> get_outgoing_edges(DiGraphView const &, std::set const &); diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h index 034294df8a..37f7a06c24 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_post_dominators_map.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::map> - get_post_dominators_map(DiGraphView const &); +std::map> get_post_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h index 4f11595430..2850296a55 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_predecessors.h @@ -5,11 +5,10 @@ namespace FlexFlow { -std::map> - get_predecessors(DiGraphView const &); +std::map> get_predecessors(DiGraphView const &); std::set get_predecessors(DiGraphView const &, Node const &); -std::map> - get_predecessors(DiGraphView const &, std::set const &); +std::map> get_predecessors(DiGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h index 0e957d2df4..9d1fd7c6d1 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::set get_strict_dominators(DiGraphView const &, - Node const &); +std::set get_strict_dominators(DiGraphView const &, Node const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h index 7fb71b9c8d..38933f4ab6 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_strict_dominators_map.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::map> - get_strict_dominators_map(DiGraphView const &); +std::map> get_strict_dominators_map(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h index 74fbf51ea0..bed41dd10b 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.h @@ -5,9 +5,8 @@ namespace FlexFlow { -std::set - get_subgraph_outgoing_edges(DiGraphView const &, - std::set const &); +std::set get_subgraph_outgoing_edges(DiGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h index 90372193fc..2e6a2b198a 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_subgraph_successors.h @@ -5,9 +5,8 @@ namespace FlexFlow { -std::set - get_subgraph_successors(DiGraphView const &, - std::set const &); +std::set get_subgraph_successors(DiGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h b/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h index 23195bb256..947dc72bbd 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_successors.h @@ -5,11 +5,10 @@ namespace FlexFlow { -std::map> - get_successors(DiGraphView const &); +std::map> get_successors(DiGraphView const &); std::set get_successors(DiGraphView const &, Node const &); -std::map> - get_successors(DiGraphView const &, std::set const &); +std::map> get_successors(DiGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h b/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h index 9d2343d923..be211a252e 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/get_weakly_connected_components.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::set> - get_weakly_connected_components(DiGraphView const &); +std::set> get_weakly_connected_components(DiGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h b/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h index 5d71737e1f..b391fb274a 100644 --- a/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h +++ b/lib/utils/include/utils/graph/digraph/algorithms/transitive_reduction.h @@ -11,8 +11,7 @@ struct DirectedEdgeMaskView final : public IDiGraphView { explicit DirectedEdgeMaskView(DiGraphView const &, std::set const &); - std::set - query_edges(DirectedEdgeQuery const &) const override; + std::set query_edges(DirectedEdgeQuery const &) const override; std::set query_nodes(NodeQuery const &) const override; DirectedEdgeMaskView *clone() const override; diff --git a/lib/utils/include/utils/graph/instances/adjacency_digraph.h b/lib/utils/include/utils/graph/instances/adjacency_digraph.h index 783d6e381c..8df3fa571e 100644 --- a/lib/utils/include/utils/graph/instances/adjacency_digraph.h +++ b/lib/utils/include/utils/graph/instances/adjacency_digraph.h @@ -17,16 +17,14 @@ class AdjacencyDiGraph : public IDiGraph { void remove_node_unsafe(Node const &) override; void add_edge(Edge const &) override; void remove_edge(Edge const &) override; - std::set - query_edges(DirectedEdgeQuery const &) const override; + std::set query_edges(DirectedEdgeQuery const &) const override; std::set query_nodes(NodeQuery const &) const override; AdjacencyDiGraph *clone() const override; private: - AdjacencyDiGraph( - NodeSource const &node_source, - std::map> const &adjacency); + AdjacencyDiGraph(NodeSource const &node_source, + std::map> const &adjacency); NodeSource node_source; std::map> adjacency; diff --git a/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h b/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h index 49ea58a659..84fa344315 100644 --- a/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h +++ b/lib/utils/include/utils/graph/instances/adjacency_multidigraph.h @@ -16,8 +16,7 @@ struct AdjacencyMultiDiGraph final : public IMultiDiGraph { void remove_node(Node const &) override; void remove_edge(MultiDiEdge const &) override; std::set query_nodes(NodeQuery const &) const override; - std::set - query_edges(MultiDiEdgeQuery const &) const override; + std::set query_edges(MultiDiEdgeQuery const &) const override; Node get_multidiedge_src(MultiDiEdge const &) const override; Node get_multidiedge_dst(MultiDiEdge const &) const override; void inplace_materialize_from(MultiDiGraphView const &) override; @@ -28,17 +27,13 @@ struct AdjacencyMultiDiGraph final : public IMultiDiGraph { AdjacencyMultiDiGraph( NodeSource const &, MultiDiEdgeSource const &, - std::map< - Node, - std::map>> const &, + std::map>> const &, std::map> const &); private: NodeSource node_source; MultiDiEdgeSource edge_source; - std::map>> - adjacency; + std::map>> adjacency; std::map> edge_nodes; }; diff --git a/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h b/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h index 189c19d5f8..5052c80eb3 100644 --- a/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h +++ b/lib/utils/include/utils/graph/instances/hashmap_undirected_graph.h @@ -14,8 +14,7 @@ class HashmapUndirectedGraph : public IUndirectedGraph { void remove_node_unsafe(Node const &) override; void add_edge(Edge const &) override; void remove_edge(Edge const &) override; - std::set - query_edges(UndirectedEdgeQuery const &) const override; + std::set query_edges(UndirectedEdgeQuery const &) const override; std::set query_nodes(NodeQuery const &) const override; friend bool operator==(HashmapUndirectedGraph const &, diff --git a/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h index c99bb47a56..a29eb92cb2 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_dataflow_graph.h @@ -39,13 +39,12 @@ struct UnorderedSetDataflowGraph final : virtual public IDataflowGraph, std::vector const &inputs, std::vector const &outputs); - UnorderedSetDataflowGraph( - NodeSource const &node_source, - DataflowGraphInputSource const &graph_input_source, - std::set const &nodes, - std::set const &edges, - std::set const &outputs, - std::set const &graph_inputs); + UnorderedSetDataflowGraph(NodeSource const &node_source, + DataflowGraphInputSource const &graph_input_source, + std::set const &nodes, + std::set const &edges, + std::set const &outputs, + std::set const &graph_inputs); private: NodeSource node_source; diff --git a/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h index d23231bd43..9c29700cae 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_kwarg_dataflow_graph.h @@ -19,26 +19,24 @@ struct UnorderedSetKwargDataflowGraph final : public IKwargDataflowGraph { UnorderedSetKwargDataflowGraph() = default; - KwargNodeAddedResult add_node( - std::map> const &inputs, - std::set const &output_slots) override { + KwargNodeAddedResult + add_node(std::map> const &inputs, + std::set const &output_slots) override { Node new_node = this->node_source.new_node(); - std::map> outputs = - generate_map( - output_slots, - [&](SlotName const &output_slot) -> KwargDataflowOutput { - KwargDataflowOutput output = - KwargDataflowOutput{ - /*node=*/new_node, - /*slot_name=*/output_slot, - }; + std::map> outputs = generate_map( + output_slots, + [&](SlotName const &output_slot) -> KwargDataflowOutput { + KwargDataflowOutput output = KwargDataflowOutput{ + /*node=*/new_node, + /*slot_name=*/output_slot, + }; - this->outputs.insert(output); + this->outputs.insert(output); - return output; - }); + return output; + }); this->add_node_unsafe(new_node, inputs, outputs); @@ -51,8 +49,8 @@ struct UnorderedSetKwargDataflowGraph final void add_node_unsafe( Node const &node, std::map> const &inputs, - std::map> const - &outputs) override { + std::map> const &outputs) + override { this->nodes.insert(node); for (auto const &[input_slot_name, src] : inputs) { diff --git a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h index 00266ca1ca..2e8bce3d8b 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_dataflow_graph.h @@ -23,7 +23,6 @@ #include "utils/graph/open_dataflow_graph/dataflow_graph_input_source.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge_query.h" -#include "utils/containers/keys.h" namespace FlexFlow { @@ -128,9 +127,8 @@ struct UnorderedSetLabelledOpenDataflowGraph final std::set nodes = get_nodes(view); std::set outputs = get_all_dataflow_outputs(view); std::set edges = get_edges(view); - std::map labelled_outputs = - generate_map(outputs, - [&](DataflowOutput const &o) { return view.at(o); }); + std::map labelled_outputs = generate_map( + outputs, [&](DataflowOutput const &o) { return view.at(o); }); this->inputs.clear(); this->nodes = diff --git a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h index b27f5cfd56..e0b32f341a 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.h @@ -5,6 +5,7 @@ #include "utils/containers/enumerate.h" #include "utils/containers/extend.h" #include "utils/containers/generate_map.h" +#include "utils/containers/keys.h" #include "utils/containers/map_values.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" @@ -17,7 +18,6 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge.h" #include "utils/overload.h" -#include "utils/containers/keys.h" namespace FlexFlow { @@ -34,10 +34,10 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final public: UnorderedSetLabelledOpenKwargDataflowGraph() = default; - KwargNodeAddedResult add_node( - NodeLabel const &node_label, - std::map> const &inputs, - std::map const &output_labels) override { + KwargNodeAddedResult + add_node(NodeLabel const &node_label, + std::map> const &inputs, + std::map const &output_labels) override { return this->add_node( node_label, map_values(inputs, @@ -49,8 +49,7 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::map> const + std::map> const &inputs, std::map const &output_labels) override { Node new_node = this->node_source.new_node(); @@ -68,25 +67,23 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final this->edges.insert(in_edge); } - std::map> outputs = - generate_map( - keys(output_labels), - [&](SlotName const &output_slot) -> KwargDataflowOutput { - ValueLabel value_label = output_labels.at(output_slot); + std::map> outputs = generate_map( + keys(output_labels), + [&](SlotName const &output_slot) -> KwargDataflowOutput { + ValueLabel value_label = output_labels.at(output_slot); - KwargDataflowOutput output = - KwargDataflowOutput{ - /*node=*/new_node, - /*slot_name=*/output_slot, - }; + KwargDataflowOutput output = KwargDataflowOutput{ + /*node=*/new_node, + /*slot_name=*/output_slot, + }; - this->outputs.insert({ - output, - value_label, - }); + this->outputs.insert({ + output, + value_label, + }); - return output; - }); + return output; + }); return KwargNodeAddedResult{ /*node=*/new_node, @@ -182,8 +179,8 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final std::set> view_inputs = get_all_kwarg_dataflow_graph_inputs(view); std::set view_nodes = get_nodes(view); - std::set> - view_edges = get_all_open_kwarg_dataflow_edges(view); + std::set> view_edges = + get_all_open_kwarg_dataflow_edges(view); std::set> view_outputs = get_all_kwarg_dataflow_outputs(view); @@ -214,21 +211,18 @@ struct UnorderedSetLabelledOpenKwargDataflowGraph final private: UnorderedSetLabelledOpenKwargDataflowGraph( NodeSource const &node_source, - std::map, - ValueLabel> const &graph_inputs, + std::map, ValueLabel> const + &graph_inputs, std::map const &nodes, - std::set> const - &edges, - std::map, ValueLabel> const - &outputs) + std::set> const &edges, + std::map, ValueLabel> const &outputs) : node_source(node_source), graph_inputs(graph_inputs), nodes(nodes), edges(edges), outputs(outputs) {} private: NodeSource node_source; - std::map, ValueLabel> - graph_inputs; + std::map, ValueLabel> graph_inputs; std::map nodes; std::set> edges; std::map, ValueLabel> outputs; diff --git a/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h index 1013f69691..518120f652 100644 --- a/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.h @@ -16,8 +16,7 @@ struct UnorderedSetOpenKwargDataflowGraph final UnorderedSetOpenKwargDataflowGraph() = default; KwargNodeAddedResult add_node( - std::map> const + std::map> const &inputs, std::set const &output_slots) override { Node new_node = this->node_source.new_node(); @@ -35,20 +34,18 @@ struct UnorderedSetOpenKwargDataflowGraph final this->edges.insert(in_edge); } - std::map> outputs = - generate_map( - output_slots, - [&](SlotName const &output_slot) -> KwargDataflowOutput { - KwargDataflowOutput output = - KwargDataflowOutput{ - /*node=*/new_node, - /*slot_name=*/output_slot, - }; + std::map> outputs = generate_map( + output_slots, + [&](SlotName const &output_slot) -> KwargDataflowOutput { + KwargDataflowOutput output = KwargDataflowOutput{ + /*node=*/new_node, + /*slot_name=*/output_slot, + }; - this->outputs.insert(output); + this->outputs.insert(output); - return output; - }); + return output; + }); return KwargNodeAddedResult{ /*node=*/new_node, @@ -107,11 +104,9 @@ struct UnorderedSetOpenKwargDataflowGraph final private: UnorderedSetOpenKwargDataflowGraph( NodeSource const &node_source, - std::set> const - &graph_inputs, + std::set> const &graph_inputs, std::set const &nodes, - std::set> const - &edges, + std::set> const &edges, std::set> const &outputs) : node_source(node_source), graph_inputs(graph_inputs), nodes(nodes), edges(edges), outputs(outputs) {} diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h index 06b19b6986..45420b4e81 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.h @@ -15,46 +15,43 @@ namespace FlexFlow { template DataflowGraphData dataflow_graph_data_from_kwarg_dataflow_graph_data( KwargDataflowGraphData const &kwarg_data, - std::function( - std::set const &)> const &order_slots) { + std::function(std::set const &)> const + &order_slots) { std::set> all_inputs = transform( kwarg_data.edges, [](KwargDataflowEdge const &e) -> KwargDataflowInput { return e.dst; }); - std::set> all_outputs = - kwarg_data.outputs; + std::set> all_outputs = kwarg_data.outputs; - std::map> - incoming_slots_by_node = map_values( - group_by(all_inputs, - [](KwargDataflowInput const &i) -> Node { - return i.node; - }) - .l_to_r(), - [](nonempty_set> const &is) - -> std::set { - return transform(is.unwrap_as_set(), - [](KwargDataflowInput const &i) { - return i.slot_name; - }); - }); + std::map> incoming_slots_by_node = + map_values(group_by(all_inputs, + [](KwargDataflowInput const &i) -> Node { + return i.node; + }) + .l_to_r(), + [](nonempty_set> const &is) + -> std::set { + return transform(is.unwrap_as_set(), + [](KwargDataflowInput const &i) { + return i.slot_name; + }); + }); - std::map> - outgoing_slots_by_node = map_values( - group_by(all_outputs, - [](KwargDataflowOutput const &o) -> Node { - return o.node; - }) - .l_to_r(), - [](nonempty_set> const &os) - -> std::set { - return transform(os.unwrap_as_set(), - [](KwargDataflowOutput const &o) { - return o.slot_name; - }); - }); + std::map> outgoing_slots_by_node = + map_values(group_by(all_outputs, + [](KwargDataflowOutput const &o) -> Node { + return o.node; + }) + .l_to_r(), + [](nonempty_set> const &os) + -> std::set { + return transform(os.unwrap_as_set(), + [](KwargDataflowOutput const &o) { + return o.slot_name; + }); + }); auto dataflow_input_from_kwarg_input = [&](KwargDataflowInput const &i) -> DataflowInput { diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h index 329b47b6b8..993183dcca 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h @@ -12,8 +12,8 @@ namespace FlexFlow { template DataflowGraphView dataflow_graph_from_kwarg_dataflow_graph( KwargDataflowGraphView const &kwarg_dg, - std::function( - std::set const &)> const &order_slots) { + std::function(std::set const &)> const + &order_slots) { KwargDataflowGraphData kwarg_data = get_kwarg_dataflow_graph_data(kwarg_dg); diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h index ac563d7363..0fbd2d562b 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h @@ -7,9 +7,8 @@ namespace FlexFlow { template -std::set> - get_all_kwarg_dataflow_outputs( - KwargDataflowGraphView const &view) { +std::set> get_all_kwarg_dataflow_outputs( + KwargDataflowGraphView const &view) { return view.query_outputs(kwarg_dataflow_output_query_all()); } diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h index ea5448c96e..eca7ad3f6c 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.h @@ -9,13 +9,12 @@ namespace FlexFlow { template -KwargDataflowGraphView get_kwarg_dataflow_graph_subgraph( - KwargDataflowGraphView const &g, - std::set const &subgraph_nodes) { +KwargDataflowGraphView + get_kwarg_dataflow_graph_subgraph(KwargDataflowGraphView const &g, + std::set const &subgraph_nodes) { KwargDataflowGraphData g_data = get_kwarg_dataflow_graph_data(g); - std::set nodes = - set_intersection(g_data.nodes, set_of(subgraph_nodes)); + std::set nodes = set_intersection(g_data.nodes, set_of(subgraph_nodes)); std::set> edges = filter(g_data.edges, [&](KwargDataflowEdge const &e) -> bool { diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h index 855333e500..92c2e9d0fc 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h @@ -1,12 +1,12 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_KWARG_DATAFLOW_GRAPH_AS_DOT_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_KWARG_DATAFLOW_GRAPH_ALGORITHMS_KWARG_DATAFLOW_GRAPH_AS_DOT_H +#include "utils/containers/set_of.h" #include "utils/graph/dataflow_graph/algorithms/dataflow_graph_as_dot.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_incoming_slots_for_node.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_outgoing_slots_for_node.h" #include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph_view.h" -#include "utils/containers/set_of.h" namespace FlexFlow { @@ -17,8 +17,8 @@ std::string kwarg_dataflow_graph_as_dot( std::function const &)> const &render_value, std::function const &render_slot_name, - std::function( - std::set const &)> const &order_slots) { + std::function(std::set const &)> const + &order_slots) { std::function get_input_label = [&](DataflowInput const &i) -> nlohmann::json { std::vector slot_ordering = diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h index 18ea5baa51..b3f2adf40f 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_data.h @@ -14,8 +14,7 @@ void require_kwarg_dataflow_graph_data_is_valid( KwargDataflowGraphData const &data) { std::set nodes_from_edges = flatmap( - data.edges, - [](KwargDataflowEdge const &e) -> std::set { + data.edges, [](KwargDataflowEdge const &e) -> std::set { return std::set{ e.src.node, e.dst.node, diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h index 0f8940abea..320d38832a 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/transitive_reduced_kwarg_dataflow_graph/get_transitive_reduced_kwarg_dataflow_edges_across_split.h @@ -16,14 +16,11 @@ std::set> TransitiveReducedKwargDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { - std::set src_subgraph = - set_of(get_leaves(split.get_left_child())); - std::set dst_subgraph = - set_of(get_leaves(split.get_right_child())); - - std::set raw_edges = - get_edges_from_subgraph_to_subgraph( - tr_g.transitive_reduction, src_subgraph, dst_subgraph); + std::set src_subgraph = set_of(get_leaves(split.get_left_child())); + std::set dst_subgraph = set_of(get_leaves(split.get_right_child())); + + std::set raw_edges = get_edges_from_subgraph_to_subgraph( + tr_g.transitive_reduction, src_subgraph, dst_subgraph); return flatmap(raw_edges, [&](DirectedEdge const &e) { return get_kwarg_dataflow_edges_from_node_to_node( diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h index a0c5b9b21b..52d6293da0 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h @@ -21,9 +21,10 @@ struct ViewFromKwargDataflowGraphData final std::set> query_edges( KwargDataflowEdgeQuery const &query) const override { - return filter(set_of(this->data.edges), [&](KwargDataflowEdge const &e) { - return kwarg_dataflow_edge_query_includes(query, e); - }); + return filter(set_of(this->data.edges), + [&](KwargDataflowEdge const &e) { + return kwarg_dataflow_edge_query_includes(query, e); + }); } std::set> query_outputs( diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h index 2ff3ddb40e..527216be45 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/i_kwarg_dataflow_graph.h @@ -9,15 +9,14 @@ namespace FlexFlow { template struct IKwargDataflowGraph : virtual public IKwargDataflowGraphView { - virtual KwargNodeAddedResult add_node( - std::map> const &inputs, - std::set const &outputs) = 0; + virtual KwargNodeAddedResult + add_node(std::map> const &inputs, + std::set const &outputs) = 0; virtual void add_node_unsafe( Node const &node, std::map> const &inputs, - std::map> const - &outputs) = 0; + std::map> const &outputs) = 0; virtual void inplace_materialize_from(KwargDataflowGraphView const &) = 0; diff --git a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h index d199a64fb4..2587e6ac03 100644 --- a/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/kwarg_dataflow_graph/kwarg_dataflow_graph.h @@ -10,17 +10,16 @@ namespace FlexFlow { template struct KwargDataflowGraph : virtual public KwargDataflowGraphView { public: - KwargNodeAddedResult add_node( - std::map> const &inputs, - std::set const &outputs) { + KwargNodeAddedResult + add_node(std::map> const &inputs, + std::set const &outputs) { return this->get_interface().add_node(inputs, outputs); } void add_node_unsafe( Node const &node, std::map> const &inputs, - std::map> const - &outputs) { + std::map> const &outputs) { return this->get_interface().add_node_unsafe(node, inputs, outputs); } diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h index a981bbd8da..2bf36e64b0 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_node_label_map.h @@ -1,17 +1,15 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_NODE_LABEL_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_NODE_LABEL_MAP_H +#include "utils/containers/generate_map.h" #include "utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" -#include "utils/containers/generate_map.h" namespace FlexFlow { template -std::map - get_labelled_kwarg_dataflow_graph_node_label_map( - LabelledKwargDataflowGraphView const - &g) { +std::map get_labelled_kwarg_dataflow_graph_node_label_map( + LabelledKwargDataflowGraphView const &g) { return generate_map(get_nodes(g), [&](Node const &n) -> NodeLabel { return g.at(n); }); } diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h index ef88f5e7aa..cd469e8db0 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_output_label_map.h @@ -1,9 +1,9 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_OUTPUT_LABEL_MAP_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_KWARG_DATAFLOW_GRAPH_OUTPUT_LABEL_MAP_H +#include "utils/containers/generate_map.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" #include "utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph_view.h" -#include "utils/containers/generate_map.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h index 9ed5b5fe44..76fe55e558 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.h @@ -24,9 +24,8 @@ LabelledKwargDataflowGraphView std::map g_node_labelling = get_labelled_kwarg_dataflow_graph_node_label_map(g); - std::map, OutputLabel> - g_output_labelling = - get_labelled_kwarg_dataflow_graph_output_label_map(g); + std::map, OutputLabel> g_output_labelling = + get_labelled_kwarg_dataflow_graph_output_label_map(g); return kwarg_dataflow_graph_view_with_labelling( unlabelled_subgraph, diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h index e09c83e91e..64ea2743dd 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.h @@ -15,8 +15,7 @@ struct KwargDataflowGraphLabellingWrapper final KwargDataflowGraphLabellingWrapper( KwargDataflowGraphView const &unlabelled, std::map const &node_labels, - std::map, OutputLabel> const - &output_labels) + std::map, OutputLabel> const &output_labels) : unlabelled(unlabelled), node_labels(node_labels), output_labels(output_labels) {} diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h index 752a0fbdc2..a7ab02bce3 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.h @@ -13,8 +13,8 @@ std::string labelled_kwarg_dataflow_graph_view_as_dot( std::function const &render_node_label, std::function const &render_value_label, std::function const &render_slot_name, - std::function( - std::set const &)> const &order_slots) { + std::function(std::set const &)> const + &order_slots) { std::function render_node = [&](Node const &n) -> nlohmann::json { return render_node_label(g.at(n)); diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h index c4f7e22276..2265eaad22 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/i_labelled_kwarg_dataflow_graph.h @@ -13,10 +13,10 @@ struct ILabelledKwargDataflowGraph OutputLabel, SlotName> { public: - virtual KwargNodeAddedResult add_node( - NodeLabel const &node_label, - std::map> const &inputs, - std::map const &output_labels) = 0; + virtual KwargNodeAddedResult + add_node(NodeLabel const &node_label, + std::map> const &inputs, + std::map const &output_labels) = 0; virtual void inplace_materialize_from( LabelledKwargDataflowGraphView const &) = 0; diff --git a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h index 3b38254f01..1ff3b4e008 100644 --- a/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_kwarg_dataflow_graph/labelled_kwarg_dataflow_graph.h @@ -18,10 +18,10 @@ struct LabelledKwargDataflowGraph LabelledKwargDataflowGraph & operator=(LabelledKwargDataflowGraph const &) = default; - KwargNodeAddedResult add_node( - NodeLabel const &node_label, - std::map> const &inputs, - std::map const &output_labels) { + KwargNodeAddedResult + add_node(NodeLabel const &node_label, + std::map> const &inputs, + std::map const &output_labels) { return this->get_interface().add_node(node_label, inputs, output_labels); } diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h index e12cb1a440..d94c8de94c 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/from_labelled_open_dataflow_graph_data.h @@ -16,8 +16,7 @@ LabelledOpenDataflowGraphView from_labelled_open_dataflow_graph_data( LabelledOpenDataflowGraphData const &data) { std::set values = keys(data.value_data); - std::set outputs = - filtrans(values, try_get_dataflow_output); + std::set outputs = filtrans(values, try_get_dataflow_output); OpenDataflowGraphData unlabelled_data = OpenDataflowGraphData{ keys(data.node_data), diff --git a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h index b94f3df126..67eff4a4b9 100644 --- a/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h +++ b/lib/utils/include/utils/graph/labelled_open_dataflow_graph/algorithms/is_isomorphic_under.h @@ -18,14 +18,15 @@ bool is_isomorphic_under( OpenDataflowGraphIsomorphism const &candidate_isomorphism) { bidict node_permutation = - bidict_transform_values(candidate_isomorphism.node_mapping, - [](Node const &dst_node) { return NewNode{dst_node}; }) + bidict_transform_values( + candidate_isomorphism.node_mapping, + [](Node const &dst_node) { return NewNode{dst_node}; }) .reversed(); bidict input_permutation = bidict_transform_values(candidate_isomorphism.input_mapping, - [](DataflowGraphInput const &dst_input) { - return NewDataflowGraphInput{dst_input}; - }) + [](DataflowGraphInput const &dst_input) { + return NewDataflowGraphInput{dst_input}; + }) .reversed(); return get_graph_data(permute_input_ids( permute_node_ids(src, node_permutation), input_permutation)) == diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h index 0f66eecc00..ee041c24df 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/get_labelled_open_kwarg_dataflow_graph_data.h @@ -1,14 +1,14 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_LABELLED_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H +#include "utils/containers/generate_map.h" +#include "utils/containers/set_of.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph_view.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_values.h" -#include "utils/containers/set_of.h" -#include "utils/containers/generate_map.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h index 50b59025f2..29d999fd2d 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.h @@ -18,8 +18,8 @@ std::string labelled_open_kwarg_dataflow_graph_view_as_dot( std::function const &render_node_label, std::function const &render_value_label, std::function const &render_slot_name, - std::function( - std::set const &)> const &order_slots) { + std::function(std::set const &)> const + &order_slots) { std::function render_node = [&](Node const &n) -> nlohmann::json { return render_node_label(g.at(n)); diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h index d82dc86e19..682818f5af 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graphs_are_isomorphic_under.h @@ -28,9 +28,9 @@ bool labelled_open_kwarg_dataflow_graphs_are_isomorphic_under( OpenKwargDataflowGraphIsomorphism const &candidate_isomorphism) { bidict new_node_to_old_node = - bidict_transform_values(candidate_isomorphism.node_mapping, [](Node const &n) { - return NewNode{n}; - }).reversed(); + bidict_transform_values(candidate_isomorphism.node_mapping, + [](Node const &n) { return NewNode{n}; }) + .reversed(); bidict, KwargDataflowGraphInput> diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h index 689a5380c0..2aa69a491e 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.h @@ -22,7 +22,7 @@ struct OpenKwargDataflowGraphLabellingWrapper final OpenKwargDataflowGraphView const &unlabelled, std::map const &node_labels, std::map, - ValueLabel> const &value_labels) + ValueLabel> const &value_labels) : unlabelled(unlabelled), node_labels(node_labels), value_labels(value_labels) {} @@ -66,8 +66,7 @@ struct OpenKwargDataflowGraphLabellingWrapper final private: OpenKwargDataflowGraphView unlabelled; std::map node_labels; - std::map, - ValueLabel> + std::map, ValueLabel> value_labels; }; @@ -83,7 +82,7 @@ LabelledOpenKwargDataflowGraphView const &g, std::map const &node_labels, std::map, - ValueLabel> const &value_labels) { + ValueLabel> const &value_labels) { return LabelledOpenKwargDataflowGraphView node_labels = generate_map(get_nodes(permuted), [&](Node const &n) { return g.at(n); }); - std::map, - ValueLabel> + std::map, ValueLabel> value_labels = generate_map( get_all_open_kwarg_dataflow_values(permuted), [&](OpenKwargDataflowValue const diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h index 29a8235900..683f934060 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/permute_labelled_open_kwarg_dataflow_graph_node_ids.h @@ -60,8 +60,7 @@ LabelledOpenKwargDataflowGraphView, - ValueLabel> + std::map, ValueLabel> value_labels = generate_map( get_all_open_kwarg_dataflow_values(permuted), [&](OpenKwargDataflowValue const diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h index 61bd3957b5..c074f12d97 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/rewrite_labelled_open_kwarg_dataflow_graph_labels.h @@ -40,8 +40,7 @@ LabelledOpenKwargDataflowGraphView node_labels = generate_map(get_nodes(g), get_new_node_label); - std::map, - NewValueLabel> + std::map, NewValueLabel> value_labels = generate_map(get_all_open_kwarg_dataflow_values(g), get_new_value_label); return open_kwarg_dataflow_graph_view_with_labelling( diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h index 45f82ad823..3c6b7c15b8 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/i_labelled_open_kwarg_dataflow_graph.h @@ -22,8 +22,7 @@ struct ILabelledOpenKwargDataflowGraph SlotName> { virtual KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::map> const + std::map> const &inputs, std::map const &output_labels) = 0; diff --git a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h index 5f92a86b48..9218e1a4d4 100644 --- a/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/labelled_open_kwarg_dataflow_graph/labelled_open_kwarg_dataflow_graph.h @@ -29,8 +29,7 @@ struct LabelledOpenKwargDataflowGraph KwargNodeAddedResult add_node( NodeLabel const &node_label, - std::map> const + std::map> const &inputs, std::map const &output_labels) { return this->get_interface().add_node(node_label, inputs, output_labels); diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h index a8683d1ac4..3b6d517243 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_incoming_edges.h @@ -6,11 +6,10 @@ namespace FlexFlow { std::set get_incoming_edges(MultiDiGraphView const &, - Node const &); + Node const &); std::map> - get_incoming_edges(MultiDiGraphView const &g, - std::set const &nodes); + get_incoming_edges(MultiDiGraphView const &g, std::set const &nodes); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h b/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h index b4d1254512..3d8d6ca667 100644 --- a/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h +++ b/lib/utils/include/utils/graph/multidigraph/algorithms/get_outgoing_edges.h @@ -7,11 +7,10 @@ namespace FlexFlow { std::set get_outgoing_edges(MultiDiGraphView const &, - Node const &); + Node const &); std::map> - get_outgoing_edges(MultiDiGraphView const &g, - std::set const &ns); + get_outgoing_edges(MultiDiGraphView const &g, std::set const &ns); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/node/node_query.h b/lib/utils/include/utils/graph/node/node_query.h index 9ab9bde9fe..973a2faf6b 100644 --- a/lib/utils/include/utils/graph/node/node_query.h +++ b/lib/utils/include/utils/graph/node/node_query.h @@ -8,8 +8,7 @@ namespace FlexFlow { NodeQuery node_query_all(); NodeQuery query_intersection(NodeQuery const &, NodeQuery const &); NodeQuery query_union(NodeQuery const &, NodeQuery const &); -std::set apply_node_query(NodeQuery const &, - std::set const &); +std::set apply_node_query(NodeQuery const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h index 4266c66e18..f71504996d 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.h @@ -5,13 +5,11 @@ namespace FlexFlow { -std::set - get_incoming_edges(OpenDataflowGraphView const &); +std::set get_incoming_edges(OpenDataflowGraphView const &); std::vector get_incoming_edges(OpenDataflowGraphView const &, Node const &); std::map> - get_incoming_edges(OpenDataflowGraphView const &, - std::set const &); + get_incoming_edges(OpenDataflowGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h index 425ae32d44..9d3bc07b24 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph.h @@ -13,8 +13,7 @@ OpenDataflowSubgraphResult get_subgraph(OpenDataflowGraphView const &, bidict get_full_graph_values_to_subgraph_inputs( - OpenDataflowGraphView const &g, - std::set const &subgraph_nodes); + OpenDataflowGraphView const &g, std::set const &subgraph_nodes); OpenDataflowGraphData get_subgraph_data(OpenDataflowGraphView const &g, diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h index 017bac26b9..93c4fc10ab 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/algorithms/get_subgraph_inputs.h @@ -6,9 +6,8 @@ namespace FlexFlow { -std::set - get_subgraph_inputs(OpenDataflowGraphView const &, - std::set const &); +std::set get_subgraph_inputs(OpenDataflowGraphView const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h index 6c0e31f4cc..c317d2d67c 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_edge_query.h @@ -15,9 +15,9 @@ OpenDataflowEdgeQuery open_dataflow_edge_query_all_outgoing_from(OpenDataflowValue const &); OpenDataflowEdgeQuery open_dataflow_edge_query_all_incoming_to(DataflowInput const &); -std::set apply_open_dataflow_edge_query( - OpenDataflowEdgeQuery const &, - std::set const &); +std::set + apply_open_dataflow_edge_query(OpenDataflowEdgeQuery const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h index f29db36920..6916e90a3e 100644 --- a/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/open_dataflow_graph/open_dataflow_graph_view.h @@ -12,8 +12,7 @@ struct OpenDataflowGraphView : virtual public DataflowGraphView { OpenDataflowGraphView &operator=(OpenDataflowGraphView const &) = default; std::set get_inputs() const; - std::set - query_edges(OpenDataflowEdgeQuery const &) const; + std::set query_edges(OpenDataflowEdgeQuery const &) const; template static diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h index cfa7c7c7de..b885eec919 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/find_isomorphisms_between_open_kwarg_dataflow_graphs.h @@ -30,15 +30,13 @@ std::optional> KwargDataflowGraphInput> const &unused_graph_inputs_mapping) { { - std::set already_mapped_src_nodes = - left_entries(sink_node_mapping); + std::set already_mapped_src_nodes = left_entries(sink_node_mapping); std::set src_g_sink_nodes = set_of(get_terminal_nodes(src_g)); ASSERT(already_mapped_src_nodes == src_g_sink_nodes); } { - std::set already_mapped_dst_nodes = - right_entries(sink_node_mapping); + std::set already_mapped_dst_nodes = right_entries(sink_node_mapping); std::set dst_g_sink_nodes = set_of(get_terminal_nodes(dst_g)); ASSERT(already_mapped_dst_nodes == dst_g_sink_nodes); } @@ -46,18 +44,16 @@ std::optional> { std::set> already_mapped_src_inputs = left_entries(unused_graph_inputs_mapping); - std::set> - src_g_unused_inputs = - set_of(get_unused_open_kwarg_dataflow_graph_inputs(src_g)); + std::set> src_g_unused_inputs = + set_of(get_unused_open_kwarg_dataflow_graph_inputs(src_g)); ASSERT(already_mapped_src_inputs == src_g_unused_inputs); } { std::set> already_mapped_dst_inputs = right_entries(unused_graph_inputs_mapping); - std::set> - dst_g_unused_inputs = - set_of(get_unused_open_kwarg_dataflow_graph_inputs(dst_g)); + std::set> dst_g_unused_inputs = + set_of(get_unused_open_kwarg_dataflow_graph_inputs(dst_g)); ASSERT(already_mapped_dst_inputs == dst_g_unused_inputs); } @@ -178,12 +174,10 @@ std::optional> result->node_mapping.equate(src_node, dst_node); - std::map> + std::map> src_incoming_edges = get_incoming_open_kwarg_dataflow_edges_for_node(src_g, src_node); - std::map> + std::map> dst_incoming_edges = get_incoming_open_kwarg_dataflow_edges_for_node(dst_g, dst_node); @@ -221,9 +215,8 @@ std::set> std::vector> src_unused_graph_inputs = vector_of(get_unused_open_kwarg_dataflow_graph_inputs(src)); - std::set> - dst_unused_graph_inputs = - get_unused_open_kwarg_dataflow_graph_inputs(dst); + std::set> dst_unused_graph_inputs = + get_unused_open_kwarg_dataflow_graph_inputs(dst); if (src_unused_graph_inputs.size() != dst_unused_graph_inputs.size()) { return {}; diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h index 10f10cc58c..9d994020f2 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h @@ -1,13 +1,13 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H +#include "utils/containers/set_of.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_graph_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" -#include "utils/containers/set_of.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h index 2c8d94be20..a7c6fa98d4 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_subgraph.h @@ -2,6 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_GET_OPEN_KWARG_DATAFLOW_GRAPH_SUBGRAPH_H #include "utils/bidict/generate_bidict.h" +#include "utils/containers/set_of.h" #include "utils/containers/set_union.h" #include "utils/containers/values.h" #include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.h" @@ -11,7 +12,6 @@ #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_subgraph_result.dtg.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h" #include "utils/overload.h" -#include "utils/containers/set_of.h" namespace FlexFlow { @@ -71,7 +71,8 @@ OpenKwargDataflowGraphData std::set> subgraph_input_edges = transform( - set_of(get_open_kwarg_dataflow_subgraph_incoming_edges(g, set_of(subgraph_nodes))), + set_of(get_open_kwarg_dataflow_subgraph_incoming_edges( + g, set_of(subgraph_nodes))), [&](OpenKwargDataflowEdge const &edge) { return edge.template visit< OpenKwargDataflowEdge>(overload{ @@ -115,7 +116,8 @@ OpenKwargDataflowGraphData }; std::set> - subgraph_interior_edges = set_of(g.query_edges(subgraph_interior_edges_query)); + subgraph_interior_edges = + set_of(g.query_edges(subgraph_interior_edges_query)); std::set> subgraph_inputs = set_of(values(full_graph_values_to_subgraph_inputs)); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h index 77cac17d1c..80c07aaa51 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_value_uses.h @@ -11,10 +11,9 @@ namespace FlexFlow { template -std::set> - get_open_kwarg_dataflow_value_uses( - OpenKwargDataflowGraphView const &g, - OpenKwargDataflowValue const &v) { +std::set> get_open_kwarg_dataflow_value_uses( + OpenKwargDataflowGraphView const &g, + OpenKwargDataflowValue const &v) { OpenKwargDataflowEdgeQuery query = v.template visit< OpenKwargDataflowEdgeQuery>(overload{ diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h index 48b5729ab0..c60eb7f078 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h @@ -2,12 +2,12 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_OPEN_KWARG_DATAFLOW_GRAPH_AS_DOT_H #include "utils/containers/filtrans.h" +#include "utils/containers/set_of.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_value.dtg.h" -#include "utils/containers/set_of.h" namespace FlexFlow { @@ -33,10 +33,8 @@ std::string open_kwarg_dataflow_graph_as_dot( return j; }; - std::function(std::set const &)> - order_slots = [](std::set const &unordered) { - return sorted(unordered); - }; + std::function(std::set const &)> order_slots = + [](std::set const &unordered) { return sorted(unordered); }; return open_kwarg_dataflow_graph_as_dot( g, render_node, render_value, render_slot_name, order_slots); @@ -50,8 +48,8 @@ std::string open_kwarg_dataflow_graph_as_dot( OpenKwargDataflowValue const &)> const &render_value, std::function const &render_slot_name, - std::function( - std::set const &)> const &order_slots) { + std::function(std::set const &)> const + &order_slots) { std::pair>, bidict, Node>> closed_g_and_mapping = diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h index b72bf95701..449d3529ab 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h @@ -13,8 +13,8 @@ namespace FlexFlow { template void require_open_kwarg_dataflow_graph_data_is_valid( OpenKwargDataflowGraphData const &data) { - std::set> - inputs_from_edges = filtrans( + std::set> inputs_from_edges = + filtrans( data.edges, [](OpenKwargDataflowEdge const &e) -> std::optional> { diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h index d3db64c5b3..76cc229765 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/try_find_isomorphism_between_open_kwarg_dataflow_graphs.h @@ -12,9 +12,8 @@ std::optional> try_find_isomorphism_between_open_kwarg_dataflow_graphs( OpenKwargDataflowGraphView const &src, OpenKwargDataflowGraphView const &dst) { - std::set> - isomorphisms = - find_isomorphisms_between_open_kwarg_dataflow_graphs(src, dst); + std::set> isomorphisms = + find_isomorphisms_between_open_kwarg_dataflow_graphs(src, dst); return try_get_one_of(isomorphisms); } diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h index 19d9f10ee7..01dd554245 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.h @@ -3,6 +3,7 @@ #include "utils/bidict/algorithms/right_entries.h" #include "utils/bidict/generate_bidict.h" +#include "utils/containers/set_of.h" #include "utils/containers/set_union.h" #include "utils/containers/transform.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" @@ -12,9 +13,8 @@ #include "utils/graph/node/node_source.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/get_open_kwarg_dataflow_graph_data.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" -#include "utils/overload.h" -#include "utils/containers/set_of.h" #include "utils/json/optional.h" +#include "utils/overload.h" namespace FlexFlow { @@ -91,16 +91,13 @@ std::pair>, KwargDataflowGraphData> closed_g_data = KwargDataflowGraphData>{ /*nodes=*/set_of( - set_union( - open_g_data.nodes, - right_entries(graph_input_nodes))), + set_union(open_g_data.nodes, right_entries(graph_input_nodes))), /*edges=*/set_of(transform(open_g_data.edges, convert_edge)), /*outputs=*/ - set_of( - set_union( - transform(open_g_data.outputs, convert_kwarg_dataflow_output), - transform(open_g_data.inputs, - kwarg_dataflow_output_for_graph_input))), + set_of(set_union( + transform(open_g_data.outputs, convert_kwarg_dataflow_output), + transform(open_g_data.inputs, + kwarg_dataflow_output_for_graph_input))), }; ASSERT(closed_g_data.edges.size() == open_g_data.edges.size()); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h index 56e43a73bd..6e5162794d 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/algorithms/view_from_open_kwarg_dataflow_graph_data.h @@ -1,13 +1,13 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_VIEW_FROM_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_GRAPH_OPEN_KWARG_DATAFLOW_GRAPH_ALGORITHMS_VIEW_FROM_OPEN_KWARG_DATAFLOW_GRAPH_DATA_H +#include "utils/containers/set_of.h" #include "utils/graph/kwarg_dataflow_graph/kwarg_dataflow_output_query.h" #include "utils/graph/node/node_query.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.dtg.h" #include "utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_data.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_edge_query.h" #include "utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h" -#include "utils/containers/set_of.h" namespace FlexFlow { @@ -27,9 +27,9 @@ struct ViewFromOpenKwargDataflowGraphData final return set_of(this->data.inputs); } - std::set> - query_edges(OpenKwargDataflowEdgeQuery const - &query) const override { + std::set> query_edges( + OpenKwargDataflowEdgeQuery const &query) + const override { return filter( set_of(this->data.edges), [&](OpenKwargDataflowEdge const &e) { diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h index 0be5f46e37..fd6da0857b 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph.h @@ -11,8 +11,7 @@ template struct IOpenKwargDataflowGraph : virtual public IOpenKwargDataflowGraphView { virtual KwargNodeAddedResult add_node( - std::map> const + std::map> const &inputs, std::set const &outputs) = 0; virtual KwargDataflowGraphInput diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h index 162224853d..cab10be854 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/i_open_kwarg_dataflow_graph_view.h @@ -14,9 +14,8 @@ struct IOpenKwargDataflowGraphView : virtual public IKwargDataflowGraphView { virtual std::set> get_inputs() const = 0; - virtual std::set> - query_edges(OpenKwargDataflowEdgeQuery const &) - const = 0; + virtual std::set> query_edges( + OpenKwargDataflowEdgeQuery const &) const = 0; std::set> query_edges( KwargDataflowEdgeQuery const &query) const override final { @@ -28,8 +27,8 @@ struct IOpenKwargDataflowGraphView /*standard_edge_query=*/query, }; - std::set> - open_edges = this->query_edges(open_query); + std::set> open_edges = + this->query_edges(open_query); return transform( open_edges, diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h index 335ec6a67d..4ed05b4115 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph.h @@ -13,8 +13,7 @@ struct OpenKwargDataflowGraph : virtual public OpenKwargDataflowGraphView { public: KwargNodeAddedResult add_node( - std::map> const + std::map> const &inputs, std::set const &outputs) { return this->get_interface().add_node(inputs, outputs); diff --git a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h index 989634bcf5..e4390f97b3 100644 --- a/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h +++ b/lib/utils/include/utils/graph/open_kwarg_dataflow_graph/open_kwarg_dataflow_graph_view.h @@ -14,14 +14,12 @@ struct OpenKwargDataflowGraphView OpenKwargDataflowGraphView & operator=(OpenKwargDataflowGraphView const &) = default; - std::set> - get_inputs() const { + std::set> get_inputs() const { return this->get_interface().get_inputs(); } - std::set> - query_edges( - OpenKwargDataflowEdgeQuery const &q) const { + std::set> query_edges( + OpenKwargDataflowEdgeQuery const &q) const { return this->get_interface().query_edges(q); } diff --git a/lib/utils/include/utils/graph/query_set.h b/lib/utils/include/utils/graph/query_set.h index a36327875f..84f15b212d 100644 --- a/lib/utils/include/utils/graph/query_set.h +++ b/lib/utils/include/utils/graph/query_set.h @@ -9,7 +9,6 @@ #include "utils/containers/set_of.h" #include "utils/containers/set_union.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/exception.h" #include "utils/fmt/set.h" #include "utils/hash-utils.h" @@ -17,7 +16,6 @@ #include "utils/optional.h" #include #include -#include namespace FlexFlow { @@ -113,8 +111,7 @@ std::set apply_query(query_set const &q, C const &c) { return set_of(c); } - return filter(set_of(c), - [&](T const &t) { return includes(q, t); }); + return filter(set_of(c), [&](T const &t) { return includes(q, t); }); } template #include +#include namespace FlexFlow { std::string escape_dot_string(std::string const &); -std::string render_dot_node_attrs( - std::map const &attrs); -std::string render_dot( - LabelledDataflowGraphView, - std::string> const &); +std::string + render_dot_node_attrs(std::map const &attrs); +std::string + render_dot(LabelledDataflowGraphView, + std::string> const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/digraph_generation.h b/lib/utils/include/utils/graph/series_parallel/digraph_generation.h index 1a8f9ff0b1..e784bb405a 100644 --- a/lib/utils/include/utils/graph/series_parallel/digraph_generation.h +++ b/lib/utils/include/utils/graph/series_parallel/digraph_generation.h @@ -6,10 +6,8 @@ namespace FlexFlow { -std::map parallel_extend(DiGraph &g, - DiGraphView const &ext); -std::map serial_extend(DiGraph &g, - DiGraphView const &ext); +std::map parallel_extend(DiGraph &g, DiGraphView const &ext); +std::map serial_extend(DiGraph &g, DiGraphView const &ext); DiGraph series_composition(DiGraphView const &g1, DiGraphView const &g2); DiGraph parallel_composition(DiGraphView const &g1, DiGraphView const &g2); DiGraph series_composition(std::vector const &graphs); diff --git a/lib/utils/include/utils/graph/series_parallel/get_ancestors.h b/lib/utils/include/utils/graph/series_parallel/get_ancestors.h index 15e3993442..0fcce08e40 100644 --- a/lib/utils/include/utils/graph/series_parallel/get_ancestors.h +++ b/lib/utils/include/utils/graph/series_parallel/get_ancestors.h @@ -43,7 +43,7 @@ namespace FlexFlow { * */ std::set get_ancestors(SeriesParallelDecomposition const &sp, - Node const &node); + Node const &node); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h b/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h index 9e448d0db6..aac8610074 100644 --- a/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h +++ b/lib/utils/include/utils/graph/series_parallel/series_parallel_decomposition.h @@ -30,8 +30,7 @@ nonnegative_int num_nodes(SeriesParallelDecomposition const &sp); SeriesParallelDecomposition series_composition( std::vector const &sp_compositions); SeriesParallelDecomposition parallel_composition( - std::multiset const - &sp_compositions); + std::multiset const &sp_compositions); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h b/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h index ca9e555803..925413f18c 100644 --- a/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h +++ b/lib/utils/include/utils/graph/series_parallel/series_parallel_metrics.h @@ -23,8 +23,7 @@ std::map float work_cost(SeriesParallelDecomposition const &sp, std::map cost_map); -float work_cost(DiGraphView const &g, - std::map const &cost_map); +float work_cost(DiGraphView const &g, std::map const &cost_map); /** * @brief Computes the total number of edges the decomposition has when viewed diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h index 4cee222872..9d791dbe1b 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/escribano_algo.h @@ -8,14 +8,12 @@ #include namespace FlexFlow { -DiGraph add_dummy_nodes(DiGraph g, - std::map &node_roles); +DiGraph add_dummy_nodes(DiGraph g, std::map &node_roles); -std::set - get_component(DiGraph const &g, - Node const &node, - std::map const &depth_map, - std::map const &node_roles); +std::set get_component(DiGraph const &g, + Node const &node, + std::map const &depth_map, + std::map const &node_roles); /** * \brief See \ref spization-escribano. diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h index d3f87c8e15..c145c14ce8 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/node_role.h @@ -8,8 +8,7 @@ namespace FlexFlow { -std::map - get_initial_node_role_map(DiGraphView const &g); +std::map get_initial_node_role_map(DiGraphView const &g); /** * @brief Contracts out nodes of a given role from the graph. diff --git a/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h b/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h index 5cc3b9bd6f..0821827383 100644 --- a/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h +++ b/lib/utils/include/utils/graph/series_parallel/sp_ization/up_down_partition.h @@ -12,14 +12,14 @@ namespace FlexFlow { * is no outgoing edge from n. */ std::set get_up_frontier(DiGraph const &sp, - UpDownPartition const &partition); + UpDownPartition const &partition); /** * @brief Returns the nodes n in the down set such that in the down subgraph, * there is no incoming edge to n. */ std::set get_down_frontier(DiGraph const &sp, - UpDownPartition const &partition); + UpDownPartition const &partition); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/traversal.h b/lib/utils/include/utils/graph/traversal.h index 8fe4040377..1231555945 100644 --- a/lib/utils/include/utils/graph/traversal.h +++ b/lib/utils/include/utils/graph/traversal.h @@ -17,8 +17,7 @@ struct unchecked_dfs_iterator { using reference = Node const &; unchecked_dfs_iterator(DiGraphView const &g, std::vector const &); - unchecked_dfs_iterator(DiGraphView const &g, - std::set const &); + unchecked_dfs_iterator(DiGraphView const &g, std::set const &); reference operator*() const; pointer operator->(); @@ -77,8 +76,7 @@ struct bfs_iterator { bfs_iterator(DiGraphView const &, std::queue const &, std::optional> const &); - bfs_iterator(DiGraphView const &, - std::set const &starting_points); + bfs_iterator(DiGraphView const &, std::set const &starting_points); reference operator*() const; pointer operator->(); @@ -126,8 +124,7 @@ struct UncheckedDFSView { struct BFSView { BFSView() = delete; - explicit BFSView(DiGraphView const &, - std::set const &starting_points); + explicit BFSView(DiGraphView const &, std::set const &starting_points); bfs_iterator begin() const; bfs_iterator end() const; @@ -175,10 +172,8 @@ UncheckedDFSView unchecked_dfs(DiGraphView const &, std::set const &starting_points); /* BoundaryDFSView boundary_dfs(IDiGraphView const &, std::set * const &starting_points); */ -CheckedDFSView dfs(DiGraphView const &, - std::set const &starting_points); -BFSView bfs(DiGraphView const &, - std::set const &starting_points); +CheckedDFSView dfs(DiGraphView const &, std::set const &starting_points); +BFSView bfs(DiGraphView const &, std::set const &starting_points); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h b/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h index 3cd3f5abf6..c702fef98f 100644 --- a/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h +++ b/lib/utils/include/utils/graph/undirected/algorithms/get_connected_components.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::set> - get_connected_components(UndirectedGraphView const &); +std::set> get_connected_components(UndirectedGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h b/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h index f98ae4b51d..dce19975c2 100644 --- a/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h +++ b/lib/utils/include/utils/graph/undirected/algorithms/get_neighboring_nodes.h @@ -5,8 +5,7 @@ namespace FlexFlow { -std::set get_neighboring_nodes(UndirectedGraphView const &, - Node const &); +std::set get_neighboring_nodes(UndirectedGraphView const &, Node const &); } // namespace FlexFlow diff --git a/lib/utils/include/utils/graph/undirected/i_undirected_graph.h b/lib/utils/include/utils/graph/undirected/i_undirected_graph.h index 24527ae98a..fbaae88f2e 100644 --- a/lib/utils/include/utils/graph/undirected/i_undirected_graph.h +++ b/lib/utils/include/utils/graph/undirected/i_undirected_graph.h @@ -12,8 +12,7 @@ struct IUndirectedGraph : public IUndirectedGraphView { virtual void add_edge(UndirectedEdge const &) = 0; virtual void remove_edge(UndirectedEdge const &) = 0; - virtual std::set - query_nodes(NodeQuery const &query) const = 0; + virtual std::set query_nodes(NodeQuery const &query) const = 0; virtual IUndirectedGraph *clone() const = 0; }; diff --git a/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h b/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h index 3e5c82b519..5df10478cc 100644 --- a/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h +++ b/lib/utils/include/utils/graph/undirected/i_undirected_graph_view.h @@ -14,8 +14,7 @@ struct IUndirectedGraphView : public IGraphView { IUndirectedGraphView(IUndirectedGraphView const &) = delete; IUndirectedGraphView &operator=(IUndirectedGraphView const &) = delete; - virtual std::set - query_edges(UndirectedEdgeQuery const &) const = 0; + virtual std::set query_edges(UndirectedEdgeQuery const &) const = 0; virtual ~IUndirectedGraphView() = default; IUndirectedGraphView *clone() const override = 0; diff --git a/lib/utils/include/utils/graph/views/views.h b/lib/utils/include/utils/graph/views/views.h index 639c8750b7..89a308596b 100644 --- a/lib/utils/include/utils/graph/views/views.h +++ b/lib/utils/include/utils/graph/views/views.h @@ -11,8 +11,7 @@ namespace FlexFlow { struct UndirectedSubgraphView : public IUndirectedGraphView { public: UndirectedSubgraphView() = delete; - UndirectedSubgraphView(UndirectedGraphView const &, - std::set const &); + UndirectedSubgraphView(UndirectedGraphView const &, std::set const &); std::set query_edges(UndirectedEdgeQuery const &) const override; @@ -30,8 +29,7 @@ struct DiSubgraphView : public IDiGraphView { DiSubgraphView() = delete; DiSubgraphView(DiGraphView const &, std::set const &); - std::set - query_edges(DirectedEdgeQuery const &) const override; + std::set query_edges(DirectedEdgeQuery const &) const override; std::set query_nodes(NodeQuery const &) const override; DiSubgraphView *clone() const override; @@ -44,16 +42,13 @@ struct DiSubgraphView : public IDiGraphView { UndirectedGraphView view_subgraph(UndirectedGraphView const &, std::set const &); -DiGraphView view_subgraph(DiGraphView const &, - std::set const &); +DiGraphView view_subgraph(DiGraphView const &, std::set const &); UndirectedEdge to_undirected_edge(DirectedEdge const &); -std::set - to_undirected_edges(std::set const &); +std::set to_undirected_edges(std::set const &); std::set to_directed_edges(UndirectedEdge const &); -std::set - to_directed_edges(std::set const &); +std::set to_directed_edges(std::set const &); struct ViewDiGraphAsUndirectedGraph : public IUndirectedGraphView { public: @@ -73,8 +68,7 @@ struct ViewUndirectedGraphAsDiGraph : public IDiGraphView { public: explicit ViewUndirectedGraphAsDiGraph(UndirectedGraphView const &); - std::set - query_edges(DirectedEdgeQuery const &) const override; + std::set query_edges(DirectedEdgeQuery const &) const override; std::set query_nodes(NodeQuery const &) const override; ViewUndirectedGraphAsDiGraph *clone() const override; diff --git a/lib/utils/include/utils/json/check_is_jsonable.h b/lib/utils/include/utils/json/check_is_jsonable.h index 8597e11c22..79331330d3 100644 --- a/lib/utils/include/utils/json/check_is_jsonable.h +++ b/lib/utils/include/utils/json/check_is_jsonable.h @@ -7,9 +7,9 @@ namespace FlexFlow { #define CHECK_IS_JSONABLE(...) \ - static_assert(::FlexFlow::is_json_serializable<__VA_ARGS__>::value, \ + static_assert(::FlexFlow::is_json_serializable<__VA_ARGS__>::value, \ #__VA_ARGS__ " should be json serializeable"); \ - static_assert(::FlexFlow::is_json_deserializable<__VA_ARGS__>::value, \ + static_assert(::FlexFlow::is_json_deserializable<__VA_ARGS__>::value, \ #__VA_ARGS__ " should be json deserializeable") } // namespace FlexFlow diff --git a/lib/utils/include/utils/many_to_one/many_to_one.h b/lib/utils/include/utils/many_to_one/many_to_one.h index 9a01d77e5c..e1306880d2 100644 --- a/lib/utils/include/utils/many_to_one/many_to_one.h +++ b/lib/utils/include/utils/many_to_one/many_to_one.h @@ -1,23 +1,23 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_MANY_TO_ONE_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_MANY_TO_ONE_MANY_TO_ONE_H +#include "utils/containers/keys.h" #include "utils/containers/require_same.h" +#include "utils/containers/set_of.h" #include "utils/containers/try_at.h" #include "utils/containers/values.h" #include "utils/exception.h" #include "utils/fmt/map.h" #include "utils/fmt/set.h" #include "utils/hash-utils.h" -#include "utils/hash/tuple.h" #include "utils/hash/map.h" +#include "utils/hash/tuple.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" +#include "utils/nonempty_set/nonempty_set.h" #include #include #include -#include "utils/containers/set_of.h" -#include "utils/nonempty_set/nonempty_set.h" -#include "utils/containers/keys.h" namespace FlexFlow { @@ -65,12 +65,11 @@ struct ManyToOne { } else if (found_r.value() == r) { return; } else { - PANIC( - "Existing mapping found for left value {}: tried to map to right " - "value {}, but is already bound to right value {}", - l, - r, - found_r.value()); + PANIC("Existing mapping found for left value {}: tried to map to right " + "value {}, but is already bound to right value {}", + l, + r, + found_r.value()); } } @@ -128,8 +127,7 @@ struct ManyToOne { }; template -std::map, R> - format_as(ManyToOne const &m) { +std::map, R> format_as(ManyToOne const &m) { std::map, R> result; for (R const &r : m.right_values()) { @@ -179,7 +177,8 @@ struct adl_serializer<::FlexFlow::ManyToOne> { CHECK_IS_JSON_SERIALIZABLE(L); CHECK_IS_JSON_SERIALIZABLE(R); - j = ::FlexFlow::set_of(::FlexFlow::unstructured_relation_from_many_to_one(m)); + j = ::FlexFlow::set_of( + ::FlexFlow::unstructured_relation_from_many_to_one(m)); } }; diff --git a/lib/utils/include/utils/nonempty_set/nonempty_set.h b/lib/utils/include/utils/nonempty_set/nonempty_set.h index 61276f4db9..9e985875c0 100644 --- a/lib/utils/include/utils/nonempty_set/nonempty_set.h +++ b/lib/utils/include/utils/nonempty_set/nonempty_set.h @@ -1,17 +1,17 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_NONEMPTY_SET_NONEMPTY_SET_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_NONEMPTY_SET_NONEMPTY_SET_H -#include -#include +#include "utils/containers/set_of.h" +#include "utils/containers/unordered_set_of.h" +#include "utils/fmt/set.h" #include "utils/hash-utils.h" #include "utils/hash/set.h" #include "utils/hash/tuple.h" -#include "utils/fmt/set.h" -#include "utils/positive_int/positive_int.h" -#include "utils/containers/set_of.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" -#include "utils/containers/unordered_set_of.h" +#include "utils/positive_int/positive_int.h" +#include +#include namespace FlexFlow { @@ -106,14 +106,12 @@ struct nonempty_set { }; template -bool operator==(std::set const &lhs, - nonempty_set const &rhs) { +bool operator==(std::set const &lhs, nonempty_set const &rhs) { return lhs == rhs.unwrap_as_set(); } template -bool operator!=(std::set const &lhs, - nonempty_set const &rhs) { +bool operator!=(std::set const &lhs, nonempty_set const &rhs) { return lhs != rhs.unwrap_as_set(); } diff --git a/lib/utils/include/utils/one_to_many/one_to_many.h b/lib/utils/include/utils/one_to_many/one_to_many.h index f46afaae42..3f7cd576ec 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many.h +++ b/lib/utils/include/utils/one_to_many/one_to_many.h @@ -5,6 +5,7 @@ #include "utils/containers/items.h" #include "utils/containers/keys.h" #include "utils/containers/require_same.h" +#include "utils/containers/set_of.h" #include "utils/containers/transform.h" #include "utils/containers/try_at.h" #include "utils/containers/values.h" @@ -12,18 +13,17 @@ #include "utils/fmt/map.h" #include "utils/fmt/set.h" #include "utils/hash-utils.h" -#include "utils/hash/tuple.h" #include "utils/hash/map.h" #include "utils/hash/set.h" +#include "utils/hash/tuple.h" #include "utils/json/check_is_json_deserializable.h" #include "utils/json/check_is_json_serializable.h" #include "utils/nonempty_set/nonempty_set.h" #include +#include #include #include -#include #include -#include "utils/containers/set_of.h" namespace FlexFlow { @@ -87,12 +87,11 @@ struct OneToMany { } else if (found_l.value() == l) { return; } else { - PANIC( - "Existing mapping found for right value {}: tried to map " - "to left value {}, but is already bound to left value {}", - r, - l, - found_l.value()); + PANIC("Existing mapping found for right value {}: tried to map " + "to left value {}, but is already bound to left value {}", + r, + l, + found_l.value()); } } @@ -149,8 +148,7 @@ struct OneToMany { }; template -std::map> - format_as(OneToMany const &m) { +std::map> format_as(OneToMany const &m) { return generate_map(m.left_values(), [&](L const &l) { return m.at_l(l); }); } @@ -197,7 +195,8 @@ struct adl_serializer<::FlexFlow::OneToMany> { CHECK_IS_JSON_SERIALIZABLE(L); CHECK_IS_JSON_SERIALIZABLE(R); - j = ::FlexFlow::set_of(::FlexFlow::unstructured_relation_from_one_to_many(m)); + j = ::FlexFlow::set_of( + ::FlexFlow::unstructured_relation_from_one_to_many(m)); } }; diff --git a/lib/utils/include/utils/one_to_many/one_to_many_filter_values.h b/lib/utils/include/utils/one_to_many/one_to_many_filter_values.h index 4694b06d65..c191ab71c4 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many_filter_values.h +++ b/lib/utils/include/utils/one_to_many/one_to_many_filter_values.h @@ -18,5 +18,4 @@ OneToMany one_to_many_filter_values(OneToMany const &m, F &&f) { } // namespace FlexFlow - #endif diff --git a/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h b/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h index eae63e6ade..03bb958cdb 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h +++ b/lib/utils/include/utils/one_to_many/one_to_many_from_l_to_r_mapping.h @@ -7,8 +7,8 @@ namespace FlexFlow { template -OneToMany one_to_many_from_l_to_r_mapping( - std::map> const &m) { +OneToMany + one_to_many_from_l_to_r_mapping(std::map> const &m) { OneToMany result; for (auto const &[l, rs] : m) { diff --git a/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h b/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h index 9d5fbb3665..8d29b22310 100644 --- a/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h +++ b/lib/utils/include/utils/one_to_many/one_to_many_transform_values.h @@ -4,7 +4,6 @@ #include "utils/containers/transform.h" #include "utils/one_to_many/one_to_many.h" - namespace FlexFlow { template > OneToMany one_to_many_transform_values(OneToMany const &input, F f) { - return one_to_many_from_unstructured_relation(transform( - set_of(input.relation()), - [&](std::pair const &p) -> std::pair { - return {p.first, f(p.second)}; - })); + return one_to_many_from_unstructured_relation( + transform(set_of(input.relation()), + [&](std::pair const &p) -> std::pair { + return {p.first, f(p.second)}; + })); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h b/lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h index 9d31dc968c..b1a244703e 100644 --- a/lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h +++ b/lib/utils/include/utils/one_to_many/require_one_to_many_is_bijection.h @@ -2,20 +2,19 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ONE_TO_MANY_REQUIRE_ONE_TO_MANY_IS_BIJECTION_H #include "utils/bidict/algorithms/bidict_from_map.h" -#include "utils/containers/map_values.h" #include "utils/containers/get_only.h" -#include "utils/one_to_many/one_to_many.h" +#include "utils/containers/map_values.h" #include "utils/nonempty_set/nonempty_set.h" +#include "utils/one_to_many/one_to_many.h" namespace FlexFlow { template bidict require_one_to_many_is_bijection(OneToMany const &otm) { return bidict_from_map( - map_values(otm.l_to_r(), - [](nonempty_set const &s) -> R { - return get_only(s.unwrap_as_set()); - })); + map_values(otm.l_to_r(), [](nonempty_set const &s) -> R { + return get_only(s.unwrap_as_set()); + })); } } // namespace FlexFlow diff --git a/lib/utils/include/utils/ord/unordered_map.h b/lib/utils/include/utils/ord/unordered_map.h index 072386b3fc..1cfbdb27b6 100644 --- a/lib/utils/include/utils/ord/unordered_map.h +++ b/lib/utils/include/utils/ord/unordered_map.h @@ -3,8 +3,8 @@ #include "utils/type_traits_core.h" #include -#include #include +#include namespace FlexFlow { diff --git a/lib/utils/include/utils/orthotope/dim_coord.h b/lib/utils/include/utils/orthotope/dim_coord.h index ccb2cb320b..9508000908 100644 --- a/lib/utils/include/utils/orthotope/dim_coord.h +++ b/lib/utils/include/utils/orthotope/dim_coord.h @@ -3,6 +3,7 @@ #include "utils/containers/all_of.h" #include "utils/containers/contains_key.h" +#include "utils/containers/generate_map.h" #include "utils/containers/get_all_assignments.h" #include "utils/containers/is_subseteq_of.h" #include "utils/containers/keys.h" @@ -12,9 +13,9 @@ #include "utils/containers/require_same.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/scanr.h" +#include "utils/containers/set_of.h" #include "utils/containers/sorted_by.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/containers/zip_with_strict.h" #include "utils/exception.h" #include "utils/nonnegative_int/nonnegative_range.h" @@ -23,8 +24,6 @@ #include "utils/orthotope/dim_domain.h" #include "utils/orthotope/minimal_dim_domain.h" #include "utils/orthotope/orthotope.h" -#include "utils/containers/set_of.h" -#include "utils/containers/generate_map.h" namespace FlexFlow { @@ -78,23 +77,19 @@ DimCoord lift_dim_coord(DimCoord const &coord, } template -std::set> - get_coords_in_dim_domain(DimDomain const &dim_domain) { - std::map> - component_possible_values = map_values( - dim_domain.dims, - [](positive_int component_size) - -> std::set { - return set_of(nonnegative_range(component_size)); - }); - - return set_of(transform( - get_all_assignments(component_possible_values), - [](std::map const &assignment) { - return DimCoord{ - assignment, - }; - })); +std::set> get_coords_in_dim_domain(DimDomain const &dim_domain) { + std::map> component_possible_values = + map_values(dim_domain.dims, + [](positive_int component_size) -> std::set { + return set_of(nonnegative_range(component_size)); + }); + + return set_of(transform(get_all_assignments(component_possible_values), + [](std::map const &assignment) { + return DimCoord{ + assignment, + }; + })); } template diff --git a/lib/utils/include/utils/orthotope/dim_domain.h b/lib/utils/include/utils/orthotope/dim_domain.h index 7c6abf509a..39c8b63822 100644 --- a/lib/utils/include/utils/orthotope/dim_domain.h +++ b/lib/utils/include/utils/orthotope/dim_domain.h @@ -2,6 +2,7 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ORTHOTOPE_DIM_DOMAIN_H #include "utils/containers/filter.h" +#include "utils/containers/keys.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/restrict_keys.h" #include "utils/containers/set_minus.h" @@ -11,7 +12,6 @@ #include "utils/orthotope/dim_domain.dtg.h" #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/orthotope.dtg.h" -#include "utils/containers/keys.h" namespace FlexFlow { diff --git a/lib/utils/include/utils/orthotope/dim_projection.h b/lib/utils/include/utils/orthotope/dim_projection.h index 44b7a7cebb..177556ec1b 100644 --- a/lib/utils/include/utils/orthotope/dim_projection.h +++ b/lib/utils/include/utils/orthotope/dim_projection.h @@ -32,8 +32,7 @@ DimProjection } template -std::set - input_dims_of_projection(DimProjection const &projection) { +std::set input_dims_of_projection(DimProjection const &projection) { return projection.template visit>(overload{ [](UpProjection const &p) { return input_dims_of_up_projection(p); @@ -48,8 +47,7 @@ std::set } template -std::set - output_dims_of_projection(DimProjection const &projection) { +std::set output_dims_of_projection(DimProjection const &projection) { return projection.template visit>(overload{ [](UpProjection const &p) { return output_dims_of_up_projection(p); @@ -102,8 +100,7 @@ DimCoord compute_dim_projection(DimProjection const &projection, { std::set nontrivial_input_domain_dims = get_nontrivial_domain_dims(input_domain); - std::set projection_input_dims = - input_dims_of_projection(projection); + std::set projection_input_dims = input_dims_of_projection(projection); std::set all_input_domain_dims = get_domain_dims(input_domain); ASSERT(is_subseteq_of(nontrivial_input_domain_dims, projection_input_dims), @@ -117,10 +114,8 @@ DimCoord compute_dim_projection(DimProjection const &projection, { std::set nontrivial_output_domain_dims = get_nontrivial_domain_dims(output_domain); - std::set projection_output_dims = - output_dims_of_projection(projection); - std::set all_output_domain_dims = - get_domain_dims(output_domain); + std::set projection_output_dims = output_dims_of_projection(projection); + std::set all_output_domain_dims = get_domain_dims(output_domain); ASSERT( is_subseteq_of(nontrivial_output_domain_dims, projection_output_dims), diff --git a/lib/utils/include/utils/orthotope/down_projection.h b/lib/utils/include/utils/orthotope/down_projection.h index 2dda306371..406c993001 100644 --- a/lib/utils/include/utils/orthotope/down_projection.h +++ b/lib/utils/include/utils/orthotope/down_projection.h @@ -44,8 +44,7 @@ DimCoord compute_down_projection(DownProjection const &projection, "compute_down_projection expected coord dimensions to match " "projection input dimensions"); - std::set output_dims = - output_dims_of_down_projection(projection); + std::set output_dims = output_dims_of_down_projection(projection); return DimCoord{ generate_map( diff --git a/lib/utils/include/utils/orthotope/eq_projection.h b/lib/utils/include/utils/orthotope/eq_projection.h index 12a0046583..7baaa4af31 100644 --- a/lib/utils/include/utils/orthotope/eq_projection.h +++ b/lib/utils/include/utils/orthotope/eq_projection.h @@ -15,14 +15,12 @@ EqProjection make_empty_eq_projection() { } template -std::set - input_dims_of_eq_projection(EqProjection const &projection) { +std::set input_dims_of_eq_projection(EqProjection const &projection) { return projection.dim_mapping.left_values(); } template -std::set - output_dims_of_eq_projection(EqProjection const &projection) { +std::set output_dims_of_eq_projection(EqProjection const &projection) { return projection.dim_mapping.right_values(); } diff --git a/lib/utils/include/utils/orthotope/minimal_dim_domain.h b/lib/utils/include/utils/orthotope/minimal_dim_domain.h index 0720d2e8ec..b83eda18d7 100644 --- a/lib/utils/include/utils/orthotope/minimal_dim_domain.h +++ b/lib/utils/include/utils/orthotope/minimal_dim_domain.h @@ -2,8 +2,10 @@ #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ORTHOTOPE_MINIMAL_DIM_DOMAIN_H #include "utils/containers/are_disjoint.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include "utils/containers/filtermap_values.h" #include "utils/containers/generate_map.h" +#include "utils/containers/keys.h" #include "utils/containers/map_from_keys_and_values.h" #include "utils/containers/map_values.h" #include "utils/containers/restrict_keys.h" @@ -14,8 +16,6 @@ #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/minimal_dim_domain.dtg.h" #include "utils/orthotope/minimal_orthotope.dtg.h" -#include "utils/containers/keys.h" -#include "utils/containers/binary_merge_disjoint_maps.h" namespace FlexFlow { @@ -60,8 +60,7 @@ template DimDomain dim_domain_from_minimal_dim_domain( MinimalDimDomain const &minimal_dim_domain, std::set const &trivial_dims) { - std::set nontrivial_dims = - get_minimal_domain_dims(minimal_dim_domain); + std::set nontrivial_dims = get_minimal_domain_dims(minimal_dim_domain); ASSERT(are_disjoint(nontrivial_dims, trivial_dims)); @@ -75,8 +74,7 @@ DimDomain dim_domain_from_minimal_dim_domain( } template -std::set - get_minimal_domain_dims(MinimalDimDomain const &domain) { +std::set get_minimal_domain_dims(MinimalDimDomain const &domain) { return keys(domain.dims); } diff --git a/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h b/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h index c15152ceef..0520499764 100644 --- a/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h +++ b/lib/utils/include/utils/orthotope/minimal_dim_domain_mapping.h @@ -1,6 +1,7 @@ #ifndef _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ORTHOTOPE_MINIMAL_DIM_DOMAIN_MAPPING_H #define _FLEXFLOW_LIB_UTILS_INCLUDE_UTILS_ORTHOTOPE_MINIMAL_DIM_DOMAIN_MAPPING_H +#include "utils/bidict/algorithms/bidict_transform_keys_and_values.h" #include "utils/bidict/algorithms/exhaustive_relational_join.h" #include "utils/bidict/algorithms/left_entries.h" #include "utils/bidict/algorithms/right_entries.h" @@ -13,7 +14,6 @@ #include "utils/orthotope/dim_ordering.dtg.h" #include "utils/orthotope/dim_projection.h" #include "utils/orthotope/minimal_dim_domain.dtg.h" -#include "utils/bidict/algorithms/bidict_transform_keys_and_values.h" namespace FlexFlow { @@ -92,24 +92,20 @@ template MinimalDimDomainMapping minimal_mapping_from_dim_domain_mapping(DimDomainMapping const &m) { - std::set l_nontrivial_dims = - get_nontrivial_domain_dims(m.l_domain); + std::set l_nontrivial_dims = get_nontrivial_domain_dims(m.l_domain); - std::set r_nontrivial_dims = - get_nontrivial_domain_dims(m.r_domain); + std::set r_nontrivial_dims = get_nontrivial_domain_dims(m.r_domain); return MinimalDimDomainMapping{ /*coord_mapping=*/ bidict_transform_keys_and_values( - m.coord_mapping, - [&](DimCoord const &l_coord) { - return restrict_coord_to_dims(l_coord, - l_nontrivial_dims); - }, - [&](DimCoord const &r_coord) { - return restrict_coord_to_dims( - r_coord, r_nontrivial_dims); - }), + m.coord_mapping, + [&](DimCoord const &l_coord) { + return restrict_coord_to_dims(l_coord, l_nontrivial_dims); + }, + [&](DimCoord const &r_coord) { + return restrict_coord_to_dims(r_coord, r_nontrivial_dims); + }), /*l_domain=*/minimal_dim_domain_from_dim_domain(m.l_domain), /*r_domain=*/minimal_dim_domain_from_dim_domain(m.r_domain), }; @@ -132,14 +128,13 @@ DimDomainMapping dim_domain_mapping_from_minimal_dim_domain( return DimDomainMapping{ /*coord_mapping=*/ bidict_transform_keys_and_values( - m.coord_mapping, - [&](DimCoord const &l_coord) { - return lift_dim_coord(l_coord, all_l_dims); - }, - [&](DimCoord const &r_coord) { - return lift_dim_coord(r_coord, - all_r_dims); - }), + m.coord_mapping, + [&](DimCoord const &l_coord) { + return lift_dim_coord(l_coord, all_l_dims); + }, + [&](DimCoord const &r_coord) { + return lift_dim_coord(r_coord, all_r_dims); + }), /*l_domain=*/l_domain, /*r_domain=*/r_domain, }; @@ -207,14 +202,12 @@ DimDomainMapping compose_dim_domain_mappings_through_minimal( MinimalDimDomainMapping minimal_lhs = minimal_mapping_from_dim_domain_mapping(lhs); - std::set t1_trivial_dims = - get_trivial_domain_dims(lhs.l_domain); + std::set t1_trivial_dims = get_trivial_domain_dims(lhs.l_domain); MinimalDimDomainMapping minimal_rhs = minimal_mapping_from_dim_domain_mapping(rhs); - std::set t3_trivial_dims = - get_trivial_domain_dims(rhs.r_domain); + std::set t3_trivial_dims = get_trivial_domain_dims(rhs.r_domain); return dim_domain_mapping_from_minimal_dim_domain( compose_minimal_dim_domain_mappings(minimal_lhs, minimal_rhs), diff --git a/lib/utils/include/utils/orthotope/orthotope.h b/lib/utils/include/utils/orthotope/orthotope.h index ef8ad75846..013e30fc1b 100644 --- a/lib/utils/include/utils/orthotope/orthotope.h +++ b/lib/utils/include/utils/orthotope/orthotope.h @@ -10,8 +10,7 @@ nonnegative_int orthotope_get_num_dims(Orthotope const &); positive_int orthotope_get_volume(Orthotope const &); -std::set - get_all_coords_in_orthotope(Orthotope const &); +std::set get_all_coords_in_orthotope(Orthotope const &); bool orthotope_contains_coord(Orthotope const &, OrthotopeCoord const &); diff --git a/lib/utils/include/utils/orthotope/up_projection.h b/lib/utils/include/utils/orthotope/up_projection.h index 19b753c69f..5f694f8e90 100644 --- a/lib/utils/include/utils/orthotope/up_projection.h +++ b/lib/utils/include/utils/orthotope/up_projection.h @@ -22,14 +22,12 @@ UpProjection make_empty_up_projection() { } template -std::set - input_dims_of_up_projection(UpProjection const &projection) { +std::set input_dims_of_up_projection(UpProjection const &projection) { return projection.dim_mapping.left_values(); } template -std::set - output_dims_of_up_projection(UpProjection const &projection) { +std::set output_dims_of_up_projection(UpProjection const &projection) { return projection.dim_mapping.right_values(); } @@ -50,12 +48,10 @@ DimCoord compute_up_projection(UpProjection const &projection, DimCoord unlifted = DimCoord{ flatmap(coord.raw, - [&](L const &input_dim, nonnegative_int input_dim_val) - -> std::map - { + [&](L const &input_dim, nonnegative_int input_dim_val) + -> std::map { std::set dst_dims = - projection.dim_mapping.at_l(input_dim) - .unwrap_as_set(); + projection.dim_mapping.at_l(input_dim).unwrap_as_set(); DimDomain dst_domain = restrict_domain_to_dims(output_domain, dst_dims); diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc index e52c8703d8..70d5ebdd1a 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_keys_and_values.cc @@ -6,9 +6,7 @@ namespace FlexFlow { using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template - bidict bidict_from_keys_and_values( - std::vector const &, - std::vector const &); +template bidict bidict_from_keys_and_values(std::vector const &, + std::vector const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc index 271fc35582..6d5ddf503a 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_pairs.cc @@ -6,7 +6,6 @@ namespace FlexFlow { using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template - bidict bidict_from_pairs(std::vector> const &); +template bidict bidict_from_pairs(std::vector> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc b/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc index 7ff4b06576..689b7ee31a 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_from_unstructured_relation.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template bidict bidict_from_unstructured_relation( - std::set> const &); +template bidict + bidict_from_unstructured_relation(std::set> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc b/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc index 8e836c2da6..70008592c0 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_transform_keys_and_values.cc @@ -10,6 +10,7 @@ using V2 = ordered_value_type<3>; using KF = std::function; using VF = std::function; -template bidict bidict_transform_keys_and_values(bidict const &, KF &&, VF &&); +template bidict + bidict_transform_keys_and_values(bidict const &, KF &&, VF &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc b/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc index 339754a8f2..b557dc18b0 100644 --- a/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc +++ b/lib/utils/src/utils/bidict/algorithms/bidict_unordered_set_of.cc @@ -6,6 +6,7 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = ordered_value_type<1>; -std::unordered_set> bidict_unordered_set_of(bidict const &); +std::unordered_set> + bidict_unordered_set_of(bidict const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/bidict/bidict.cc b/lib/utils/src/utils/bidict/bidict.cc index 534b99f039..5a79f6e71f 100644 --- a/lib/utils/src/utils/bidict/bidict.cc +++ b/lib/utils/src/utils/bidict/bidict.cc @@ -1,8 +1,8 @@ #include "utils/bidict/bidict.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/jsonable_value_type.h" #include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/rapidcheckable_value_type.h" -#include "utils/archetypes/jsonable_ordered_value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc index 073265d427..9aac84d22f 100644 --- a/lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_left.cc @@ -8,7 +8,7 @@ using R = ordered_value_type<1>; using L2 = ordered_value_type<2>; using F = std::function; -template - BinaryRelation binary_relation_transform_left(BinaryRelation const &, F &&); +template BinaryRelation + binary_relation_transform_left(BinaryRelation const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc index a26b8d1ee4..1bda449114 100644 --- a/lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_left2.cc @@ -8,7 +8,7 @@ using L2 = ordered_value_type<1>; using R = ordered_value_type<2>; using F = std::function; -template - BinaryRelation binary_relation_transform_left2(BinaryRelation const &, F &&); +template BinaryRelation + binary_relation_transform_left2(BinaryRelation const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc index 5274340885..a8c12bbbae 100644 --- a/lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_right.cc @@ -8,7 +8,7 @@ using R = ordered_value_type<1>; using R2 = ordered_value_type<2>; using F = std::function; -template - BinaryRelation binary_relation_transform_right(BinaryRelation const &, F &&); +template BinaryRelation + binary_relation_transform_right(BinaryRelation const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc b/lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc index 9f76943eae..2a12e255d8 100644 --- a/lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc +++ b/lib/utils/src/utils/binary_relation/binary_relation_transform_right2.cc @@ -8,7 +8,7 @@ using R = ordered_value_type<1>; using R2 = ordered_value_type<2>; using F = std::function; -template - BinaryRelation binary_relation_transform_right2(BinaryRelation const &, F &&); +template BinaryRelation + binary_relation_transform_right2(BinaryRelation const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/filter_binary_relation.cc b/lib/utils/src/utils/binary_relation/filter_binary_relation.cc index fd0d497b98..83bf2fd301 100644 --- a/lib/utils/src/utils/binary_relation/filter_binary_relation.cc +++ b/lib/utils/src/utils/binary_relation/filter_binary_relation.cc @@ -7,6 +7,7 @@ using L = ordered_value_type<0>; using R = ordered_value_type<1>; using F = std::function; -template BinaryRelation filter_binary_relation(BinaryRelation const &, F &&); +template BinaryRelation + filter_binary_relation(BinaryRelation const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc b/lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc index 09b1c1d539..48bb5c6541 100644 --- a/lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc +++ b/lib/utils/src/utils/binary_relation/require_binary_relation_is_left_unique.cc @@ -6,8 +6,7 @@ namespace FlexFlow { using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template - OneToMany require_binary_relation_is_left_unique(BinaryRelation const &); - +template OneToMany + require_binary_relation_is_left_unique(BinaryRelation const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc b/lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc index aebdc9895c..b6000f932d 100644 --- a/lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc +++ b/lib/utils/src/utils/binary_relation/require_binary_relation_is_right_unique.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template - ManyToOne require_binary_relation_is_right_unique(BinaryRelation const &); +template ManyToOne + require_binary_relation_is_right_unique(BinaryRelation const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/are_disjoint.cc b/lib/utils/src/utils/containers/are_disjoint.cc index b71cb3730b..a72d56012a 100644 --- a/lib/utils/src/utils/containers/are_disjoint.cc +++ b/lib/utils/src/utils/containers/are_disjoint.cc @@ -1,12 +1,13 @@ #include "utils/containers/are_disjoint.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { using T = value_type<0>; -template bool are_disjoint(std::unordered_set const &, std::unordered_set const &); +template bool are_disjoint(std::unordered_set const &, + std::unordered_set const &); using R = ordered_value_type<0>; diff --git a/lib/utils/src/utils/containers/argmax.cc b/lib/utils/src/utils/containers/argmax.cc index f84c33e8d4..8b253659a1 100644 --- a/lib/utils/src/utils/containers/argmax.cc +++ b/lib/utils/src/utils/containers/argmax.cc @@ -3,7 +3,6 @@ #include "utils/archetypes/value_type.h" #include #include -#include #include namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/argmin.cc b/lib/utils/src/utils/containers/argmin.cc index 9f1434861e..5af603796a 100644 --- a/lib/utils/src/utils/containers/argmin.cc +++ b/lib/utils/src/utils/containers/argmin.cc @@ -3,7 +3,6 @@ #include "utils/archetypes/value_type.h" #include #include -#include #include namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/at_idx.cc b/lib/utils/src/utils/containers/at_idx.cc index c0f315d766..01bb0a6189 100644 --- a/lib/utils/src/utils/containers/at_idx.cc +++ b/lib/utils/src/utils/containers/at_idx.cc @@ -1,6 +1,6 @@ #include "utils/containers/at_idx.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/binary_cartesian_product.cc b/lib/utils/src/utils/containers/binary_cartesian_product.cc index 70195e2a1f..de3d3c398c 100644 --- a/lib/utils/src/utils/containers/binary_cartesian_product.cc +++ b/lib/utils/src/utils/containers/binary_cartesian_product.cc @@ -7,7 +7,6 @@ using A = ordered_value_type<0>; using B = ordered_value_type<1>; template std::set> - binary_cartesian_product(std::set const &, - std::set const &); + binary_cartesian_product(std::set const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc b/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc index 95d19083ac..cc428d7fd6 100644 --- a/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc +++ b/lib/utils/src/utils/containers/binary_merge_disjoint_maps.cc @@ -1,14 +1,13 @@ #include "utils/containers/binary_merge_disjoint_maps.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<1>; -template std::map - binary_merge_disjoint_maps(std::map const &, - std::map const &); +template std::map binary_merge_disjoint_maps(std::map const &, + std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_maps_with.cc b/lib/utils/src/utils/containers/binary_merge_maps_with.cc index 4679d21227..36f269461d 100644 --- a/lib/utils/src/utils/containers/binary_merge_maps_with.cc +++ b/lib/utils/src/utils/containers/binary_merge_maps_with.cc @@ -1,6 +1,6 @@ #include "utils/containers/binary_merge_maps_with.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,7 +8,8 @@ using K = ordered_value_type<0>; using V = value_type<1>; using F = std::function; -template std::map binary_merge_maps_with( - std::map const &, std::map const &, F &&); +template std::map binary_merge_maps_with(std::map const &, + std::map const &, + F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc b/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc index d5b4f6cebe..75bce861f0 100644 --- a/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_maps_with_left_dominating.cc @@ -1,6 +1,6 @@ #include "utils/containers/binary_merge_maps_with_left_dominating.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc index bbc799150e..48caff1db5 100644 --- a/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_maps_with_right_dominating.cc @@ -1,6 +1,6 @@ #include "utils/containers/binary_merge_maps_with_right_dominating.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc index d2f58b90eb..405dabeea2 100644 --- a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc @@ -1,6 +1,6 @@ #include "utils/containers/binary_merge_unordered_maps_with_left_dominating.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,7 +8,7 @@ using K = ordered_value_type<0>; using V = value_type<1>; template std::unordered_map - binary_merge_unordered_maps_with_left_dominating(std::unordered_map const &, - std::unordered_map const &); + binary_merge_unordered_maps_with_left_dominating( + std::unordered_map const &, std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc index 62382bfa03..a5ac198aa2 100644 --- a/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc +++ b/lib/utils/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc @@ -1,14 +1,14 @@ -#include "utils/containers/binary_merge_maps_with_right_dominating.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" +#include "utils/containers/binary_merge_maps_with_right_dominating.h" namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<1>; -template - std::map binary_merge_maps_with_right_dominating( - std::map const &, std::map const &); +template std::map + binary_merge_maps_with_right_dominating(std::map const &, + std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/contains_duplicates.cc b/lib/utils/src/utils/containers/contains_duplicates.cc index 757882319e..66987bd5cd 100644 --- a/lib/utils/src/utils/containers/contains_duplicates.cc +++ b/lib/utils/src/utils/containers/contains_duplicates.cc @@ -1,6 +1,6 @@ #include "utils/containers/contains_duplicates.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/contains_value.cc b/lib/utils/src/utils/containers/contains_value.cc index 4a8332d631..a63fad2b2b 100644 --- a/lib/utils/src/utils/containers/contains_value.cc +++ b/lib/utils/src/utils/containers/contains_value.cc @@ -1,6 +1,6 @@ #include "utils/containers/contains_value.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -13,5 +13,4 @@ using O_K = ordered_value_type<0>; template bool contains_value(std::map const &, V const &); - } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/filter.cc b/lib/utils/src/utils/containers/filter.cc index 4931d97704..eb0538dc0e 100644 --- a/lib/utils/src/utils/containers/filter.cc +++ b/lib/utils/src/utils/containers/filter.cc @@ -1,6 +1,6 @@ #include "utils/containers/filter.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,38 +8,29 @@ using VT0 = value_type<0>; using VT1 = value_type<1>; using OVT0 = ordered_value_type<0>; -template - std::vector filter(std::vector const &, - std::function const &); - -template - std::unordered_set filter(std::unordered_set const &, +template std::vector filter(std::vector const &, std::function const &); -template - std::unordered_map +template std::unordered_set + filter(std::unordered_set const &, + std::function const &); + +template std::unordered_map filter(std::unordered_map const &, std::function const &)> const &); -template - std::set filter( - std::set const &, - std::function const &); +template std::set filter(std::set const &, + std::function const &); -template - std::map +template std::map filter(std::map const &, std::function const &)> const &); -template - std::multiset - filter(std::multiset const &, - std::function const &); +template std::multiset filter(std::multiset const &, + std::function const &); -template - std::unordered_multiset +template std::unordered_multiset filter(std::unordered_multiset const &, std::function const &); - } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/filtrans.cc b/lib/utils/src/utils/containers/filtrans.cc index 6bb5bbbb85..bfd8ade53d 100644 --- a/lib/utils/src/utils/containers/filtrans.cc +++ b/lib/utils/src/utils/containers/filtrans.cc @@ -1,6 +1,6 @@ #include "utils/containers/filtrans.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -10,7 +10,8 @@ using F = std::function(In const &)>; template std::vector filtrans(std::vector const &, F &&); template std::unordered_set filtrans(std::unordered_set const &, F &&); -template std::unordered_multiset filtrans(std::unordered_multiset const &, F &&); +template std::unordered_multiset + filtrans(std::unordered_multiset const &, F &&); using O_In = ordered_value_type<0>; using O_Out = ordered_value_type<0>; diff --git a/lib/utils/src/utils/containers/flatmap.cc b/lib/utils/src/utils/containers/flatmap.cc index 2f71264c2e..e34240178f 100644 --- a/lib/utils/src/utils/containers/flatmap.cc +++ b/lib/utils/src/utils/containers/flatmap.cc @@ -38,8 +38,8 @@ using O_OutK = ordered_value_type<2>; using O_OutV = value_type<3>; using O_F3 = std::function(O_InK, O_InV)>; -template std::map - flatmap(std::map const &, O_F3 &&); +template std::map flatmap(std::map const &, + O_F3 &&); using F4 = std::function(In const &)>; diff --git a/lib/utils/src/utils/containers/generate_unordered_map.cc b/lib/utils/src/utils/containers/generate_unordered_map.cc index 73ecfb48ff..49fe19ae6e 100644 --- a/lib/utils/src/utils/containers/generate_unordered_map.cc +++ b/lib/utils/src/utils/containers/generate_unordered_map.cc @@ -8,8 +8,7 @@ using K = value_type<0>; using V = value_type<1>; using F = std::function; -template - std::unordered_map generate_unordered_map( - std::unordered_set const &, F &&); +template std::unordered_map + generate_unordered_map(std::unordered_set const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/get_all_assignments.cc b/lib/utils/src/utils/containers/get_all_assignments.cc index fbe9eaddb0..4d12448da3 100644 --- a/lib/utils/src/utils/containers/get_all_assignments.cc +++ b/lib/utils/src/utils/containers/get_all_assignments.cc @@ -1,6 +1,6 @@ #include "utils/containers/get_all_assignments.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" #include "utils/hash/unordered_map.h" namespace FlexFlow { @@ -11,6 +11,7 @@ using V = ordered_value_type<1>; template std::unordered_set> get_all_assignments(std::unordered_map> const &); -template std::set> get_all_assignments(std::map> const &); +template std::set> + get_all_assignments(std::map> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/get_element_counts.cc b/lib/utils/src/utils/containers/get_element_counts.cc index c640d51d35..ed44beee23 100644 --- a/lib/utils/src/utils/containers/get_element_counts.cc +++ b/lib/utils/src/utils/containers/get_element_counts.cc @@ -1,13 +1,15 @@ #include "utils/containers/get_element_counts.h" -#include "utils/containers/vector_of.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/containers/vector_of.h" namespace FlexFlow { using O_T = ordered_value_type<0>; -template std::map get_element_counts(std::vector const &); -template std::map get_element_counts(std::multiset const &); +template std::map + get_element_counts(std::vector const &); +template std::map + get_element_counts(std::multiset const &); std::map get_element_counts(std::string const &s) { return get_element_counts(vector_of(s)); diff --git a/lib/utils/src/utils/containers/get_only.cc b/lib/utils/src/utils/containers/get_only.cc index 8c24aa77e5..e9d890386c 100644 --- a/lib/utils/src/utils/containers/get_only.cc +++ b/lib/utils/src/utils/containers/get_only.cc @@ -1,6 +1,6 @@ #include "utils/containers/get_only.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/invert_map.cc b/lib/utils/src/utils/containers/invert_map.cc index 575f0854d9..dbb7468c03 100644 --- a/lib/utils/src/utils/containers/invert_map.cc +++ b/lib/utils/src/utils/containers/invert_map.cc @@ -6,7 +6,6 @@ namespace FlexFlow { using O_K = ordered_value_type<0>; using O_V = ordered_value_type<1>; -template std::map> - invert_map(std::map const &); +template std::map> invert_map(std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/invert_unordered_map.cc b/lib/utils/src/utils/containers/invert_unordered_map.cc index 5cf85480ab..fd1e13eba8 100644 --- a/lib/utils/src/utils/containers/invert_unordered_map.cc +++ b/lib/utils/src/utils/containers/invert_unordered_map.cc @@ -9,5 +9,4 @@ using V = value_type<1>; template std::unordered_map> invert_unordered_map(std::unordered_map const &); - } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/is_submapeq_of.cc b/lib/utils/src/utils/containers/is_submapeq_of.cc index db0c349557..0c70cf59de 100644 --- a/lib/utils/src/utils/containers/is_submapeq_of.cc +++ b/lib/utils/src/utils/containers/is_submapeq_of.cc @@ -8,5 +8,4 @@ using V = value_type<1>; bool is_submapeq_of(std::map const &, std::map const &); - } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/is_subseteq_of.cc b/lib/utils/src/utils/containers/is_subseteq_of.cc index c3ab1f8ef2..9ca2c6d42a 100644 --- a/lib/utils/src/utils/containers/is_subseteq_of.cc +++ b/lib/utils/src/utils/containers/is_subseteq_of.cc @@ -9,7 +9,6 @@ template bool is_subseteq_of(std::unordered_set const &, std::unordered_set const &); using T2 = ordered_value_type<0>; -template bool is_subseteq_of(std::set const &, - std::set const &); +template bool is_subseteq_of(std::set const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/lookup_in_map.cc b/lib/utils/src/utils/containers/lookup_in_map.cc index c8f3e2feab..9c596eca95 100644 --- a/lib/utils/src/utils/containers/lookup_in_map.cc +++ b/lib/utils/src/utils/containers/lookup_in_map.cc @@ -1,13 +1,12 @@ #include "utils/containers/lookup_in_map.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<1>; -template std::function - lookup_in_map(std::map const &map); +template std::function lookup_in_map(std::map const &map); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_from_keys_and_values.cc b/lib/utils/src/utils/containers/map_from_keys_and_values.cc index 0c94aace3c..21ef533d81 100644 --- a/lib/utils/src/utils/containers/map_from_keys_and_values.cc +++ b/lib/utils/src/utils/containers/map_from_keys_and_values.cc @@ -1,13 +1,13 @@ #include "utils/containers/map_from_keys_and_values.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { using K1 = ordered_value_type<0>; using V1 = value_type<1>; -template std::map - map_from_keys_and_values(std::vector const &, std::vector const &); +template std::map map_from_keys_and_values(std::vector const &, + std::vector const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_from_pairs.cc b/lib/utils/src/utils/containers/map_from_pairs.cc index 47936dc3da..d2f408c8e2 100644 --- a/lib/utils/src/utils/containers/map_from_pairs.cc +++ b/lib/utils/src/utils/containers/map_from_pairs.cc @@ -9,13 +9,11 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = ordered_value_type<1>; -template std::map - map_from_pairs(std::set> const &); +template std::map map_from_pairs(std::set> const &); template std::map map_from_pairs(std::unordered_set> const &); -template std::map - map_from_pairs(std::vector> const &); +template std::map map_from_pairs(std::vector> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_from_unordered.cc b/lib/utils/src/utils/containers/map_from_unordered.cc index 11558af765..ac8900ce29 100644 --- a/lib/utils/src/utils/containers/map_from_unordered.cc +++ b/lib/utils/src/utils/containers/map_from_unordered.cc @@ -7,7 +7,6 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<1>; -template - std::map map_from_unordered(std::unordered_map const &); +template std::map map_from_unordered(std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_keys.cc b/lib/utils/src/utils/containers/map_keys.cc index f5c07b91ec..2e363ea299 100644 --- a/lib/utils/src/utils/containers/map_keys.cc +++ b/lib/utils/src/utils/containers/map_keys.cc @@ -1,6 +1,6 @@ #include "utils/containers/map_keys.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,15 +8,13 @@ using K1 = value_type<0>; using K2 = value_type<1>; using V = value_type<2>; -template - std::unordered_map map_keys(std::unordered_map const &, - std::function &&); +template std::unordered_map map_keys(std::unordered_map const &, + std::function &&); using O_K1 = ordered_value_type<0>; using O_K2 = ordered_value_type<1>; -template - std::map map_keys(std::map const &m, - std::function &&); +template std::map map_keys(std::map const &m, + std::function &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_keys2.cc b/lib/utils/src/utils/containers/map_keys2.cc index 2401d64e81..e95cbdc7d5 100644 --- a/lib/utils/src/utils/containers/map_keys2.cc +++ b/lib/utils/src/utils/containers/map_keys2.cc @@ -1,6 +1,6 @@ #include "utils/containers/map_keys2.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -9,7 +9,6 @@ using V = value_type<1>; using K2 = ordered_value_type<2>; using F = std::function; -template std::map map_keys2(std::map const &, - F const &); +template std::map map_keys2(std::map const &, F const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_keys_and_values.cc b/lib/utils/src/utils/containers/map_keys_and_values.cc index 95608a09bd..fafe36044d 100644 --- a/lib/utils/src/utils/containers/map_keys_and_values.cc +++ b/lib/utils/src/utils/containers/map_keys_and_values.cc @@ -1,6 +1,6 @@ #include "utils/containers/map_keys_and_values.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -18,7 +18,7 @@ using OK = ordered_value_type<0>; using OK2 = ordered_value_type<1>; using OFK = std::function; -template std::map map_keys_and_values( - std::map const &, OFK const &, FV const &); +template std::map + map_keys_and_values(std::map const &, OFK const &, FV const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_keys_with_value_merging.cc b/lib/utils/src/utils/containers/map_keys_with_value_merging.cc index 7a6fd94e3f..a88ae50563 100644 --- a/lib/utils/src/utils/containers/map_keys_with_value_merging.cc +++ b/lib/utils/src/utils/containers/map_keys_with_value_merging.cc @@ -19,7 +19,7 @@ using O_K2 = ordered_value_type<2>; using O_F = std::function; -template std::map map_keys_with_value_merging( - std::map const &, O_F &&, MergeF &&); +template std::map + map_keys_with_value_merging(std::map const &, O_F &&, MergeF &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/map_values.cc b/lib/utils/src/utils/containers/map_values.cc index e850ecf31e..48adce3801 100644 --- a/lib/utils/src/utils/containers/map_values.cc +++ b/lib/utils/src/utils/containers/map_values.cc @@ -1,6 +1,6 @@ #include "utils/containers/map_values.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/map_values2.cc b/lib/utils/src/utils/containers/map_values2.cc index 9f95d32324..edf8debbff 100644 --- a/lib/utils/src/utils/containers/map_values2.cc +++ b/lib/utils/src/utils/containers/map_values2.cc @@ -1,6 +1,6 @@ #include "utils/containers/map_values2.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,14 +8,14 @@ using K = value_type<0>; using V1 = value_type<1>; using V2 = value_type<2>; -template std::unordered_map map_values2( - std::unordered_map const &, - std::function &&); +template std::unordered_map + map_values2(std::unordered_map const &, + std::function &&); using O_K = ordered_value_type<0>; -template std::map map_values2( - std::map const &, - std::function &&); +template std::map + map_values2(std::map const &, + std::function &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_disjoint_maps.cc b/lib/utils/src/utils/containers/merge_disjoint_maps.cc index e810b6b4f0..cf2d0fd3ec 100644 --- a/lib/utils/src/utils/containers/merge_disjoint_maps.cc +++ b/lib/utils/src/utils/containers/merge_disjoint_maps.cc @@ -1,6 +1,6 @@ #include "utils/containers/merge_disjoint_maps.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc b/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc index 6257356457..1056e08d2f 100644 --- a/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc +++ b/lib/utils/src/utils/containers/merge_disjoint_unordered_maps.cc @@ -6,8 +6,7 @@ namespace FlexFlow { using K = value_type<0>; using V = value_type<1>; -template - std::unordered_map merge_disjoint_unordered_maps( - std::vector> const &); +template std::unordered_map merge_disjoint_unordered_maps( + std::vector> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_in_map.cc b/lib/utils/src/utils/containers/merge_in_map.cc index 618128ff3b..3443dd8323 100644 --- a/lib/utils/src/utils/containers/merge_in_map.cc +++ b/lib/utils/src/utils/containers/merge_in_map.cc @@ -6,7 +6,6 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = ordered_value_type<1>; -template void merge_in_map(std::map const &, - std::map &); +template void merge_in_map(std::map const &, std::map &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_maps_with.cc b/lib/utils/src/utils/containers/merge_maps_with.cc index b5b471a1f8..0282265720 100644 --- a/lib/utils/src/utils/containers/merge_maps_with.cc +++ b/lib/utils/src/utils/containers/merge_maps_with.cc @@ -1,6 +1,6 @@ #include "utils/containers/merge_maps_with.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -8,7 +8,7 @@ using K = ordered_value_type<0>; using V = value_type<1>; using F = std::function; -template std::map - merge_maps_with(std::vector> const &, F &&); +template std::map merge_maps_with(std::vector> const &, + F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc index f33b46780c..dff4708c81 100644 --- a/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc +++ b/lib/utils/src/utils/containers/merge_maps_with_right_dominating.cc @@ -1,6 +1,6 @@ #include "utils/containers/merge_maps_with_right_dominating.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/merge_unordered_maps_with.cc b/lib/utils/src/utils/containers/merge_unordered_maps_with.cc index c9f09e61f0..3c1b9245ba 100644 --- a/lib/utils/src/utils/containers/merge_unordered_maps_with.cc +++ b/lib/utils/src/utils/containers/merge_unordered_maps_with.cc @@ -1,5 +1,5 @@ -#include "utils/containers/merge_maps_with.h" #include "utils/archetypes/value_type.h" +#include "utils/containers/merge_maps_with.h" namespace FlexFlow { @@ -7,8 +7,6 @@ using K = value_type<0>; using V = value_type<1>; using F = std::function; -std::map - merge_maps_with(std::vector> const &, - F &&); +std::map merge_maps_with(std::vector> const &, F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc b/lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc index 1dd7da70a3..9208df8f6a 100644 --- a/lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc +++ b/lib/utils/src/utils/containers/merge_unordered_maps_with_right_dominating.cc @@ -7,6 +7,7 @@ using K = value_type<0>; using V = value_type<1>; using C = std::vector>; -template std::unordered_map merge_unordered_maps_with_right_dominating(C const &); +template std::unordered_map + merge_unordered_maps_with_right_dominating(C const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/multiset_union.cc b/lib/utils/src/utils/containers/multiset_union.cc index 980f9ea262..5b5bba2a6d 100644 --- a/lib/utils/src/utils/containers/multiset_union.cc +++ b/lib/utils/src/utils/containers/multiset_union.cc @@ -1,24 +1,21 @@ #include "utils/containers/multiset_union.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { using T = value_type<0>; -template - std::unordered_multiset - multiset_union(std::unordered_multiset const &, - std::unordered_multiset const &); +template std::unordered_multiset + multiset_union(std::unordered_multiset const &, + std::unordered_multiset const &); using O_T = ordered_value_type<0>; -template - std::multiset - multiset_union(std::multiset const &, - std::multiset const &); +template std::multiset multiset_union(std::multiset const &, + std::multiset const &); -template - std::multiset multiset_union(std::vector> const &); +template std::multiset + multiset_union(std::vector> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/require_all_of.cc b/lib/utils/src/utils/containers/require_all_of.cc index 5d85ed3d1c..5865abf061 100644 --- a/lib/utils/src/utils/containers/require_all_of.cc +++ b/lib/utils/src/utils/containers/require_all_of.cc @@ -2,7 +2,6 @@ #include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/value_type.h" #include -#include namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/require_only_key.cc b/lib/utils/src/utils/containers/require_only_key.cc index 26ec81528a..d72f62ccb3 100644 --- a/lib/utils/src/utils/containers/require_only_key.cc +++ b/lib/utils/src/utils/containers/require_only_key.cc @@ -1,6 +1,6 @@ #include "utils/containers/require_only_key.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/require_two_keys.cc b/lib/utils/src/utils/containers/require_two_keys.cc index 30cc7419a3..ab19f4ba50 100644 --- a/lib/utils/src/utils/containers/require_two_keys.cc +++ b/lib/utils/src/utils/containers/require_two_keys.cc @@ -1,6 +1,6 @@ #include "utils/containers/require_two_keys.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/restrict_keys.cc b/lib/utils/src/utils/containers/restrict_keys.cc index 7c314733ba..812b64cf90 100644 --- a/lib/utils/src/utils/containers/restrict_keys.cc +++ b/lib/utils/src/utils/containers/restrict_keys.cc @@ -1,22 +1,19 @@ #include "utils/containers/restrict_keys.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { using VT0 = value_type<0>; using VT1 = value_type<1>; -template - std::unordered_map +template std::unordered_map restrict_keys(std::unordered_map const &, std::unordered_set const &); using OV0 = ordered_value_type<0>; -template - std::map restrict_keys(std::map const &, - std::set const &); - +template std::map restrict_keys(std::map const &, + std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/set_of.cc b/lib/utils/src/utils/containers/set_of.cc index f35f45a0ed..fa3e70f453 100644 --- a/lib/utils/src/utils/containers/set_of.cc +++ b/lib/utils/src/utils/containers/set_of.cc @@ -15,5 +15,4 @@ using V = ordered_value_type<1>; template std::set> set_of(std::map const &); - } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/set_union.cc b/lib/utils/src/utils/containers/set_union.cc index bb2ae9be76..b56e539940 100644 --- a/lib/utils/src/utils/containers/set_union.cc +++ b/lib/utils/src/utils/containers/set_union.cc @@ -6,13 +6,13 @@ namespace FlexFlow { using T = value_type<0>; -template std::unordered_set set_union(std::unordered_set const &, std::unordered_set const &); +template std::unordered_set set_union(std::unordered_set const &, + std::unordered_set const &); using O_T = ordered_value_type<0>; template std::set set_union(std::set const &, std::set const &); -template std::set - set_union(std::vector> const &); +template std::set set_union(std::vector> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/transform.cc b/lib/utils/src/utils/containers/transform.cc index 55255cbc1e..15427c801b 100644 --- a/lib/utils/src/utils/containers/transform.cc +++ b/lib/utils/src/utils/containers/transform.cc @@ -1,6 +1,6 @@ #include "utils/containers/transform.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -10,9 +10,11 @@ using F = std::function; template std::vector transform(std::vector const &, F const &); -template std::unordered_set transform(std::unordered_set const &, F const &); +template std::unordered_set transform(std::unordered_set const &, + F const &); -template std::unordered_multiset transform(std::unordered_multiset const &, F const &); +template std::unordered_multiset + transform(std::unordered_multiset const &, F const &); using In2 = ordered_value_type<0>; using Out2 = ordered_value_type<1>; @@ -20,7 +22,8 @@ using F2 = std::function; template std::set transform(std::set const &, F2 const &); -template std::multiset transform(std::multiset const &v, F2 const &f); +template std::multiset transform(std::multiset const &v, + F2 const &f); using F3 = std::function; @@ -31,16 +34,18 @@ using V = value_type<4>; using K2 = value_type<5>; using V2 = value_type<6>; -template std::unordered_map transform(std::unordered_map const &, - std::function(K const &, V const &)> const &); +template std::unordered_map + transform(std::unordered_map const &, + std::function(K const &, V const &)> const &); using K3 = ordered_value_type<3>; using V3 = value_type<4>; using K4 = ordered_value_type<5>; using V4 = value_type<6>; -template std::map transform(std::map const &, - std::function(K3 const &, V3 const &)> const &); +template std::map + transform(std::map const &, + std::function(K3 const &, V3 const &)> const &); template std::optional transform(std::optional const &o, std::function const &); diff --git a/lib/utils/src/utils/containers/transform_pairs.cc b/lib/utils/src/utils/containers/transform_pairs.cc index 1c6deb907d..f4564d224d 100644 --- a/lib/utils/src/utils/containers/transform_pairs.cc +++ b/lib/utils/src/utils/containers/transform_pairs.cc @@ -1,6 +1,6 @@ #include "utils/containers/transform_pairs.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -17,7 +17,7 @@ using O_R = ordered_value_type<1>; using O_Out = ordered_value_type<2>; using O_F = std::function; -template std::set - transform_pairs(std::set> const &, O_F &&); +template std::set transform_pairs(std::set> const &, + O_F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc b/lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc index c97da474e5..cbf83a70d6 100644 --- a/lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc +++ b/lib/utils/src/utils/containers/try_merge_nondisjoint_maps.cc @@ -7,9 +7,7 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<1>; -template - std::optional> - try_merge_nondisjoint_maps(std::map const &, - std::map const &); +template std::optional> + try_merge_nondisjoint_maps(std::map const &, std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc b/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc index 3576b1b7b5..015eb4a367 100644 --- a/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc +++ b/lib/utils/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc @@ -6,9 +6,8 @@ namespace FlexFlow { using K = value_type<0>; using V = value_type<1>; -template - std::optional> - try_merge_nondisjoint_unordered_maps(std::unordered_map const &, - std::unordered_map const &); +template std::optional> + try_merge_nondisjoint_unordered_maps(std::unordered_map const &, + std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_items.cc b/lib/utils/src/utils/containers/unordered_items.cc index 0ec4d51b15..a4f26fe1e0 100644 --- a/lib/utils/src/utils/containers/unordered_items.cc +++ b/lib/utils/src/utils/containers/unordered_items.cc @@ -7,8 +7,7 @@ namespace FlexFlow { using K = value_type<0>; using V = value_type<1>; -template - std::unordered_set> +template std::unordered_set> unordered_items(std::unordered_map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unordered_keys.cc b/lib/utils/src/utils/containers/unordered_keys.cc index 57ad2efa1c..d072220f8f 100644 --- a/lib/utils/src/utils/containers/unordered_keys.cc +++ b/lib/utils/src/utils/containers/unordered_keys.cc @@ -1,6 +1,6 @@ -#include "utils/containers/keys.h" #include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/value_type.h" +#include "utils/containers/keys.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc b/lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc index 9a56082145..728906722e 100644 --- a/lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc +++ b/lib/utils/src/utils/containers/unordered_map_from_keys_and_values.cc @@ -6,8 +6,7 @@ namespace FlexFlow { using K = value_type<0>; using V = value_type<1>; -template - std::unordered_map +template std::unordered_map unordered_map_from_keys_and_values(std::vector const &, std::vector const &); diff --git a/lib/utils/src/utils/containers/unordered_map_from_map.cc b/lib/utils/src/utils/containers/unordered_map_from_map.cc index a0ffa034b5..4ebc42f142 100644 --- a/lib/utils/src/utils/containers/unordered_map_from_map.cc +++ b/lib/utils/src/utils/containers/unordered_map_from_map.cc @@ -7,6 +7,7 @@ namespace FlexFlow { using K = ordered_value_type<0>; using V = value_type<0>; -template std::unordered_map unordered_map_from_map(std::map const &); +template std::unordered_map + unordered_map_from_map(std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc b/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc index 1c975c36e7..69ff87c404 100644 --- a/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc +++ b/lib/utils/src/utils/containers/unstructured_exhaustive_relational_join.cc @@ -8,8 +8,7 @@ using C = ordered_value_type<1>; using R = ordered_value_type<2>; template std::set> - unstructured_exhaustive_relational_join( - std::set> const &, - std::set> const &); + unstructured_exhaustive_relational_join(std::set> const &, + std::set> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/vector_of.cc b/lib/utils/src/utils/containers/vector_of.cc index 0c23655b0f..75c26b35d4 100644 --- a/lib/utils/src/utils/containers/vector_of.cc +++ b/lib/utils/src/utils/containers/vector_of.cc @@ -1,8 +1,8 @@ #include "utils/containers/vector_of.h" +#include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/value_type.h" #include #include -#include "utils/archetypes/ordered_value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/containers/without_nullopts.cc b/lib/utils/src/utils/containers/without_nullopts.cc index 25a3c85526..343d8cc86f 100644 --- a/lib/utils/src/utils/containers/without_nullopts.cc +++ b/lib/utils/src/utils/containers/without_nullopts.cc @@ -2,8 +2,7 @@ namespace FlexFlow { -template std::set - without_nullopts(std::set> const &); +template std::set without_nullopts(std::set> const &); template std::vector without_nullopts(std::vector> const &); diff --git a/lib/utils/src/utils/containers/zip_values_strict.cc b/lib/utils/src/utils/containers/zip_values_strict.cc index c1710129a5..5514fb13e7 100644 --- a/lib/utils/src/utils/containers/zip_values_strict.cc +++ b/lib/utils/src/utils/containers/zip_values_strict.cc @@ -1,6 +1,6 @@ #include "utils/containers/zip_values_strict.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -15,9 +15,6 @@ template std::unordered_map> using OV0 = ordered_value_type<0>; template std::map> - zip_values_strict(std::map const &, - std::map const &); - - + zip_values_strict(std::map const &, std::map const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/containers/zip_values_strict_with.cc b/lib/utils/src/utils/containers/zip_values_strict_with.cc index d7614ab177..4a1b7d357f 100644 --- a/lib/utils/src/utils/containers/zip_values_strict_with.cc +++ b/lib/utils/src/utils/containers/zip_values_strict_with.cc @@ -1,6 +1,6 @@ #include "utils/containers/zip_values_strict_with.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -10,7 +10,8 @@ using V2 = value_type<2>; using Out = value_type<3>; using F = std::function; -template std::map zip_values_strict_with( - std::map const &, std::map const &, F &&); +template std::map zip_values_strict_with(std::map const &, + std::map const &, + F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/full_binary_tree/get_leaves.cc b/lib/utils/src/utils/full_binary_tree/get_leaves.cc index 9e16012a97..a5dccf12a4 100644 --- a/lib/utils/src/utils/full_binary_tree/get_leaves.cc +++ b/lib/utils/src/utils/full_binary_tree/get_leaves.cc @@ -1,6 +1,6 @@ #include "utils/full_binary_tree/get_leaves.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/graph/algorithms.cc b/lib/utils/src/utils/graph/algorithms.cc index 80e073d6f3..7bef583579 100644 --- a/lib/utils/src/utils/graph/algorithms.cc +++ b/lib/utils/src/utils/graph/algorithms.cc @@ -5,7 +5,6 @@ #include "utils/containers/set_difference.h" #include "utils/containers/set_of.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/containers/values.h" #include "utils/exception.h" #include "utils/graph/digraph/algorithms/get_incoming_edges.h" @@ -55,8 +54,7 @@ struct GetNodesFunctor { } }; -std::set query_nodes(GraphView const &g, - std::set const &nodes) { +std::set query_nodes(GraphView const &g, std::set const &nodes) { NodeQuery query = NodeQuery{ query_set::match_values_in(set_of(nodes)), }; @@ -127,8 +125,7 @@ void add_edges(DiGraph &g, std::set const &edges) { } } -void add_edges(UndirectedGraph &g, - std::set const &edges) { +void add_edges(UndirectedGraph &g, std::set const &edges) { for (UndirectedEdge const &e : edges) { g.add_edge(e); } @@ -157,8 +154,7 @@ void remove_edges(DiGraph &g, std::set const &edges) { } } -void remove_edges(UndirectedGraph &g, - std::set const &edges) { +void remove_edges(UndirectedGraph &g, std::set const &edges) { for (UndirectedEdge const &e : edges) { ASSERT(contains_edge(g, e), "remove_edges expected edge to exist in UndirectedGraph"); @@ -167,7 +163,7 @@ void remove_edges(UndirectedGraph &g, } std::set get_node_edges(UndirectedGraphView const &g, - Node const &n) { + Node const &n) { UndirectedEdgeQuery query = UndirectedEdgeQuery{ query_set::match_single_value(n), }; @@ -175,22 +171,21 @@ std::set get_node_edges(UndirectedGraphView const &g, return g.query_edges(query); } -std::vector get_unchecked_dfs_ordering( - DiGraphView const &g, std::set const &starting_points) { +std::vector + get_unchecked_dfs_ordering(DiGraphView const &g, + std::set const &starting_points) { UncheckedDFSView dfs_view = unchecked_dfs(g, starting_points); return {dfs_view.begin(), dfs_view.end()}; } -std::vector - get_dfs_ordering(DiGraphView const &g, - std::set const &starting_points) { +std::vector get_dfs_ordering(DiGraphView const &g, + std::set const &starting_points) { CheckedDFSView dfs_view = dfs(g, starting_points); return {dfs_view.begin(), dfs_view.end()}; } -std::vector - get_bfs_ordering(DiGraphView const &g, - std::set const &starting_points) { +std::vector get_bfs_ordering(DiGraphView const &g, + std::set const &starting_points) { BFSView bfs_view = bfs(g, starting_points); return {bfs_view.begin(), bfs_view.end()}; } @@ -200,8 +195,7 @@ std::set get_neighbors(DiGraphView const &g, Node const &n) { return get_neighbors(undirected, n); } -std::set get_neighbors(UndirectedGraphView const &g, - Node const &n) { +std::set get_neighbors(UndirectedGraphView const &g, Node const &n) { return flatmap(get_node_edges(g, n), [&](UndirectedEdge const &edge) { return set_difference(get_endpoints(edge), {n}); }); @@ -212,8 +206,7 @@ UndirectedGraphView get_subgraph(UndirectedGraphView const &g, return UndirectedGraphView::create(g, nodes); } -DiGraphView get_subgraph(DiGraphView const &g, - std::set const &nodes) { +DiGraphView get_subgraph(DiGraphView const &g, std::set const &nodes) { return DiGraphView::create(g, nodes); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc index 35f67ac808..423e2698ad 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms.cc @@ -34,8 +34,7 @@ std::vector get_outputs(DataflowGraphView const &g, }); } -std::set - get_all_dataflow_outputs(DataflowGraphView const &g) { +std::set get_all_dataflow_outputs(DataflowGraphView const &g) { return g.query_outputs(dataflow_output_query_all()); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc index 241ad7efc7..9408ea4dc4 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/find_isomorphisms.cc @@ -8,9 +8,8 @@ namespace FlexFlow { std::set find_isomorphisms(DataflowGraphView const &src, DataflowGraphView const &dst) { - std::set open_isomorphisms = - find_isomorphisms(view_as_open_dataflow_graph(src), - view_as_open_dataflow_graph(dst)); + std::set open_isomorphisms = find_isomorphisms( + view_as_open_dataflow_graph(src), view_as_open_dataflow_graph(dst)); return transform(open_isomorphisms, [](OpenDataflowGraphIsomorphism const &open) { diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc index 6c6f2fc3db..9a85108270 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_incoming_edges.cc @@ -17,9 +17,8 @@ std::vector get_incoming_edges(DataflowGraphView const &g, }); } -std::set - get_incoming_edges(DataflowGraphView const &g, - std::set const &ns) { +std::set get_incoming_edges(DataflowGraphView const &g, + std::set const &ns) { DataflowEdgeQuery query = DataflowEdgeQuery{ query_set::matchall(), query_set::matchall(), diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc index d64ce3c17c..913c680892 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_outgoing_edges.cc @@ -5,7 +5,7 @@ namespace FlexFlow { std::set get_outgoing_edges(DataflowGraphView const &g, - Node const &n) { + Node const &n) { return g.query_edges(DataflowEdgeQuery{ query_set::match_single_value(n), query_set::matchall(), @@ -14,9 +14,8 @@ std::set get_outgoing_edges(DataflowGraphView const &g, }); } -std::set - get_outgoing_edges(DataflowGraphView const &g, - std::set const &ns) { +std::set get_outgoing_edges(DataflowGraphView const &g, + std::set const &ns) { DataflowEdgeQuery query = DataflowEdgeQuery{ query_set::match_values_in(set_of(ns)), query_set::matchall(), diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc index 8d322c3c1d..ea4e53c0f8 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_incoming_edges.cc @@ -5,9 +5,8 @@ namespace FlexFlow { -std::set - get_subgraph_incoming_edges(DataflowGraphView const &g, - std::set const &ns) { +std::set get_subgraph_incoming_edges(DataflowGraphView const &g, + std::set const &ns) { std::set all_nodes = get_nodes(g); query_set src_query = diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc index 6dfed0a9af..02f4f11d56 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/get_subgraph_outgoing_edges.cc @@ -5,9 +5,8 @@ namespace FlexFlow { -std::set - get_subgraph_outgoing_edges(DataflowGraphView const &g, - std::set const &ns) { +std::set get_subgraph_outgoing_edges(DataflowGraphView const &g, + std::set const &ns) { std::set all_nodes = get_nodes(g); query_set dst_query = diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc index 4c02f9c54e..d091c9aad4 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/transitive_reduced_dataflow_graph/get_transitive_reduced_edges_across_split.cc @@ -9,14 +9,11 @@ namespace FlexFlow { std::set get_transitive_reduced_edges_across_split( TransitiveReducedDataflowGraphView const &tr_g, BinarySeriesSplit const &split) { - std::set src_subgraph = - set_of(get_leaves(split.get_left_child())); - std::set dst_subgraph = - set_of(get_leaves(split.get_right_child())); + std::set src_subgraph = set_of(get_leaves(split.get_left_child())); + std::set dst_subgraph = set_of(get_leaves(split.get_right_child())); - std::set raw_edges = - get_edges_from_subgraph_to_subgraph( - tr_g.transitive_reduction, src_subgraph, dst_subgraph); + std::set raw_edges = get_edges_from_subgraph_to_subgraph( + tr_g.transitive_reduction, src_subgraph, dst_subgraph); return flatmap(raw_edges, [&](DirectedEdge const &e) { return get_dataflow_edges_from_node_to_node( diff --git a/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc b/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc index 3c8fabb51e..a54db45a8a 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/algorithms/view_as_open_dataflow_graph.cc @@ -12,9 +12,8 @@ std::set ViewDataflowGraphAsOpenDataflowGraph::query_nodes( return this->g.query_nodes(q); } -std::set - ViewDataflowGraphAsOpenDataflowGraph::query_edges( - OpenDataflowEdgeQuery const &q) const { +std::set ViewDataflowGraphAsOpenDataflowGraph::query_edges( + OpenDataflowEdgeQuery const &q) const { std::set closed_edges = this->g.query_edges(q.standard_edge_query); @@ -22,9 +21,8 @@ std::set [](DataflowEdge const &e) { return OpenDataflowEdge{e}; }); } -std::set - ViewDataflowGraphAsOpenDataflowGraph::query_outputs( - DataflowOutputQuery const &q) const { +std::set ViewDataflowGraphAsOpenDataflowGraph::query_outputs( + DataflowOutputQuery const &q) const { return this->g.query_outputs(q); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc b/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc index 6fa21a5a54..d27596ad8d 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/dataflow_graph_view.cc @@ -2,8 +2,7 @@ namespace FlexFlow { -std::set - DataflowGraphView::query_nodes(NodeQuery const &q) const { +std::set DataflowGraphView::query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } diff --git a/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc b/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc index 628679ce30..0abf6847a9 100644 --- a/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/dataflow_graph/i_dataflow_graph_view.cc @@ -11,8 +11,7 @@ std::set q.dsts, matchall(), }; - std::set dataflow_edges = - this->query_edges(dataflow_query); + std::set dataflow_edges = this->query_edges(dataflow_query); return transform(dataflow_edges, [](DataflowEdge const &e) { return DirectedEdge{e.src.node, e.dst.node}; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc index a42feea337..010cd14e8b 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/complete_bipartite_composite/get_cbc_decomposition.cc @@ -53,11 +53,10 @@ std::optional return std::nullopt; } - std::set from_head_to_tail = - g.query_edges(DirectedEdgeQuery{ - query_set::match_values_in(set_of(head)), - query_set::match_values_in(set_of(tail)), - }); + std::set from_head_to_tail = g.query_edges(DirectedEdgeQuery{ + query_set::match_values_in(set_of(head)), + query_set::match_values_in(set_of(tail)), + }); DiGraphView subgraph = get_subgraph(g, set_union(head, tail)); if (!is_complete_bipartite_digraph(subgraph, head)) { diff --git a/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc b/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc index d192dc8cdf..66e74f17d1 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/contract_node.cc @@ -17,8 +17,7 @@ std::set }); } -std::set - ContractNodeView::query_nodes(NodeQuery const &q) const { +std::set ContractNodeView::query_nodes(NodeQuery const &q) const { return transform(g.query_nodes(q), [&](Node const &n) { if (n == this->from) { return this->to; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc b/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc index 61f735db24..04bad14b96 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/flipped.cc @@ -13,8 +13,7 @@ std::set result, [](DirectedEdge const &e) { return flipped_directed_edge(e); }); } -std::set - FlippedView::query_nodes(NodeQuery const &query) const { +std::set FlippedView::query_nodes(NodeQuery const &query) const { return this->g.query_nodes(query); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc index 4af4620070..edc45fe171 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_ancestors.cc @@ -4,8 +4,7 @@ #include "utils/graph/digraph/algorithms/is_acyclic.h" namespace FlexFlow { -std::set get_ancestors(DiGraphView const &g, - Node const &starting_node) { +std::set get_ancestors(DiGraphView const &g, Node const &starting_node) { assert(is_acyclic(g)); return get_descendants(flipped(g), starting_node); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc index 6004b1a4d1..888382ee0d 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_descendants.cc @@ -9,12 +9,11 @@ namespace FlexFlow { std::set get_descendants(DiGraphView const &g, - Node const &starting_node) { + Node const &starting_node) { assert(is_acyclic(g)); assert(contains(get_nodes(g), starting_node)); - return set_of( - get_bfs_ordering(g, get_successors(g, starting_node))); + return set_of(get_bfs_ordering(g, get_successors(g, starting_node))); }; } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc index aaf0a850ef..6cf9b2fee1 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators.cc @@ -13,8 +13,7 @@ std::set get_dominators(DiGraphView const &g, Node const &n) { return get_dominators_map(g).at(n); } -std::set get_dominators(DiGraphView const &g, - std::set const &n) { +std::set get_dominators(DiGraphView const &g, std::set const &n) { ASSERT(n.size() > 0, "Cannot find dominators of no nodes"); std::optional> result = diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc index aaaf2f47a6..9e09da3494 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_dominators_map.cc @@ -15,8 +15,7 @@ namespace FlexFlow { -std::map> - get_dominators_map(DiGraphView const &g) { +std::map> get_dominators_map(DiGraphView const &g) { std::set initial_nodes = get_initial_nodes(g); std::queue queue; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc index f7e07905aa..29bda61c78 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_edges_from_subgraph_to_subgraph.cc @@ -4,10 +4,10 @@ namespace FlexFlow { -std::set get_edges_from_subgraph_to_subgraph( - DiGraphView const &g, - std::set const &src_subgraph, - std::set const &dst_subgraph) { +std::set + get_edges_from_subgraph_to_subgraph(DiGraphView const &g, + std::set const &src_subgraph, + std::set const &dst_subgraph) { if (!are_disjoint(src_subgraph, dst_subgraph)) { throw mk_runtime_error( fmt::format("get_edges_from_subgraph_to_subgraph(DiGraphView, ...) " diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc index 21dbb737e0..7420bb4328 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc @@ -15,8 +15,7 @@ namespace FlexFlow { std::map> get_imm_dominators_map(DiGraphView const &g) { - std::map> node_to_its_dominators = - get_dominators_map(g); + std::map> node_to_its_dominators = get_dominators_map(g); auto get_imm_dominator = [&](Node const &n) { std::set n_dominators = node_to_its_dominators.at(n); @@ -27,8 +26,8 @@ std::map> })); std::map dominator_counts = get_element_counts(recursive_dominator_list); - std::set imm_dominators = keys( - filter_values(dominator_counts, [](positive_int count) { return count <= 1; })); + std::set imm_dominators = keys(filter_values( + dominator_counts, [](positive_int count) { return count <= 1; })); ASSERT(imm_dominators.size() <= 1); return maybe_get_only(imm_dominators); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc index ea95b6829f..573e747cb8 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_imm_post_dominator.cc @@ -18,9 +18,8 @@ std::optional get_imm_post_dominator(DiGraphView const &g, return get_imm_post_dominators_map(g).at(n); } -std::optional - get_imm_post_dominator(DiGraphView const &g, - std::set const &nodes) { +std::optional get_imm_post_dominator(DiGraphView const &g, + std::set const &nodes) { if (nodes.empty()) { throw mk_runtime_error("Cannot get imm_post_dominator of no nodes"); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc index ad0f8cffaa..1f6e5ea9c5 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_incoming_edges.cc @@ -6,8 +6,7 @@ namespace FlexFlow { -std::set get_incoming_edges(DiGraphView const &g, - Node const &n) { +std::set get_incoming_edges(DiGraphView const &g, Node const &n) { return g.query_edges(DirectedEdgeQuery{ query_set::matchall(), query_set::match_single_value(n), @@ -15,23 +14,21 @@ std::set get_incoming_edges(DiGraphView const &g, } std::map> - get_incoming_edges(DiGraphView const &g, - std::set const &ns) { + get_incoming_edges(DiGraphView const &g, std::set const &ns) { std::map> by_dst = - group_by(g.query_edges(DirectedEdgeQuery{ - query_set::matchall(), - query_set::match_values_in(set_of(ns)), - }), - [](DirectedEdge const &e) { return e.dst; }) - .l_to_r(); - - std::map> result = - map_values(by_dst, - [](nonempty_set const &s) - -> std::set { - return s.unwrap_as_set(); - }); + group_by(g.query_edges(DirectedEdgeQuery{ + query_set::matchall(), + query_set::match_values_in(set_of(ns)), + }), + [](DirectedEdge const &e) { return e.dst; }) + .l_to_r(); + + std::map> result = map_values( + by_dst, + [](nonempty_set const &s) -> std::set { + return s.unwrap_as_set(); + }); for (Node const &n : ns) { result[n]; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc index a87170397c..ed74fdc543 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc @@ -22,12 +22,10 @@ std::optional> if (num_nodes(g) == 0 || nodes.size() == 0) { return std::nullopt; } - std::set> ancestors = - transform(nodes, [&](Node const &n) { - return set_union(get_ancestors(g, n), {n}); - }); - std::set common_ancestors = - set_intersection(ancestors).value(); + std::set> ancestors = transform(nodes, [&](Node const &n) { + return set_union(get_ancestors(g, n), {n}); + }); + std::set common_ancestors = set_intersection(ancestors).value(); if (common_ancestors.empty()) { return std::set{}; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc index ba42fca1b5..4ffaf03b2a 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_outgoing_edges.cc @@ -7,22 +7,21 @@ namespace FlexFlow { std::map> - get_outgoing_edges(DiGraphView const &g, - std::set const &ns) { + get_outgoing_edges(DiGraphView const &g, std::set const &ns) { std::map> by_src = - group_by(g.query_edges(DirectedEdgeQuery{ - query_set::match_values_in(set_of(ns)), - query_set::matchall(), - }), - [](DirectedEdge const &e) { return e.src; }) - .l_to_r(); - - std::map> result = - map_values(by_src, - [](nonempty_set const &s) -> std::set { - return s.unwrap_as_set(); - }); + group_by(g.query_edges(DirectedEdgeQuery{ + query_set::match_values_in(set_of(ns)), + query_set::matchall(), + }), + [](DirectedEdge const &e) { return e.src; }) + .l_to_r(); + + std::map> result = map_values( + by_src, + [](nonempty_set const &s) -> std::set { + return s.unwrap_as_set(); + }); for (Node const &n : ns) { result[n]; @@ -31,8 +30,7 @@ std::map> return result; } -std::set get_outgoing_edges(DiGraphView const &g, - Node const &n) { +std::set get_outgoing_edges(DiGraphView const &g, Node const &n) { return g.query_edges(DirectedEdgeQuery{ query_set::match_single_value(n), query_set::matchall(), diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc index 2b0b43c9a8..3bcc72ee67 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators.cc @@ -4,8 +4,7 @@ namespace FlexFlow { -std::set get_post_dominators(DiGraphView const &g, - Node const &n) { +std::set get_post_dominators(DiGraphView const &g, Node const &n) { return get_post_dominators_map(g).at(n); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc index c48e69cdbd..13363cd55c 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc @@ -4,8 +4,7 @@ namespace FlexFlow { -std::map> - get_post_dominators_map(DiGraphView const &g) { +std::map> get_post_dominators_map(DiGraphView const &g) { return get_dominators_map(flipped(g)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc index d7f99477d2..51d3e58722 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_predecessors.cc @@ -6,8 +6,7 @@ namespace FlexFlow { -std::map> - get_predecessors(DiGraphView const &g) { +std::map> get_predecessors(DiGraphView const &g) { return get_predecessors(g, get_nodes(g)); } @@ -15,13 +14,12 @@ std::set get_predecessors(DiGraphView const &g, Node const &n) { return get_predecessors(g, std::set{n}).at(n); } -std::map> - get_predecessors(DiGraphView const &g, std::set const &ns) { - return map_values(get_incoming_edges(g, ns), - [](std::set const &es) { - return transform( - es, [](DirectedEdge const &e) { return e.src; }); - }); +std::map> get_predecessors(DiGraphView const &g, + std::set const &ns) { + return map_values( + get_incoming_edges(g, ns), [](std::set const &es) { + return transform(es, [](DirectedEdge const &e) { return e.src; }); + }); } } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc index 9dce4f8472..9029cfaaa3 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators.cc @@ -4,8 +4,7 @@ namespace FlexFlow { -std::set get_strict_dominators(DiGraphView const &g, - Node const &n) { +std::set get_strict_dominators(DiGraphView const &g, Node const &n) { std::set result = get_dominators(g, {n}); result.erase(n); return result; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc index 07acbd274e..34498608e5 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_strict_dominators_map.cc @@ -4,8 +4,7 @@ namespace FlexFlow { -std::map> - get_strict_dominators_map(DiGraphView const &g) { +std::map> get_strict_dominators_map(DiGraphView const &g) { return transform(get_dominators_map(g), [](Node const &n, std::set const &doms) { std::set result = doms; diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc index 6889fc10f6..6996ccf7b0 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_outgoing_edges.cc @@ -5,10 +5,10 @@ namespace FlexFlow { -std::set get_subgraph_outgoing_edges( - DiGraphView const &g, std::set const &subgraph_nodes) { - std::set external_nodes = - set_minus(get_nodes(g), subgraph_nodes); +std::set + get_subgraph_outgoing_edges(DiGraphView const &g, + std::set const &subgraph_nodes) { + std::set external_nodes = set_minus(get_nodes(g), subgraph_nodes); DirectedEdgeQuery query = DirectedEdgeQuery{ query_set::match_values_in(set_of(subgraph_nodes)), query_set::match_values_in(set_of(external_nodes)), diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc index ac525353c0..c6fa456157 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_subgraph_successors.cc @@ -3,9 +3,8 @@ namespace FlexFlow { -std::set - get_subgraph_successors(DiGraphView const &g, - std::set const &subgraph_nodes) { +std::set get_subgraph_successors(DiGraphView const &g, + std::set const &subgraph_nodes) { std::set successors = transform(get_subgraph_outgoing_edges(g, subgraph_nodes), [](DirectedEdge const &e) { return e.dst; }); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc index 8864e27417..00c52e4977 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_successors.cc @@ -4,8 +4,7 @@ namespace FlexFlow { -std::map> - get_successors(DiGraphView const &g) { +std::map> get_successors(DiGraphView const &g) { return get_predecessors(flipped(g)); } @@ -13,8 +12,8 @@ std::set get_successors(DiGraphView const &g, Node const &n) { return get_predecessors(flipped(g), n); } -std::map> - get_successors(DiGraphView const &g, std::set const &ns) { +std::map> get_successors(DiGraphView const &g, + std::set const &ns) { return get_predecessors(flipped(g), ns); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc b/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc index 48c1e87a45..1925acd0aa 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc @@ -4,8 +4,7 @@ namespace FlexFlow { -std::set> - get_weakly_connected_components(DiGraphView const &g) { +std::set> get_weakly_connected_components(DiGraphView const &g) { return get_connected_components(as_undirected(g)); } diff --git a/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc b/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc index ce4b029f4b..9a986372c5 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/transitive_closure.cc @@ -17,9 +17,9 @@ DiGraphView transitive_closure(DiGraphView const &g) { // incredibly slow (> minutes) for even moderately sized graphs // (i.e., 200 nodes) without optimization enabled. - bidict nodes = - bidict_transform_keys(bidict_from_enumerating(get_nodes(g)), - [](nonnegative_int x) { return x.unwrap_nonnegative(); }); + bidict nodes = bidict_transform_keys( + bidict_from_enumerating(get_nodes(g)), + [](nonnegative_int x) { return x.unwrap_nonnegative(); }); std::set edges = get_edges(g); int num_nodes = nodes.size(); diff --git a/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc b/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc index 2c3b7f2b0c..3bfc74c376 100644 --- a/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc +++ b/lib/utils/src/utils/graph/digraph/algorithms/transitive_reduction.cc @@ -22,8 +22,7 @@ std::set return set_intersection(g.query_edges(q), this->edge_mask); } -std::set - DirectedEdgeMaskView::query_nodes(NodeQuery const &q) const { +std::set DirectedEdgeMaskView::query_nodes(NodeQuery const &q) const { return g.query_nodes(q); } @@ -40,9 +39,9 @@ DiGraph transitive_reduction(DiGraphView const &g) { // transitive_closure inlined to avoid any drifts in node numbering // between transitive_closure and transitive_reduction - bidict nodes = - bidict_transform_keys(bidict_from_enumerating(get_nodes(g)), - [](nonnegative_int x) { return x.unwrap_nonnegative(); }); + bidict nodes = bidict_transform_keys( + bidict_from_enumerating(get_nodes(g)), + [](nonnegative_int x) { return x.unwrap_nonnegative(); }); int num_nodes = nodes.size(); std::vector edge_matrix(num_nodes * num_nodes, false); diff --git a/lib/utils/src/utils/graph/digraph/digraph.cc b/lib/utils/src/utils/graph/digraph/digraph.cc index 41070d73b5..be9cfde326 100644 --- a/lib/utils/src/utils/graph/digraph/digraph.cc +++ b/lib/utils/src/utils/graph/digraph/digraph.cc @@ -26,8 +26,7 @@ std::set DiGraph::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } -std::set - DiGraph::query_edges(DirectedEdgeQuery const &q) const { +std::set DiGraph::query_edges(DirectedEdgeQuery const &q) const { return this->get_ptr().query_edges(q); } diff --git a/lib/utils/src/utils/graph/digraph/digraph_view.cc b/lib/utils/src/utils/graph/digraph/digraph_view.cc index 53cf868514..36f80b2758 100644 --- a/lib/utils/src/utils/graph/digraph/digraph_view.cc +++ b/lib/utils/src/utils/graph/digraph/digraph_view.cc @@ -6,8 +6,7 @@ std::set DiGraphView::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } -std::set - DiGraphView::query_edges(EdgeQuery const &query) const { +std::set DiGraphView::query_edges(EdgeQuery const &query) const { return get_ptr().query_edges(query); } diff --git a/lib/utils/src/utils/graph/instances/adjacency_digraph.cc b/lib/utils/src/utils/graph/instances/adjacency_digraph.cc index dd4e6672d1..e8a7770f00 100644 --- a/lib/utils/src/utils/graph/instances/adjacency_digraph.cc +++ b/lib/utils/src/utils/graph/instances/adjacency_digraph.cc @@ -52,8 +52,7 @@ std::set return result; } -std::set - AdjacencyDiGraph::query_nodes(NodeQuery const &query) const { +std::set AdjacencyDiGraph::query_nodes(NodeQuery const &query) const { return apply_query(query.nodes, keys(this->adjacency)); } diff --git a/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc b/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc index faf4fea868..bf6eacb939 100644 --- a/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc +++ b/lib/utils/src/utils/graph/instances/adjacency_multidigraph.cc @@ -2,11 +2,11 @@ #include "utils/containers/contains_key.h" #include "utils/containers/extend.h" #include "utils/containers/generate_map.h" +#include "utils/containers/keys.h" #include "utils/containers/values.h" #include "utils/graph/multidigraph/algorithms/get_edges.h" #include "utils/graph/node/algorithms.h" #include "utils/hash/set.h" -#include "utils/containers/keys.h" namespace FlexFlow { @@ -15,21 +15,16 @@ AdjacencyMultiDiGraph::AdjacencyMultiDiGraph() {} AdjacencyMultiDiGraph::AdjacencyMultiDiGraph( NodeSource const &node_source, MultiDiEdgeSource const &edge_source, - std::map< - Node, - std::map>> const - &adjacency, + std::map>> const &adjacency, std::map> const &edge_nodes) : node_source(node_source), edge_source(edge_source), adjacency(adjacency), edge_nodes(edge_nodes) {} Node AdjacencyMultiDiGraph::add_node() { Node new_node = this->node_source.new_node(); - std::set all_nodes = - set_union(keys(this->adjacency), {new_node}); - this->adjacency[new_node] = generate_map(all_nodes, [](Node const &) { - return std::set{}; - }); + std::set all_nodes = set_union(keys(this->adjacency), {new_node}); + this->adjacency[new_node] = generate_map( + all_nodes, [](Node const &) { return std::set{}; }); for (Node const &n : all_nodes) { this->adjacency.at(n)[new_node] = {}; @@ -49,8 +44,7 @@ MultiDiEdge AdjacencyMultiDiGraph::add_edge(Node const &src, Node const &dst) { void AdjacencyMultiDiGraph::remove_node(Node const &n) { assert(contains_key(this->adjacency, n)); - std::set outgoing = - set_union(values(this->adjacency.at(n))); + std::set outgoing = set_union(values(this->adjacency.at(n))); std::set incoming; for (auto const &[k, v] : this->adjacency) { if (k != n) { @@ -75,8 +69,7 @@ void AdjacencyMultiDiGraph::remove_edge(MultiDiEdge const &e) { this->adjacency.at(src).at(dst).erase(e); } -std::set - AdjacencyMultiDiGraph::query_nodes(NodeQuery const &q) const { +std::set AdjacencyMultiDiGraph::query_nodes(NodeQuery const &q) const { return apply_query(q.nodes, keys(this->adjacency)); } @@ -109,8 +102,8 @@ void AdjacencyMultiDiGraph::inplace_materialize_from( std::set edges = get_edges(g); this->adjacency = generate_map(nodes, [&](Node const &) { - return generate_map( - nodes, [&](Node const &) { return std::set{}; }); + return generate_map(nodes, + [&](Node const &) { return std::set{}; }); }); this->edge_nodes.clear(); diff --git a/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc b/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc index 9da333ced0..faf8c1fd79 100644 --- a/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/instances/unordered_set_dataflow_graph.cc @@ -2,8 +2,8 @@ #include "utils/containers/are_disjoint.h" #include "utils/containers/enumerate_vector.h" #include "utils/containers/extend.h" -#include "utils/containers/transform.h" #include "utils/containers/set_of.h" +#include "utils/containers/transform.h" #include "utils/graph/dataflow_graph/algorithms.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_dataflow_graph/open_dataflow_edge.h" @@ -73,8 +73,7 @@ std::set UnorderedSetDataflowGraph::query_outputs( }); } -std::set - UnorderedSetDataflowGraph::get_inputs() const { +std::set UnorderedSetDataflowGraph::get_inputs() const { return this->graph_inputs; } diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc index 18ee3e677a..06dd9f1e1e 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc @@ -7,7 +7,6 @@ using SlotName = jsonable_ordered_value_type<0>; template DataflowGraphData dataflow_graph_data_from_kwarg_dataflow_graph_data( KwargDataflowGraphData const &, - std::function< - std::vector(std::set const &)> const &); + std::function(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc index 927c6e2d82..bf83174793 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc @@ -7,7 +7,6 @@ using SlotName = jsonable_ordered_value_type<0>; template DataflowGraphView dataflow_graph_from_kwarg_dataflow_graph( KwargDataflowGraphView const &, - std::function< - std::vector(std::set const &)> const &); + std::function(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc index d9a278099d..81429e8014 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_incoming_edges.cc @@ -7,7 +7,6 @@ using SlotName = ordered_value_type<0>; template std::set> get_kwarg_dataflow_subgraph_incoming_edges( - KwargDataflowGraphView const &, - std::set const &); + KwargDataflowGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc index 3333408b2a..014482e18f 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_subgraph_outgoing_edges.cc @@ -7,7 +7,6 @@ using SlotName = ordered_value_type<0>; template std::set> get_kwarg_dataflow_subgraph_outgoing_edges( - KwargDataflowGraphView const &, - std::set const &); + KwargDataflowGraphView const &, std::set const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc index 6f6d9bacad..aab8e8cb3a 100644 --- a/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc +++ b/lib/utils/src/utils/graph/kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_as_dot.cc @@ -11,7 +11,6 @@ template std::string kwarg_dataflow_graph_as_dot( std::function const &)> const &, std::function const &, - std::function< - std::vector(std::set const &)> const &); + std::function(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc index 85abc54c05..25de65a98c 100644 --- a/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc +++ b/lib/utils/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/labelled_kwarg_dataflow_graph_view_as_dot.cc @@ -13,7 +13,6 @@ template std::string labelled_kwarg_dataflow_graph_view_as_dot( std::function const &, std::function const &, std::function const &, - std::function< - std::vector(std::set const &)> const &); + std::function(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc index 1dcd6c9f2d..29f8acc46f 100644 --- a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc +++ b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/labelled_open_kwarg_dataflow_graph_view_as_dot.cc @@ -17,7 +17,6 @@ template std::string labelled_open_kwarg_dataflow_graph_view_as_dot( std::function const &, std::function const &, std::function const &, - std::function< - std::vector(std::set const &)> const &); + std::function(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc index 1ac5c9f239..effcb294c1 100644 --- a/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc +++ b/lib/utils/src/utils/graph/labelled_open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_view_with_labelling.cc @@ -27,6 +27,6 @@ template LabelledOpenKwargDataflowGraphView const &, std::map const &, std::map, - ValueLabel> const &); + ValueLabel> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc index 8d98e1c50d..08115b9915 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc @@ -11,7 +11,7 @@ namespace FlexFlow { std::set get_incoming_edges(MultiDiGraphView const &g, - Node const &n) { + Node const &n) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::matchall(), query_set::match_single_value(n), @@ -21,8 +21,7 @@ std::set get_incoming_edges(MultiDiGraphView const &g, } std::map> - get_incoming_edges(MultiDiGraphView const &g, - std::set const &ns) { + get_incoming_edges(MultiDiGraphView const &g, std::set const &ns) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::matchall(), query_set::match_values_in(set_of(ns)), @@ -32,8 +31,7 @@ std::map> group_by(g.query_edges(query), [&](MultiDiEdge const &e) { return g.get_multidiedge_dst(e); }) .l_to_r(), - [](nonempty_set const &s) - -> std::set { + [](nonempty_set const &s) -> std::set { return s.unwrap_as_set(); }); diff --git a/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc b/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc index c66eed3dca..3dab3ec975 100644 --- a/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc @@ -9,7 +9,7 @@ namespace FlexFlow { std::set get_outgoing_edges(MultiDiGraphView const &g, - Node const &n) { + Node const &n) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::match_single_value(n), query_set::matchall(), @@ -19,8 +19,7 @@ std::set get_outgoing_edges(MultiDiGraphView const &g, } std::map> - get_outgoing_edges(MultiDiGraphView const &g, - std::set const &ns) { + get_outgoing_edges(MultiDiGraphView const &g, std::set const &ns) { MultiDiEdgeQuery query = MultiDiEdgeQuery{ query_set::match_values_in(set_of(ns)), query_set::matchall(), @@ -30,8 +29,7 @@ std::map> group_by(g.query_edges(query), [&](MultiDiEdge const &e) { return g.get_multidiedge_src(e); }) .l_to_r(), - [](nonempty_set const &s) - -> std::set { + [](nonempty_set const &s) -> std::set { return s.unwrap_as_set(); }); diff --git a/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc b/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc index b308293a49..50b38ec436 100644 --- a/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc +++ b/lib/utils/src/utils/graph/multidigraph/multidigraph_view.cc @@ -2,8 +2,7 @@ namespace FlexFlow { -std::set - MultiDiGraphView::query_nodes(NodeQuery const &q) const { +std::set MultiDiGraphView::query_nodes(NodeQuery const &q) const { return this->get_interface().query_nodes(q); } diff --git a/lib/utils/src/utils/graph/node/node_query.cc b/lib/utils/src/utils/graph/node/node_query.cc index 79b53cb686..c76d123944 100644 --- a/lib/utils/src/utils/graph/node/node_query.cc +++ b/lib/utils/src/utils/graph/node/node_query.cc @@ -29,7 +29,7 @@ NodeQuery query_union(NodeQuery const &lhs, NodeQuery const &rhs) { } std::set apply_node_query(NodeQuery const &query, - std::set const &ns) { + std::set const &ns) { return apply_query(query.nodes, ns); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc index 78a160c390..fc47776fe6 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/find_isomorphisms.cc @@ -32,15 +32,13 @@ static std::optional bidict const &unused_graph_inputs_mapping) { { - std::set already_mapped_src_nodes = - left_entries(sink_node_mapping); + std::set already_mapped_src_nodes = left_entries(sink_node_mapping); std::set src_g_sink_nodes = set_of(get_terminal_nodes(src_g)); ASSERT(already_mapped_src_nodes == src_g_sink_nodes); } { - std::set already_mapped_dst_nodes = - right_entries(sink_node_mapping); + std::set already_mapped_dst_nodes = right_entries(sink_node_mapping); std::set dst_g_sink_nodes = set_of(get_terminal_nodes(dst_g)); ASSERT(already_mapped_dst_nodes == dst_g_sink_nodes); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc index c36bfee6ab..7baa8f5083 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/from_open_dataflow_graph_data.cc @@ -24,8 +24,7 @@ std::set FromOpenDataflowGraphDataView::query_outputs( return apply_dataflow_output_query(q, this->data.outputs); } -std::set - FromOpenDataflowGraphDataView::get_inputs() const { +std::set FromOpenDataflowGraphDataView::get_inputs() const { return this->data.inputs; } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc index 72de3b1b4d..84c1dff0e4 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_incoming_edges.cc @@ -8,13 +8,11 @@ namespace FlexFlow { -std::set - get_incoming_edges(OpenDataflowGraphView const &g) { - std::set raw_edges = - g.query_edges(OpenDataflowEdgeQuery{ - dataflow_input_edge_query_all(), - dataflow_edge_query_none(), - }); +std::set get_incoming_edges(OpenDataflowGraphView const &g) { + std::set raw_edges = g.query_edges(OpenDataflowEdgeQuery{ + dataflow_input_edge_query_all(), + dataflow_edge_query_none(), + }); return transform(raw_edges, [](OpenDataflowEdge const &e) { return e.get(); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc index a88dcdb4cb..4413cc268f 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc @@ -17,9 +17,8 @@ namespace FlexFlow { -OpenDataflowSubgraphResult - get_subgraph(OpenDataflowGraphView const &g, - std::set const &subgraph_nodes) { +OpenDataflowSubgraphResult get_subgraph(OpenDataflowGraphView const &g, + std::set const &subgraph_nodes) { bidict full_graph_values_to_subgraph_inputs = get_full_graph_values_to_subgraph_inputs(g, subgraph_nodes); @@ -34,8 +33,7 @@ OpenDataflowSubgraphResult bidict get_full_graph_values_to_subgraph_inputs( - OpenDataflowGraphView const &g, - std::set const &subgraph_nodes) { + OpenDataflowGraphView const &g, std::set const &subgraph_nodes) { DataflowGraphInputSource input_source; return generate_bidict(get_subgraph_inputs(g, subgraph_nodes), [&](OpenDataflowValue const &v) -> DataflowGraphInput { diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc index d73a582252..8dbf23ada0 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/algorithms/is_isomorphic_under.cc @@ -14,14 +14,15 @@ bool is_isomorphic_under( OpenDataflowGraphIsomorphism const &candidate_isomorphism) { bidict node_permutation = - bidict_transform_values(candidate_isomorphism.node_mapping, - [](Node const &dst_node) { return NewNode{dst_node}; }) + bidict_transform_values( + candidate_isomorphism.node_mapping, + [](Node const &dst_node) { return NewNode{dst_node}; }) .reversed(); bidict input_permutation = bidict_transform_values(candidate_isomorphism.input_mapping, - [](DataflowGraphInput const &dst_input) { - return NewDataflowGraphInput{dst_input}; - }) + [](DataflowGraphInput const &dst_input) { + return NewDataflowGraphInput{dst_input}; + }) .reversed(); return get_graph_data(permute_input_ids( permute_node_ids(src, node_permutation), input_permutation)) == diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc b/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc index 5fb5d4ce4e..14bb153f5a 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/i_open_dataflow_graph_view.cc @@ -11,8 +11,7 @@ std::set q, }; - std::set open_edges = - this->query_edges(open_query); + std::set open_edges = this->query_edges(open_query); return transform(open_edges, [](OpenDataflowEdge const &e) { return e.get(); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc index 7899ad1c36..b959268d51 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_edge_query.cc @@ -58,9 +58,9 @@ OpenDataflowEdgeQuery }; } -std::set apply_open_dataflow_edge_query( - OpenDataflowEdgeQuery const &q, - std::set const &es) { +std::set + apply_open_dataflow_edge_query(OpenDataflowEdgeQuery const &q, + std::set const &es) { return filter(es, [&](OpenDataflowEdge const &e) { return open_dataflow_edge_query_includes(q, e); }); diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc index 339087c1bd..42c37c283f 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/open_dataflow_graph_view.cc @@ -2,8 +2,7 @@ namespace FlexFlow { -std::set - OpenDataflowGraphView::get_inputs() const { +std::set OpenDataflowGraphView::get_inputs() const { return this->get_interface().get_inputs(); } diff --git a/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc b/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc index be1478d38b..7bfb5781a0 100644 --- a/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc +++ b/lib/utils/src/utils/graph/open_dataflow_graph/unordered_set_open_dataflow_graph.cc @@ -57,8 +57,7 @@ std::set UnorderedSetOpenDataflowGraph::query_outputs( }); } -std::set - UnorderedSetOpenDataflowGraph::get_inputs() const { +std::set UnorderedSetOpenDataflowGraph::get_inputs() const { return this->graph_inputs; } diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc index 291188723b..0d847dd681 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_all_open_kwarg_dataflow_edges.cc @@ -6,9 +6,8 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template - std::set> - get_all_open_kwarg_dataflow_edges( - OpenKwargDataflowGraphView const &); +template std::set> + get_all_open_kwarg_dataflow_edges( + OpenKwargDataflowGraphView const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc index 3bffb4f9f5..4cbdd34531 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_edges_for_node.cc @@ -6,8 +6,7 @@ namespace FlexFlow { using GraphInputName = ordered_value_type<0>; using SlotName = ordered_value_type<1>; -template std::map> +template std::map> get_incoming_open_kwarg_dataflow_edges_for_node( OpenKwargDataflowGraphView const &, Node const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc index be7affcded..077e2c9673 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/get_incoming_open_kwarg_dataflow_values_for_node.cc @@ -6,8 +6,7 @@ namespace FlexFlow { using SlotName = ordered_value_type<0>; using GraphInputName = ordered_value_type<1>; -template std::map> +template std::map> get_incoming_open_kwarg_dataflow_values_for_node( OpenKwargDataflowGraphView const &, Node const &); diff --git a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc index 643e600ba5..c05e88bc80 100644 --- a/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc +++ b/lib/utils/src/utils/graph/open_kwarg_dataflow_graph/algorithms/open_kwarg_dataflow_graph_as_dot.cc @@ -13,7 +13,6 @@ template std::string open_kwarg_dataflow_graph_as_dot( std::function const &)> const &, std::function const &, - std::function< - std::vector(std::set const &)> const &); + std::function(std::set const &)> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/graph/render_dot.cc b/lib/utils/src/utils/graph/render_dot.cc index 6f163f73d8..958cbfef07 100644 --- a/lib/utils/src/utils/graph/render_dot.cc +++ b/lib/utils/src/utils/graph/render_dot.cc @@ -59,9 +59,9 @@ std::string render_node_label( return oss.str(); } -std::string render_dot( - LabelledDataflowGraphView, - std::string> const &g) { +std::string + render_dot(LabelledDataflowGraphView, + std::string> const &g) { std::vector lines; lines.push_back("digraph {"); diff --git a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc index 9a76b940d5..ed74151edd 100644 --- a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc +++ b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/balanced_binary_sp_tree_from_nary.cc @@ -1,8 +1,8 @@ #include "utils/containers/foldl1.h" #include "utils/containers/get_only.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/slice.h" #include "utils/containers/transform.h" -#include "utils/containers/multiset_of.h" #include "utils/containers/vector_of.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_parallel_split.dtg.h" #include "utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.dtg.h" @@ -13,7 +13,6 @@ #include "utils/overload.h" #include #include -#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -45,8 +44,7 @@ BinarySPDecompositionTree } auto s1 = multiset_of(slice(children, 0, children.size() / 2)); - auto s2 = multiset_of( - slice(children, children.size() / 2, std::nullopt)); + auto s2 = multiset_of(slice(children, children.size() / 2, std::nullopt)); return BinarySPDecompositionTree{BinaryParallelSplit{ from_parallel(ParallelSplit{s1}), from_parallel(ParallelSplit{s2})}}; diff --git a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc index cb0505a6a9..cef30e0ffe 100644 --- a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc +++ b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/binary_sp_decomposition_tree.cc @@ -66,8 +66,7 @@ bool is_binary_sp_tree_right_associative( generic_impl_for_binary_sp_tree()); } -std::multiset - get_leaves(BinarySPDecompositionTree const &tree) { +std::multiset get_leaves(BinarySPDecompositionTree const &tree) { return get_leaves(tree, generic_impl_for_binary_sp_tree()); } diff --git a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc index 3929f35f94..06ab4e7f20 100644 --- a/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc +++ b/lib/utils/src/utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.cc @@ -1,6 +1,6 @@ #include "utils/graph/series_parallel/binary_sp_decomposition_tree/generic_binary_sp_decomposition_tree/get_leaves.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/graph/series_parallel/digraph_generation.cc b/lib/utils/src/utils/graph/series_parallel/digraph_generation.cc index 5d19f2032a..93671ad35f 100644 --- a/lib/utils/src/utils/graph/series_parallel/digraph_generation.cc +++ b/lib/utils/src/utils/graph/series_parallel/digraph_generation.cc @@ -14,8 +14,7 @@ namespace FlexFlow { -std::map parallel_extend(DiGraph &g, - DiGraphView const &ext) { +std::map parallel_extend(DiGraph &g, DiGraphView const &ext) { std::map node_map; for (Node const &node : get_nodes(ext)) { node_map.emplace(node, g.add_node()); @@ -26,8 +25,7 @@ std::map parallel_extend(DiGraph &g, return node_map; } -std::map serial_extend(DiGraph &g, - DiGraphView const &ext) { +std::map serial_extend(DiGraph &g, DiGraphView const &ext) { std::set original_sinks = get_terminal_nodes(g); std::map node_map = parallel_extend(g, ext); for (Node const &node1 : original_sinks) { diff --git a/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc b/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc index baffa96c53..4c21102d01 100644 --- a/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc +++ b/lib/utils/src/utils/graph/series_parallel/get_ancestors.cc @@ -2,9 +2,9 @@ #include "utils/containers/contains.h" #include "utils/containers/filter.h" #include "utils/containers/get_only.h" +#include "utils/containers/set_of.h" #include "utils/containers/set_union.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/variant.h" #include @@ -12,14 +12,14 @@ namespace FlexFlow { std::set get_ancestors(SeriesParallelDecomposition const &sp, - Node const &node); + Node const &node); static std::set get_ancestors(Node const &, Node const &node) { return {}; } static std::set get_ancestors(SeriesSplit const &serial, - Node const &node) { + Node const &node) { std::set ancestors{}; for (std::variant const &child : serial.children) { SeriesParallelDecomposition child_sp = @@ -33,7 +33,7 @@ static std::set get_ancestors(SeriesSplit const &serial, } static std::set get_ancestors(ParallelSplit const ¶llel, - Node const &node) { + Node const &node) { SeriesParallelDecomposition branch = get_only(filter(transform(parallel.get_children(), [](std::variant const &c) { @@ -46,7 +46,7 @@ static std::set get_ancestors(ParallelSplit const ¶llel, } std::set get_ancestors(SeriesParallelDecomposition const &sp, - Node const &node) { + Node const &node) { assert(contains(get_nodes(sp), node)); return sp.visit>( [&](auto const &t) { return get_ancestors(t, node); }); diff --git a/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc index 5391fafd56..baa9efddde 100644 --- a/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/get_series_parallel_decomposition.cc @@ -1,8 +1,8 @@ #include "utils/graph/series_parallel/get_series_parallel_decomposition.h" #include "utils/containers/get_only.h" #include "utils/containers/map_values.h" -#include "utils/containers/transform.h" #include "utils/containers/multiset_of.h" +#include "utils/containers/transform.h" #include "utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.h" #include "utils/graph/digraph/algorithms/transitive_reduction.h" #include "utils/graph/instances/adjacency_multidigraph.h" @@ -37,10 +37,9 @@ std::optional MultiDiGraph ttsp = MultiDiGraph::materialize_copy_of( inverse_line_graph_result.graph); - std::map - ttsp_edge_to_sp_tree = map_values( - inverse_line_graph_result.inverse_edge_to_line_node_bidict - .as_map(), + std::map ttsp_edge_to_sp_tree = + map_values( + inverse_line_graph_result.inverse_edge_to_line_node_bidict.as_map(), [](Node const &n) { return SeriesParallelDecomposition{n}; }); auto perform_extended_parallel_reduction = @@ -132,10 +131,9 @@ std::optional MultiDiGraph ttsp = MultiDiGraph::materialize_copy_of( inverse_line_graph_result.graph); - std::map - ttsp_edge_to_sp_tree = map_values( - inverse_line_graph_result.inverse_edge_to_line_node_bidict - .as_map(), + std::map ttsp_edge_to_sp_tree = + map_values( + inverse_line_graph_result.inverse_edge_to_line_node_bidict.as_map(), [](Node const &n) { return BinarySPDecompositionTree{n}; }); while (true) { diff --git a/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc index 624f06d760..e5ad6e876a 100644 --- a/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/non_normal_sp_decomposition.cc @@ -1,6 +1,7 @@ #include "utils/graph/series_parallel/non_normal_sp_decomposition.h" #include "utils/containers/all_of.h" #include "utils/containers/extend.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/multiset_union.h" #include "utils/containers/transform.h" #include "utils/containers/vector_of.h" @@ -11,8 +12,6 @@ #include "utils/graph/series_parallel/series_split.dtg.h" #include "utils/overload.h" #include "utils/variant.h" -#include "utils/containers/multiset_of.h" -#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -45,7 +44,8 @@ NonNormalSPDecomposition non_normal_parallel_composition( for (NonNormalSPDecomposition const &sp_comp : sp_compositions) { if (sp_comp.has()) { composition = multiset_union( - composition, multiset_of(sp_comp.get().get_children())); + composition, + multiset_of(sp_comp.get().get_children())); } else if (sp_comp.has()) { composition.insert(sp_comp.get()); } else { @@ -53,7 +53,8 @@ NonNormalSPDecomposition non_normal_parallel_composition( composition.insert(sp_comp.get()); } } - return NonNormalSPDecomposition(NonNormalParallelSplit{multiset_of(composition)}); + return NonNormalSPDecomposition( + NonNormalParallelSplit{multiset_of(composition)}); } static Node as_non_normal(Node const &n) { @@ -72,11 +73,12 @@ static NonNormalSeriesSplit as_non_normal(SeriesSplit const &s) { static NonNormalParallelSplit as_non_normal(ParallelSplit const &p) { return non_normal_parallel_composition( - multiset_of(transform(p.get_children(), - [](std::variant const &child) { - return as_non_normal( - widen(child)); - }))) + multiset_of( + transform(p.get_children(), + [](std::variant const &child) { + return as_non_normal( + widen(child)); + }))) .get(); } diff --git a/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc index 25a8ca1cfa..e8cfe6c0b8 100644 --- a/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/normalize_sp_decomposition.cc @@ -1,12 +1,12 @@ #include "utils/graph/series_parallel/normalize_sp_decomposition.h" #include "utils/containers/filter.h" #include "utils/containers/get_only.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/transform.h" #include "utils/exception.h" #include "utils/graph/series_parallel/non_normal_sp_decomposition.h" #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/variant.h" -#include "utils/containers/multiset_of.h" namespace FlexFlow { diff --git a/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc b/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc index 69acec155b..0c6b6553d4 100644 --- a/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc +++ b/lib/utils/src/utils/graph/series_parallel/parallel_reduction.cc @@ -3,8 +3,8 @@ #include "utils/containers/contains_key.h" #include "utils/containers/get_one_of.h" #include "utils/containers/group_by.h" -#include "utils/containers/transform.h" #include "utils/containers/set_of.h" +#include "utils/containers/transform.h" #include "utils/containers/values.h" #include "utils/graph/digraph/directed_edge.dtg.h" #include "utils/graph/multidigraph/algorithms/get_directed_edge.h" @@ -40,20 +40,18 @@ std::optional std::set find_all_extended_parallel_reductions(MultiDiGraphView const &g) { - std::map> - reduction_groups; + std::map> reduction_groups; for (MultiDiEdge const &edge : get_edges(g)) { reduction_groups[get_directed_edge(g, edge)].insert(edge); } - std::set> reductions = filter( - set_of(values(reduction_groups)), - [](std::set const &s) { return s.size() > 1; }); + std::set> reductions = + filter(set_of(values(reduction_groups)), + [](std::set const &s) { return s.size() > 1; }); - return transform(reductions, - [&](std::set const &edges) { - return ExtendedParallelReduction{edges}; - }); + return transform(reductions, [&](std::set const &edges) { + return ExtendedParallelReduction{edges}; + }); } MultiDiEdge apply_parallel_reduction(MultiDiGraph &g, diff --git a/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc b/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc index 918c13d57d..d041b9656e 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_parallel_decomposition.cc @@ -2,11 +2,11 @@ #include "utils/containers/all_of.h" #include "utils/containers/extend.h" #include "utils/containers/get_only.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/multiset_union.h" #include "utils/containers/set_union.h" #include "utils/containers/sum.h" #include "utils/containers/transform.h" -#include "utils/containers/multiset_of.h" #include "utils/containers/values.h" #include "utils/containers/vector_of.h" #include "utils/exception.h" @@ -16,7 +16,6 @@ #include "utils/nonnegative_int/nonnegative_int.h" #include "utils/variant.h" #include -#include "utils/containers/multiset_of.h" namespace FlexFlow { @@ -59,8 +58,7 @@ SeriesParallelDecomposition to_final_ast( } std::multiset get_nodes(SeriesParallelDecomposition const &sp) { - return sp.visit>( - [](auto &&t) { return get_nodes(t); }); + return sp.visit>([](auto &&t) { return get_nodes(t); }); } std::multiset get_nodes(SeriesSplit const &serial) { @@ -122,8 +120,7 @@ SeriesParallelDecomposition series_composition( } SeriesParallelDecomposition parallel_composition( - std::multiset const - &sp_compositions) { + std::multiset const &sp_compositions) { ASSERT(sp_compositions.size() > 0, "Cannot create parallel composition with zero elements"); @@ -132,13 +129,13 @@ SeriesParallelDecomposition parallel_composition( return get_only(sp_compositions); } - std::multiset< - std::variant<::FlexFlow::SeriesSplit, ::FlexFlow::Node>> + std::multiset> composition{}; for (SeriesParallelDecomposition const &sp_comp : sp_compositions) { if (sp_comp.has()) { - composition = multiset_union(composition, - multiset_of(sp_comp.get().get_children())); + composition = multiset_union( + composition, + multiset_of(sp_comp.get().get_children())); } else if (sp_comp.has()) { composition.insert(sp_comp.get()); } else { diff --git a/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc b/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc index 12f5eb2582..b4167dff8f 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_parallel_metrics.cc @@ -5,7 +5,6 @@ #include "utils/containers/values.h" #include "utils/containers/vector_of.h" #include "utils/fmt/multiset.h" -#include "utils/fmt/multiset.h" #include "utils/graph/digraph/algorithms/get_edges.h" #include "utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h" #include "utils/graph/digraph/digraph_view.h" @@ -55,21 +54,18 @@ float work_cost(SeriesParallelDecomposition const &sp, [&](Node const &node) { return cost_map.at(node); })); } -float work_cost(DiGraphView const &g, - std::map const &cost_map) { +float work_cost(DiGraphView const &g, std::map const &cost_map) { return sum(transform(vector_of(get_nodes(g)), [&](Node const &node) { return cost_map.at(node); })); } -static float - critical_path_cost(Node const &node, - std::map const &cost_map) { +static float critical_path_cost(Node const &node, + std::map const &cost_map) { return cost_map.at(node); } -static float - critical_path_cost(SeriesSplit const &serial, - std::map const &cost_map) { +static float critical_path_cost(SeriesSplit const &serial, + std::map const &cost_map) { return sum(transform( serial.children, [&](std::variant const &child) { return critical_path_cost(widen(child), @@ -77,9 +73,8 @@ static float })); } -static float - critical_path_cost(ParallelSplit const ¶llel, - std::map const &cost_map) { +static float critical_path_cost(ParallelSplit const ¶llel, + std::map const &cost_map) { return maximum(transform(parallel.get_children(), [&](std::variant const &child) { return critical_path_cost( diff --git a/lib/utils/src/utils/graph/series_parallel/series_reduction.cc b/lib/utils/src/utils/graph/series_parallel/series_reduction.cc index c8aa6db3c8..6e34306f26 100644 --- a/lib/utils/src/utils/graph/series_parallel/series_reduction.cc +++ b/lib/utils/src/utils/graph/series_parallel/series_reduction.cc @@ -3,8 +3,8 @@ #include "utils/containers/contains_key.h" #include "utils/containers/get_only.h" #include "utils/containers/require_same.h" -#include "utils/containers/slice.h" #include "utils/containers/set_of.h" +#include "utils/containers/slice.h" #include "utils/containers/values.h" #include "utils/graph/digraph/algorithms/get_predecessors.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc index 1721d16083..7702f3dbd6 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc @@ -38,9 +38,9 @@ namespace FlexFlow { -static std::set filter_out_sync_nodes( - std::set const &nodes, - std::map const &node_roles) { +static std::set + filter_out_sync_nodes(std::set const &nodes, + std::map const &node_roles) { return filter( nodes, [&](Node const &n) { return node_roles.at(n) != NodeRole::SYNC; }); } @@ -52,8 +52,7 @@ static nonnegative_int depth_map, [&](Node const &n) { return contains(get_nodes(sp), n); }))); } -DiGraph add_dummy_nodes(DiGraph g, - std::map &node_roles) { +DiGraph add_dummy_nodes(DiGraph g, std::map &node_roles) { std::map depth_map = get_longest_path_lengths_from_root(g); @@ -87,11 +86,10 @@ DiGraph add_dummy_nodes(DiGraph g, return g; } -std::set - get_component(DiGraph const &g, - Node const &node, - std::map const &depth_map, - std::map const &node_roles) { +std::set get_component(DiGraph const &g, + Node const &node, + std::map const &depth_map, + std::map const &node_roles) { nonnegative_int max_depth = get_max_depth(g, depth_map); auto is_in_last_2_strata = [&](Node const &n) { @@ -140,18 +138,16 @@ static std::set return set_intersection(subtree, component).size() > 0; }); - std::set forest = - set_union(subtrees_overlapping_with_component); + std::set forest = set_union(subtrees_overlapping_with_component); forest.insert(handle); return filter_out_sync_nodes(forest, node_roles); } static std::pair, nonempty_set> - get_up_and_down_sets( - DiGraph const &g, - std::set const &forest, - std::map const &depth_map) { + get_up_and_down_sets(DiGraph const &g, + std::set const &forest, + std::map const &depth_map) { nonnegative_int max_depth = get_max_depth(g, depth_map); @@ -163,10 +159,9 @@ static std::pair, nonempty_set> grouped_by_depth.at_l(max_depth)); } -static std::set - edges_to_remove(DiGraph const &g, - std::set const &up, - std::set const &down) { +static std::set edges_to_remove(DiGraph const &g, + std::set const &up, + std::set const &down) { std::set to_remove; for (Node const &u : up) { @@ -179,10 +174,9 @@ static std::set return to_remove; } -static std::set - edges_to_add_escribano(std::set const &up, - std::set const &down, - Node const &sync_node) { +static std::set edges_to_add_escribano(std::set const &up, + std::set const &down, + Node const &sync_node) { return set_union(transform(up, [&](Node const &u) { return DirectedEdge{u, sync_node}; @@ -192,8 +186,7 @@ static std::set })); } -static Node add_sync_node(DiGraph &sp, - std::map &node_roles) { +static Node add_sync_node(DiGraph &sp, std::map &node_roles) { Node sync_node = sp.add_node(); node_roles[sync_node] = NodeRole::SYNC; return sync_node; @@ -223,18 +216,16 @@ SeriesParallelDecomposition escribano_sp_ization(DiGraph g) { sp.add_node_unsafe(node); add_edges(sp, get_incoming_edges(g, node)); - std::set component = - get_component(sp, node, depth_map, node_roles); + std::set component = get_component(sp, node, depth_map, node_roles); Node handle = get_only(get_lowest_common_ancestors(sp, component).value()); std::set forest = get_forest_escribano(sp, handle, component, node_roles); - std::pair, nonempty_set> - up_down_sets = get_up_and_down_sets(sp, forest, depth_map); + std::pair, nonempty_set> up_down_sets = + get_up_and_down_sets(sp, forest, depth_map); std::set up = up_down_sets.first.unwrap_as_set(); - std::set down = - up_down_sets.second.unwrap_as_set(); + std::set down = up_down_sets.second.unwrap_as_set(); remove_edges(sp, edges_to_remove(sp, up, down)); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc index 9da2d523a6..7ee1d0e81d 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/flexible_algo.cc @@ -43,8 +43,8 @@ namespace FlexFlow { -static std::set - get_component(DiGraph const &sp, std::set const &nodes) { +static std::set get_component(DiGraph const &sp, + std::set const &nodes) { std::set parents = set_union( transform(nodes, [&](Node const &n) { return get_predecessors(sp, n); })); std::set children = set_union(transform( @@ -141,8 +141,7 @@ static UpDownPartition partitions.insert( UpDownPartition{base_up, set_union(base_down, assignable_nodes)}); - std::set valid_partitions = - filter(partitions, is_valid); + std::set valid_partitions = filter(partitions, is_valid); ASSERT(!valid_partitions.empty()); auto partition_cost = [&](UpDownPartition const &p) { @@ -155,11 +154,11 @@ static UpDownPartition return argmin(valid_partitions, partition_cost); } -static std::set edges_to_remove_flexible( - DiGraph const &sp, - std::set const &up, - std::set const &down, - std::map const &node_roles) { +static std::set + edges_to_remove_flexible(DiGraph const &sp, + std::set const &up, + std::set const &down, + std::map const &node_roles) { std::set to_remove; // from up to down @@ -210,10 +209,9 @@ static Node add_sync_node(DiGraph &sp, return sync_node; } -static std::set - get_next_nodes(DiGraph const &sp, - DiGraph const &g, - std::map const &cost_map) { +static std::set get_next_nodes(DiGraph const &sp, + DiGraph const &g, + std::map const &cost_map) { std::map sp_longest_paths = get_weighted_longest_path_lengths_from_root(sp, cost_map); @@ -221,14 +219,13 @@ static std::set std::set g_nodes = get_nodes(g); // candidate nodes: not in sp but all predecessors in sp - std::set candidate_nodes = - filter(g_nodes, [&](Node const &node) { - if (contains(sp_nodes, node)) { - return false; - } - std::set preds = get_predecessors(g, node); - return is_subseteq_of(preds, sp_nodes); - }); + std::set candidate_nodes = filter(g_nodes, [&](Node const &node) { + if (contains(sp_nodes, node)) { + return false; + } + std::set preds = get_predecessors(g, node); + return is_subseteq_of(preds, sp_nodes); + }); ASSERT(!candidate_nodes.empty()); @@ -265,8 +262,7 @@ SeriesParallelDecomposition DiGraph g_reduced = materialize_digraph_view(transitive_reduction(g)); - std::map node_roles = - get_initial_node_role_map(g_reduced); + std::map node_roles = get_initial_node_role_map(g_reduced); DiGraph sp = DiGraph::create(); Node root = get_only(get_initial_nodes(g_reduced)); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc index cdc209f9d1..cc15f64ca2 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc @@ -1,9 +1,10 @@ #include "utils/graph/series_parallel/sp_ization/naive_stratum_sync.h" #include "utils/containers/group_by.h" +#include "utils/containers/keys.h" #include "utils/containers/maximum.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/range.h" #include "utils/containers/transform.h" -#include "utils/containers/multiset_of.h" #include "utils/fmt/multiset.h" #include "utils/graph/digraph/algorithms/get_longest_path_lengths_from_root.h" #include "utils/graph/digraph/algorithms/is_acyclic.h" @@ -12,7 +13,6 @@ #include "utils/graph/series_parallel/series_parallel_decomposition.h" #include "utils/graph/series_parallel/sp_ization/dependencies_are_maintained.h" #include -#include "utils/containers/keys.h" namespace FlexFlow { @@ -27,24 +27,22 @@ std::vector> nonnegative_int num_strata = maximum(strata_to_nodes.left_values()); - return transform(range(1, num_strata.unwrap_nonnegative() + 1), - [&](int depth) { - return multiset_of( - strata_to_nodes.at_l(nonnegative_int{depth})); - }); + return transform( + range(1, num_strata.unwrap_nonnegative() + 1), [&](int depth) { + return multiset_of(strata_to_nodes.at_l(nonnegative_int{depth})); + }); } -static SeriesParallelDecomposition naive_stratum_merge( - std::vector> stratum_split) { +static SeriesParallelDecomposition + naive_stratum_merge(std::vector> stratum_split) { - auto merge_one_stratum = - [&](std::multiset const &stratum_nodes) { - auto as_singleton_sp = [](Node const &node) { - return NonNormalSPDecomposition{node}; - }; - return non_normal_parallel_composition( - transform(stratum_nodes, as_singleton_sp)); - }; + auto merge_one_stratum = [&](std::multiset const &stratum_nodes) { + auto as_singleton_sp = [](Node const &node) { + return NonNormalSPDecomposition{node}; + }; + return non_normal_parallel_composition( + transform(stratum_nodes, as_singleton_sp)); + }; std::vector parallel_strata = transform(stratum_split, merge_one_stratum); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc index d70abfac58..9dd1142adc 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/node_role.cc @@ -8,8 +8,7 @@ namespace FlexFlow { -std::map - get_initial_node_role_map(DiGraphView const &g) { +std::map get_initial_node_role_map(DiGraphView const &g) { return generate_map(get_nodes(g), [](Node const &) { return NodeRole::PURE; }); } diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc index d4c62007dd..63230a72c7 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/up_down_partition.cc @@ -7,7 +7,7 @@ namespace FlexFlow { std::set get_up_frontier(DiGraph const &sp, - UpDownPartition const &partition) { + UpDownPartition const &partition) { DiGraphView up_subgraph = get_subgraph(sp, partition.up); return filter(partition.up, [&](Node const &node) { return get_outgoing_edges(up_subgraph, node).empty(); @@ -15,7 +15,7 @@ std::set get_up_frontier(DiGraph const &sp, } std::set get_down_frontier(DiGraph const &sp, - UpDownPartition const &partition) { + UpDownPartition const &partition) { DiGraphView down_subgraph = get_subgraph(sp, partition.down); return filter(partition.down, [&](Node const &node) { return get_incoming_edges(down_subgraph, node).empty(); diff --git a/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc b/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc index db28ee9970..68036d3f15 100644 --- a/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc +++ b/lib/utils/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc @@ -2,9 +2,9 @@ #include "utils/containers/filter.h" #include "utils/containers/get_only.h" #include "utils/containers/group_by.h" +#include "utils/containers/multiset_of.h" #include "utils/containers/slice.h" #include "utils/containers/transform.h" -#include "utils/containers/multiset_of.h" #include "utils/fmt/variant.h" #include "utils/graph/digraph/algorithms/get_initial_nodes.h" #include "utils/graph/digraph/algorithms/get_predecessors.h" @@ -113,9 +113,8 @@ static SeriesParallelDecomposition for (Node const &node : get_topological_ordering(g)) { std::multiset predecessors_as_sp = - multiset_of( - transform(get_predecessors(g, node), - [&](Node const &p) { return node_to_sp.at(p); })); + multiset_of(transform(get_predecessors(g, node), + [&](Node const &p) { return node_to_sp.at(p); })); NonNormalSPDecomposition parallel_comp = non_normal_parallel_composition(predecessors_as_sp); diff --git a/lib/utils/src/utils/graph/traversal.cc b/lib/utils/src/utils/graph/traversal.cc index beffcf90d2..a27e8db2aa 100644 --- a/lib/utils/src/utils/graph/traversal.cc +++ b/lib/utils/src/utils/graph/traversal.cc @@ -129,8 +129,7 @@ bfs_iterator &bfs_iterator::operator++() { this->seen.value().insert(current); this->q.pop(); - std::set outgoing = - get_outgoing_edges(graph, {current}); + std::set outgoing = get_outgoing_edges(graph, {current}); for (DirectedEdge const &e : outgoing) { if (!contains(this->seen.value(), e.dst)) { this->q.push(e.dst); @@ -189,8 +188,8 @@ CheckedDFSView dfs(DiGraphView const &g, return CheckedDFSView(g, starting_points); } -UncheckedDFSView::UncheckedDFSView( - DiGraphView const &g, std::set const &starting_points) +UncheckedDFSView::UncheckedDFSView(DiGraphView const &g, + std::set const &starting_points) : graph(g), starting_points(starting_points) {} unchecked_dfs_iterator UncheckedDFSView::cbegin() const { @@ -209,14 +208,12 @@ unchecked_dfs_iterator UncheckedDFSView::end() const { return this->cend(); } -UncheckedDFSView - unchecked_dfs(DiGraphView const &g, - std::set const &starting_points) { +UncheckedDFSView unchecked_dfs(DiGraphView const &g, + std::set const &starting_points) { return UncheckedDFSView(g, starting_points); } -BFSView::BFSView(DiGraphView const &g, - std::set const &starting_points) +BFSView::BFSView(DiGraphView const &g, std::set const &starting_points) : graph(g), starting_points(starting_points) {} bfs_iterator BFSView::cbegin() const { @@ -235,8 +232,7 @@ bfs_iterator BFSView::end() const { return this->cend(); } -BFSView bfs(DiGraphView const &g, - std::set const &starting_points) { +BFSView bfs(DiGraphView const &g, std::set const &starting_points) { return BFSView(g, starting_points); } diff --git a/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc b/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc index 562fc1f906..69c1949afc 100644 --- a/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc +++ b/lib/utils/src/utils/graph/undirected/algorithms/get_connected_components.cc @@ -11,8 +11,7 @@ std::set> std::set visited; for (Node const &node : get_nodes(g)) { - std::set component = - set_of(get_bfs_ordering(as_digraph(g), {node})); + std::set component = set_of(get_bfs_ordering(as_digraph(g), {node})); components.insert(component); visited = set_union(visited, component); } diff --git a/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc b/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc index 8bf76965a3..3acf7af7c9 100644 --- a/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc +++ b/lib/utils/src/utils/graph/undirected/algorithms/get_neighboring_nodes.cc @@ -4,7 +4,7 @@ namespace FlexFlow { std::set get_neighboring_nodes(UndirectedGraphView const &g, - Node const &n) { + Node const &n) { std::set edges = g.query_edges( UndirectedEdgeQuery{query_set::match_single_value(n)}); diff --git a/lib/utils/src/utils/graph/undirected/undirected_graph.cc b/lib/utils/src/utils/graph/undirected/undirected_graph.cc index 12662fd6f0..6a409f8012 100644 --- a/lib/utils/src/utils/graph/undirected/undirected_graph.cc +++ b/lib/utils/src/utils/graph/undirected/undirected_graph.cc @@ -37,8 +37,7 @@ std::set return this->get_ptr().query_edges(q); } -std::set - UndirectedGraph::query_nodes(NodeQuery const &q) const { +std::set UndirectedGraph::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } diff --git a/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc b/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc index 2daa8dcbfd..14ba4d3ba8 100644 --- a/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc +++ b/lib/utils/src/utils/graph/undirected/undirected_graph_view.cc @@ -7,8 +7,7 @@ std::set return this->get_ptr().query_edges(q); } -std::set - UndirectedGraphView::query_nodes(NodeQuery const &q) const { +std::set UndirectedGraphView::query_nodes(NodeQuery const &q) const { return this->get_ptr().query_nodes(q); } diff --git a/lib/utils/src/utils/graph/views/views.cc b/lib/utils/src/utils/graph/views/views.cc index 5977d800af..2f418bae21 100644 --- a/lib/utils/src/utils/graph/views/views.cc +++ b/lib/utils/src/utils/graph/views/views.cc @@ -11,8 +11,7 @@ namespace FlexFlow { UndirectedSubgraphView::UndirectedSubgraphView( - UndirectedGraphView const &g, - std::set const &subgraph_nodes) + UndirectedGraphView const &g, std::set const &subgraph_nodes) : g(g), subgraph_nodes(subgraph_nodes) {} UndirectedSubgraphView *UndirectedSubgraphView::clone() const { @@ -49,8 +48,7 @@ std::set return this->g.query_edges(query_intersection(query, subgraph_query)); } -std::set - DiSubgraphView::query_nodes(NodeQuery const &query) const { +std::set DiSubgraphView::query_nodes(NodeQuery const &query) const { NodeQuery subgraph_query = NodeQuery{ query_set::match_values_in(set_of(this->subgraph_nodes)), }; @@ -62,9 +60,8 @@ DiSubgraphView *DiSubgraphView::clone() const { return new DiSubgraphView(g, subgraph_nodes); } -UndirectedGraphView - view_subgraph(UndirectedGraphView const &g, - std::set const &subgraph_nodes) { +UndirectedGraphView view_subgraph(UndirectedGraphView const &g, + std::set const &subgraph_nodes) { return UndirectedGraphView::create(g, subgraph_nodes); } @@ -77,8 +74,8 @@ UndirectedEdge to_undirected_edge(DirectedEdge const &e) { return make_undirected_edge(e.src, e.dst); } -std::set to_undirected_edges( - std::set const &directed_edges) { +std::set + to_undirected_edges(std::set const &directed_edges) { return transform(directed_edges, [](DirectedEdge const &e) { return to_undirected_edge(e); }); } @@ -89,8 +86,8 @@ std::set to_directed_edges(UndirectedEdge const &e) { DirectedEdge{e.endpoints.max(), e.endpoints.min()}}; } -std::set to_directed_edges( - std::set const &undirected_edges) { +std::set + to_directed_edges(std::set const &undirected_edges) { return flatmap(undirected_edges, [](UndirectedEdge const &e) { return to_directed_edges(e); }); } diff --git a/lib/utils/src/utils/many_to_one/many_to_one.cc b/lib/utils/src/utils/many_to_one/many_to_one.cc index 933d32a176..c5d0fcf2cf 100644 --- a/lib/utils/src/utils/many_to_one/many_to_one.cc +++ b/lib/utils/src/utils/many_to_one/many_to_one.cc @@ -1,7 +1,7 @@ #include "utils/many_to_one/many_to_one.h" #include "utils/archetypes/jsonable_ordered_value_type.h" -#include "utils/archetypes/rapidcheckable_value_type.h" #include "utils/archetypes/ordered_value_type.h" +#include "utils/archetypes/rapidcheckable_value_type.h" using namespace ::FlexFlow; @@ -12,16 +12,15 @@ using R = ordered_value_type<1>; template struct ManyToOne; -template std::map, R> - format_as(ManyToOne const &); +template std::map, R> format_as(ManyToOne const &); template std::ostream &operator<<(std::ostream &, ManyToOne const &); template std::set> unstructured_relation_from_many_to_one(ManyToOne const &); -template ManyToOne many_to_one_from_unstructured_relation( - std::set> const &); +template ManyToOne + many_to_one_from_unstructured_relation(std::set> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/nonempty_set/nonempty_set.cc b/lib/utils/src/utils/nonempty_set/nonempty_set.cc index a7bd7f7c5a..e239ee81b5 100644 --- a/lib/utils/src/utils/nonempty_set/nonempty_set.cc +++ b/lib/utils/src/utils/nonempty_set/nonempty_set.cc @@ -1,6 +1,6 @@ #include "utils/nonempty_set/nonempty_set.h" -#include "utils/archetypes/ordered_value_type.h" #include "utils/archetypes/jsonable_ordered_value_type.h" +#include "utils/archetypes/ordered_value_type.h" using T = ::FlexFlow::ordered_value_type<0>; using J = ::FlexFlow::jsonable_ordered_value_type<0>; diff --git a/lib/utils/src/utils/one_to_many/one_to_many.cc b/lib/utils/src/utils/one_to_many/one_to_many.cc index 2094a65ece..108b3fd4cc 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many.cc @@ -1,8 +1,8 @@ #include "utils/one_to_many/one_to_many.h" +#include "utils/archetypes/jsonable_ordered_value_type.h" #include "utils/archetypes/jsonable_value_type.h" -#include "utils/archetypes/rapidcheckable_value_type.h" #include "utils/archetypes/ordered_value_type.h" -#include "utils/archetypes/jsonable_ordered_value_type.h" +#include "utils/archetypes/rapidcheckable_value_type.h" using namespace ::FlexFlow; @@ -13,8 +13,7 @@ using R = ordered_value_type<1>; template struct OneToMany; -template std::map> - format_as(OneToMany const &); +template std::map> format_as(OneToMany const &); template std::ostream &operator<<(std::ostream &, OneToMany const &); diff --git a/lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc b/lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc index fa11ccf56f..1787a879dc 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many_filter_values.cc @@ -8,6 +8,7 @@ using R = ordered_value_type<1>; using F = std::function; -template OneToMany one_to_many_filter_values(OneToMany const &, F &&); +template OneToMany one_to_many_filter_values(OneToMany const &, + F &&); } // namespace FlexFlow diff --git a/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc b/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc index eddc23fe5f..f2521e9f2b 100644 --- a/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc +++ b/lib/utils/src/utils/one_to_many/one_to_many_from_l_to_r_mapping.cc @@ -6,7 +6,7 @@ namespace FlexFlow { using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template OneToMany one_to_many_from_l_to_r_mapping( - std::map> const &); +template OneToMany + one_to_many_from_l_to_r_mapping(std::map> const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc b/lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc index 60a78bb123..973b7a30a8 100644 --- a/lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc +++ b/lib/utils/src/utils/one_to_many/require_one_to_many_is_bijection.cc @@ -6,7 +6,6 @@ namespace FlexFlow { using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template - bidict require_one_to_many_is_bijection(OneToMany const &); +template bidict require_one_to_many_is_bijection(OneToMany const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/orthotope/dim_coord.cc b/lib/utils/src/utils/orthotope/dim_coord.cc index dbe5273929..e8e7fc84ee 100644 --- a/lib/utils/src/utils/orthotope/dim_coord.cc +++ b/lib/utils/src/utils/orthotope/dim_coord.cc @@ -13,16 +13,13 @@ template DimCoord restrict_coord_to_dims(DimCoord const &, template OrthotopeCoord orthotope_coord_from_dim_coord(DimCoord const &, DimOrdering const &); -template DimCoord - dim_coord_from_orthotope_coord(OrthotopeCoord const &, - std::set const &, - DimOrdering const &); +template DimCoord dim_coord_from_orthotope_coord(OrthotopeCoord const &, + std::set const &, + DimOrdering const &); -template DimCoord lift_dim_coord(DimCoord const &, - std::set const &); +template DimCoord lift_dim_coord(DimCoord const &, std::set const &); -template std::set> - get_coords_in_dim_domain(DimDomain const &); +template std::set> get_coords_in_dim_domain(DimDomain const &); template std::set> get_coords_in_minimal_dim_domain(MinimalDimDomain const &); diff --git a/lib/utils/src/utils/orthotope/dim_projection.cc b/lib/utils/src/utils/orthotope/dim_projection.cc index 64f172ab10..0f67e48586 100644 --- a/lib/utils/src/utils/orthotope/dim_projection.cc +++ b/lib/utils/src/utils/orthotope/dim_projection.cc @@ -1,6 +1,6 @@ #include "utils/orthotope/dim_projection.h" -#include "utils/archetypes/value_type.h" #include "utils/archetypes/jsonable_ordered_value_type.h" +#include "utils/archetypes/value_type.h" namespace FlexFlow { @@ -13,11 +13,9 @@ template DimProjection DimOrdering const &, DimOrdering const &); -template std::set - input_dims_of_projection(DimProjection const &); +template std::set input_dims_of_projection(DimProjection const &); -template std::set - output_dims_of_projection(DimProjection const &); +template std::set output_dims_of_projection(DimProjection const &); template DimProjection invert_dim_projection(DimProjection const &); diff --git a/lib/utils/src/utils/orthotope/down_projection.cc b/lib/utils/src/utils/orthotope/down_projection.cc index d92969165f..9cf9885bd3 100644 --- a/lib/utils/src/utils/orthotope/down_projection.cc +++ b/lib/utils/src/utils/orthotope/down_projection.cc @@ -19,9 +19,8 @@ template DimCoord compute_down_projection(DownProjection const &, DimDomain const &, DimOrdering const &); -template void project_dims(DownProjection &, - std::set const &, - R const &); +template void + project_dims(DownProjection &, std::set const &, R const &); template UpProjection invert_down_projection(DownProjection const &); diff --git a/lib/utils/src/utils/orthotope/eq_projection.cc b/lib/utils/src/utils/orthotope/eq_projection.cc index 7deffc7711..a0553e9a89 100644 --- a/lib/utils/src/utils/orthotope/eq_projection.cc +++ b/lib/utils/src/utils/orthotope/eq_projection.cc @@ -8,11 +8,9 @@ using R = ordered_value_type<1>; template EqProjection make_empty_eq_projection(); -template std::set - input_dims_of_eq_projection(EqProjection const &); +template std::set input_dims_of_eq_projection(EqProjection const &); -template std::set - output_dims_of_eq_projection(EqProjection const &); +template std::set output_dims_of_eq_projection(EqProjection const &); template void project_dims(EqProjection &, L const &, R const &); diff --git a/lib/utils/src/utils/orthotope/minimal_dim_domain.cc b/lib/utils/src/utils/orthotope/minimal_dim_domain.cc index 5bff87fc27..e1aa9e4369 100644 --- a/lib/utils/src/utils/orthotope/minimal_dim_domain.cc +++ b/lib/utils/src/utils/orthotope/minimal_dim_domain.cc @@ -22,8 +22,7 @@ template DimDomain dim_domain_from_minimal_dim_domain(MinimalDimDomain const &, std::set const &); -template std::set - get_minimal_domain_dims(MinimalDimDomain const &); +template std::set get_minimal_domain_dims(MinimalDimDomain const &); template MinimalDimDomain restrict_minimal_domain_to_dims(MinimalDimDomain const &, @@ -33,9 +32,7 @@ template MinimalOrthotope minimal_orthotope_from_minimal_dim_domain(MinimalDimDomain const &, DimOrdering const &); -template MinimalDimDomain - minimal_dim_domain_from_minimal_orthotope(MinimalOrthotope const &, - std::set const &, - DimOrdering const &); +template MinimalDimDomain minimal_dim_domain_from_minimal_orthotope( + MinimalOrthotope const &, std::set const &, DimOrdering const &); } // namespace FlexFlow diff --git a/lib/utils/src/utils/orthotope/orthotope.cc b/lib/utils/src/utils/orthotope/orthotope.cc index 0a53a113d6..1e4e6e79eb 100644 --- a/lib/utils/src/utils/orthotope/orthotope.cc +++ b/lib/utils/src/utils/orthotope/orthotope.cc @@ -6,10 +6,10 @@ #include "utils/containers/filter_idxs.h" #include "utils/containers/product.h" #include "utils/containers/scanr.h" +#include "utils/containers/set_of.h" #include "utils/containers/slice.h" #include "utils/containers/sum.h" #include "utils/containers/transform.h" -#include "utils/containers/set_of.h" #include "utils/containers/zip3_with_strict.h" #include "utils/containers/zip_strict.h" #include "utils/containers/zip_with_strict.h" diff --git a/lib/utils/src/utils/orthotope/up_projection.cc b/lib/utils/src/utils/orthotope/up_projection.cc index 587ccdeb19..d4d25267b1 100644 --- a/lib/utils/src/utils/orthotope/up_projection.cc +++ b/lib/utils/src/utils/orthotope/up_projection.cc @@ -14,11 +14,9 @@ template UpProjection using L = ordered_value_type<0>; using R = ordered_value_type<1>; -template std::set - input_dims_of_up_projection(UpProjection const &); +template std::set input_dims_of_up_projection(UpProjection const &); -template std::set - output_dims_of_up_projection(UpProjection const &); +template std::set output_dims_of_up_projection(UpProjection const &); template DimCoord compute_up_projection(UpProjection const &, DimCoord const &, @@ -27,9 +25,8 @@ template DimCoord compute_up_projection(UpProjection const &, template UpProjection make_empty_up_projection(); -template void project_dims(UpProjection &, - L const &, - std::set const &); +template void + project_dims(UpProjection &, L const &, std::set const &); template DownProjection invert_up_projection(UpProjection const &); diff --git a/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h b/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h index 659c23814c..badd536060 100644 --- a/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h +++ b/lib/utils/test/common/include/test/utils/doctest/check_without_stringify.h @@ -1,10 +1,10 @@ #include "utils/fmt/expected.h" #include #include -#include -#include #include #include +#include +#include #include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_values.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_values.cc index 5c67e11c4f..f5f009e91d 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_values.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_filter_values.cc @@ -10,8 +10,8 @@ TEST_SUITE(FF_TEST_SUITE) { {2, "two"}, }; - bidict result = - bidict_filter_values(dict, [](std::string const &v) { return v == "two"; }); + bidict result = bidict_filter_values( + dict, [](std::string const &v) { return v == "two"; }); bidict correct = { {2, "two"}, }; diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_values.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_values.cc index 8687d539bc..a4c7cfbaf7 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_values.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_filtrans_values.cc @@ -10,8 +10,8 @@ TEST_SUITE(FF_TEST_SUITE) { {2, "two"}, }; - bidict result = - bidict_filtrans_values(dict, [](std::string const &v) -> std::optional { + bidict result = bidict_filtrans_values( + dict, [](std::string const &v) -> std::optional { if (v == "two") { return std::nullopt; } else { diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc index 4d7f0ad495..175aacb5dc 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_from_enumerating.cc @@ -36,13 +36,11 @@ TEST_SUITE(FF_TEST_SUITE) { bidict result = bidict_from_enumerating(input); - std::set result_left_entries = - left_entries(result); + std::set result_left_entries = left_entries(result); std::set correct_left_entries = {0_n, 1_n, 2_n}; CHECK(result_left_entries == correct_left_entries); - std::set result_right_entries = - right_entries(result); + std::set result_right_entries = right_entries(result); std::set correct_right_entries = input; CHECK(result_right_entries == correct_right_entries); } diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_keys.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_keys.cc index 994db63aa6..ab8b0ceb74 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_keys.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_keys.cc @@ -10,11 +10,12 @@ TEST_SUITE(FF_TEST_SUITE) { {2, "two"}, }; - bidict result = bidict_transform_keys(dict, [](int k) { - std::ostringstream oss; - oss << k; - return oss.str(); - }); + bidict result = + bidict_transform_keys(dict, [](int k) { + std::ostringstream oss; + oss << k; + return oss.str(); + }); bidict correct = { {"1", "one"}, {"2", "two"}, diff --git a/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_values.cc b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_values.cc index 606823df87..491d916dde 100644 --- a/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_values.cc +++ b/lib/utils/test/src/utils/bidict/algorithms/bidict_transform_values.cc @@ -10,8 +10,8 @@ TEST_SUITE(FF_TEST_SUITE) { {2, "two"}, }; - bidict result = - bidict_transform_values(dict, [](std::string const &v) { return v + "a"; }); + bidict result = bidict_transform_values( + dict, [](std::string const &v) { return v + "a"; }); bidict correct = { {1, "onea"}, {2, "twoa"}, diff --git a/lib/utils/test/src/utils/binary_relation/binary_relation.cc b/lib/utils/test/src/utils/binary_relation/binary_relation.cc index 26634dad31..a06391df01 100644 --- a/lib/utils/test/src/utils/binary_relation/binary_relation.cc +++ b/lib/utils/test/src/utils/binary_relation/binary_relation.cc @@ -1,8 +1,8 @@ -#include #include "utils/binary_relation/binary_relation.h" -#include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/pair.h" #include "test/utils/doctest/fmt/multiset.h" +#include "test/utils/doctest/fmt/pair.h" +#include "test/utils/doctest/fmt/set.h" +#include using namespace ::FlexFlow; @@ -22,22 +22,22 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("initializer_list constuctor") { BinaryRelation b = BinaryRelation{ - { - 2, - "even", - }, - { - 2, - "EVEN", - }, - { - 3, - "odd", - }, - { - 1, - "odd", - }, + { + 2, + "even", + }, + { + 2, + "EVEN", + }, + { + 3, + "odd", + }, + { + 1, + "odd", + }, }; CHECK(b.size() == 4); @@ -45,10 +45,10 @@ TEST_SUITE(FF_TEST_SUITE) { std::set> raw = b.unwrap_as_set(); std::set> correct_raw = { - {2, "even"}, - {2, "EVEN"}, - {3, "odd"}, - {1, "odd"}, + {2, "even"}, + {2, "EVEN"}, + {3, "odd"}, + {1, "odd"}, }; CHECK(raw == correct_raw); @@ -57,22 +57,22 @@ TEST_SUITE(FF_TEST_SUITE) { BinaryRelation empty_rel; BinaryRelation b = BinaryRelation{ - { - 2, - "even", - }, - { - 2, - "EVEN", - }, - { - 3, - "odd", - }, - { - 1, - "odd", - }, + { + 2, + "even", + }, + { + 2, + "EVEN", + }, + { + 3, + "odd", + }, + { + 1, + "odd", + }, }; SUBCASE("left_values") { @@ -99,10 +99,10 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("right_value_occurences") { std::multiset result = b.right_value_occurences(); std::multiset correct = { - "odd", - "odd", - "even", - "EVEN", + "odd", + "odd", + "even", + "EVEN", }; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc b/lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc index 5a470b60d8..5028390154 100644 --- a/lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc +++ b/lib/utils/test/src/utils/binary_relation/filter_binary_relation.cc @@ -1,45 +1,43 @@ -#include #include "utils/binary_relation/filter_binary_relation.h" +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("filter_binary_relation") { BinaryRelation rel = BinaryRelation{ - { - 2, - "even", - }, - { - 2, - "EVEN", - }, - { - 3, - "odd", - }, - { - 1, - "odd", - }, + { + 2, + "even", + }, + { + 2, + "EVEN", + }, + { + 3, + "odd", + }, + { + 1, + "odd", + }, }; - BinaryRelation result - = filter_binary_relation( - rel, - [](int l, std::string const &r) -> bool { - return l > 1 && r != "EVEN"; - }); + BinaryRelation result = + filter_binary_relation(rel, [](int l, std::string const &r) -> bool { + return l > 1 && r != "EVEN"; + }); BinaryRelation correct = BinaryRelation{ - { - 2, - "even", - }, - { - 3, - "odd", - }, + { + 2, + "even", + }, + { + 3, + "odd", + }, }; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc index e90f1eb633..ba1fcd75c5 100644 --- a/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc +++ b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_left_unique.cc @@ -1,5 +1,5 @@ -#include #include "utils/binary_relation/require_binary_relation_is_left_unique.h" +#include using namespace ::FlexFlow; @@ -7,24 +7,25 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("require_binary_relation_is_left_unique") { SUBCASE("relation is left unique") { BinaryRelation rel = BinaryRelation{ - { - 1, - "one", - }, - { - 2, - "two", - }, - { - 2, - "TWO", - }, + { + 1, + "one", + }, + { + 2, + "two", + }, + { + 2, + "TWO", + }, }; - OneToMany result = require_binary_relation_is_left_unique(rel); + OneToMany result = + require_binary_relation_is_left_unique(rel); OneToMany correct = { - {1, {"one"}}, - {2, {"two", "TWO"}}, + {1, {"one"}}, + {2, {"two", "TWO"}}, }; CHECK(result == correct); @@ -32,18 +33,18 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("relation is not left unique") { BinaryRelation rel = BinaryRelation{ - { - 1, - "odd", - }, - { - 2, - "even", - }, - { - 3, - "odd", - }, + { + 1, + "odd", + }, + { + 2, + "even", + }, + { + 3, + "odd", + }, }; CHECK_THROWS(require_binary_relation_is_left_unique(rel)); diff --git a/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc index f2951d257f..8f515ce805 100644 --- a/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc +++ b/lib/utils/test/src/utils/binary_relation/require_binary_relation_is_right_unique.cc @@ -1,5 +1,5 @@ -#include #include "utils/binary_relation/require_binary_relation_is_right_unique.h" +#include using namespace ::FlexFlow; @@ -7,24 +7,25 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("require_binary_relation_is_right_unique") { SUBCASE("relation is right unique") { BinaryRelation rel = BinaryRelation{ - { - 1, - "odd", - }, - { - 2, - "even", - }, - { - 3, - "odd", - }, + { + 1, + "odd", + }, + { + 2, + "even", + }, + { + 3, + "odd", + }, }; - ManyToOne result = require_binary_relation_is_right_unique(rel); + ManyToOne result = + require_binary_relation_is_right_unique(rel); ManyToOne correct = { - {{1, 3}, "odd"}, - {{2}, "even"}, + {{1, 3}, "odd"}, + {{2}, "even"}, }; CHECK(result == correct); @@ -32,18 +33,18 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("relation is not right unique") { BinaryRelation rel = BinaryRelation{ - { - 1, - "one", - }, - { - 2, - "two", - }, - { - 2, - "TWO", - }, + { + 1, + "one", + }, + { + 2, + "two", + }, + { + 2, + "TWO", + }, }; CHECK_THROWS(require_binary_relation_is_right_unique(rel)); diff --git a/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc index d4487343f2..8b75e5b9c8 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_disjoint_unordered_maps.cc @@ -1,5 +1,5 @@ -#include "utils/containers/binary_merge_disjoint_maps.h" #include "test/utils/doctest/fmt/map.h" +#include "utils/containers/binary_merge_disjoint_maps.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc index 6a848565e1..75b3e64e77 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with.cc @@ -1,5 +1,5 @@ -#include "utils/containers/binary_merge_maps_with.h" #include "test/utils/doctest/fmt/map.h" +#include "utils/containers/binary_merge_maps_with.h" #include #include diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc index fe152dd832..aab3cf696f 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_left_dominating.cc @@ -1,5 +1,5 @@ -#include "utils/containers/binary_merge_maps_with_left_dominating.h" #include "test/utils/doctest/fmt/map.h" +#include "utils/containers/binary_merge_maps_with_left_dominating.h" #include #include diff --git a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc index c107f2b7ff..83d40de49b 100644 --- a/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc +++ b/lib/utils/test/src/utils/containers/binary_merge_unordered_maps_with_right_dominating.cc @@ -1,5 +1,5 @@ -#include "utils/containers/binary_merge_maps_with_right_dominating.h" #include "test/utils/doctest/fmt/map.h" +#include "utils/containers/binary_merge_maps_with_right_dominating.h" #include #include diff --git a/lib/utils/test/src/utils/containers/cartesian_product.cc b/lib/utils/test/src/utils/containers/cartesian_product.cc index e25ffd67f7..372a5dcb74 100644 --- a/lib/utils/test/src/utils/containers/cartesian_product.cc +++ b/lib/utils/test/src/utils/containers/cartesian_product.cc @@ -12,40 +12,35 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("empty") { std::vector> containers = {}; - std::multiset> result = - cartesian_product(containers); + std::multiset> result = cartesian_product(containers); std::multiset> correct = {{}}; CHECK(result == correct); } SUBCASE("single container, one element") { std::vector> containers = {{1}}; - std::multiset> result = - cartesian_product(containers); + std::multiset> result = cartesian_product(containers); std::multiset> correct = {{1}}; CHECK(result == correct); } SUBCASE("single container, multiple elements") { std::vector> containers = {{1, 2, 3}}; - std::multiset> result = - cartesian_product(containers); + std::multiset> result = cartesian_product(containers); std::multiset> correct = {{1}, {2}, {3}}; CHECK(result == correct); } SUBCASE("multiple containers, one element each") { std::vector> containers = {{1}, {2}, {3}}; - std::multiset> result = - cartesian_product(containers); + std::multiset> result = cartesian_product(containers); std::multiset> correct = {{1, 2, 3}}; CHECK(result == correct); } SUBCASE("multiple containers, multiple elements") { std::vector> containers = {{1, 2}, {3, 4}}; - std::multiset> result = - cartesian_product(containers); + std::multiset> result = cartesian_product(containers); std::multiset> correct = { {1, 3}, {1, 4}, {2, 3}, {2, 4}}; CHECK(result == correct); @@ -53,8 +48,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("multiple containers, duplicate elements") { std::vector> containers = {{1, 1}, {2, 3}}; - std::multiset> result = - cartesian_product(containers); + std::multiset> result = cartesian_product(containers); std::multiset> correct = { {1, 2}, {1, 3}, {1, 3}, {1, 2}}; CHECK(result == correct); @@ -62,8 +56,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("1 empty container, 1 non-empty container") { std::vector> containers = {{}, {2, 3}}; - std::multiset> result = - cartesian_product(containers); + std::multiset> result = cartesian_product(containers); std::multiset> correct = {}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/contains_key.cc b/lib/utils/test/src/utils/containers/contains_key.cc index 48933e3fab..e48eff904d 100644 --- a/lib/utils/test/src/utils/containers/contains_key.cc +++ b/lib/utils/test/src/utils/containers/contains_key.cc @@ -2,7 +2,6 @@ #include #include #include -#include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/enumerate.cc b/lib/utils/test/src/utils/containers/enumerate.cc index 17fcdbc046..dca8307669 100644 --- a/lib/utils/test/src/utils/containers/enumerate.cc +++ b/lib/utils/test/src/utils/containers/enumerate.cc @@ -1,7 +1,7 @@ #include "utils/containers/enumerate.h" #include "test/utils/doctest/fmt/map.h" -#include "test/utils/doctest/fmt/pair.h" #include "test/utils/doctest/fmt/multiset.h" +#include "test/utils/doctest/fmt/pair.h" #include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "utils/containers/keys.h" diff --git a/lib/utils/test/src/utils/containers/filter.cc b/lib/utils/test/src/utils/containers/filter.cc index cffd3676b7..33c47ce57e 100644 --- a/lib/utils/test/src/utils/containers/filter.cc +++ b/lib/utils/test/src/utils/containers/filter.cc @@ -1,7 +1,5 @@ #include "utils/containers/filter.h" #include "test/utils/doctest/fmt/map.h" -#include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" diff --git a/lib/utils/test/src/utils/containers/filter_keys.cc b/lib/utils/test/src/utils/containers/filter_keys.cc index 6de906c491..8fabd1386f 100644 --- a/lib/utils/test/src/utils/containers/filter_keys.cc +++ b/lib/utils/test/src/utils/containers/filter_keys.cc @@ -1,15 +1,14 @@ #include "utils/containers/filter_keys.h" #include "test/utils/doctest/fmt/map.h" #include -#include #include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("filter_keys") { - std::map m = { - {1, "one"}, {2, "two"}, {3, "three"}}; + std::map m = {{1, "one"}, {2, "two"}, {3, "three"}}; auto f = [](int x) { return x % 2 == 1; }; std::map result = filter_keys(m, f); std::map correct = {{1, "one"}, {3, "three"}}; diff --git a/lib/utils/test/src/utils/containers/filtermap_keys.cc b/lib/utils/test/src/utils/containers/filtermap_keys.cc index b71bb8a052..f28789f616 100644 --- a/lib/utils/test/src/utils/containers/filtermap_keys.cc +++ b/lib/utils/test/src/utils/containers/filtermap_keys.cc @@ -1,6 +1,5 @@ #include "utils/containers/filtermap_keys.h" #include "test/utils/doctest/fmt/map.h" -#include "test/utils/doctest/fmt/map.h" #include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/containers/filtermap_values.cc b/lib/utils/test/src/utils/containers/filtermap_values.cc index f6d4335405..6e52ca4418 100644 --- a/lib/utils/test/src/utils/containers/filtermap_values.cc +++ b/lib/utils/test/src/utils/containers/filtermap_values.cc @@ -1,6 +1,5 @@ #include "utils/containers/filtermap_values.h" #include "test/utils/doctest/fmt/map.h" -#include "test/utils/doctest/fmt/map.h" #include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/containers/filtrans.cc b/lib/utils/test/src/utils/containers/filtrans.cc index 4aab65d528..ccee148407 100644 --- a/lib/utils/test/src/utils/containers/filtrans.cc +++ b/lib/utils/test/src/utils/containers/filtrans.cc @@ -1,6 +1,5 @@ #include "utils/containers/filtrans.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include diff --git a/lib/utils/test/src/utils/containers/find.cc b/lib/utils/test/src/utils/containers/find.cc index b3fc17eb82..6b6f02025e 100644 --- a/lib/utils/test/src/utils/containers/find.cc +++ b/lib/utils/test/src/utils/containers/find.cc @@ -3,7 +3,6 @@ #include #include #include -#include #include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/containers/flatmap.cc b/lib/utils/test/src/utils/containers/flatmap.cc index 18e0cf88a3..d9a48a5d38 100644 --- a/lib/utils/test/src/utils/containers/flatmap.cc +++ b/lib/utils/test/src/utils/containers/flatmap.cc @@ -1,6 +1,6 @@ #include "utils/containers/flatmap.h" -#include "test/utils/doctest/fmt/pair.h" #include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/pair.h" #include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "utils/containers/map_keys.h" @@ -57,8 +57,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::set input = {"hello", " ", "", "world", "!"}; std::set result = flatmap(input, get_chars); - std::set correct = { - 'h', 'e', 'l', 'o', ' ', 'w', 'r', 'd', '!'}; + std::set correct = {'h', 'e', 'l', 'o', ' ', 'w', 'r', 'd', '!'}; CHECK(result == correct); } @@ -106,8 +105,7 @@ TEST_SUITE(FF_TEST_SUITE) { } TEST_CASE("flatmap(std::map, F)") { - auto de_nest_keys = [](int k1, - std::map const &v) { + auto de_nest_keys = [](int k1, std::map const &v) { return map_keys(v, [&](int k2) { return std::pair{k1, k2}; }); }; diff --git a/lib/utils/test/src/utils/containers/get_all_assignments.cc b/lib/utils/test/src/utils/containers/get_all_assignments.cc index bcf69cdc87..1e2a8ac2dc 100644 --- a/lib/utils/test/src/utils/containers/get_all_assignments.cc +++ b/lib/utils/test/src/utils/containers/get_all_assignments.cc @@ -10,8 +10,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("empty input") { std::map> input = {}; - std::set> result = - get_all_assignments(input); + std::set> result = get_all_assignments(input); std::set> correct = {{}}; CHECK(result == correct); @@ -23,8 +22,7 @@ TEST_SUITE(FF_TEST_SUITE) { {"b", {2, 3}}, }; - std::set> result = - get_all_assignments(input); + std::set> result = get_all_assignments(input); std::set> correct = { {{"a", 1}, {"b", 2}}, {{"a", 1}, {"b", 3}}, @@ -43,8 +41,7 @@ TEST_SUITE(FF_TEST_SUITE) { {"b", {2, 3}}, }; - std::set> result = - get_all_assignments(input); + std::set> result = get_all_assignments(input); std::set> correct = {}; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc b/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc index 5e75a7a25d..9f5f931126 100644 --- a/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc +++ b/lib/utils/test/src/utils/containers/get_all_permutations_with_repetition.cc @@ -60,14 +60,14 @@ TEST_SUITE(FF_TEST_SUITE) { std::multiset> result = get_all_permutations_with_repetition(input, 2_n); std::multiset> correct = {{1, 1}, - {1, 2}, - {1, 2}, - {2, 1}, - {2, 1}, - {2, 2}, - {2, 2}, - {2, 2}, - {2, 2}}; + {1, 2}, + {1, 2}, + {2, 1}, + {2, 1}, + {2, 2}, + {2, 2}, + {2, 2}, + {2, 2}}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/group_by.cc b/lib/utils/test/src/utils/containers/group_by.cc index 1f9848260a..7fb2da7864 100644 --- a/lib/utils/test/src/utils/containers/group_by.cc +++ b/lib/utils/test/src/utils/containers/group_by.cc @@ -1,5 +1,4 @@ #include "utils/containers/group_by.h" -#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" diff --git a/lib/utils/test/src/utils/containers/inplace_filter.cc b/lib/utils/test/src/utils/containers/inplace_filter.cc index 43c3c01444..f8f9dcb377 100644 --- a/lib/utils/test/src/utils/containers/inplace_filter.cc +++ b/lib/utils/test/src/utils/containers/inplace_filter.cc @@ -1,8 +1,6 @@ #include "utils/containers/inplace_filter.h" #include "test/utils/doctest/fmt/map.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/map.h" -#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "test/utils/rapidcheck.h" #include diff --git a/lib/utils/test/src/utils/containers/is_submapeq_of.cc b/lib/utils/test/src/utils/containers/is_submapeq_of.cc index f84672c38b..4be71e51af 100644 --- a/lib/utils/test/src/utils/containers/is_submapeq_of.cc +++ b/lib/utils/test/src/utils/containers/is_submapeq_of.cc @@ -1,14 +1,13 @@ #include "utils/containers/is_submapeq_of.h" #include -#include #include +#include using namespace ::FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("is_submapeq_of") { - std::map super = { - {1, "one"}, {2, "two"}, {3, "three"}}; + std::map super = {{1, "one"}, {2, "two"}, {3, "three"}}; SUBCASE("keys and values match") { std::map sub = {{1, "one"}, {2, "two"}}; @@ -21,8 +20,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("keys match but values don't") { - std::map sub = {{1, "wrong_value"}, - {2, "two"}}; + std::map sub = {{1, "wrong_value"}, {2, "two"}}; CHECK_FALSE(is_submapeq_of(sub, super)); } diff --git a/lib/utils/test/src/utils/containers/keys.cc b/lib/utils/test/src/utils/containers/keys.cc index ce455a608c..01434bc253 100644 --- a/lib/utils/test/src/utils/containers/keys.cc +++ b/lib/utils/test/src/utils/containers/keys.cc @@ -1,16 +1,15 @@ #include "utils/containers/keys.h" #include "test/utils/doctest/fmt/set.h" #include -#include #include #include +#include using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("keys") { - std::map m = { - {1, "one"}, {2, "two"}, {3, "three"}}; + std::map m = {{1, "one"}, {2, "two"}, {3, "three"}}; std::set result = keys(m); std::set expected = {1, 2, 3}; CHECK(result == expected); diff --git a/lib/utils/test/src/utils/containers/lift_optional_through_map.cc b/lib/utils/test/src/utils/containers/lift_optional_through_map.cc index e9a84ca392..6ae7170a52 100644 --- a/lib/utils/test/src/utils/containers/lift_optional_through_map.cc +++ b/lib/utils/test/src/utils/containers/lift_optional_through_map.cc @@ -1,6 +1,6 @@ #include "utils/containers/lift_optional_through_map.h" -#include "test/utils/doctest/fmt/optional.h" #include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/optional.h" #include using namespace ::FlexFlow; @@ -25,8 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::optional> result = lift_optional_through_map(input); - std::optional> correct = - std::nullopt; + std::optional> correct = std::nullopt; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/map_from_pairs.cc b/lib/utils/test/src/utils/containers/map_from_pairs.cc index 48e8b9fe05..05476eed91 100644 --- a/lib/utils/test/src/utils/containers/map_from_pairs.cc +++ b/lib/utils/test/src/utils/containers/map_from_pairs.cc @@ -1,8 +1,8 @@ #include "utils/containers/map_from_pairs.h" #include "test/utils/doctest/fmt/map.h" #include -#include #include +#include using namespace ::FlexFlow; @@ -16,11 +16,10 @@ TEST_SUITE(FF_TEST_SUITE) { std::map result = map_from_pairs(input); - std::map correct = - std::map{ - {1, "one"}, - {2, "two"}, - }; + std::map correct = std::map{ + {1, "one"}, + {2, "two"}, + }; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/map_keys.cc b/lib/utils/test/src/utils/containers/map_keys.cc index 47ba041b18..569e9bc9dc 100644 --- a/lib/utils/test/src/utils/containers/map_keys.cc +++ b/lib/utils/test/src/utils/containers/map_keys.cc @@ -1,8 +1,8 @@ #include "utils/containers/map_keys.h" #include "test/utils/doctest/fmt/map.h" #include -#include #include +#include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/containers/map_values.cc b/lib/utils/test/src/utils/containers/map_values.cc index 6ed3e8cb81..02cdb6b87e 100644 --- a/lib/utils/test/src/utils/containers/map_values.cc +++ b/lib/utils/test/src/utils/containers/map_values.cc @@ -1,8 +1,8 @@ #include "utils/containers/map_values.h" #include "test/utils/doctest/fmt/map.h" #include -#include #include +#include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc index bf4b2202d7..3e358381cf 100644 --- a/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc +++ b/lib/utils/test/src/utils/containers/merge_disjoint_unordered_maps.cc @@ -1,5 +1,5 @@ -#include "utils/containers/merge_disjoint_maps.h" #include "test/utils/doctest/fmt/map.h" +#include "utils/containers/merge_disjoint_maps.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/merge_maps_with.cc b/lib/utils/test/src/utils/containers/merge_maps_with.cc index fd73a9345e..60b4c2ac9f 100644 --- a/lib/utils/test/src/utils/containers/merge_maps_with.cc +++ b/lib/utils/test/src/utils/containers/merge_maps_with.cc @@ -12,18 +12,17 @@ TEST_SUITE(FF_TEST_SUITE) { return l + r; }; - RC_SUBCASE( - "with two inputs, matches binary_merge_maps_with", - [&](std::map const &lhs, - std::map const &rhs) { - std::map from_merge_maps_with = - merge_maps_with(std::vector{lhs, rhs}, string_concat); + RC_SUBCASE("with two inputs, matches binary_merge_maps_with", + [&](std::map const &lhs, + std::map const &rhs) { + std::map from_merge_maps_with = + merge_maps_with(std::vector{lhs, rhs}, string_concat); - std::map from_binary_merge_maps_with = - binary_merge_maps_with(lhs, rhs, string_concat); + std::map from_binary_merge_maps_with = + binary_merge_maps_with(lhs, rhs, string_concat); - CHECK(from_merge_maps_with == from_binary_merge_maps_with); - }); + CHECK(from_merge_maps_with == from_binary_merge_maps_with); + }); SUBCASE("maps overlap") { std::map map1 = { @@ -89,8 +88,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("no maps are provided") { std::vector> maps = {}; - std::map result = - merge_maps_with(maps, fail_if_called); + std::map result = merge_maps_with(maps, fail_if_called); std::map correct = {}; diff --git a/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc b/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc index fd73a9345e..0a0604b17a 100644 --- a/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc +++ b/lib/utils/test/src/utils/containers/merge_unordered_maps_with.cc @@ -1,7 +1,7 @@ -#include "utils/containers/merge_maps_with.h" #include "test/utils/doctest/fmt/map.h" #include "test/utils/rapidcheck.h" #include "utils/containers/binary_merge_maps_with.h" +#include "utils/containers/merge_maps_with.h" #include using namespace ::FlexFlow; @@ -12,18 +12,17 @@ TEST_SUITE(FF_TEST_SUITE) { return l + r; }; - RC_SUBCASE( - "with two inputs, matches binary_merge_maps_with", - [&](std::map const &lhs, - std::map const &rhs) { - std::map from_merge_maps_with = - merge_maps_with(std::vector{lhs, rhs}, string_concat); + RC_SUBCASE("with two inputs, matches binary_merge_maps_with", + [&](std::map const &lhs, + std::map const &rhs) { + std::map from_merge_maps_with = + merge_maps_with(std::vector{lhs, rhs}, string_concat); - std::map from_binary_merge_maps_with = - binary_merge_maps_with(lhs, rhs, string_concat); + std::map from_binary_merge_maps_with = + binary_merge_maps_with(lhs, rhs, string_concat); - CHECK(from_merge_maps_with == from_binary_merge_maps_with); - }); + CHECK(from_merge_maps_with == from_binary_merge_maps_with); + }); SUBCASE("maps overlap") { std::map map1 = { @@ -89,8 +88,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("no maps are provided") { std::vector> maps = {}; - std::map result = - merge_maps_with(maps, fail_if_called); + std::map result = merge_maps_with(maps, fail_if_called); std::map correct = {}; diff --git a/lib/utils/test/src/utils/containers/multiset_union.cc b/lib/utils/test/src/utils/containers/multiset_union.cc index 7eeb24df99..f4534d5e20 100644 --- a/lib/utils/test/src/utils/containers/multiset_union.cc +++ b/lib/utils/test/src/utils/containers/multiset_union.cc @@ -1,6 +1,5 @@ #include "utils/containers/multiset_union.h" #include "test/utils/doctest/fmt/multiset.h" -#include "test/utils/doctest/fmt/multiset.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/product.cc b/lib/utils/test/src/utils/containers/product.cc index ff6c515870..bf86f8e178 100644 --- a/lib/utils/test/src/utils/containers/product.cc +++ b/lib/utils/test/src/utils/containers/product.cc @@ -3,7 +3,6 @@ #include #include #include -#include #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/require_all_same1.cc b/lib/utils/test/src/utils/containers/require_all_same1.cc index 6e4c24a3cd..ba3f6f3796 100644 --- a/lib/utils/test/src/utils/containers/require_all_same1.cc +++ b/lib/utils/test/src/utils/containers/require_all_same1.cc @@ -3,14 +3,11 @@ #include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/optional.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/multiset.h" -#include "test/utils/doctest/fmt/set.h" #include "test/utils/doctest/fmt/vector.h" #include "utils/expected.h" #include #include #include -#include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/require_no_duplicates.cc b/lib/utils/test/src/utils/containers/require_no_duplicates.cc index b84c093e0c..e4a1fa3136 100644 --- a/lib/utils/test/src/utils/containers/require_no_duplicates.cc +++ b/lib/utils/test/src/utils/containers/require_no_duplicates.cc @@ -1,8 +1,6 @@ #include "utils/containers/require_no_duplicates.h" #include "test/utils/doctest/fmt/multiset.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/multiset.h" -#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/restrict_keys.cc b/lib/utils/test/src/utils/containers/restrict_keys.cc index 2b78a59468..bc45e9ba77 100644 --- a/lib/utils/test/src/utils/containers/restrict_keys.cc +++ b/lib/utils/test/src/utils/containers/restrict_keys.cc @@ -7,8 +7,7 @@ using namespace FlexFlow; TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("restrict_keys") { - std::map m = { - {1, "one"}, {2, "two"}, {3, "three"}}; + std::map m = {{1, "one"}, {2, "two"}, {3, "three"}}; std::set mask = {2, 3, 4}; std::map result = restrict_keys(m, mask); std::map correct = {{2, "two"}, {3, "three"}}; diff --git a/lib/utils/test/src/utils/containers/set_intersection.cc b/lib/utils/test/src/utils/containers/set_intersection.cc index 69a3845ef8..7f2680d637 100644 --- a/lib/utils/test/src/utils/containers/set_intersection.cc +++ b/lib/utils/test/src/utils/containers/set_intersection.cc @@ -1,7 +1,6 @@ #include "utils/containers/set_intersection.h" #include "test/utils/doctest/fmt/optional.h" #include "test/utils/doctest/fmt/set.h" -#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc b/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc index 6804ea0243..d222202cfe 100644 --- a/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc +++ b/lib/utils/test/src/utils/containers/try_merge_nondisjoint_unordered_maps.cc @@ -1,6 +1,6 @@ -#include "utils/containers/try_merge_nondisjoint_maps.h" -#include "test/utils/doctest/fmt/optional.h" #include "test/utils/doctest/fmt/map.h" +#include "test/utils/doctest/fmt/optional.h" +#include "utils/containers/try_merge_nondisjoint_maps.h" #include using namespace ::FlexFlow; @@ -32,8 +32,7 @@ TEST_SUITE(FF_TEST_SUITE) { d1.insert({2, "three"}); std::optional> result = try_merge_nondisjoint_maps(d1, d2); - std::optional> correct = - std::nullopt; + std::optional> correct = std::nullopt; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc b/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc index 13d9dabeb9..b62369d3bf 100644 --- a/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc +++ b/lib/utils/test/src/utils/containers/unordered_map_from_pairs.cc @@ -1,6 +1,6 @@ -#include "utils/containers/map_from_pairs.h" #include "test/utils/doctest/fmt/map.h" #include "utils/containers/contains.h" +#include "utils/containers/map_from_pairs.h" #include #include #include @@ -15,8 +15,7 @@ TEST_SUITE(FF_TEST_SUITE) { {3, "world"}, }; - std::map result = - map_from_pairs(input); + std::map result = map_from_pairs(input); std::map correct = { {1, "hello"}, {3, "world"}, @@ -28,8 +27,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("empty input") { std::vector> input = {}; - std::map result = - map_from_pairs(input); + std::map result = map_from_pairs(input); std::map correct = {}; CHECK(result == correct); @@ -42,14 +40,12 @@ TEST_SUITE(FF_TEST_SUITE) { {1, "b"}, }; - std::map result = - map_from_pairs(input); + std::map result = map_from_pairs(input); - std::vector> - possible_correct_values = { - {{1, "a"}, {2, "c"}}, - {{1, "b"}, {2, "c"}}, - }; + std::vector> possible_correct_values = { + {{1, "a"}, {2, "c"}}, + {{1, "b"}, {2, "c"}}, + }; CHECK(contains(possible_correct_values, result)); } diff --git a/lib/utils/test/src/utils/containers/unordered_multiset_of.cc b/lib/utils/test/src/utils/containers/unordered_multiset_of.cc index d44979f655..334be12359 100644 --- a/lib/utils/test/src/utils/containers/unordered_multiset_of.cc +++ b/lib/utils/test/src/utils/containers/unordered_multiset_of.cc @@ -1,5 +1,5 @@ -#include "utils/containers/multiset_of.h" #include "test/utils/doctest/fmt/multiset.h" +#include "utils/containers/multiset_of.h" #include #include diff --git a/lib/utils/test/src/utils/containers/unordered_set_of.cc b/lib/utils/test/src/utils/containers/unordered_set_of.cc index 69b762c6ab..d5b459fbee 100644 --- a/lib/utils/test/src/utils/containers/unordered_set_of.cc +++ b/lib/utils/test/src/utils/containers/unordered_set_of.cc @@ -1,5 +1,5 @@ -#include "utils/containers/set_of.h" #include "test/utils/doctest/fmt/set.h" +#include "utils/containers/set_of.h" #include #include diff --git a/lib/utils/test/src/utils/containers/values.cc b/lib/utils/test/src/utils/containers/values.cc index e98882e4e9..204ee27a2a 100644 --- a/lib/utils/test/src/utils/containers/values.cc +++ b/lib/utils/test/src/utils/containers/values.cc @@ -1,8 +1,8 @@ #include "utils/containers/values.h" #include "test/utils/doctest/fmt/multiset.h" #include -#include #include +#include #include using namespace FlexFlow; @@ -12,8 +12,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::map m = { {1, "one"}, {2, "two"}, {3, "three"}, {33, "three"}}; std::multiset result = values(m); - std::multiset correct = { - "one", "two", "three", "three"}; + std::multiset correct = {"one", "two", "three", "three"}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/fmt/unordered_map.cc b/lib/utils/test/src/utils/fmt/unordered_map.cc index 7f43728df7..b69fe89b1b 100644 --- a/lib/utils/test/src/utils/fmt/unordered_map.cc +++ b/lib/utils/test/src/utils/fmt/unordered_map.cc @@ -1,6 +1,6 @@ -#include "utils/fmt/map.h" #include "test/utils/doctest/fmt/map.h" #include "utils/containers/get_element_counts.h" +#include "utils/fmt/map.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/graph/algorithms.cc b/lib/utils/test/src/utils/graph/algorithms.cc index 78ee4e704d..2e592cd71e 100644 --- a/lib/utils/test/src/utils/graph/algorithms.cc +++ b/lib/utils/test/src/utils/graph/algorithms.cc @@ -75,8 +75,8 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(2)}, DirectedEdge{n.at(2), n.at(0)}, DirectedEdge{n.at(2), n.at(1)}}); - std::set> corrects = { - {n.at(0), n.at(1), n.at(2)}, {n.at(0), n.at(2), n.at(1)}}; + std::set> corrects = {{n.at(0), n.at(1), n.at(2)}, + {n.at(0), n.at(2), n.at(1)}}; std::vector result = get_bfs_ordering(g, {n.at(0)}); CHECK(contains(corrects, result)); } diff --git a/lib/utils/test/src/utils/graph/cow_ptr_t.cc b/lib/utils/test/src/utils/graph/cow_ptr_t.cc index 6feba34dab..f7758bc210 100644 --- a/lib/utils/test/src/utils/graph/cow_ptr_t.cc +++ b/lib/utils/test/src/utils/graph/cow_ptr_t.cc @@ -1,7 +1,7 @@ #include "utils/graph/cow_ptr_t.h" #include -#include #include +#include #include using namespace FlexFlow; diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc index 78be249a07..4b756cfd1c 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/complete_bipartite_composite/complete_bipartite_composite_decomposition.cc @@ -48,18 +48,14 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_head_subcomponents") { - std::set> result = - get_head_subcomponents(cbc); - std::set> correct = {bc1.head_nodes, - bc2.head_nodes}; + std::set> result = get_head_subcomponents(cbc); + std::set> correct = {bc1.head_nodes, bc2.head_nodes}; CHECK(result == correct); } SUBCASE("get_tail_subcomponents") { - std::set> result = - get_tail_subcomponents(cbc); - std::set> correct = {bc1.tail_nodes, - bc2.tail_nodes}; + std::set> result = get_tail_subcomponents(cbc); + std::set> correct = {bc1.tail_nodes, bc2.tail_nodes}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc index d4c1bdd0a0..66945af531 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/contract_node.cc @@ -29,8 +29,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("nodes") { std::set result_nodes = get_nodes(result); - std::set correct_nodes = { - n.at(1), n.at(2), n.at(3), n.at(4)}; + std::set correct_nodes = {n.at(1), n.at(2), n.at(3), n.at(4)}; CHECK(result_nodes == correct_nodes); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc index c2eeef7a1a..9f039ad6b5 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_dominators_map.cc @@ -34,8 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {n.at(0), n.at(1), n.at(5)}}, }; - std::map> result = - get_dominators_map(g); + std::map> result = get_dominators_map(g); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc index 3d71bc836d..df19acdcb2 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_imm_dominators_map.cc @@ -34,8 +34,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), n.at(1)}, }; - std::map> result = - get_imm_dominators_map(g); + std::map> result = get_imm_dominators_map(g); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc index 69ccfadc5c..d1348e8687 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_lowest_common_ancestors.cc @@ -32,8 +32,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("trees") { SUBCASE("single node") { std::vector n = add_nodes(g, 1); - std::optional> correct = - std::set{n.at(0)}; + std::optional> correct = std::set{n.at(0)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(0)}); CHECK(correct == result); @@ -46,24 +45,21 @@ TEST_SUITE(FF_TEST_SUITE) { {DirectedEdge{n.at(0), n.at(1)}, DirectedEdge{n.at(0), n.at(2)}}); SUBCASE("LCA of siblings is parent") { - std::optional> correct = - std::set{n.at(0)}; + std::optional> correct = std::set{n.at(0)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(1), n.at(2)}); CHECK(correct == result); } SUBCASE("LCA of a single node is itself") { - std::optional> correct = - std::set{n.at(1)}; + std::optional> correct = std::set{n.at(1)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(1)}); CHECK(correct == result); } SUBCASE("LCA of another single node is itself") { - std::optional> correct = - std::set{n.at(2)}; + std::optional> correct = std::set{n.at(2)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(2)}); CHECK(correct == result); @@ -80,35 +76,30 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(3), n.at(5)}}); SUBCASE("LCA of nodes at different depths (root is LCA)") { - std::optional> correct = - std::set{n.at(0)}; + std::optional> correct = std::set{n.at(0)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(5), n.at(2)}); CHECK(correct == result); } SUBCASE("LCA of node and its ancestor is the ancestor") { - std::optional> correct = - std::set{n.at(3)}; + std::optional> correct = std::set{n.at(3)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(5), n.at(3)}); CHECK(correct == result); } SUBCASE("LCA of siblings at depth 2") { - std::optional> correct = - std::set{n.at(1)}; + std::optional> correct = std::set{n.at(1)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(3), n.at(4)}); CHECK(correct == result); } SUBCASE("LCA of multiple nodes across different branches") { - std::optional> correct = - std::set{n.at(0)}; - std::optional> result = - get_lowest_common_ancestors( - g, {n.at(1), n.at(2), n.at(3), n.at(4), n.at(5)}); + std::optional> correct = std::set{n.at(0)}; + std::optional> result = get_lowest_common_ancestors( + g, {n.at(1), n.at(2), n.at(3), n.at(4), n.at(5)}); CHECK(correct == result); } } @@ -121,24 +112,21 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(2), n.at(3)}}); SUBCASE("LCA of adjacent nodes in a path") { - std::optional> correct = - std::set{n.at(2)}; + std::optional> correct = std::set{n.at(2)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(2), n.at(3)}); CHECK(correct == result); } SUBCASE("LCA of non-adjacent nodes in a path") { - std::optional> correct = - std::set{n.at(1)}; + std::optional> correct = std::set{n.at(1)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(1), n.at(3)}); CHECK(correct == result); } SUBCASE("LCA of multiple nodes in a path") { - std::optional> correct = - std::set{n.at(1)}; + std::optional> correct = std::set{n.at(1)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(1), n.at(2), n.at(3)}); CHECK(correct == result); @@ -154,8 +142,7 @@ TEST_SUITE(FF_TEST_SUITE) { g, {DirectedEdge{n.at(0), n.at(2)}, DirectedEdge{n.at(1), n.at(2)}}); - std::optional> correct = - std::set{}; + std::optional> correct = std::set{}; std::optional> result = get_lowest_common_ancestors(g, {n.at(0), n.at(1)}); CHECK(correct == result); @@ -187,8 +174,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(3), n.at(5)}, DirectedEdge{n.at(1), n.at(5)}}); - std::optional> correct = - std::set{n.at(3)}; + std::optional> correct = std::set{n.at(3)}; std::optional> result = get_lowest_common_ancestors(g, {n.at(4), n.at(5)}); CHECK(correct == result); diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc index 0b21b7dfa9..b9d1730553 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_post_dominators_map.cc @@ -14,8 +14,7 @@ TEST_SUITE(FF_TEST_SUITE) { g.add_edge(DirectedEdge{n.at(0), n.at(1)}); - std::map> result = - get_post_dominators_map(g); + std::map> result = get_post_dominators_map(g); std::map> correct = { {n.at(0), {n.at(0), n.at(1)}}, {n.at(1), {n.at(1)}}, @@ -41,8 +40,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(8), n.at(9)}, }); - std::map> result = - get_post_dominators_map(g); + std::map> result = get_post_dominators_map(g); std::map> correct = { {n.at(0), {n.at(0), n.at(9)}}, {n.at(1), {n.at(1), n.at(7), n.at(9)}}, @@ -85,8 +83,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {n.at(5)}}, }; - std::map> result = - get_post_dominators_map(g); + std::map> result = get_post_dominators_map(g); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc index 8241b8891f..cc8538e68f 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_predecessors.cc @@ -31,8 +31,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {n.at(1)}}, }; - std::map> result = - get_predecessors(g); + std::map> result = get_predecessors(g); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc index 5688494233..1f092c77e0 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_successors.cc @@ -31,8 +31,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), {}}, }; - std::map> result = - get_successors(g); + std::map> result = get_successors(g); CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc index ee407b2179..06daee46e6 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/get_weakly_connected_components.cc @@ -13,8 +13,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("single node") { std::vector n = add_nodes(g, 1); - std::set> result = - get_weakly_connected_components(g); + std::set> result = get_weakly_connected_components(g); std::set> correct = {{n.at(0)}}; CHECK(result == correct); } @@ -26,8 +25,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(0), n.at(0)}, }); - std::set> result = - get_weakly_connected_components(g); + std::set> result = get_weakly_connected_components(g); std::set> correct = {{n.at(0)}}; CHECK(result == correct); } @@ -40,10 +38,8 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(1)}, }); - std::set> result = - get_weakly_connected_components(g); - std::set> correct = {{n.at(0)}, - {n.at(1)}}; + std::set> result = get_weakly_connected_components(g); + std::set> correct = {{n.at(0)}, {n.at(1)}}; CHECK(result == correct); } @@ -54,10 +50,8 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(0), n.at(1)}, }); - std::set> result = - get_weakly_connected_components(g); - std::set> correct = { - {n.at(0), n.at(1)}}; + std::set> result = get_weakly_connected_components(g); + std::set> correct = {{n.at(0), n.at(1)}}; CHECK(result == correct); } @@ -69,10 +63,8 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(1), n.at(0)}, }); - std::set> result = - get_weakly_connected_components(g); - std::set> correct = { - {n.at(0), n.at(1)}}; + std::set> result = get_weakly_connected_components(g); + std::set> correct = {{n.at(0), n.at(1)}}; CHECK(result == correct); } @@ -90,8 +82,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(4), n.at(3)}, }); - std::set> result = - get_weakly_connected_components(g); + std::set> result = get_weakly_connected_components(g); std::set> correct = { {n.at(0), n.at(1), n.at(2)}, {n.at(3), n.at(4)}, diff --git a/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc b/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc index 3947aab336..9894a57037 100644 --- a/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc +++ b/lib/utils/test/src/utils/graph/digraph/algorithms/inverse_line_graph/get_inverse_line_graph.cc @@ -76,12 +76,11 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("inverse_edge_to_line_node_bidict") { - std::map result_bidict = - map_values(result.inverse_edge_to_line_node_bidict.reversed() - .as_map(), - [&](MultiDiEdge const &e) { - return get_directed_edge(result.graph, e); - }); + std::map result_bidict = map_values( + result.inverse_edge_to_line_node_bidict.reversed().as_map(), + [&](MultiDiEdge const &e) { + return get_directed_edge(result.graph, e); + }); std::map correct_bidict = { {n.at(0), DirectedEdge{inv.at(0), inv.at(1)}}, {n.at(1), DirectedEdge{inv.at(1), inv.at(2)}}, @@ -134,12 +133,11 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("inverse_edge_to_line_node_bidict") { - std::map result_bidict = - map_values(result.inverse_edge_to_line_node_bidict.reversed() - .as_map(), - [&](MultiDiEdge const &e) { - return get_directed_edge(result.graph, e); - }); + std::map result_bidict = map_values( + result.inverse_edge_to_line_node_bidict.reversed().as_map(), + [&](MultiDiEdge const &e) { + return get_directed_edge(result.graph, e); + }); std::map correct_bidict = { {n.at(0), DirectedEdge{inv.at(0), inv.at(1)}}, {n.at(1), DirectedEdge{inv.at(0), inv.at(1)}}, diff --git a/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc b/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc index 9ab7ee8d2d..426043ac9b 100644 --- a/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc +++ b/lib/utils/test/src/utils/graph/instances/adjacency_digraph.cc @@ -75,8 +75,7 @@ TEST_SUITE(FF_TEST_SUITE) { }; std::set result = g.query_edges(query); - std::set correct = - std::set{e[0]}; + std::set correct = std::set{e[0]}; CHECK(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc b/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc index 39230fb090..5f23ce02c7 100644 --- a/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc +++ b/lib/utils/test/src/utils/graph/instances/adjacency_multidigraph.cc @@ -11,22 +11,20 @@ TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("AdjacencyMultiDiGraph") { MultiDiGraph g = MultiDiGraph::create(); - auto check_state = - [&](std::set const &correct_nodes, - std::set const &correct_edges) { - { - std::set result = g.query_nodes(node_query_all()); - std::set correct = correct_nodes; - REQUIRE(result == correct); - } - - { - std::set result = - g.query_edges(multidiedge_query_all()); - std::set correct = correct_edges; - REQUIRE(result == correct); - } - }; + auto check_state = [&](std::set const &correct_nodes, + std::set const &correct_edges) { + { + std::set result = g.query_nodes(node_query_all()); + std::set correct = correct_nodes; + REQUIRE(result == correct); + } + + { + std::set result = g.query_edges(multidiedge_query_all()); + std::set correct = correct_edges; + REQUIRE(result == correct); + } + }; check_state({}, {}); @@ -138,8 +136,7 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("edges") { g.add_edge(n1, n2); - std::set result = - g2.query_edges(multidiedge_query_all()); + std::set result = g2.query_edges(multidiedge_query_all()); std::set correct = {e1, e2, e3, e4}; CHECK(result == correct); } diff --git a/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc b/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc index cc1142bdcd..17421f434b 100644 --- a/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/instances/unordered_set_dataflow_graph.cc @@ -18,8 +18,7 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set result = - g.query_edges(dataflow_edge_query_all()); + std::set result = g.query_edges(dataflow_edge_query_all()); std::set correct = {}; REQUIRE(result == correct); } @@ -40,8 +39,7 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set result = - g.query_edges(dataflow_edge_query_all()); + std::set result = g.query_edges(dataflow_edge_query_all()); std::set correct = {}; REQUIRE(result == correct); } @@ -49,8 +47,7 @@ TEST_SUITE(FF_TEST_SUITE) { { std::set result = g.query_outputs(dataflow_output_query_all()); - std::set correct = - set_of(added.outputs); + std::set correct = set_of(added.outputs); REQUIRE(result == correct); } @@ -63,8 +60,7 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set result = - g.query_edges(dataflow_edge_query_all()); + std::set result = g.query_edges(dataflow_edge_query_all()); std::set correct = { DataflowEdge{added.outputs.at(0), DataflowInput{added2.node, 0_n}}, DataflowEdge{added.outputs.at(1), DataflowInput{added2.node, 1_n}}, @@ -75,8 +71,8 @@ TEST_SUITE(FF_TEST_SUITE) { { std::set result = g.query_outputs(dataflow_output_query_all()); - std::set correct = set_union( - set_of(added.outputs), set_of(added2.outputs)); + std::set correct = + set_union(set_of(added.outputs), set_of(added2.outputs)); REQUIRE(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc b/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc index 5448a81760..6ad8593774 100644 --- a/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/instances/unordered_set_labelled_open_kwarg_dataflow_graph.cc @@ -25,11 +25,10 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set> - result = g.query_edges( + std::set> result = + g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::set> - correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } @@ -75,11 +74,10 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set> - result = g.query_edges( + std::set> result = + g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::set> - correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } @@ -130,8 +128,8 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set> - result = g.query_edges( + std::set> result = + g.query_edges( open_kwarg_dataflow_edge_query_all()); auto internal_edge = [](KwargDataflowOutput const &src, @@ -164,11 +162,10 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::set> - correct = { - internal_edge(added_output_1, added2.node, "input_1"), - internal_edge(added_output_3, added2.node, "input_2"), - }; + std::set> correct = { + internal_edge(added_output_1, added2.node, "input_1"), + internal_edge(added_output_3, added2.node, "input_2"), + }; REQUIRE(result == correct); } diff --git a/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc b/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc index 354ea4ac85..e3289316ff 100644 --- a/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/instances/unordered_set_open_kwarg_dataflow_graph.cc @@ -19,11 +19,10 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set> - result = g.query_edges( + std::set> result = + g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::set> - correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } @@ -58,11 +57,10 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set> - result = g.query_edges( + std::set> result = + g.query_edges( open_kwarg_dataflow_edge_query_all()); - std::set> - correct = {}; + std::set> correct = {}; REQUIRE(result == correct); } @@ -109,8 +107,8 @@ TEST_SUITE(FF_TEST_SUITE) { } { - std::set> - result = g.query_edges( + std::set> result = + g.query_edges( open_kwarg_dataflow_edge_query_all()); auto internal_edge = [](KwargDataflowOutput const &src, @@ -143,11 +141,10 @@ TEST_SUITE(FF_TEST_SUITE) { }; }; - std::set> - correct = { - internal_edge(added_output_1, added2.node, "input_1"), - internal_edge(added_output_3, added2.node, "input_2"), - }; + std::set> correct = { + internal_edge(added_output_1, added2.node, "input_1"), + internal_edge(added_output_3, added2.node, "input_2"), + }; REQUIRE(result == correct); } diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc index 5d8983ff4e..6984fb9859 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_data_from_kwarg_dataflow_graph_data.cc @@ -66,10 +66,11 @@ TEST_SUITE(FF_TEST_SUITE) { }, }; - std::function( - std::set const &)> - slot_ordering = [](std::set const &slots) - -> std::vector { return reversed(sorted(slots)); }; + std::function(std::set const &)> + slot_ordering = + [](std::set const &slots) -> std::vector { + return reversed(sorted(slots)); + }; DataflowGraphData result = dataflow_graph_data_from_kwarg_dataflow_graph_data(input, diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc index f46e62c588..12a536bb8c 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/dataflow_graph_from_kwarg_dataflow_graph.cc @@ -22,8 +22,7 @@ TEST_SUITE(FF_TEST_SUITE) { UnorderedSetKwargDataflowGraph>(); KwargNodeAddedResult n0_added = g.add_node( - /*inputs=*/std::map>{}, + /*inputs=*/std::map>{}, /*outputs=*/std::set{ "a", }); @@ -32,8 +31,7 @@ TEST_SUITE(FF_TEST_SUITE) { require_only_key(n0_added.outputs, std::string{"a"}); KwargNodeAddedResult n1_added = g.add_node( - /*inputs=*/std::map>{}, + /*inputs=*/std::map>{}, /*outputs=*/std::set{ "b", "c", @@ -56,10 +54,11 @@ TEST_SUITE(FF_TEST_SUITE) { return g; }(); - std::function( - std::set const &)> - slot_ordering = [](std::set const &slots) - -> std::vector { return reversed(sorted(slots)); }; + std::function(std::set const &)> + slot_ordering = + [](std::set const &slots) -> std::vector { + return reversed(sorted(slots)); + }; DataflowGraphView result = dataflow_graph_from_kwarg_dataflow_graph(input, slot_ordering); diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc index 42be45f161..30e42914bc 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/get_kwarg_dataflow_graph_subgraph.cc @@ -52,8 +52,8 @@ TEST_SUITE(FF_TEST_SUITE) { KwargDataflowGraphView g = view_from_kwarg_dataflow_graph_data(g_data); SUBCASE("node set is contains all graph nodes") { - KwargDataflowGraphView result = get_kwarg_dataflow_graph_subgraph( - g, std::set{n1, n2, n3, n4, n5}); + KwargDataflowGraphView result = + get_kwarg_dataflow_graph_subgraph(g, std::set{n1, n2, n3, n4, n5}); KwargDataflowGraphData result_data = get_kwarg_dataflow_graph_data(result); diff --git a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc index 8d7342e120..d08310a588 100644 --- a/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc +++ b/lib/utils/test/src/utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.cc @@ -1,9 +1,9 @@ #include "utils/graph/kwarg_dataflow_graph/algorithms/view_from_kwarg_dataflow_graph_data.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_edges.h" #include "utils/graph/kwarg_dataflow_graph/algorithms/get_all_kwarg_dataflow_outputs.h" #include "utils/graph/node/algorithms.h" #include -#include "test/utils/doctest/fmt/set.h" using namespace ::FlexFlow; @@ -75,16 +75,14 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_all_kwarg_dataflow_edges") { std::set>> result = get_all_kwarg_dataflow_edges(g); - std::set>> correct = - all_edges; + std::set>> correct = all_edges; ASSERT(result == correct); } SUBCASE("get_all_kwarg_dataflow_outputs") { std::set>> result = get_all_kwarg_dataflow_outputs(g); - std::set>> correct = - all_outputs; + std::set>> correct = all_outputs; ASSERT(result == correct); } } diff --git a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc index d24aed9903..87b490deca 100644 --- a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc +++ b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/get_labelled_kwarg_dataflow_graph_subgraph.cc @@ -120,8 +120,8 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set includes only some graph nodes") { LabelledKwargDataflowGraphView result = - get_labelled_kwarg_dataflow_graph_subgraph( - input, std::set{n2, n3}); + get_labelled_kwarg_dataflow_graph_subgraph(input, + std::set{n2, n3}); LabelledKwargDataflowGraphData result_data = get_labelled_kwarg_dataflow_graph_data(result); @@ -148,8 +148,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("node set includes no graph nodes") { LabelledKwargDataflowGraphView result = - get_labelled_kwarg_dataflow_graph_subgraph( - input, std::set{}); + get_labelled_kwarg_dataflow_graph_subgraph(input, std::set{}); LabelledKwargDataflowGraphData result_data = get_labelled_kwarg_dataflow_graph_data(result); diff --git a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc index 98b61d8542..a189d92556 100644 --- a/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc +++ b/lib/utils/test/src/utils/graph/labelled_kwarg_dataflow_graph/algorithms/kwarg_dataflow_graph_view_with_labelling.cc @@ -80,13 +80,12 @@ TEST_SUITE(FF_TEST_SUITE) { std::string n3_1_label = "c"; std::string n5_0_label = "d"; - std::map, std::string> value_labelling = - { - {n1_0, n1_0_label}, - {n2_3, n2_3_label}, - {n3_1, n3_1_label}, - {n5_0, n5_0_label}, - }; + std::map, std::string> value_labelling = { + {n1_0, n1_0_label}, + {n2_3, n2_3_label}, + {n3_1, n3_1_label}, + {n5_0, n5_0_label}, + }; LabelledKwargDataflowGraphView result = kwarg_dataflow_graph_view_with_labelling( diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc index 6c16277914..61557d23d6 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/add_edges.cc @@ -26,8 +26,7 @@ TEST_SUITE(FF_TEST_SUITE) { auto dst = [&](MultiDiEdge const &e) { return g.get_multidiedge_dst(e); }; SUBCASE("adds only those edges") { - std::set added = - g.query_edges(multidiedge_query_all()); + std::set added = g.query_edges(multidiedge_query_all()); std::set returned = set_of(result); CHECK(returned == added); } diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc index 4380c1c76e..637e57d6f9 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_incoming_edges.cc @@ -40,8 +40,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_incoming_edges(MultiDiGraphView, std::set)") { std::set ns = {n.at(0), n.at(2)}; - std::map> result = - get_incoming_edges(g, ns); + std::map> result = get_incoming_edges(g, ns); std::map> correct = { {n.at(0), {edges.at(0), edges.at(3), edges.at(4)}}, {n.at(2), {}}}; diff --git a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc index 09ba24f997..32942ab07f 100644 --- a/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc +++ b/lib/utils/test/src/utils/graph/multidigraph/algorithms/get_outgoing_edges.cc @@ -43,8 +43,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("get_outgoing_edges(MultiDiGraphView, std::set)") { std::set ns = {n.at(0), n.at(1)}; - std::map> result = - get_outgoing_edges(g, ns); + std::map> result = get_outgoing_edges(g, ns); std::map> correct = { {n.at(0), {edges.at(0), edges.at(1), edges.at(2), edges.at(3)}}, diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc index 6a5c9fbc2a..b5196afbc7 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_graph_inputs.cc @@ -15,8 +15,7 @@ TEST_SUITE(FF_TEST_SUITE) { NodeAddedResult n0_added = g.add_node({}, 1_n); - std::set result = - get_open_dataflow_graph_inputs(g); + std::set result = get_open_dataflow_graph_inputs(g); std::set correct = {i0, i1}; CHECK(result == correct); diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc index e29788c138..816ef2ae6b 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/get_subgraph.cc @@ -1,9 +1,9 @@ #include "utils/graph/open_dataflow_graph/algorithms/get_subgraph.h" +#include "test/utils/doctest/fmt/set.h" #include "utils/bidict/algorithms/left_entries.h" #include "utils/containers/contains.h" #include "utils/containers/get_only.h" #include "utils/graph/instances/unordered_set_dataflow_graph.h" -#include "test/utils/doctest/fmt/set.h" #include "utils/graph/node/algorithms.h" #include "utils/graph/open_dataflow_graph/algorithms/get_open_dataflow_values.h" #include "utils/graph/open_dataflow_graph/open_dataflow_graph.h" @@ -59,9 +59,8 @@ TEST_SUITE(FF_TEST_SUITE) { } } - TEST_CASE( - "get_subgraph_data(OpenDataflowGraphView, std::set, " - "bidict)") { + TEST_CASE("get_subgraph_data(OpenDataflowGraphView, std::set, " + "bidict)") { SUBCASE("2-node graph without inputs") { OpenDataflowGraph graph = OpenDataflowGraph::create(); diff --git a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc index 127f03c0d9..74492cec49 100644 --- a/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc +++ b/lib/utils/test/src/utils/graph/open_dataflow_graph/algorithms/permute_node_ids.cc @@ -119,8 +119,7 @@ TEST_SUITE(FF_TEST_SUITE) { dataflow_edge_query_for_edge( DataflowEdge{n0_output, DataflowInput{n1, 1_n}}), }; - std::set result_nodes = - result.query_edges(query); + std::set result_nodes = result.query_edges(query); std::set correct = {}; CHECK(result_nodes == correct); } @@ -139,8 +138,7 @@ TEST_SUITE(FF_TEST_SUITE) { dataflow_edge_query_for_edge(new_standard_edge), }; - std::set result_nodes = - result.query_edges(query); + std::set result_nodes = result.query_edges(query); std::set correct = { OpenDataflowEdge{new_standard_edge}, OpenDataflowEdge{new_input_edge}, @@ -156,8 +154,7 @@ TEST_SUITE(FF_TEST_SUITE) { DataflowOutputQuery query = dataflow_output_query_for_output(old_output); - std::set result_outputs = - result.query_outputs(query); + std::set result_outputs = result.query_outputs(query); std::set correct = {}; @@ -169,8 +166,7 @@ TEST_SUITE(FF_TEST_SUITE) { DataflowOutputQuery query = dataflow_output_query_for_output(new_output); - std::set result_outputs = - result.query_outputs(query); + std::set result_outputs = result.query_outputs(query); std::set correct = {new_output}; diff --git a/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc b/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc index e17079b770..23897b06a8 100644 --- a/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc +++ b/lib/utils/test/src/utils/graph/open_kwarg_dataflow_graph/algorithms/view_as_closed_kwarg_dataflow_graph_by_materializing_inputs.cc @@ -98,8 +98,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult> n1_added = g.add_node( /*inputs=*/ - std::map, - KwargDataflowOutput>>{ + std::map, KwargDataflowOutput>>{ { 1, input1, @@ -122,8 +121,7 @@ TEST_SUITE(FF_TEST_SUITE) { KwargNodeAddedResult> n2_added = g.add_node( /*inputs=*/ - std::map, - KwargDataflowOutput>>{ + std::map, KwargDataflowOutput>>{ { 4, input2, diff --git a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc index 1589f682dc..033d90d405 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/left_associative_binary_sp_tree_from_nary.cc @@ -97,8 +97,7 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(is_binary_sp_tree_left_associative(result)); std::multiset result_nodes = get_leaves(result); - std::multiset correct_nodes = { - n1, n2, n3, n3, n5, n6, n4, n5}; + std::multiset correct_nodes = {n1, n2, n3, n3, n5, n6, n4, n5}; CHECK(result_nodes == correct_nodes); } diff --git a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc index bb4aa5ec54..7a575a949c 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/binary_sp_decomposition_tree/right_associative_binary_sp_tree_from_nary.cc @@ -95,8 +95,7 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(is_binary_sp_tree_right_associative(result)); std::multiset result_nodes = get_nodes(input); - std::multiset correct_nodes = { - n1, n2, n3, n3, n5, n6, n4, n5}; + std::multiset correct_nodes = {n1, n2, n3, n3, n5, n6, n4, n5}; CHECK(result_nodes == correct_nodes); } diff --git a/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc b/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc index 4242131d2e..0ab6fa4ccc 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/parallel_reduction.cc @@ -178,7 +178,7 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge e = get_directed_edge(g, reduction_e1); new_edge_counts.at(e) = positive_int{ - new_edge_counts.at(e).int_from_positive_int() - 1, + new_edge_counts.at(e).int_from_positive_int() - 1, }; return new_edge_counts; }(); diff --git a/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc b/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc index d6b06d28d4..f3b3007650 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/series_reduction.cc @@ -213,8 +213,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("nodes") { std::set result_nodes = get_nodes(g); - std::set correct_nodes = - set_minus(set_of(n), {n.at(4)}); + std::set correct_nodes = set_minus(set_of(n), {n.at(4)}); CHECK(result_nodes == correct_nodes); } @@ -367,8 +366,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("nodes") { std::set result_nodes = get_nodes(g); - std::set correct_nodes = - set_minus(set_of(n), {n.at(4), n.at(3)}); + std::set correct_nodes = set_minus(set_of(n), {n.at(4), n.at(3)}); CHECK(result_nodes == correct_nodes); } diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc index 002e77828b..19c95137d0 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/escribano_algo.cc @@ -47,13 +47,12 @@ TEST_SUITE(FF_TEST_SUITE) { CHECK(get_outgoing_edges(g, n.at(0)).size() == 2); CHECK(node_types.size() == 6); - CHECK(values(node_types) == - std::multiset{NodeRole::PURE, - NodeRole::PURE, - NodeRole::PURE, - NodeRole::PURE, - NodeRole::DUMMY, - NodeRole::DUMMY}); + CHECK(values(node_types) == std::multiset{NodeRole::PURE, + NodeRole::PURE, + NodeRole::PURE, + NodeRole::PURE, + NodeRole::DUMMY, + NodeRole::DUMMY}); DiGraph restored = contract_out_nodes_of_given_role(result, NodeRole::DUMMY, node_types); @@ -109,15 +108,13 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(5), 1_n}, }; SUBCASE("n.at(4)'s component") { - std::set correct = { - n.at(0), n.at(1), n.at(4), n.at(5)}; + std::set correct = {n.at(0), n.at(1), n.at(4), n.at(5)}; std::set result = get_component(g, n.at(4), depth_map, node_roles); CHECK(correct == result); } SUBCASE("n.at(5)'s component") { - std::set correct = { - n.at(0), n.at(1), n.at(4), n.at(5)}; + std::set correct = {n.at(0), n.at(1), n.at(4), n.at(5)}; std::set result = get_component(g, n.at(5), depth_map, node_roles); CHECK(correct == result); @@ -134,31 +131,28 @@ TEST_SUITE(FF_TEST_SUITE) { DirectedEdge{n.at(3), n.at(4)}, DirectedEdge{n.at(4), n.at(5)}, DirectedEdge{n.at(4), n.at(6)}}); - std::map node_roles = { - {n.at(0), NodeRole::PURE}, - {n.at(1), NodeRole::SYNC}, - {n.at(2), NodeRole::PURE}, - {n.at(3), NodeRole::PURE}, - {n.at(4), NodeRole::SYNC}, - {n.at(5), NodeRole::PURE}, - {n.at(6), NodeRole::PURE}}; + std::map node_roles = {{n.at(0), NodeRole::PURE}, + {n.at(1), NodeRole::SYNC}, + {n.at(2), NodeRole::PURE}, + {n.at(3), NodeRole::PURE}, + {n.at(4), NodeRole::SYNC}, + {n.at(5), NodeRole::PURE}, + {n.at(6), NodeRole::PURE}}; std::map depth_map = {{n.at(0), 0_n}, - {n.at(2), 1_n}, - {n.at(3), 1_n}, - {n.at(5), 2_n}, - {n.at(6), 2_n}}; + {n.at(2), 1_n}, + {n.at(3), 1_n}, + {n.at(5), 2_n}, + {n.at(6), 2_n}}; SUBCASE("n.at(5)'s component") { - std::set correct = { - n.at(2), n.at(3), n.at(5), n.at(6)}; + std::set correct = {n.at(2), n.at(3), n.at(5), n.at(6)}; std::set result = get_component(g, n.at(5), depth_map, node_roles); CHECK(correct == result); } SUBCASE("n.at(6)'s component") { - std::set correct = { - n.at(2), n.at(3), n.at(5), n.at(6)}; + std::set correct = {n.at(2), n.at(3), n.at(5), n.at(6)}; std::set result = get_component(g, n.at(6), depth_map, node_roles); CHECK(correct == result); @@ -194,12 +188,12 @@ TEST_SUITE(FF_TEST_SUITE) { }; std::map depth_map = {{n.at(0), 0_n}, - {n.at(2), 1_n}, - {n.at(3), 1_n}, - {n.at(4), 1_n}, - {n.at(7), 2_n}, - {n.at(8), 2_n}, - {n.at(9), 2_n}}; + {n.at(2), 1_n}, + {n.at(3), 1_n}, + {n.at(4), 1_n}, + {n.at(7), 2_n}, + {n.at(8), 2_n}, + {n.at(9), 2_n}}; SUBCASE("n.at(7)'s component") { std::set correct = {n.at(2), n.at(7), n.at(8)}; std::set result = diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc index 5679d53f08..b58a037389 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/naive_stratum_sync.cc @@ -44,10 +44,10 @@ TEST_SUITE(FF_TEST_SUITE) { }); std::map cost_map = {{n.at(0), 1.0f}, - {n.at(1), 2.0f}, - {n.at(2), 3.0f}, - {n.at(3), 4.0f}, - {n.at(4), 5.0f}}; + {n.at(1), 2.0f}, + {n.at(2), 3.0f}, + {n.at(3), 4.0f}, + {n.at(4), 5.0f}}; SeriesParallelDecomposition sp = naive_stratum_sync_sp_ization(g); @@ -74,10 +74,10 @@ TEST_SUITE(FF_TEST_SUITE) { }); std::map cost_map = {{n.at(0), 2.0f}, - {n.at(1), 3.0f}, - {n.at(2), 5.0f}, - {n.at(3), 7.0f}, - {n.at(4), 11.0f}}; + {n.at(1), 3.0f}, + {n.at(2), 5.0f}, + {n.at(3), 7.0f}, + {n.at(4), 11.0f}}; SeriesParallelDecomposition sp = naive_stratum_sync_sp_ization(g); @@ -106,11 +106,11 @@ TEST_SUITE(FF_TEST_SUITE) { }); std::map cost_map = {{n.at(0), 1.0f}, - {n.at(1), 1.0f}, - {n.at(2), 10.0f}, - {n.at(3), 1.0f}, - {n.at(4), 1.0f}, - {n.at(5), 1.0f}}; + {n.at(1), 1.0f}, + {n.at(2), 10.0f}, + {n.at(3), 1.0f}, + {n.at(4), 1.0f}, + {n.at(5), 1.0f}}; SeriesParallelDecomposition sp = naive_stratum_sync_sp_ization(g); diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc index df7e980db1..a834877c32 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/node_role.cc @@ -36,12 +36,11 @@ TEST_SUITE(FF_TEST_SUITE) { DiGraph result = contract_out_nodes_of_given_role(g, NodeRole::DUMMY, node_roles); - CHECK(get_nodes(result) == - std::set{n.at(0), n.at(3), n.at(4)}); + CHECK(get_nodes(result) == std::set{n.at(0), n.at(3), n.at(4)}); CHECK(get_edges(result) == std::set{DirectedEdge{n.at(0), n.at(4)}, - DirectedEdge{n.at(0), n.at(3)}, - DirectedEdge{n.at(3), n.at(4)}}); + DirectedEdge{n.at(0), n.at(3)}, + DirectedEdge{n.at(3), n.at(4)}}); } SUBCASE("graph is unchanged when no node has the target role") { diff --git a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc index 6ebbe5ab42..5014981bc4 100644 --- a/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc +++ b/lib/utils/test/src/utils/graph/series_parallel/sp_ization/work_duplicating_sp_ization.cc @@ -45,10 +45,9 @@ static std::pair> } } - std::map cost_map = - generate_map(get_nodes(g), [](Node const &) { - return static_cast(*rc::gen::inRange(1, 101)); - }); + std::map cost_map = generate_map(get_nodes(g), [](Node const &) { + return static_cast(*rc::gen::inRange(1, 101)); + }); return {g, cost_map}; } diff --git a/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc b/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc index bb03895540..56fce6ca98 100644 --- a/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc +++ b/lib/utils/test/src/utils/graph/undirected/algorithms/get_connected_components.cc @@ -20,8 +20,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(1)}, {n.at(2)}, }; - std::set> result = - get_connected_components(g); + std::set> result = get_connected_components(g); CHECK(correct == result); } @@ -39,8 +38,7 @@ TEST_SUITE(FF_TEST_SUITE) { std::set> correct = { {n.at(0), n.at(1), n.at(2), n.at(3)}, }; - std::set> result = - get_connected_components(g); + std::set> result = get_connected_components(g); CHECK(correct == result); } @@ -57,8 +55,7 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(0), n.at(1), n.at(2)}, {n.at(3)}, }; - std::set> result = - get_connected_components(g); + std::set> result = get_connected_components(g); CHECK(correct == result); } @@ -78,16 +75,14 @@ TEST_SUITE(FF_TEST_SUITE) { {n.at(3), n.at(4)}, {n.at(5)}, }; - std::set> result = - get_connected_components(g); + std::set> result = get_connected_components(g); CHECK(correct == result); } SUBCASE("empty graph") { std::set> correct = {}; - std::set> result = - get_connected_components(g); + std::set> result = get_connected_components(g); CHECK(correct == result); } diff --git a/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc b/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc index df4d28dfd5..f42b08101c 100644 --- a/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc +++ b/lib/utils/test/src/utils/graph/undirected/undirected_graph.cc @@ -75,10 +75,9 @@ TEST_SUITE(FF_TEST_SUITE) { }; std::set result = g.query_edges(query); - std::set correct = - std::set{ - e.at(0), - }; + std::set correct = std::set{ + e.at(0), + }; CHECK(result == correct); } @@ -107,11 +106,9 @@ TEST_SUITE(FF_TEST_SUITE) { g.remove_edge(e.at(0)); CHECK(g.query_edges(undirected_edge_query_all()) == - std::set{ - e.at(1), e.at(2), e.at(3), e.at(4)}); + std::set{e.at(1), e.at(2), e.at(3), e.at(4)}); CHECK(g.query_nodes(node_query_all()) == - std::set{ - n.at(0), n.at(1), n.at(2), n.at(3), n.at(4)}); + std::set{n.at(0), n.at(1), n.at(2), n.at(3), n.at(4)}); g.remove_edge(e.at(1)); g.remove_edge(e.at(3)); diff --git a/lib/utils/test/src/utils/graph/views/views.cc b/lib/utils/test/src/utils/graph/views/views.cc index d363264cc1..1950fadafb 100644 --- a/lib/utils/test/src/utils/graph/views/views.cc +++ b/lib/utils/test/src/utils/graph/views/views.cc @@ -1,6 +1,6 @@ #include "utils/graph/views/views.h" -#include "utils/containers/set_union.h" #include "utils/containers/set_of.h" +#include "utils/containers/set_union.h" #include "utils/fmt/map.h" #include "utils/fmt/set.h" #include "utils/graph/algorithms.h" @@ -138,14 +138,13 @@ TEST_SUITE(FF_TEST_SUITE) { } SUBCASE("get_edges") { - std::set expected = { - DirectedEdge{n.at(0), n.at(0)}, - DirectedEdge{n.at(0), n.at(1)}, - DirectedEdge{n.at(1), n.at(0)}, - DirectedEdge{n.at(1), n.at(2)}, - DirectedEdge{n.at(2), n.at(1)}, - DirectedEdge{n.at(2), n.at(0)}, - DirectedEdge{n.at(0), n.at(2)}}; + std::set expected = {DirectedEdge{n.at(0), n.at(0)}, + DirectedEdge{n.at(0), n.at(1)}, + DirectedEdge{n.at(1), n.at(0)}, + DirectedEdge{n.at(1), n.at(2)}, + DirectedEdge{n.at(2), n.at(1)}, + DirectedEdge{n.at(2), n.at(0)}, + DirectedEdge{n.at(0), n.at(2)}}; std::set result = get_edges(view); diff --git a/lib/utils/test/src/utils/one_to_many/one_to_many.cc b/lib/utils/test/src/utils/one_to_many/one_to_many.cc index f6ffabc698..4a1bfe7efe 100644 --- a/lib/utils/test/src/utils/one_to_many/one_to_many.cc +++ b/lib/utils/test/src/utils/one_to_many/one_to_many.cc @@ -1,11 +1,9 @@ #include "utils/one_to_many/one_to_many.h" #include "test/utils/doctest/fmt/multiset.h" -#include "test/utils/doctest/fmt/set.h" +#include "test/utils/doctest/fmt/pair.h" #include "test/utils/doctest/fmt/set.h" #include "utils/containers/multiset_of.h" #include "utils/one_to_many/one_to_many_from_l_to_r_mapping.h" -#include "test/utils/doctest/fmt/pair.h" -#include "test/utils/doctest/fmt/set.h" #include using namespace ::FlexFlow; diff --git a/lib/utils/test/src/utils/orthotope/dim_coord.cc b/lib/utils/test/src/utils/orthotope/dim_coord.cc index 0766f1067a..0ed73ab503 100644 --- a/lib/utils/test/src/utils/orthotope/dim_coord.cc +++ b/lib/utils/test/src/utils/orthotope/dim_coord.cc @@ -104,8 +104,7 @@ TEST_SUITE(FF_TEST_SUITE) { {7, 2_p}, }}; - std::set> result = - get_coords_in_dim_domain(dim_domain); + std::set> result = get_coords_in_dim_domain(dim_domain); std::set> correct = { DimCoord{{ @@ -125,8 +124,7 @@ TEST_SUITE(FF_TEST_SUITE) { {2, 3_p}, }}; - std::set> result = - get_coords_in_dim_domain(dim_domain); + std::set> result = get_coords_in_dim_domain(dim_domain); auto mk_dim_coord = [](nonnegative_int dim7, nonnegative_int dim2) { return DimCoord{{ @@ -153,8 +151,7 @@ TEST_SUITE(FF_TEST_SUITE) { {2, 3_p}, }}; - std::set> result = - get_coords_in_dim_domain(dim_domain); + std::set> result = get_coords_in_dim_domain(dim_domain); auto mk_dim_coord = [](nonnegative_int dim7, nonnegative_int dim2) { return DimCoord{{ @@ -175,8 +172,7 @@ TEST_SUITE(FF_TEST_SUITE) { SUBCASE("zero-dimensional dim domain") { DimDomain dim_domain = DimDomain{{}}; - std::set> result = - get_coords_in_dim_domain(dim_domain); + std::set> result = get_coords_in_dim_domain(dim_domain); std::set> correct = { DimCoord{{}}, From 7b806a9a3e6fdebcfffdf574a866799c024c5694 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Fri, 26 Jun 2026 19:05:28 -0700 Subject: [PATCH 32/35] Remove dead code discovered in PR review --- .../src/task-spec/dynamic_graph/shard_expansion.cc | 6 ------ 1 file changed, 6 deletions(-) diff --git a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc index c566a30f8d..98df030438 100644 --- a/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc +++ b/lib/task-spec/src/task-spec/dynamic_graph/shard_expansion.cc @@ -105,12 +105,6 @@ static DynamicNodeInvocationShardingInfo invocation_sharding_info_for_binding( }; }; - DynamicNodeAttrs expanded_node_attrs = [&]() { - DynamicNodeAttrs result = i.node_attrs; - result.device_ids = nonempty_set{device_id}; - return result; - }(); - DynamicNodeInvocationShardingInfo result = DynamicNodeInvocationShardingInfo{ /*device_coord=*/nonempty_set{device_id}, /*value_sharding=*/ From 3629063da017f021367408a40bceeee5f6bacaf4 Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Tue, 30 Jun 2026 13:11:41 -0700 Subject: [PATCH 33/35] Fix incorrect proj flake uri --- flake.lock | 4 ++-- flake.nix | 3 +-- 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/flake.lock b/flake.lock index c06618c181..233e2265e2 100644 --- a/flake.lock +++ b/flake.lock @@ -73,11 +73,11 @@ "rev": "f2bdddab299b98fb67612e3c04136949b2339d74", "revCount": 161, "type": "git", - "url": "file:///home/lockshaw/x/ff/proj/proj" + "url": "https://git.sr.ht/~lockshaw/proj" }, "original": { "type": "git", - "url": "file:///home/lockshaw/x/ff/proj/proj" + "url": "https://git.sr.ht/~lockshaw/proj" } }, "python38-nixpkgs": { diff --git a/flake.nix b/flake.nix index ccb0402498..ad71cbefb4 100644 --- a/flake.nix +++ b/flake.nix @@ -18,8 +18,7 @@ flake-utils.url = "github:numtide/flake-utils"; proj-repo = { - # url = "git+https://git.sr.ht/~lockshaw/proj"; - url = "git+file:///home/lockshaw/x/ff/proj/proj"; + url = "git+https://git.sr.ht/~lockshaw/proj"; inputs.nixpkgs.follows = "nixpkgs"; inputs.flake-utils.follows = "flake-utils"; }; From 2488ff43f2788b2f90d54c1bf8164f4ce09223b0 Mon Sep 17 00:00:00 2001 From: Elliott Slaughter Date: Tue, 30 Jun 2026 18:28:02 -0700 Subject: [PATCH 34/35] Properly abstract and implement collective interface. --- .../src/realm-execution/pcg_instance.cc | 162 ++++++++++++------ 1 file changed, 108 insertions(+), 54 deletions(-) diff --git a/lib/realm-execution/src/realm-execution/pcg_instance.cc b/lib/realm-execution/src/realm-execution/pcg_instance.cc index 1814af9cb0..1d102f8cca 100644 --- a/lib/realm-execution/src/realm-execution/pcg_instance.cc +++ b/lib/realm-execution/src/realm-execution/pcg_instance.cc @@ -25,8 +25,10 @@ #include "utils/containers/transform.h" #include "utils/containers/try_at.h" #include "utils/containers/values.h" +#include "utils/containers/vector_of.h" #include "utils/graph/digraph/algorithms/get_topological_ordering.h" #include "utils/optional.h" +#include namespace FlexFlow { @@ -171,6 +173,87 @@ PCGInstance create_pcg_instance( }; } +static Realm::Event + issue_p2p_copy(RealmContext &ctx, + DynamicValueAttrs const &input, + DynamicValueAttrs const &output, + TensorInstanceBacking const &tensor_instance_backing, + Realm::Event precondition) { + Realm::RegionInstance src_inst = + tensor_instance_backing.backing.at(input).first; + Realm::RegionInstance dst_inst = + tensor_instance_backing.backing.at(output).first; + return ctx.issue_copy(assert_unwrap(input.parallel_tensor_shape), + src_inst, + assert_unwrap(output.parallel_tensor_shape), + dst_inst, + Realm::ProfilingRequestSet{}, + precondition); +} + +static Realm::Event + issue_p2p_reduction(RealmContext &ctx, + DynamicValueAttrs const &input, + DynamicValueAttrs const &output, + TensorInstanceBacking const &tensor_instance_backing, + redop_id_t redop_id, + bool is_fold, + bool exclusive, + Realm::Event precondition) { + Realm::RegionInstance src_inst = + tensor_instance_backing.backing.at(input).first; + Realm::RegionInstance dst_inst = + tensor_instance_backing.backing.at(output).first; + return ctx.issue_reduction(assert_unwrap(input.parallel_tensor_shape), + src_inst, + assert_unwrap(output.parallel_tensor_shape), + dst_inst, + redop_id, + is_fold, + exclusive, + Realm::ProfilingRequestSet{}, + precondition); +} + +static Realm::Event issue_collective_broadcast( + RealmContext &ctx, + DynamicValueAttrs const &input, + std::vector const &outputs, + TensorInstanceBacking const &tensor_instance_backing, + Realm::Event precondition) { + // For now we just implement this as the naive set of N p2p copies. + std::vector result = + transform(outputs, [&](DynamicValueAttrs const &output) { + return issue_p2p_copy( + ctx, input, output, tensor_instance_backing, precondition); + }); + return Realm::Event::merge_events(result); +} + +static Realm::Event issue_collective_reduction( + RealmContext &ctx, + std::vector const &inputs, + DynamicValueAttrs const &output, + TensorInstanceBacking const &tensor_instance_backing, + redop_id_t redop_id, + Realm::Event precondition) { + // For now we just implement this as a naive set of N p2p reductions. Because + // we're launching them in parallel they cannot be exclusive (i.e., they need + // to use per-element atomics to update the output tensor) + std::vector result = + transform(inputs, [&](DynamicValueAttrs const &input) { + return issue_p2p_reduction(ctx, + input, + output, + tensor_instance_backing, + redop_id, + /*is_fold*/ false, + /*exclusive*/ false, + precondition); + }); + return Realm::Event::merge_events(result); +} + /** * \brief Spawn the Realm operations (tasks, copies, etc.) for a given \ref * DynamicNodeInvocation, given the specified dependencies, instances, etc. Note @@ -212,59 +295,26 @@ static Realm::Event spawn_dynamic_node_invocation( auto issue_copy = [&]() { DynamicValueAttrs const &input = get_only(invocation.inputs).second; DynamicValueAttrs const &output = get_only(invocation.outputs).second; - Realm::RegionInstance src_inst = - tensor_instance_backing.backing.at(input).first; - Realm::RegionInstance dst_inst = - tensor_instance_backing.backing.at(output).first; - return ctx.issue_copy(assert_unwrap(input.parallel_tensor_shape), - src_inst, - assert_unwrap(output.parallel_tensor_shape), - dst_inst, - Realm::ProfilingRequestSet{}, - precondition); + return issue_p2p_copy( + ctx, input, output, tensor_instance_backing, precondition); }; - auto issue_replicate_bwd = [&]() { - DynamicValueAttrs output_grad = get_only(values( - filter_keys(invocation.inputs, [](DynamicTensorSlot const &s) -> bool { - return s.slot_tensor_role == - DynamicTensorRole{FwbTensorType::GRADIENT}; - }))); - - DynamicValueAttrs input_grad = get_only(values(invocation.outputs)); - - Realm::RegionInstance dst_inst = - tensor_instance_backing.backing.at(input_grad).first; + auto issue_replicate = [&]() { + DynamicValueAttrs const &input = get_only(invocation.inputs).second; + std::vector outputs = + vector_of(values(invocation.outputs)); + return issue_collective_broadcast( + ctx, input, outputs, tensor_instance_backing, precondition); + }; + auto issue_reduction = [&]() { + std::vector inputs = + vector_of(values(invocation.inputs)); + DynamicValueAttrs const &output = get_only(invocation.outputs).second; redop_id_t redop_id = get_sum_redop_id_for_data_type( - assert_unwrap(output_grad.parallel_tensor_shape).data_type); - - // chain reductions sequentially to avoid write races on dst - Realm::Event result = precondition; - for (auto const &[p, d] : assert_unwrap(output_grad.mapping).raw) { - DynamicValueAttrs replica_key = output_grad; - replica_key.mapping = ParallelTensorMapping{ - bidict{ - {p, d}, - }, - }; - replica_key.shard_coord = p; - - Realm::RegionInstance src_inst = - tensor_instance_backing.backing.at(replica_key).first; - - result = ctx.issue_reduction( - /*src_shape=*/assert_unwrap(output_grad.parallel_tensor_shape), - /*src_inst=*/src_inst, - /*dst_shape=*/assert_unwrap(input_grad.parallel_tensor_shape), - /*dst_inst=*/dst_inst, - /*redop_id=*/redop_id, - /*is_fold=*/false, - /*exlusive=*/false, - /*requests=*/Realm::ProfilingRequestSet{}, - /*wait_on=*/result); - } - return result; + assert_unwrap(output.parallel_tensor_shape).data_type); + return issue_collective_reduction( + ctx, inputs, output, tensor_instance_backing, redop_id, precondition); }; TrainingOperationAttrs op_attrs = @@ -275,12 +325,16 @@ static Realm::Event spawn_dynamic_node_invocation( [&](InputAttrs const &) { return Realm::Event::NO_EVENT; }, [&](WeightAttrs const &) { return Realm::Event::NO_EVENT; }, [&](ReplicateAttrs const &) { - if (invocation.node_attrs.task_type.has_value() && - invocation.node_attrs.task_type.value() == - DynamicTaskType::BWD) { - return issue_replicate_bwd(); + DynamicTaskType task_type = + assert_unwrap(invocation.node_attrs.task_type); + switch (task_type) { + case DynamicTaskType::FWD: + return issue_replicate(); + case DynamicTaskType::BWD: + return issue_reduction(); + default: + PANIC("Unhandled replicate task type ", task_type); } - return issue_copy(); // forward }, [&](auto const &) { return spawn_task(); }, }); From 880a7dc8a0969739e96423d3fa439709ecd149fc Mon Sep 17 00:00:00 2001 From: Colin Unger Date: Wed, 1 Jul 2026 13:36:43 -0700 Subject: [PATCH 35/35] Minor fixes to pass cpu-ci checks --- .../test/src/realm-execution/test_e2e.cc | 302 +++++++++++++++ .../src/realm-execution/test_op_replicate.cc | 348 ------------------ .../dynamic_value_attrs.dtg.toml | 2 +- .../include/task-spec/dynamic_graph/index.dox | 2 +- .../transform_binary_relation.h | 4 +- lib/utils/include/utils/fmt/unordered_map.h | 10 +- .../transform_binary_relation.cc | 15 + lib/utils/src/utils/fmt/unordered_map.cc | 2 +- 8 files changed, 327 insertions(+), 358 deletions(-) delete mode 100644 lib/realm-execution/test/src/realm-execution/test_op_replicate.cc create mode 100644 lib/utils/src/utils/binary_relation/transform_binary_relation.cc diff --git a/lib/realm-execution/test/src/realm-execution/test_e2e.cc b/lib/realm-execution/test/src/realm-execution/test_e2e.cc index 5f681a6a3d..9ba4886b4b 100644 --- a/lib/realm-execution/test/src/realm-execution/test_e2e.cc +++ b/lib/realm-execution/test/src/realm-execution/test_e2e.cc @@ -4,6 +4,7 @@ #include "kernels/copy_tensor_accessor.h" #include "kernels/format_accessor_contents.h" #include "kernels/tensor_accessor_reductions.h" +#include "op-attrs/ops/element_unary.h" #include "op-attrs/parallel_tensor_shape.h" #include "op-attrs/tensor_shape.dtg.h" #include "op-attrs/tensor_slot_name.dtg.h" @@ -29,6 +30,14 @@ namespace test { using namespace ::FlexFlow; namespace Realm = ::FlexFlow::Realm; +template +static ParallelLayerAttrs make_layer_attrs(T const &op_attrs) { + return ParallelLayerAttrs{ + /*op_attrs=*/PCGOperatorAttrs{op_attrs}, + /*name=*/std::nullopt, + }; +}; + static bool did_loss_decrease(GenericTensorAccessorR const &first_epoch, GenericTensorAccessorR const &last_epoch, Allocator &allocator) { @@ -214,6 +223,194 @@ static E2ETrainingConfig create_e2e_test_case() { }; } +MappedParallelComputationGraph + make_test_replicate_mpcg_for_device_type(DeviceType device_type) { + positive_int batch_size = 10_p; + positive_int data_dim = 16_p; + positive_int hidden_dim = 32_p; + positive_int output_dim = 1_p; + + TensorShape output_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + TensorShape label_tensor_shape = TensorShape{ + TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; + + ParallelComputationGraph pcg = empty_parallel_computation_graph(); + + TensorShape input_tensor_shape = + TensorShape{TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; + + ParallelLayerAddedResult inputs_layer = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input = + require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult inputs_layer_2 = + pcg_add_input_layer(pcg, input_tensor_shape); + parallel_tensor_guid_t t_input_2 = + require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); + + ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ + OperatorType::EW_ADD, + DataType::FLOAT, + false, + false, + }; + + ParallelLayerAddedResult add_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(add_attrs), + { + { + TensorSlotName::LHS_INPUT, + t_input, + }, + { + TensorSlotName::RHS_INPUT, + t_input_2, + }, + }, + /*weights=*/{}); + + parallel_tensor_guid_t t_add_1 = + require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); + + positive_int replicate_degree = 2_p; + ReplicateAttrs repl_attrs = ReplicateAttrs{replicate_degree}; + ParallelLayerAddedResult repl_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(repl_attrs), + { + { + TensorSlotName::INPUT, + t_add_1, + }, + }, + /*weight=*/{}); + + parallel_tensor_guid_t t_repl_1 = + require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); + + ParallelLayerAddedResult relu_operator_1 = + add_parallel_layer(pcg, + make_layer_attrs(make_relu_attrs()), + /*inputs=*/ + { + { + TensorSlotName::INPUT, + t_repl_1, + }, + }, + /*weights=*/{}); + + parallel_tensor_guid_t t_relu_1 = + require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); + + MachineSpaceCoordinate mc0{0_n, 0_n}; + MachineSpaceCoordinate mc1{0_n, 1_n}; + + ParallelTensorSpaceCoordinate tensor_coord0{ + /*sum_component=*/0_n, + /*discard_copy_component=*/0_n, + /*shard_component=*/FFOrdered{0_n}}; + ParallelTensorSpaceCoordinate tensor_coord1{ + /*sum_component=*/0_n, + /*discard_copy_component=*/1_n, + /*shard_component=*/FFOrdered{0_n}}; + + MappedParallelComputationGraph mpcg = + mapped_pcg_from_pcg_and_mapped_op_task_groups( + /*pcg=*/pcg, + /*mapped_op_task_groups=*/{ + { + inputs_layer.parallel_layer, + MappedOperatorTaskGroup{ + { + { + mc0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + inputs_layer_2.parallel_layer, + MappedOperatorTaskGroup{ + { + { + mc0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + add_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + mc0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::LHS_INPUT, tensor_coord0}, + {TensorSlotName::RHS_INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + }, + }, + }, + { + repl_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + mc0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + mc1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, + }, + { + relu_operator_1.parallel_layer, + MappedOperatorTaskGroup{ + { + { + mc0, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord0}, + {TensorSlotName::OUTPUT, tensor_coord0}, + }}, + }, + { + mc1, + OperatorAtomicTaskShardBinding{{ + {TensorSlotName::INPUT, tensor_coord1}, + {TensorSlotName::OUTPUT, tensor_coord1}, + }}, + }, + }, + }, + }, + }); + + return mpcg; +} + TEST_SUITE(FF_TEST_SUITE) { TEST_CASE("RealmBackend e2e Training (CPU Model Parallelism)") { std::vector fake_args = @@ -288,6 +485,58 @@ TEST_SUITE(FF_TEST_SUITE) { format_accessor_r_contents(last_epoch_loss))); }); } + + TEST_CASE("RealmBackend e2e Training Replicate Op (CPU Model Parallelism)") { + std::vector fake_args = + make_fake_realm_args(/*num_cpus=*/2_p, /*num_gpus=*/0_n); + int fake_argc = fake_args.size(); + char **fake_argv = fake_args.data(); + + RealmManager manager = RealmManager{&fake_argc, &fake_argv}; + ControllerTaskResult result = + manager.start_controller([](RealmContext &ctx) { + Allocator allocator = ctx.get_current_device_allocator(); + + MappedParallelComputationGraph mpcg = + make_test_replicate_mpcg_for_device_type(DeviceType::CPU); + + std::map input_tensors; + + OptimizerAttrs optimizer_attrs = OptimizerAttrs{ + SGDOptimizerAttrs{ + /*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001, + }, + }; + + DistributedFfHandle device_handle = create_distributed_ff_handle( + ctx, + /*workSpaceSize=*/1024 * 1024, + /*allowTensorOpMathConversion=*/true); + + PCGInstance pcg_instance = create_pcg_instance( + /*ctx=*/ctx, + /*mpcg=*/mpcg, + /*optimizer=*/optimizer_attrs, + /*loss=*/std::nullopt, + /*input_tensors=*/input_tensors, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::CPU); + + // begin training loop + int num_epochs = 1; + for (int i = 0; i < num_epochs; i++) { + perform_all_passes_for_pcg_instance( + /*instance=*/pcg_instance, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle); + } + }); + result.wait(); + } } TEST_SUITE(FF_CUDA_TEST_SUITE) { @@ -370,6 +619,59 @@ TEST_SUITE(FF_CUDA_TEST_SUITE) { result.wait(); //! [realm-execution example] } + + TEST_CASE("RealmBackend e2e Training Replicate Op (GPU Model Parallelism)") { + std::vector fake_args = + make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/2_n); + int fake_argc = fake_args.size(); + char **fake_argv = fake_args.data(); + + RealmManager manager = RealmManager{&fake_argc, &fake_argv}; + + ControllerTaskResult result = + manager.start_controller([](RealmContext &ctx) { + Allocator allocator = ctx.get_current_device_allocator(); + + MappedParallelComputationGraph mpcg = + make_test_replicate_mpcg_for_device_type(DeviceType::GPU); + + OptimizerAttrs optimizer_attrs = OptimizerAttrs{ + SGDOptimizerAttrs{ + /*lr=*/0.001, + /*momentum=*/0.9, + /*nesterov=*/false, + /*weight_decay=*/0.001, + }, + }; + + std::map input_tensors; + + DistributedFfHandle device_handle = create_distributed_ff_handle( + ctx, + /*workSpaceSize=*/1024 * 1024, + /*allowTensorOpMathConversion=*/true); + + PCGInstance pcg_instance = create_pcg_instance( + /*ctx=*/ctx, + /*mpcg=*/mpcg, + /*optimizer=*/optimizer_attrs, + /*loss=*/std::nullopt, + /*input_tensors=*/input_tensors, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle, + /*device_type=*/DeviceType::GPU); + + // begin training loop + int num_epochs = 1; + for (int i = 0; i < num_epochs; i++) { + perform_all_passes_for_pcg_instance( + /*instance=*/pcg_instance, + /*profiling_settings=*/ProfilingSettings{0, 0}, + /*device_handle=*/device_handle); + } + }); + result.wait(); + } } } // namespace test diff --git a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc b/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc deleted file mode 100644 index 955e7d73d2..0000000000 --- a/lib/realm-execution/test/src/realm-execution/test_op_replicate.cc +++ /dev/null @@ -1,348 +0,0 @@ -#include "internal/realm_test_utils.h" -#include "kernels/allocation.h" -#include "kernels/compare_tensor_accessors.h" -#include "kernels/copy_tensor_accessor.h" -#include "kernels/format_accessor_contents.h" -#include "kernels/tensor_accessor_reductions.h" -#include "op-attrs/operator_task_space_to_operator_task_space_mapping.h" -#include "op-attrs/ops/element_unary.h" -#include "op-attrs/ops/linear.h" -#include "op-attrs/ops/replicate.h" -#include "op-attrs/parallel_tensor_shape.h" -#include "op-attrs/tensor_shape.dtg.h" -#include "op-attrs/tensor_slot_name.dtg.h" -#include "pcg/device_type.dtg.h" -#include "pcg/machine_space_coordinate.dtg.h" -#include "pcg/mapped_parallel_computation_graph/mapped_parallel_computation_graph.h" -#include "pcg/mapped_parallel_computation_graph/operator_atomic_task_shard_binding.dtg.h" -#include "pcg/parallel_computation_graph/parallel_computation_graph.h" -#include "pcg/parallel_computation_graph/parallel_computation_graph_builder.h" -#include "pcg/parallel_computation_graph/parallel_layer_guid_t.dtg.h" -#include "pcg/parallel_computation_graph/parallel_tensor_guid_t.dtg.h" -#include "realm-execution/distributed_ff_handle.h" -#include "realm-execution/dynamic_tensor_accessor_from_instance.h" -#include "realm-execution/pcg_instance.h" -#include "realm-execution/realm_context.h" -#include "realm-execution/realm_manager.h" -#include "task-spec/permissions.h" -#include "test/utils/doctest/check_kv.h" -#include "utils/containers/require_only_key.h" -#include - -namespace test { - -using namespace ::FlexFlow; -namespace Realm = ::FlexFlow::Realm; - -template -static ParallelLayerAttrs make_layer_attrs(T const &op_attrs) { - return ParallelLayerAttrs{ - /*op_attrs=*/PCGOperatorAttrs{op_attrs}, - /*name=*/std::nullopt, - }; -}; - -static bool did_loss_decrease(GenericTensorAccessorR const &first_epoch, - GenericTensorAccessorR const &last_epoch, - Allocator &allocator) { - return tensor_accessor_all( - compare_tensor_accessors_le(last_epoch, first_epoch, allocator)); -} - -MappedParallelComputationGraph - make_test_mpcg_for_device_type(DeviceType device_type) { - positive_int batch_size = 10_p; - positive_int data_dim = 16_p; - positive_int hidden_dim = 32_p; - positive_int output_dim = 1_p; - - TensorShape output_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - TensorShape label_tensor_shape = TensorShape{ - TensorDims{FFOrdered{batch_size, output_dim}}, DataType::FLOAT}; - - ParallelComputationGraph pcg = empty_parallel_computation_graph(); - - TensorShape input_tensor_shape = - TensorShape{TensorDims{FFOrdered{batch_size, data_dim}}, DataType::FLOAT}; - - ParallelLayerAddedResult inputs_layer = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input = - require_only_key(inputs_layer.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult inputs_layer_2 = - pcg_add_input_layer(pcg, input_tensor_shape); - parallel_tensor_guid_t t_input_2 = - require_only_key(inputs_layer_2.outputs, TensorSlotName::OUTPUT); - - ElementBinaryAttrs add_attrs = ElementBinaryAttrs{ - OperatorType::EW_ADD, - DataType::FLOAT, - false, - false, - }; - - ParallelLayerAddedResult add_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(add_attrs), - { - { - TensorSlotName::LHS_INPUT, - t_input, - }, - { - TensorSlotName::RHS_INPUT, - t_input_2, - }, - }, - /*weights=*/{}); - - parallel_tensor_guid_t t_add_1 = - require_only_key(add_operator_1.outputs, TensorSlotName::OUTPUT); - - positive_int replicate_degree = 2_p; - ReplicateAttrs repl_attrs = ReplicateAttrs{replicate_degree}; - ParallelLayerAddedResult repl_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(repl_attrs), - { - { - TensorSlotName::INPUT, - t_add_1, - }, - }, - /*weight=*/{}); - - parallel_tensor_guid_t t_repl_1 = - require_only_key(repl_operator_1.outputs, TensorSlotName::OUTPUT); - - ParallelLayerAddedResult relu_operator_1 = - add_parallel_layer(pcg, - make_layer_attrs(make_relu_attrs()), - /*inputs=*/ - { - { - TensorSlotName::INPUT, - t_repl_1, - }, - }, - /*weights=*/{}); - - parallel_tensor_guid_t t_relu_1 = - require_only_key(relu_operator_1.outputs, TensorSlotName::OUTPUT); - - MachineSpaceCoordinate mc0{0_n, 0_n}; - MachineSpaceCoordinate mc1{0_n, 1_n}; - - ParallelTensorSpaceCoordinate tensor_coord0{ - /*sum_component=*/0_n, - /*discard_copy_component=*/0_n, - /*shard_component=*/FFOrdered{0_n}}; - ParallelTensorSpaceCoordinate tensor_coord1{ - /*sum_component=*/0_n, - /*discard_copy_component=*/1_n, - /*shard_component=*/FFOrdered{0_n}}; - - MappedParallelComputationGraph mpcg = - mapped_pcg_from_pcg_and_mapped_op_task_groups( - /*pcg=*/pcg, - /*mapped_op_task_groups=*/{ - { - inputs_layer.parallel_layer, - MappedOperatorTaskGroup{ - { - { - mc0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - }, - }, - }, - { - inputs_layer_2.parallel_layer, - MappedOperatorTaskGroup{ - { - { - mc0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - }, - }, - }, - { - add_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { - { - mc0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::LHS_INPUT, tensor_coord0}, - {TensorSlotName::RHS_INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - }, - }, - }, - { - repl_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { - { - mc0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - { - mc1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, - }, - }, - }, - }, - { - relu_operator_1.parallel_layer, - MappedOperatorTaskGroup{ - { - { - mc0, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord0}, - {TensorSlotName::OUTPUT, tensor_coord0}, - }}, - }, - { - mc1, - OperatorAtomicTaskShardBinding{{ - {TensorSlotName::INPUT, tensor_coord1}, - {TensorSlotName::OUTPUT, tensor_coord1}, - }}, - }, - }, - }, - }, - }); - - return mpcg; -} - -TEST_SUITE(FF_TEST_SUITE) { - TEST_CASE("RealmBackend e2e Training Replicate Op (CPU Model Parallelism)") { - std::vector fake_args = - make_fake_realm_args(/*num_cpus=*/2_p, /*num_gpus=*/0_n); - int fake_argc = fake_args.size(); - char **fake_argv = fake_args.data(); - - RealmManager manager = RealmManager{&fake_argc, &fake_argv}; - ControllerTaskResult result = - manager.start_controller([](RealmContext &ctx) { - Allocator allocator = ctx.get_current_device_allocator(); - - MappedParallelComputationGraph mpcg = - make_test_mpcg_for_device_type(DeviceType::CPU); - - std::map input_tensors; - - OptimizerAttrs optimizer_attrs = OptimizerAttrs{ - SGDOptimizerAttrs{ - /*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001, - }, - }; - - DistributedFfHandle device_handle = create_distributed_ff_handle( - ctx, - /*workSpaceSize=*/1024 * 1024, - /*allowTensorOpMathConversion=*/true); - - PCGInstance pcg_instance = create_pcg_instance( - /*ctx=*/ctx, - /*mpcg=*/mpcg, - /*optimizer=*/optimizer_attrs, - /*loss=*/std::nullopt, - /*input_tensors=*/input_tensors, - /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*device_type=*/DeviceType::CPU); - - // begin training loop - int num_epochs = 1; - for (int i = 0; i < num_epochs; i++) { - perform_all_passes_for_pcg_instance( - /*instance=*/pcg_instance, - /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle); - } - }); - result.wait(); - } -} - -TEST_SUITE(FF_CUDA_TEST_SUITE) { - TEST_CASE("RealmBackend e2e Training Replicate Op (GPU Model Parallelism)") { - std::vector fake_args = - make_fake_realm_args(/*num_cpus=*/1_p, /*num_gpus=*/2_n); - int fake_argc = fake_args.size(); - char **fake_argv = fake_args.data(); - - RealmManager manager = RealmManager{&fake_argc, &fake_argv}; - - ControllerTaskResult result = - manager.start_controller([](RealmContext &ctx) { - Allocator allocator = ctx.get_current_device_allocator(); - - MappedParallelComputationGraph mpcg = - make_test_mpcg_for_device_type(DeviceType::GPU); - - OptimizerAttrs optimizer_attrs = OptimizerAttrs{ - SGDOptimizerAttrs{ - /*lr=*/0.001, - /*momentum=*/0.9, - /*nesterov=*/false, - /*weight_decay=*/0.001, - }, - }; - - std::map input_tensors; - - DistributedFfHandle device_handle = create_distributed_ff_handle( - ctx, - /*workSpaceSize=*/1024 * 1024, - /*allowTensorOpMathConversion=*/true); - - PCGInstance pcg_instance = create_pcg_instance( - /*ctx=*/ctx, - /*mpcg=*/mpcg, - /*optimizer=*/optimizer_attrs, - /*loss=*/std::nullopt, - /*input_tensors=*/input_tensors, - /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle, - /*device_type=*/DeviceType::GPU); - - // begin training loop - int num_epochs = 1; - for (int i = 0; i < num_epochs; i++) { - perform_all_passes_for_pcg_instance( - /*instance=*/pcg_instance, - /*profiling_settings=*/ProfilingSettings{0, 0}, - /*device_handle=*/device_handle); - } - }); - result.wait(); - } -} -} // namespace test diff --git a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml index 2c10f23f22..5fa9cb63b5 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml +++ b/lib/task-spec/include/task-spec/dynamic_graph/dynamic_value_attrs.dtg.toml @@ -30,7 +30,7 @@ type = "::FlexFlow::dynamic_tensor_guid_t" docstring = ''' \brief The \ref tensor_guid_t or \ref parallel_tensor_guid_t of the (usually parallel) tensor this value originates from. Also allows representing tensors for computing the loss that lie outside of the scope of the \ref ComputationGraph or \ref ParallelComputationGraph, e.g., the label tensor. -For a \ref DynamicOpenDataflowGraph originating from a \ref MapepdParallelComputationGraph, this field is filled in by \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h. +For a \ref DynamicOpenDataflowGraph originating from a \ref MappedParallelComputationGraph, this field is filled in by \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h. ''' [[fields]] diff --git a/lib/task-spec/include/task-spec/dynamic_graph/index.dox b/lib/task-spec/include/task-spec/dynamic_graph/index.dox index 97b72f8553..d256da328d 100644 --- a/lib/task-spec/include/task-spec/dynamic_graph/index.dox +++ b/lib/task-spec/include/task-spec/dynamic_graph/index.dox @@ -7,7 +7,7 @@ namespace FlexFlow { \section task-spec-lowering-passes Lowering Passes -- \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h: Embeds a \ref MappedParallelComputationGraph as a \ref DynamicOpenDataflowGraph. The first of the lowering passes. +- \ref make_dynamic_open_dataflow_graph_from_mapped_pcg.h "": Embeds a \ref MappedParallelComputationGraph as a \ref DynamicOpenDataflowGraph. The first of the lowering passes. - \ref pass_expansion.h - \ref shard_expansion.h - \ref update_insertion.h diff --git a/lib/utils/include/utils/binary_relation/transform_binary_relation.h b/lib/utils/include/utils/binary_relation/transform_binary_relation.h index cf69a97d0a..11141e263e 100644 --- a/lib/utils/include/utils/binary_relation/transform_binary_relation.h +++ b/lib/utils/include/utils/binary_relation/transform_binary_relation.h @@ -9,8 +9,8 @@ namespace FlexFlow { template ::first_type, - typename R2 = std::invoke_result_t::second_type> + typename L2 = typename std::invoke_result_t::first_type, + typename R2 = typename std::invoke_result_t::second_type> BinaryRelation binary_relation_transform_left(BinaryRelation const &rel, F &&f) { BinaryRelation result; diff --git a/lib/utils/include/utils/fmt/unordered_map.h b/lib/utils/include/utils/fmt/unordered_map.h index ecd0257f4e..12faa64e32 100644 --- a/lib/utils/include/utils/fmt/unordered_map.h +++ b/lib/utils/include/utils/fmt/unordered_map.h @@ -6,19 +6,19 @@ #include "utils/join_strings.h" #include #include -#include +#include #include namespace fmt { template struct formatter< - ::std::map, + ::std::unordered_map, Char, - std::enable_if_t>::value>> + std::enable_if_t>::value>> : formatter<::std::string> { template - auto format(::std::map const &m, FormatContext &ctx) const + auto format(::std::unordered_map const &m, FormatContext &ctx) const -> decltype(ctx.out()) { CHECK_FMTABLE(K); CHECK_FMTABLE(V); @@ -38,7 +38,7 @@ struct formatter< namespace FlexFlow { template -std::ostream &operator<<(std::ostream &s, std::map const &m) { +std::ostream &operator<<(std::ostream &s, std::unordered_map const &m) { CHECK_FMTABLE(K); CHECK_FMTABLE(V); diff --git a/lib/utils/src/utils/binary_relation/transform_binary_relation.cc b/lib/utils/src/utils/binary_relation/transform_binary_relation.cc new file mode 100644 index 0000000000..8a83b950cd --- /dev/null +++ b/lib/utils/src/utils/binary_relation/transform_binary_relation.cc @@ -0,0 +1,15 @@ +#include "utils/binary_relation/transform_binary_relation.h" +#include "utils/archetypes/ordered_value_type.h" + +namespace FlexFlow { + +using L = ordered_value_type<0>; +using R = ordered_value_type<1>; +using L2 = ordered_value_type<2>; +using R2 = ordered_value_type<3>; +using F = std::function(L const &, R const &)>; + +template BinaryRelation + binary_relation_transform_left(BinaryRelation const &, F &&); + +} // namespace FlexFlow diff --git a/lib/utils/src/utils/fmt/unordered_map.cc b/lib/utils/src/utils/fmt/unordered_map.cc index 21db320044..f8746e85a0 100644 --- a/lib/utils/src/utils/fmt/unordered_map.cc +++ b/lib/utils/src/utils/fmt/unordered_map.cc @@ -1 +1 @@ -#include "utils/fmt/map.h" +#include "utils/fmt/unordered_map.h"