From faef96b375a6aa1900442769b2a04d72e4f6bbfa Mon Sep 17 00:00:00 2001 From: Liu Yiqun Date: Thu, 17 Apr 2025 12:20:54 +0800 Subject: [PATCH 1/3] Support non-3d tensors. --- tests/ap/index_code_gen_value_util.py | 5 +++-- tests/ap/index_program_translator_util.py | 17 ++++++++++++----- tests/ap/matmul_binary_tpl.py | 8 ++++++++ tests/ap/test_matmul_binary.py | 4 +++- 4 files changed, 26 insertions(+), 8 deletions(-) diff --git a/tests/ap/index_code_gen_value_util.py b/tests/ap/index_code_gen_value_util.py index 5e30317..457e514 100644 --- a/tests/ap/index_code_gen_value_util.py +++ b/tests/ap/index_code_gen_value_util.py @@ -1,4 +1,5 @@ class IndexCodeGenValue: - def __init__(self, iter_var_names): + def __init__(self, iter_var_names, iter_dim_splits): self.iter_var_names = iter_var_names - self.const_data = None \ No newline at end of file + self.iter_dim_splits = iter_dim_splits + self.const_data = None diff --git a/tests/ap/index_program_translator_util.py b/tests/ap/index_program_translator_util.py index 6e85f1f..b952fa7 100644 --- a/tests/ap/index_program_translator_util.py +++ b/tests/ap/index_program_translator_util.py @@ -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( @@ -39,12 +41,14 @@ def make_translator(self, program_id, index_program): pass_manager.add_pass(ir_tools.create_access_topo_drr_one_step_pass(drr_pass)) pass_manager.add_pass(ir_tools.create_dce_pass()) pass_manager.run(index_program) + print(f"index_program_with_reshape: {index_program}") return IndexProgramTranslator( 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, ) @@ -56,13 +60,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) @@ -85,7 +91,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( @@ -102,4 +109,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] \ No newline at end of file + self.ir_value_index2translated_value[pair[0]] = pair[1] diff --git a/tests/ap/matmul_binary_tpl.py b/tests/ap/matmul_binary_tpl.py index 696d4ac..64465b0 100644 --- a/tests/ap/matmul_binary_tpl.py +++ b/tests/ap/matmul_binary_tpl.py @@ -10,6 +10,14 @@ def get_anchor_iter_var_names(): return ["coord.batch", "coord.row", "coord.column"] +def get_anchor_iter_dim_splits(symbolic_shape): + num_anchor_iters = len(get_anchor_iter_var_names()) + diff = len(symbolic_shape) - num_anchor_iters + def get_dim_split(i): + return diff + 1 if i < diff else 1 + return map(lambda i: get_dim_split(i), range(num_anchor_iters)) + + class MatmulBinaryTemplate: def __init__( self, diff --git a/tests/ap/test_matmul_binary.py b/tests/ap/test_matmul_binary.py index a4e4ed1..3634789 100644 --- a/tests/ap/test_matmul_binary.py +++ b/tests/ap/test_matmul_binary.py @@ -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_binary_tpl.get_anchor_iter_var_names() + anchor_iter_var_names=matmul_binary_tpl.get_anchor_iter_var_names(), + anchor_iter_dim_splits=matmul_binary_tpl.get_anchor_iter_dim_splits(mm_out_symbolic_shape), ) self._replace_with_load_from_register( mut_program, From 68ef38360474f8beb823866e8b1954b7a23e2c41 Mon Sep 17 00:00:00 2001 From: Liu Yiqun Date: Thu, 17 Apr 2025 13:38:42 +0800 Subject: [PATCH 2/3] Support 2-D tensors. --- tests/ap/index_program_translator_util.py | 1 - tests/ap/matmul_binary_tpl.py | 15 +++- tests/ap/op_index_translator_util.py | 94 ++++++++++++++++++----- tests/ap/test_matmul_binary.py | 2 +- tests/ap/test_matmul_epilogue.py | 6 +- 5 files changed, 88 insertions(+), 30 deletions(-) diff --git a/tests/ap/index_program_translator_util.py b/tests/ap/index_program_translator_util.py index b952fa7..1cc45d6 100644 --- a/tests/ap/index_program_translator_util.py +++ b/tests/ap/index_program_translator_util.py @@ -41,7 +41,6 @@ def make_translator(self, program_id, index_program): pass_manager.add_pass(ir_tools.create_access_topo_drr_one_step_pass(drr_pass)) pass_manager.add_pass(ir_tools.create_dce_pass()) pass_manager.run(index_program) - print(f"index_program_with_reshape: {index_program}") return IndexProgramTranslator( index_program, program_id=program_id, diff --git a/tests/ap/matmul_binary_tpl.py b/tests/ap/matmul_binary_tpl.py index 64465b0..eeea609 100644 --- a/tests/ap/matmul_binary_tpl.py +++ b/tests/ap/matmul_binary_tpl.py @@ -6,15 +6,22 @@ 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()) + num_anchor_iters = len(get_anchor_iter_var_names(symbolic_shape)) diff = len(symbolic_shape) - num_anchor_iters + def get_dim_split(i): - return diff + 1 if i < diff else 1 + # 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)) diff --git a/tests/ap/op_index_translator_util.py b/tests/ap/op_index_translator_util.py index bea4d80..ab0a4b7 100644 --- a/tests/ap/op_index_translator_util.py +++ b/tests/ap/op_index_translator_util.py @@ -1,5 +1,6 @@ import index_code_gen_value_util + class PdOpDataCodeGen: def __init__(self, index_program_id, @@ -7,16 +8,21 @@ 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 [index_code_gen_value_util.IndexCodeGenValue(self.anchor_iter_var_names)] + return [index_code_gen_value_util.IndexCodeGenValue( + self.anchor_iter_var_names, + self.anchor_iter_dim_splits + )] class PdOpFullIntArrayCodeGen: @@ -26,16 +32,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): @@ -45,6 +53,7 @@ def convert_list(lst): ) return [out] + class PdOpSumCodeGen: def __init__(self, index_program_id, @@ -52,28 +61,51 @@ 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): input_iter_var_names = inputs[0].iter_var_names - reduced_axes_set = OrderedDict( - map(lambda x: [int(x), True], inputs[1].const_data) + 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_reduced_axes_set = OrderedDict( + map(lambda i: [i, is_anchor_reduced_axes(i)], + range(len(input_iter_dim_splits))) ) - non_reduced_axes = filter( - lambda x: reduced_axes_set.contains(x) == False, + anchor_non_reduced_axes = filter( + lambda x: anchor_reduced_axes_set[x] == 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 + ) + 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)] + return [index_code_gen_value_util.IndexCodeGenValue( + output_iter_var_names, + output_iter_dim_splits + )] class CinnOpReshapeCodeGen: @@ -83,25 +115,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( @@ -115,7 +160,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: @@ -125,13 +173,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 [] @@ -153,7 +203,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, @@ -162,6 +213,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): diff --git a/tests/ap/test_matmul_binary.py b/tests/ap/test_matmul_binary.py index 3634789..79a2a23 100644 --- a/tests/ap/test_matmul_binary.py +++ b/tests/ap/test_matmul_binary.py @@ -227,7 +227,7 @@ def _get_program_translator(self, ctx, o, t): 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_binary_tpl.get_anchor_iter_var_names(), + anchor_iter_var_names=matmul_binary_tpl.get_anchor_iter_var_names(mm_out_symbolic_shape), anchor_iter_dim_splits=matmul_binary_tpl.get_anchor_iter_dim_splits(mm_out_symbolic_shape), ) self._replace_with_load_from_register( diff --git a/tests/ap/test_matmul_epilogue.py b/tests/ap/test_matmul_epilogue.py index 9c005dd..84147ee 100644 --- a/tests/ap/test_matmul_epilogue.py +++ b/tests/ap/test_matmul_epilogue.py @@ -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()) @@ -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 ) @@ -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_binary_tpl.get_anchor_iter_var_names() + anchor_iter_var_names=matmul_binary_tpl.get_anchor_iter_var_names(mm_out_symbolic_shape), + anchor_iter_dim_splits=matmul_binary_tpl.get_anchor_iter_dim_splits(mm_out_symbolic_shape), ) self._replace_with_load_from_register( mut_program, From 8611b4101d4529a6c843c3b3f189dd39d31d8d09 Mon Sep 17 00:00:00 2001 From: Liu Yiqun Date: Fri, 18 Apr 2025 10:11:37 +0800 Subject: [PATCH 3/3] Simplify the codes. --- tests/ap/op_index_translator_util.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/tests/ap/op_index_translator_util.py b/tests/ap/op_index_translator_util.py index ab0a4b7..0666dbf 100644 --- a/tests/ap/op_index_translator_util.py +++ b/tests/ap/op_index_translator_util.py @@ -19,6 +19,7 @@ def __init__(self, self.anchor_iter_dim_splits = anchor_iter_dim_splits def __call__(self, inputs, mut_kernel_arg_id_registry, mut_lir_code_gen_ctx): + 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 @@ -86,12 +87,8 @@ def is_anchor_reduced_axes(i): 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_reduced_axes_set = OrderedDict( - map(lambda i: [i, is_anchor_reduced_axes(i)], - range(len(input_iter_dim_splits))) - ) anchor_non_reduced_axes = filter( - lambda x: anchor_reduced_axes_set[x] == False, + lambda i: is_anchor_reduced_axes(i) == False, range(len(input_iter_var_names)) ) output_iter_var_names = map(