Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions include/tvm/relax/distributed/transform.h
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,13 @@ using DataflowBlock = tvm::relax::DataflowBlock;
*/
TVM_DLL Pass PropagateSharding();

/*!
* \brief Legalize redistribute op to ccl op.
*
* \return The Pass.
*/
TVM_DLL Pass LegalizeRedistribute();

} // namespace transform
} // namespace distributed
} // namespace relax
Expand Down
2 changes: 1 addition & 1 deletion python/tvm/relax/distributed/transform/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,4 +16,4 @@
# under the License.
"""Relax distributed-related transformations. """

from .transform import PropagateSharding
from .transform import PropagateSharding, LegalizeRedistribute
13 changes: 13 additions & 0 deletions python/tvm/relax/distributed/transform/transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,3 +30,16 @@ def PropagateSharding() -> tvm.ir.transform.Pass:
The registered pass
"""
return _ffi_api.PropagateSharding() # type: ignore


def LegalizeRedistribute() -> tvm.ir.transform.Pass:
"""Legalize redistribute op to ccl op.
Comment thread
jinhongyii marked this conversation as resolved.
S->R: R.ccl.allgather
R->S: R.dist.redistribute_replica_to_shard

Returns
-------
ret : tvm.transform.Pass
The registered pass
"""
return _ffi_api.LegalizeRedistribute() # type: ignore
2 changes: 1 addition & 1 deletion python/tvm/relax/op/distributed/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,4 +16,4 @@
# under the License.
"""Operators serving for distributed Relax."""

from .distributed import annotate_sharding, redistribute
from .distributed import annotate_sharding, redistribute, redistribute_replica_to_shard
26 changes: 26 additions & 0 deletions python/tvm/relax/op/distributed/distributed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think we should change the type of num_workers from int to Expr. That allows the number of workers to be a symbolic variable. It doesn't require any runtime support, as the symbolic variable would be specialized later on, but this is very useful when writing generic implementations.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The 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 num_workers when converting to call_dps_packed(ccl, ...). This behavior needs to be changed if we want to make num_worker a Expr

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The 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.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The 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 num_workers is still dynamic after being lowered to either the disco runtime or the ccl op legalization. It is easier to write a single dynamic implementation then specialize to a variety of static cases, than it is to write several distinct static implementations. However, the initial writing of the dynamic implementation requires that it be expressible, even though it will be specialized out later in lowering.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The 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 num_workers as a global constant?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The 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 num_workers.

Effectively, having num_workers: Expr de-couples the communication mechanics from the choice of how many workers to use. It still can be explicitly specified at any stage of lowering, and it still must be statically-known after lowering, but we gain flexibility before then. This has a number of benefits, for example:

  • Predicted memory usage. If the number of workers is stored symbolically, adding up the size of all live values at any point gives the memory footprint as a function of the number of workers. Requiring the number of workers to be static at all points of lowering prevents this analysis.

  • Consistent optimization. If an optimization is applicable regardless of the number of workers, the optimization should be applied at a point when the number of workers is unknown. This prevents a developer from accidentally making a less general optimization. (e.g. By using a sharded tensor shape in the pattern-matching.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The 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.

@jinhongyii jinhongyii Nov 14, 2023 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The 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?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The 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 relax.build for lowering/compilation.

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 IRModule much earlier.

"""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
Comment thread
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)
1 change: 1 addition & 0 deletions python/tvm/relax/transform/legalize_ops/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from . import ccl
from . import create
from . import datatype
from . import distributed
from . import grad
from . import image
from . import index
Expand Down
43 changes: 43 additions & 0 deletions python/tvm/relax/transform/legalize_ops/distributed.py
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
Comment thread
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],
)
3 changes: 2 additions & 1 deletion python/tvm/script/ir_builder/relax/distributed/ir.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
# pylint: disable=redefined-builtin, wrong-import-order, no-member, invalid-name
# pylint: disable=redefined-builtin, wrong-import-order, no-member, invalid-name, unused-import

"""IRBuilder for distributed Relax dialect"""
from typing import Union, List, Tuple, Optional
Expand All @@ -32,6 +32,7 @@
from tvm.relax.op.distributed import (
redistribute as _redistribute,
annotate_sharding as _annotate_sharding,
redistribute_replica_to_shard,
)
from tvm.relax.distributed import DeviceMesh, Placement
from . import _ffi_api
Expand Down
27 changes: 12 additions & 15 deletions python/tvm/script/parser/relax/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,18 +34,15 @@
else:
from .entry import function, macro

__all__ = (
_relax.__all__
+ dist.__all__
+ [
"Callable",
"Object",
"Prim",
"Shape",
"Tensor",
"Tuple",
"function",
"macro",
"match_cast",
]
)
__all__ = _relax.__all__ + [
"dist",
"Callable",
"Object",
"Prim",
"Shape",
"Tensor",
"Tuple",
"function",
"macro",
"match_cast",
]
8 changes: 7 additions & 1 deletion python/tvm/script/parser/relax/dist.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,13 @@
from tvm.relax.distributed import DeviceMesh, Placement, DTensorStructInfo, device_mesh
from tvm.script.ir_builder import IRBuilder
from tvm.script.ir_builder.ir import IRModuleFrame
from tvm.script.ir_builder.relax.distributed import call_tir, const, annotate_sharding, redistribute
from tvm.script.ir_builder.relax.distributed import (
call_tir,
const,
annotate_sharding,
redistribute,
redistribute_replica_to_shard,
)
from .entry import StructInfoProxy, TensorProxy


Expand Down
123 changes: 123 additions & 0 deletions src/relax/distributed/transform/legalize_redistribute.cc
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],
Comment thread
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
Loading