Repository navigation
[Unity][DistIR] Legalize redistribute #16098
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
12da87b
792e8ea
99b5a5e
08cd0d3
6bc265b
2e72711
03b6142
af83e28
f47fbd2
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -59,3 +59,29 @@ def redistribute(input: Expr, device_mesh: DeviceMesh, placement: Placement) -> | |
| The tensor after redistribution. | ||
| """ | ||
| return _ffi_api.redistribute(input, device_mesh, placement) # type: ignore | ||
|
|
||
|
|
||
| def redistribute_replica_to_shard(input: Expr, num_workers: int, axis: int) -> Expr: | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I think we should change the type of
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. We do need runtime support. Currently Disco runtime treats num_workers as a constant, and legalize_ops of ccl op will throw away
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I would recommend we do a refactor on the whole stack to support dynamic num_workers if there is really a need in the future.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. That is correct, that we would require runtime support, but only if the
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I don't understand the case you are talking about. Why do we want to use a symbolic var for a value that is constant after lowering? And why will we have a variety of static case to specialize, given disco runtime regards
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. One example is if you are comparing how performance scales with the number of GPUs, then you need a way to specify the number of GPUs. Each data point collected would be for a specialized value of Effectively, having
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. got it. This makes sense. Since redistribute_R_to_S shares the attribute with scatter_from_worker0, I'd like to open up a followup PR for this.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. @Lunderberg I just come up with an additional question on the attrs->args change: Do we need to recompile each time for different num_workers after this change with specialization? If yes, then what's the difference between define-symbolic->specialize flow and assume-constant flow?
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The main difference occurs when there are additional optimization steps that occur after a module definition but before the module is handed off to If sharding (and propagation of sharding) is done early in optimization, then propagation of sharding produces large portions of the compute graph where no communication steps are required. These portions can be optimized as if they were single-GPU modules, with no specific handling of multi-GPU setups required. If we can write modules with a dynamic number of gpus, we can write the lowering steps as a single optimization pipeline. # With specialization occurring late in the pipeline.
mod = Sequential([
pre_sharding_optimizations,
shard_across_multiple_gpus,
propagate_sharding,
convert_to_local_view,
single_gpu_optimizations,
])(mod)
built_modules = [relax.build(specialize(mod, num_gpus)) for num_gpus in num_gpu_list]If we can only write modules with a static number of gpus, we cannot write an optimization pipeline, as the optimization pipeline # With specialization occurring at the start of the pipeline.
mod = pre_sharding_optimizations(mod)
mods = [shard_across_multiple_gpus(mod, num_gpus) for num_gpus in num_gpu_list]
pipeline = Sequential([
propagate_sharding,
convert_to_local_view,
single_gpu_optimizations,
])
mods = [pipeline(mod) for mod in mods]
built_modules = [relax.build(mod) for mod in mods]It's not that it's impossible by any means, but that the restricted expressability in an early step means that a user must leave the world of a single |
||
| """Slice tensor into several parts along one axis, | ||
| and each worker takes one part. | ||
| input.struct_info.shape[axis] % num_workers == 0 is required. | ||
| Each worker must have an identical copy of the input. | ||
| This is a specialized version of redistribute op. | ||
|
|
||
| Parameters | ||
| ---------- | ||
| input : relax.Expr | ||
| The buffer to be sliced into equal parts. | ||
|
|
||
| num_worker : int | ||
|
jinhongyii marked this conversation as resolved.
|
||
| The number of workers, i.e. the number of parts the given buffer should be sliced into. | ||
|
|
||
| axis : int | ||
| The axis of the tensor to be sliced. | ||
|
|
||
| Returns | ||
| ------- | ||
| result : relax.Expr | ||
| Sliced Tensor kept by each device. | ||
| """ | ||
| return _ffi_api.redistribute_replica_to_shard(input, num_workers, axis) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,43 @@ | ||
| # Licensed to the Apache Software Foundation (ASF) under one | ||
| # or more contributor license agreements. See the NOTICE file | ||
| # distributed with this work for additional information | ||
| # regarding copyright ownership. The ASF licenses this file | ||
| # to you under the Apache License, Version 2.0 (the | ||
| # "License"); you may not use this file except in compliance | ||
| # with the License. You may obtain a copy of the License at | ||
| # | ||
| # http://www.apache.org/licenses/LICENSE-2.0 | ||
| # | ||
| # Unless required by applicable law or agreed to in writing, | ||
| # software distributed under the License is distributed on an | ||
| # "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY | ||
| # KIND, either express or implied. See the License for the | ||
| # specific language governing permissions and limitations | ||
| # under the License. | ||
| # pylint: disable=invalid-name | ||
| """Default legalization function for distir-related operators.""" | ||
| from tvm import tir, relax | ||
| from ...block_builder import BlockBuilder | ||
| from ...expr import Call, Expr | ||
| from ...op import call_pure_packed | ||
| from ...struct_info import ShapeStructInfo | ||
| from .common import register_legalize | ||
|
|
||
|
|
||
| @register_legalize("relax.dist.redistribute_replica_to_shard") | ||
| def _redistribute_replica_to_shard(_bb: BlockBuilder, call: Call) -> Expr: | ||
| num_workers = call.attrs.num_workers | ||
|
jinhongyii marked this conversation as resolved.
|
||
| axis = call.attrs.axis | ||
| worker_id_symbol = tir.Var("worker_id", "int64") | ||
| worker_id_var = _bb.emit( | ||
| call_pure_packed("runtime.disco.worker_id", sinfo_args=[ShapeStructInfo(None)]) | ||
| ) | ||
| _bb.match_cast(worker_id_var, ShapeStructInfo([worker_id_symbol])) | ||
|
|
||
| split_axis_size = call.args[0].struct_info.shape[axis] | ||
| return relax.op.strided_slice( | ||
| call.args[0], | ||
| axes=[axis], | ||
| begin=[worker_id_symbol * split_axis_size // num_workers], | ||
| end=[(worker_id_symbol + 1) * split_axis_size // num_workers], | ||
| ) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,123 @@ | ||
| /* | ||
| * Licensed to the Apache Software Foundation (ASF) under one | ||
| * or more contributor license agreements. See the NOTICE file | ||
| * distributed with this work for additional information | ||
| * regarding copyright ownership. The ASF licenses this file | ||
| * to you under the Apache License, Version 2.0 (the | ||
| * "License"); you may not use this file except in compliance | ||
| * with the License. You may obtain a copy of the License at | ||
| * | ||
| * http://www.apache.org/licenses/LICENSE-2.0 | ||
| * | ||
| * Unless required by applicable law or agreed to in writing, | ||
| * software distributed under the License is distributed on an | ||
| * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY | ||
| * KIND, either express or implied. See the License for the | ||
| * specific language governing permissions and limitations | ||
| * under the License. | ||
| */ | ||
|
|
||
| /*! | ||
| * \file tvm/relax/distributed/transform/legalize_redistribute.cc | ||
| * \brief Pass for legalizing redistribute op to ccl op. | ||
| */ | ||
|
|
||
| #include <tvm/relax/attrs/ccl.h> | ||
| #include <tvm/relax/attrs/distributed.h> | ||
| #include <tvm/relax/distributed/axis_group_graph.h> | ||
| #include <tvm/relax/distributed/transform.h> | ||
| #include <tvm/relax/expr_functor.h> | ||
| #include <tvm/tir/stmt_functor.h> | ||
|
|
||
| #include "../../../tir/schedule/transform.h" | ||
| #include "../../op/ccl/ccl.h" | ||
| #include "../../op/distributed/distributed.h" | ||
|
|
||
| namespace tvm { | ||
| namespace relax { | ||
| namespace distributed { | ||
|
|
||
| class RedistributeLegalizer : public ExprMutator { | ||
| public: | ||
| static IRModule LegalizeRedistribute(IRModule mod) { | ||
| return RedistributeLegalizer(mod).Legalize(); | ||
| } | ||
|
|
||
| private: | ||
| explicit RedistributeLegalizer(IRModule mod) : ExprMutator(mod) {} | ||
|
|
||
| IRModule Legalize() { | ||
| auto mod = builder_->GetContextIRModule(); | ||
| for (const auto& [gv, base_func] : mod->functions) { | ||
| const auto* func_ = base_func.as<FunctionNode>(); | ||
| if (func_ == nullptr) { | ||
| continue; | ||
| } | ||
| Expr new_func_body = VisitExpr(func_->body); | ||
| auto new_func = make_object<FunctionNode>(*func_); | ||
| new_func->body = new_func_body; | ||
| builder_->UpdateFunction(gv, Function(new_func)); | ||
| } | ||
| return builder_->GetContextIRModule(); | ||
| } | ||
| using ExprMutator::VisitExpr_; | ||
| Expr VisitExpr_(const CallNode* op) final { | ||
| Call call = Downcast<Call>(ExprMutator::VisitExpr_(op)); | ||
| static Op redistribute_op = Op::Get("relax.dist.redistribute"); | ||
| if (call->op.same_as(redistribute_op)) { | ||
| const auto* attrs = call->attrs.as<DistributionAttrs>(); | ||
| ICHECK(attrs); | ||
| const auto* input_sinfo = call->args[0]->struct_info_.as<DTensorStructInfoNode>(); | ||
| ICHECK(input_sinfo); | ||
| // As the first step, we only support redistribute in the same device mesh, | ||
| // and the device mesh must be 1d | ||
| // todo: extend the ccl ops so that it can support 2d device mesh, and different sharding | ||
| // dimension | ||
| ICHECK(StructuralEqual()(input_sinfo->device_mesh, attrs->device_mesh)); | ||
| ICHECK(input_sinfo->device_mesh->shape.size() == 1); | ||
| // only support "S[x]"-> "R" and "R" -> "S[x]" | ||
| PlacementSpec input_spec = input_sinfo->placement->dim_specs[0]; | ||
| PlacementSpec output_spec = attrs->placement->dim_specs[0]; | ||
| if (input_spec->kind == PlacementSpecKind::kReplica && | ||
| output_spec->kind == PlacementSpecKind::kReplica) { | ||
| // "R" -> "R" | ||
| return call->args[0]; | ||
| } else if (input_spec->kind == PlacementSpecKind::kSharding && | ||
| output_spec->kind == PlacementSpecKind::kSharding) { | ||
| // "S[x]" -> "S[y]" | ||
| if (input_spec->axis != output_spec->axis) { | ||
| LOG(FATAL) << "AlltoAll not implemented yet"; | ||
| } else { | ||
| return call->args[0]; | ||
| } | ||
| } else if (input_spec->kind == PlacementSpecKind::kSharding && | ||
| output_spec->kind == PlacementSpecKind::kReplica) { | ||
| // "S[x]" -> "R" | ||
| LOG(FATAL) << "Allgather not implemented yet"; | ||
| } else if (input_spec->kind == PlacementSpecKind::kReplica && | ||
| output_spec->kind == PlacementSpecKind::kSharding) { | ||
| // "R" -> "S[x]" | ||
| return redistribute_replica_to_shard(call->args[0], attrs->device_mesh->shape[0], | ||
|
jinhongyii marked this conversation as resolved.
|
||
| output_spec->axis); | ||
| } else { | ||
| LOG(FATAL) << "Unsupported redistribute op"; | ||
| } | ||
| } | ||
| return call; | ||
| } | ||
| }; | ||
|
|
||
| namespace transform { | ||
|
|
||
| Pass LegalizeRedistribute() { | ||
| runtime::TypedPackedFunc<IRModule(IRModule, PassContext)> pass_func = | ||
| [=](IRModule m, PassContext pc) { return RedistributeLegalizer::LegalizeRedistribute(m); }; | ||
| return CreateModulePass(pass_func, 1, "LegalizeRedistribute", {}); | ||
| } | ||
| TVM_REGISTER_GLOBAL("relax.distributed.transform.LegalizeRedistribute") | ||
| .set_body_typed(LegalizeRedistribute); | ||
| } // namespace transform | ||
|
|
||
| } // namespace distributed | ||
| } // namespace relax | ||
| } // namespace tvm | ||
Uh oh!
There was an error while loading. Please reload this page.