[Device] Add dynamic fetch/reduce pipelining for reduction collectives - Simple protocol (#1861)
* Support pipelining codegen and template specialization * Support ReduceCopy pipelining for AllReduce, ReduceScatter, and Reduce (currently enabled for bfloat16) * Remove need for FUNC_INDEX_TOTAL * Add pipeline field to device function key construction logic * Avoid unneeded codegen for LL/LL64 kernels * Modify conditions and add pipeline dtypes env * Optimize selection for both gfx942 and gfx950 * Increase pipeline bitfield width * Use __forceinline__ for all device functions * Realign reduceCopy with original form * Add opt-out option to enable perf debugs * Remove force-reduce-pipelining option from README * Update CHANGELOG.md --------- Co-authored-by: Jeffrey Novotny <jnovotny@amd.com>
This commit is contained in:
committed by
GitHub
parent
b882af9ffd
commit
277747c199
+59
-31
@@ -11,7 +11,13 @@ all_protos = ["LL","LL128","SIMPLE"]
|
||||
all_algos = ["TREE","RING", "", "", "", "", "PAT"]
|
||||
all_unroll = ["1", "2", "4"]
|
||||
use_acc = ["0", "1"]
|
||||
all_params = [all_colls, all_algos, all_protos, all_redops, all_tys, use_acc, all_unroll]
|
||||
|
||||
# Pipelining is not supported for LL/LL64 prims, so "1" is not a valid value for low latency protocols.
|
||||
# However, if it needs to be supported, equivalent_primary() can be modified to avoid the "non-zero"->"0" mapping.
|
||||
all_pipeline = ["0", "1"]
|
||||
pipelined_types = ["bf16"]
|
||||
all_params = [all_colls, all_algos, all_protos, all_redops, all_tys, use_acc, all_pipeline, all_unroll]
|
||||
|
||||
|
||||
################################################################################
|
||||
# The first command line argument is the path to the directory to generate and
|
||||
@@ -114,14 +120,25 @@ redops_of_coll = {
|
||||
}
|
||||
|
||||
tys_of_coll = {
|
||||
"AllGather": ["i8"],
|
||||
"AllReduce": all_tys,
|
||||
"AllGather": ["i8"],
|
||||
"AllReduce": all_tys,
|
||||
"AllReduceWithBias": all_tys,
|
||||
"AllToAllPivot": ["i8"],
|
||||
"Broadcast": ["i8"],
|
||||
"Reduce": all_tys,
|
||||
"ReduceScatter": all_tys,
|
||||
"SendRecv": ["i8"]
|
||||
"AllToAllPivot": ["i8"],
|
||||
"Broadcast": ["i8"],
|
||||
"Reduce": all_tys,
|
||||
"ReduceScatter": all_tys,
|
||||
"SendRecv": ["i8"]
|
||||
}
|
||||
|
||||
pipelines_of_coll = {
|
||||
"AllGather": ["0"],
|
||||
"AllReduce": all_pipeline,
|
||||
"AllReduceWithBias": ["0"],
|
||||
"AllToAllPivot": ["0"],
|
||||
"Broadcast": ["0"],
|
||||
"Reduce": all_pipeline,
|
||||
"ReduceScatter": all_pipeline,
|
||||
"SendRecv": ["0"]
|
||||
}
|
||||
|
||||
coll_camel_to_lower = {
|
||||
@@ -179,7 +196,7 @@ def calc_unroll_for_local_arch():
|
||||
return all_unroll
|
||||
|
||||
# Helper function to check if the conditions for the collective is being met
|
||||
def func_validate(coll, algo, proto, redop, ty, acc, unroll):
|
||||
def func_validate(coll, algo, proto, redop, ty, acc, pipeline, unroll):
|
||||
if acc == "1" and coll != "AllReduceWithBias":
|
||||
return False
|
||||
if acc == "0" and coll == "AllReduceWithBias":
|
||||
@@ -188,7 +205,7 @@ def func_validate(coll, algo, proto, redop, ty, acc, unroll):
|
||||
return False
|
||||
if coll == "" or algo == "":
|
||||
return False
|
||||
if algo not in algos_of_coll[coll] or proto not in protos_of_coll[coll] or redop not in redops_of_coll[coll] or ty not in tys_of_coll[coll] or acc not in use_acc or unroll not in all_unroll:
|
||||
if algo not in algos_of_coll[coll] or proto not in protos_of_coll[coll] or redop not in redops_of_coll[coll] or ty not in tys_of_coll[coll] or acc not in use_acc or unroll not in all_unroll or pipeline not in pipelines_of_coll[coll] or (pipeline in ["1"] and ty not in pipelined_types):
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -233,10 +250,10 @@ def func_filter(function_params, current_idx, item_list=None):
|
||||
# For each loop layer remove the last element in item_list
|
||||
item_list.pop()
|
||||
else:
|
||||
coll, algo, proto, redop, ty, acc, unroll = item_list
|
||||
coll, algo, proto, redop, ty, acc, pipeline, unroll = item_list
|
||||
if func_validate(coll, algo, proto, redop, ty, acc, pipeline, unroll):
|
||||
yield(coll, algo, proto, redop, ty, acc, pipeline, unroll)
|
||||
|
||||
if func_validate(coll, algo, proto, redop, ty, acc, unroll):
|
||||
yield(coll, algo, proto, redop, ty, acc, unroll)
|
||||
|
||||
# Parse ONLY_FUNCS input and feed it to func_filter
|
||||
def parse_input(func_pattern):
|
||||
@@ -256,7 +273,7 @@ def parse_input(func_pattern):
|
||||
|
||||
# Maps functions to the chosen representative for the equivalence class it
|
||||
# belongs to. For instance (sum, signed int) maps to (sum, unsigned int).
|
||||
def equivalent_primary(coll, algo, proto, redop, ty, acc, unroll):
|
||||
def equivalent_primary(coll, algo, proto, redop, ty, acc, pipeline, unroll):
|
||||
if coll in ("AllReduce", "AllReduceWithBias", "Reduce", "ReduceScatter"):
|
||||
# map signed integer sum/prod to unsigned
|
||||
if redop in ("Sum","Prod","PreMulSum","SumPostDiv") and ty[0]=="i":
|
||||
@@ -264,7 +281,11 @@ def equivalent_primary(coll, algo, proto, redop, ty, acc, unroll):
|
||||
# map signed integer min/max to unsigned for non-NVLS
|
||||
elif redop=="MinMax" and ty[0]=="i" and ("NVLS" not in algo):
|
||||
ty = "u"+ty[1:]
|
||||
return (coll, algo, proto, redop, ty, acc, unroll)
|
||||
# map pipelined to non-pipelined for LL/LL128 to avoid extra device codegen
|
||||
if (pipeline != "0" and proto != "SIMPLE"):
|
||||
pipeline = "0"
|
||||
|
||||
return (coll, algo, proto, redop, ty, acc, pipeline, unroll)
|
||||
|
||||
# Order rows are enumerated must match formula of `ncclDevFuncId()`:
|
||||
# outermost loop should be for unroll factor; refer to host_table section
|
||||
@@ -276,12 +297,12 @@ def enumerate_func_rows():
|
||||
for proto in all_protos:
|
||||
for redop in all_redops:
|
||||
for ty in all_tys:
|
||||
if func_validate(coll, algo, proto, redop, ty, acc, unroll):
|
||||
yield (coll, algo, proto, redop, ty, acc, unroll)
|
||||
|
||||
for pipeline in all_pipeline:
|
||||
if func_validate(coll, algo, proto, redop, ty, acc, pipeline, unroll):
|
||||
yield (coll, algo, proto, redop, ty, acc, pipeline, unroll)
|
||||
# Sort the hashmap based on custom key <coll> <algo> <proto> <redop> <ty>
|
||||
def custom_sort_key(fn):
|
||||
coll, algo, proto, redop, ty, acc, unroll = fn
|
||||
coll, algo, proto, redop, ty, acc, pipeline, unroll = fn
|
||||
return (
|
||||
all_unroll.index(unroll),
|
||||
use_acc.index(acc),
|
||||
@@ -289,7 +310,8 @@ def custom_sort_key(fn):
|
||||
all_algos.index(algo),
|
||||
all_protos.index(proto),
|
||||
all_redops.index(redop),
|
||||
all_tys.index(ty)
|
||||
all_tys.index(ty),
|
||||
all_pipeline.index(pipeline)
|
||||
)
|
||||
|
||||
################################################################################
|
||||
@@ -333,7 +355,7 @@ with open(os.path.join(gensrc, "device_table.h"), "w") as f:
|
||||
out("__device__ ncclDevFuncPtr_t const ncclDevFuncTable_1[] = {\n")
|
||||
index1 = 0
|
||||
for fn in primary_funcs:
|
||||
coll, algo, proto, redop, ty, acc, unroll = fn
|
||||
coll, algo, proto, redop, ty, acc, pipeline, unroll = fn
|
||||
if unroll != "1": continue
|
||||
sym = paste("_", "ncclDevFunc", *fn)
|
||||
if fn[2] == "LL128":
|
||||
@@ -350,7 +372,7 @@ with open(os.path.join(gensrc, "device_table.h"), "w") as f:
|
||||
out("__device__ ncclDevFuncPtr_t const ncclDevFuncTable_2[] = {\n")
|
||||
index2 = 0
|
||||
for fn in primary_funcs:
|
||||
coll, algo, proto, redop, ty, acc, unroll = fn
|
||||
coll, algo, proto, redop, ty, acc, pipeline, unroll = fn
|
||||
if unroll != "2": continue
|
||||
sym = paste("_", "ncclDevFunc", *fn)
|
||||
if fn[2] == "LL128":
|
||||
@@ -367,7 +389,7 @@ with open(os.path.join(gensrc, "device_table.h"), "w") as f:
|
||||
out("__device__ ncclDevFuncPtr_t const ncclDevFuncTable_4[] = {\n")
|
||||
index4 = 0
|
||||
for fn in primary_funcs:
|
||||
coll, algo, proto, redop, ty, acc, unroll = fn
|
||||
coll, algo, proto, redop, ty, acc, pipeline, unroll = fn
|
||||
if unroll != "4": continue
|
||||
sym = paste("_", "ncclDevFunc", *fn)
|
||||
if fn[2] == "LL128":
|
||||
@@ -448,7 +470,7 @@ if is_colltrace:
|
||||
out("\n")
|
||||
|
||||
seen_fns = set()
|
||||
out("const char* funcNames[FUNC_INDEX_TOTAL] = {\n")
|
||||
out("const char* funcNames[] = {\n")
|
||||
for fn in primary_funcs:
|
||||
fn_no_unroll = fn[:-1]
|
||||
if fn_no_unroll not in seen_fns:
|
||||
@@ -466,6 +488,7 @@ with open(os.path.join(gensrc, "host_table.cpp"), "w") as f:
|
||||
out('#include "device.h"\n')
|
||||
out("\n")
|
||||
out("// The key for the ncclDevFuncNameToId map is a 64-bit unsigned integer.\n")
|
||||
out("// Each field (coll, algo, proto, redop, ty, pipeline) is packed into 4 bits,\n")
|
||||
out("// Each field (coll, algo, proto, redop, ty) is packed into 4 bits,\n")
|
||||
out("// This allows up to 16 unique values per field. The layout is:\n")
|
||||
out("// bits 0-3: coll index\n")
|
||||
@@ -473,8 +496,9 @@ with open(os.path.join(gensrc, "host_table.cpp"), "w") as f:
|
||||
out("// bits 8-11: proto index\n")
|
||||
out("// bits 12-15: redop index\n")
|
||||
out("// bits 16-19: ty index\n")
|
||||
out("// bits 20-23: pipeline index\n")
|
||||
out("#include <unordered_map>\n")
|
||||
out("extern std::unordered_map<uint64_t, int> ncclDevFuncNameToId = {\n")
|
||||
out("std::unordered_map<uint64_t, int> ncclDevFuncNameToId = {\n")
|
||||
|
||||
# host_table entries map device functions based on collective, algorithm, protocol, redop, and datatype
|
||||
# For GPU targets that support multiple unrolls, e.g., gfx950
|
||||
@@ -485,17 +509,20 @@ with open(os.path.join(gensrc, "host_table.cpp"), "w") as f:
|
||||
fn_id = primary_to_index[equivalent_primary(*fn)]
|
||||
comment = " // " + paste(" ", *fn[:-1])
|
||||
# Build the function signature string: "<coll> <algo> <proto> <redop> <ty>"
|
||||
# get parts indexes in order (coll, algo, proto, redop, ty, acc, pipeline, unroll)
|
||||
coll_idx = all_colls.index(fn[0])
|
||||
algo_idx = all_algos.index(fn[1])
|
||||
proto_idx = all_protos.index(fn[2])
|
||||
redop_idx = all_redops.index(fn[3])
|
||||
ty_idx = all_tys.index(fn[4])
|
||||
pipeline_idx = all_pipeline.index(fn[6])
|
||||
# Assert that 4 bits (16 values) is enough to map all_colls, all_algos, etc.
|
||||
assert len(all_colls) <= 16, "Error: all_colls has more than 16 values, which exceeds 4-bit capacity."
|
||||
assert len(all_algos) <= 16, "Error: all_algos has more than 16 values, which exceeds 4-bit capacity."
|
||||
assert len(all_protos) <= 16, "Error: all_protos has more than 16 values, which exceeds 4-bit capacity."
|
||||
assert len(all_redops) <= 16, "Error: all_redops has more than 16 values, which exceeds 4-bit capacity."
|
||||
assert len(all_tys) <= 16, "Error: all_tys has more than 16 values, which exceeds 4-bit capacity."
|
||||
assert len(all_pipeline) <= 16, "Error: all_pipeline has more than 16 values, which exceeds 4-bit capacity."
|
||||
# Create a 64-bit unsigned integer key and pack the indices into 4 bits each
|
||||
key = (
|
||||
(coll_idx & 0xF)
|
||||
@@ -503,8 +530,9 @@ with open(os.path.join(gensrc, "host_table.cpp"), "w") as f:
|
||||
| ((proto_idx & 0xF) << 8)
|
||||
| ((redop_idx & 0xF) << 12)
|
||||
| ((ty_idx & 0xF) << 16)
|
||||
| ((pipeline_idx & 0xF) << 20)
|
||||
)
|
||||
fn_str = f"{coll_idx} {algo_idx} {proto_idx} {redop_idx} {ty_idx}"
|
||||
fn_str = f"{coll_idx} {algo_idx} {proto_idx} {redop_idx} {ty_idx} {pipeline_idx}"
|
||||
if fn[0] == "Broadcast":
|
||||
key = ((coll_idx & 0x3F) | ((proto_idx & 0x3F) << 8))
|
||||
if fn[0] in ["SendRecv", "AllToAllPivot"]:
|
||||
@@ -515,7 +543,7 @@ with open(os.path.join(gensrc, "host_table.cpp"), "w") as f:
|
||||
# Maps to .cu filename which implements this func. The only constraint is that
|
||||
# "coll" is reflected in the name: formally that no two funcs having different
|
||||
# coll's map to the same filename.
|
||||
def impl_filename(coll, algo, proto, redop, ty, acc, unroll):
|
||||
def impl_filename(coll, algo, proto, redop, ty, acc, pipeline, unroll):
|
||||
return "%s.cpp" % paste("_", coll_camel_to_lower[coll], redop and redop.lower(), ty)
|
||||
|
||||
# Partition the functions and kernels to the .cu filenames. The partition is
|
||||
@@ -573,14 +601,14 @@ for name in name_to_funcs.keys():
|
||||
)
|
||||
|
||||
for fn in fns:
|
||||
(coll, algo, proto, redop, ty, acc, unroll) = fn
|
||||
sym = paste("_", coll, algo, proto, redop, ty, acc, unroll)
|
||||
(coll, algo, proto, redop, ty, acc, pipeline, unroll) = fn
|
||||
sym = paste("_", coll, algo, proto, redop, ty, acc, pipeline, unroll)
|
||||
if proto == "LL128":
|
||||
out("#if (defined(__gfx90a__) || defined(__gfx942__) || defined(__gfx950__)) && defined(ENABLE_LL128)\n")
|
||||
out(
|
||||
"DEFINE_ncclDevFunc({sym}, ncclFunc{coll}, {redop_cxx}, {ty_cxx}, NCCL_ALGO_{algo}, NCCL_PROTO_{proto}, {acc}, {unroll})\n"
|
||||
"DEFINE_ncclDevFunc({sym}, ncclFunc{coll}, {redop_cxx}, {ty_cxx}, NCCL_ALGO_{algo}, NCCL_PROTO_{proto}, {acc}, {pipeline}, {unroll})\n"
|
||||
.format(sym=sym, coll=coll, redop_cxx=redop_to_cxx[redop], ty_cxx=ty_to_cxx[ty],
|
||||
algo=(algo or "RING"), proto=(proto or "SIMPLE"), acc=acc, unroll=unroll)
|
||||
algo=(algo or "RING"), proto=(proto or "SIMPLE"), acc=acc, pipeline=pipeline, unroll=unroll)
|
||||
)
|
||||
if proto == "LL128":
|
||||
out("#endif\n")
|
||||
|
||||
Reference in New Issue
Block a user