Skip to content
Open
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
5 changes: 3 additions & 2 deletions tests/ap/index_code_gen_value_util.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
class IndexCodeGenValue:
def __init__(self, iter_var_names):
def __init__(self, iter_var_names, iter_dim_splits):

Copy link
Copy Markdown
Owner

Choose a reason for hiding this comment

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

iter_dim_splits超出了IndexCodeGenValue的语义

self.iter_var_names = iter_var_names
self.const_data = None
self.iter_dim_splits = iter_dim_splits
self.const_data = None
16 changes: 11 additions & 5 deletions tests/ap/index_program_translator_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,12 @@ def __init__(
self,
index_func_unique_id2index_program,
kernel_arg_translator,
anchor_iter_var_names
anchor_iter_var_names,
anchor_iter_dim_splits
):
self.kernel_arg_translator = kernel_arg_translator
self.anchor_iter_var_names = anchor_iter_var_names
self.anchor_iter_dim_splits = anchor_iter_dim_splits
items = index_func_unique_id2index_program.items()
self.index_func_unique_id2translator = OrderedDict(
map(
Expand Down Expand Up @@ -44,7 +46,8 @@ def make_translator(self, program_id, index_program):
program_id=program_id,
kernel_arg_translator=self.kernel_arg_translator,
index_op_translator_maker=op_index_translator_util.OpIndexTranslatorFactory(),
anchor_iter_var_names=self.anchor_iter_var_names
anchor_iter_var_names=self.anchor_iter_var_names,
anchor_iter_dim_splits=self.anchor_iter_dim_splits,
)


Expand All @@ -56,13 +59,15 @@ def __init__(
program_id,
kernel_arg_translator,
index_op_translator_maker,
anchor_iter_var_names
anchor_iter_var_names,
anchor_iter_dim_splits
):
self.program_id = program_id
self.program_property = index_program.copy_to_const_program_data()
self.kernel_arg_translator = kernel_arg_translator
self.index_op_translator_maker = index_op_translator_maker
self.anchor_iter_var_names = anchor_iter_var_names
self.anchor_iter_dim_splits = anchor_iter_dim_splits
self.ir_value_index2translated_value = MutableList()
def PushNone(x):
self.ir_value_index2translated_value.append(None)
Expand All @@ -85,7 +90,8 @@ def _translate_op(self, op_property, mut_kernel_arg_id_registry, mut_lir_code_ge
input_properties=map(self._get_value_property, op_property.input_value_indexes),
output_properties=map(self._get_value_property, op_property.output_value_indexes),
kernel_arg_translator=self.kernel_arg_translator,
anchor_iter_var_names=self.anchor_iter_var_names
anchor_iter_var_names=self.anchor_iter_var_names,
anchor_iter_dim_splits=self.anchor_iter_dim_splits
)
inputs = map(self._get_translated_value, op_property.input_value_indexes)
outputs = index_op_translator(
Expand All @@ -102,4 +108,4 @@ def _get_translated_value(self, i):
return self.ir_value_index2translated_value[i]

def _set_translated_value(self, pair):
self.ir_value_index2translated_value[pair[0]] = pair[1]
self.ir_value_index2translated_value[pair[0]] = pair[1]
19 changes: 17 additions & 2 deletions tests/ap/matmul_variadic_tpl.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,23 @@ def make_kernel_arg_translator():
return kernel_arg_translator_util.KernelArgTranslator(param_struct_name="args")


def get_anchor_iter_var_names():
return ["coord.batch", "coord.row", "coord.column"]
def get_anchor_iter_var_names(symbolic_shape):
return (
["coord.batch", "coord.row", "coord.column"]
if len(symbolic_shape) >= 3
else ["coord.row", "coord.column"]
)


def get_anchor_iter_dim_splits(symbolic_shape):
num_anchor_iters = len(get_anchor_iter_var_names(symbolic_shape))
diff = len(symbolic_shape) - num_anchor_iters

def get_dim_split(i):
# 0 is batch which may have no dimension or multiple dimensions in symbolic_shape
return diff + 1 if i < num_anchor_iters - 2 else 1

return map(lambda i: get_dim_split(i), range(num_anchor_iters))


class MatmulVariadicTemplate:
Expand Down
93 changes: 71 additions & 22 deletions tests/ap/op_index_translator_util.py
Original file line number Diff line number Diff line change
@@ -1,22 +1,29 @@
import index_code_gen_value_util


class PdOpDataCodeGen:
def __init__(self,
index_program_id,
op_property,
input_properties,
output_properties,
kernel_arg_translator,
anchor_iter_var_names):
anchor_iter_var_names,
anchor_iter_dim_splits):
self.index_program_id = index_program_id
self.op_property = op_property
self.input_properties = input_properties
self.output_properties = output_properties
self.kernel_arg_translator = kernel_arg_translator
self.anchor_iter_var_names = anchor_iter_var_names
self.anchor_iter_dim_splits = anchor_iter_dim_splits

def __call__(self, inputs, mut_kernel_arg_id_registry, mut_lir_code_gen_ctx):
return [index_code_gen_value_util.IndexCodeGenValue(self.anchor_iter_var_names)]
assert len(self.anchor_iter_var_names) == len(self.anchor_iter_dim_splits), "The length of anchor_iter_var_names and anchor_iter_dim_splits is expected to be the same."
return [index_code_gen_value_util.IndexCodeGenValue(
self.anchor_iter_var_names,
self.anchor_iter_dim_splits
)]


class PdOpFullIntArrayCodeGen:
Expand All @@ -26,16 +33,18 @@ def __init__(self,
input_properties,
output_properties,
kernel_arg_translator,
anchor_iter_var_names):
anchor_iter_var_names,
anchor_iter_dim_splits):
self.index_program_id = index_program_id
self.op_property = op_property
self.input_properties = input_properties
self.output_properties = output_properties
self.kernel_arg_translator = kernel_arg_translator
self.anchor_iter_var_names = anchor_iter_var_names
self.anchor_iter_dim_splits = anchor_iter_dim_splits

def __call__(self, inputs, mut_kernel_arg_id_registry, mut_lir_code_gen_ctx):
out = index_code_gen_value_util.IndexCodeGenValue(None)
out = index_code_gen_value_util.IndexCodeGenValue(None, None)
def get_int64(attr):
return attr.match(a_i64=lambda x:x)
def convert_list(lst):
Expand All @@ -45,35 +54,55 @@ def convert_list(lst):
)
return [out]


class PdOpSumCodeGen:
def __init__(self,
index_program_id,
op_property,
input_properties,
output_properties,
kernel_arg_translator,
anchor_iter_var_names):
anchor_iter_var_names,
anchor_iter_dim_splits):
self.index_program_id = index_program_id
self.op_property = op_property
self.input_properties = input_properties
self.output_properties = output_properties
self.kernel_arg_translator = kernel_arg_translator
self.anchor_iter_var_names = anchor_iter_var_names
self.anchor_iter_dim_splits = anchor_iter_dim_splits

def __call__(self, inputs, mut_kernel_arg_id_registry, mut_lir_code_gen_ctx):
input_iter_var_names = inputs[0].iter_var_names
reduced_axes_set = OrderedDict(
map(lambda x: [int(x), True], inputs[1].const_data)
)
non_reduced_axes = filter(
lambda x: reduced_axes_set.contains(x) == False,
input_iter_dim_splits = inputs[0].iter_dim_splits
def is_reduced_axes(dim):
return False if len(filter(lambda x: int(x) == dim, inputs[1].const_data)) == 0 else True
input_dim_split_starts = MutableList()
input_dim_split_starts.append(0)
def is_anchor_reduced_axes(i):
dim_split_start = int(input_dim_split_starts[i])
dim_split_stop = dim_split_start + input_iter_dim_splits[i]
input_dim_split_starts.append(dim_split_stop)
is_reduced_axes_result = filter(
lambda dim: is_reduced_axes(dim), range(dim_split_start, dim_split_stop)
)
return False if len(is_reduced_axes_result) == 0 else True
anchor_non_reduced_axes = filter(
lambda i: is_anchor_reduced_axes(i) == False,
range(len(input_iter_var_names))
)
output_iter_var_names = map(
lambda i: input_iter_var_names[i],
non_reduced_axes
anchor_non_reduced_axes
)
return [index_code_gen_value_util.IndexCodeGenValue(output_iter_var_names)]
output_iter_dim_splits = map(
lambda i: input_iter_dim_splits[i],
anchor_non_reduced_axes
)
return [index_code_gen_value_util.IndexCodeGenValue(
output_iter_var_names,
output_iter_dim_splits
)]


class CinnOpReshapeCodeGen:
Expand All @@ -83,25 +112,38 @@ def __init__(self,
input_properties,
output_properties,
kernel_arg_translator,
anchor_iter_var_names):
anchor_iter_var_names,
anchor_iter_dim_splits):
self.index_program_id = index_program_id
self.op_property = op_property
self.input_properties = input_properties
self.output_properties = output_properties
self.kernel_arg_translator = kernel_arg_translator
self.anchor_iter_var_names = anchor_iter_var_names
self.anchor_iter_dim_splits = anchor_iter_dim_splits

def __call__(self, inputs, mut_kernel_arg_id_registry, mut_lir_code_gen_ctx):
symbolic_shape = self.input_properties[0].symbolic_shape
def get_or_create_dim_var_name(dim_expr):
arg_var_name = mut_kernel_arg_id_registry.get_dim_expr_var_name(dim_expr)
return self.kernel_arg_translator.get_use_name(arg_var_name)
input_iter_var_names = inputs[0].iter_var_names
input_iter_dim_splits = inputs[0].iter_dim_splits
def get_dim_var_name(i):
dim_expr = symbolic_shape[i]
return get_or_create_dim_var_name(dim_expr)
rank = len(symbolic_shape)
arg_var_name = mut_kernel_arg_id_registry.get_dim_expr_var_name(dim_expr)
return self.kernel_arg_translator.get_use_name(arg_var_name)
input_dim_split_starts = MutableList()
input_dim_split_starts.append(0)
def get_anchor_iter_dims(i):
dim_split_start = int(input_dim_split_starts[i])
dim_split_stop = dim_split_start + input_iter_dim_splits[i]
input_dim_split_starts.append(dim_split_stop)
current_dim_names = map(
lambda idx: get_dim_var_name(idx), range(dim_split_start, dim_split_stop)
)
return " * ".join(current_dim_names)
rank = len(input_iter_dim_splits)
anchor_iter_dims = map(lambda i : get_anchor_iter_dims(i), range(rank))
stride_dims_list = map(
lambda num_dims: map(lambda i: get_dim_var_name(num_dims + i + 1), range(rank - 1 - num_dims)),
lambda num_dims: map(lambda i: anchor_iter_dims[num_dims + i + 1], range(rank - 1 - num_dims)),
range(rank)
)
var_name_and_dims_list = map(
Expand All @@ -115,7 +157,10 @@ def get_dim_var_name(i):
)
)
assert len(self.output_properties[0].symbolic_shape) == 1, "len(self.output_properties[0]) should be 1"
return [index_code_gen_value_util.IndexCodeGenValue([f"({offset_expr})"])]
return [index_code_gen_value_util.IndexCodeGenValue(
[f"({offset_expr})"],
input_iter_dim_splits)
]


class CfYieldCodeGen:
Expand All @@ -125,13 +170,15 @@ def __init__(self,
input_properties,
output_properties,
kernel_arg_translator,
anchor_iter_var_names):
anchor_iter_var_names,
anchor_iter_dim_splits):
self.index_program_id = index_program_id
self.op_property = op_property
self.input_properties = input_properties
self.output_properties = output_properties
self.kernel_arg_translator = kernel_arg_translator
self.anchor_iter_var_names = anchor_iter_var_names
self.anchor_iter_dim_splits = anchor_iter_dim_splits

def __call__(self, inputs, mut_kernel_arg_id_registry, mut_lir_code_gen_ctx):
return []
Expand All @@ -153,7 +200,8 @@ def __call__(self,
input_properties,
output_properties,
kernel_arg_translator,
anchor_iter_var_names):
anchor_iter_var_names,
anchor_iter_dim_splits):
cls = self._get_class(op_property.op_name)
return cls(
index_program_id=index_program_id,
Expand All @@ -162,6 +210,7 @@ def __call__(self,
output_properties=output_properties,
kernel_arg_translator=kernel_arg_translator,
anchor_iter_var_names=anchor_iter_var_names,
anchor_iter_dim_splits=anchor_iter_dim_splits,
)

def _get_class(self, op_name):
Expand Down
4 changes: 3 additions & 1 deletion tests/ap/test_matmul_binary.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,10 +223,12 @@ def _get_program_translator(self, ctx, o, t):
output_names=[],
)
print("index_func_unique_id2index_program:\n", index_func_unique_id2index_program)
mm_out_symbolic_shape = t.mm_out.symbolic_shape_to_list()
index_program_translator_map = index_program_translator_util.IndexProgramTranslatorMap(
index_func_unique_id2index_program=index_func_unique_id2index_program,
kernel_arg_translator=kernel_arg_translator,
anchor_iter_var_names=matmul_variadic_tpl.get_anchor_iter_var_names()
anchor_iter_var_names=matmul_variadic_tpl.get_anchor_iter_var_names(mm_out_symbolic_shape),
anchor_iter_dim_splits=matmul_variadic_tpl.get_anchor_iter_dim_splits(mm_out_symbolic_shape),
)
self._replace_with_load_from_register(
mut_program,
Expand Down
6 changes: 3 additions & 3 deletions tests/ap/test_matmul_epilogue.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,6 @@ def result_pattern(self, o, t):
def constraint(self, o, t):
program = ir_tools.copy_fused_ops_to_program(o.trivial_op, tensor_match_ctx=t)
print("before-umprime: ", program)
# umprime passes
pass_manager = ir_tools.create_pass_manager()
pass_manager.add_pass(ir_tools.create_access_topo_drr_pass("umprime"))
pass_manager.add_pass(ir_tools.create_dce_pass())
Expand All @@ -56,7 +55,6 @@ def constraint(self, o, t):
)
outputs_name_list = map(lambda i: f"output{i}", range(self.number_of_outputs()))
inputs_name_list = map(lambda i: f"input{i+2}", range(self.number_of_inputs() - 2)) if self.number_of_inputs() > 2 else []
print('inputs_name_list: ', ', '.join(inputs_name_list))
init_fake_data_for_yield_input = topo_drr_pass.FakeDataForYieldAccessTopoPass(
outputs_name_list
)
Expand Down Expand Up @@ -243,10 +241,12 @@ def _get_program_translator(self, ctx, o, t):
output_names=other_outputs_name_list,
)
print("index_func_unique_id2index_program:\n", index_func_unique_id2index_program)
mm_out_symbolic_shape = t.mm_out.symbolic_shape_to_list()
index_program_translator_map = index_program_translator_util.IndexProgramTranslatorMap(
index_func_unique_id2index_program=index_func_unique_id2index_program,
kernel_arg_translator=kernel_arg_translator,
anchor_iter_var_names=matmul_variadic_tpl.get_anchor_iter_var_names()
anchor_iter_var_names=matmul_variadic_tpl.get_anchor_iter_var_names(mm_out_symbolic_shape),
anchor_iter_dim_splits=matmul_variadic_tpl.get_anchor_iter_dim_splits(mm_out_symbolic_shape),
)
self._replace_with_load_from_register(
mut_program,
Expand Down