211namespace distributed {
219 Axis(
const ExprNode* tensor,
int dim,
int tuple_index = 0)
220 : tensor(tensor), dim(dim), tuple_index(tuple_index) {
225 return tensor ==
other.tensor && dim ==
other.dim && tuple_index ==
other.tuple_index;
231 size_t operator()(
const Axis& axis)
const {
232 size_t const h1(std::hash<const ExprNode*>()(axis.tensor));
233 size_t const h2(std::hash<int>()(axis.dim));
234 size_t const h3(std::hash<int>()(axis.tuple_index));
235 return h1 ^ (
h2 << 1) ^ (
h3 << 2);
239using AxisGroup = std::unordered_set<Axis, AxisHash>;
243 size_t operator()(
const AxisGroup&
axis_group)
const {
246 seed ^= AxisHash()(axis) + 0x9e3779b9 + (
seed << 6) + (
seed >> 2);
252using ShardingSpec = std::pair<DeviceMesh, Placement>;
255using AxisShardingSpec = std::pair<DeviceMesh, int>;
256class AxisShardingSpecEqual {
258 bool operator()(
const AxisShardingSpec& lhs,
const AxisShardingSpec& rhs)
const {
259 return ffi::StructuralEqual()(lhs.first, rhs.first) && lhs.second == rhs.second;
263class AxisShardingSpecHash {
265 size_t operator()(
const AxisShardingSpec&
sharding_spec)
const {
278class AxisGroupGraph {
280 enum class EdgeType { kAscend, kDescend, kSimbling };
283 static EdgeType ReverseEdgeType(EdgeType type) {
285 case EdgeType::kAscend:
286 return EdgeType::kDescend;
287 case EdgeType::kDescend:
288 return EdgeType::kAscend;
289 case EdgeType::kSimbling:
290 return EdgeType::kSimbling;
296 static int GetEdgePriority(EdgeType type) {
298 case EdgeType::kAscend:
300 case EdgeType::kDescend:
302 case EdgeType::kSimbling:
309 struct AxisGraphEdge {
327 Path AddEdge(EdgeType type) {
return {direction |= (1 << GetEdgePriority(type))}; }
329 int GetPriority()
const {
344 AxisGroupGraph() =
default;
355 void JoinAxis(Axis
axis1, Axis
axis2, EdgeType type) {
365 void AddSrcShardingPoint(Axis axis, AxisShardingSpec
spec) {
366 src_axis_sharding_spec_[axis] =
spec;
372 void PropagateShardingSpec() {
373 axis_sharding_specs_priority_.clear();
374 for (
const auto&
pr : src_axis_sharding_spec_) {
375 std::unordered_set<Axis, AxisHash>
visited;
376 PropagateShardingSpec(
pr.first,
pr.second, Path(), &
visited);
378 ChooseAxisShardingSpec();
387 void AddPropagationCutPoint(Axis axis, AxisShardingSpec
spec) {
388 cutpoint_axis_sharding_spec_[axis] =
spec;
398 std::tuple<AxisShardingSpec, bool> GetAxisShardingSpec(Axis axis) {
399 if (axis_sharding_specs_priority_.count(axis)) {
400 return {axis_sharding_specs_priority_[axis].begin()->first,
true};
407 void AddEdge(Axis src, Axis dst, EdgeType type) {
408 if (!graph_.count(src)) {
411 graph_[src].push_back({src, dst, type});
414 void PropagateShardingSpec(Axis axis, AxisShardingSpec
spec, Path path,
415 std::unordered_set<Axis, AxisHash>*
visited) {
416 if (cutpoint_axis_sharding_spec_.count(axis) ||
417 (src_axis_sharding_spec_.count(axis) &&
418 !AxisShardingSpecEqual()(src_axis_sharding_spec_[axis],
spec)) ||
423 if (!axis_sharding_specs_priority_.count(axis)) {
424 axis_sharding_specs_priority_[axis] = {};
426 axis_sharding_specs_priority_[axis][
spec] = path.GetPriority();
427 for (
auto edge : graph_[axis]) {
432 void ChooseAxisShardingSpec() {
433 for (
auto&
pr : axis_sharding_specs_priority_) {
434 auto& axis =
pr.first;
448 <<
"multiple possible sharding for axis: (" << ffi::GetRef<Expr>(axis.tensor) <<
", "
454 std::unordered_map<Axis, std::vector<AxisGraphEdge>, AxisHash> graph_;
455 std::unordered_map<Axis, AxisShardingSpec, AxisHash> src_axis_sharding_spec_;
456 std::unordered_map<Axis, AxisShardingSpec, AxisHash> cutpoint_axis_sharding_spec_;
458 Axis, std::unordered_map<AxisShardingSpec, int, AxisShardingSpecHash, AxisShardingSpecEqual>,
460 axis_sharding_specs_priority_;