---
title: kutacc_af2_outer_product_mean_chunk
description: "outer_product_mean中用于分块计算结果的计算函数。"
url: https://www.hikunpeng.com/document/detail/zh/kunpenghpcs/hpckit/devg/KunpengHPCKit_developer_139.html
sourcePath: /source/zh/kunpenghpcs/hpckit/devg/KunpengHPCKit_developer_139.html
indexId: ee99c2aeecce11c4c4d93500191a1f7be5d8f1fc631faf64d90a1ff22aa4161d67
---
# kutacc_af2_outer_product_mean_chunk

outer_product_mean中用于分块计算结果的计算函数。

#### 接口定义

kutacc_export void kutacc_af2_outer_product_mean_chunk(kutacc_af2_opm_act_inputs_t *opm_acts_ptr, kutacc_af2_opm_mask_inputs_t *opm_masks_ptr, kutacc_af2_opm_weights_t *opm_weights_ptr,kutacc_tensor_h out, int64_t left_block_size, int64_t right_block_size);


#### 参数

表1 入参定义

| 参数名 | 类型 | 描述 | 输入/输出 |
| --- | --- | --- | --- |
| opm\_acts\_ptr | kutacc\_af2\_opm\_act\_inputs\_t \* | kutacc\_af2\_opm\_act\_inputs\_t类型的指针，具体数据结构定义见下方表2 kutacc\_af2\_opm\_act\_inputs\_t 数据结构表定义 | 输入 |
| opm\_masks\_ptr | kutacc\_af2\_opm\_mask\_inputs\_t \* | kutacc\_af2\_opm\_mask\_inputs\_t类型的指针，具体数据结构定义见下方表3 kutacc\_af2\_opm\_mask\_inputs\_t 数据结构表定义 | 输入 |
| opm\_weights\_ptr | kutacc\_af2\_opm\_weights\_t \* | kutacc\_af2\_opm\_weights\_t类型的指针，具体数据结构定义见下方表4 kutacc\_af2\_opm\_weights\_t数据结构表定义 | 输入 |
| out | kutacc\_tensor\_h | 输出数据 | 输出 |
| left\_block\_size | int64\_t | 左分块大小 | 输入 |
| right\_block\_size | int64\_t | 右分块大小 | 输入 |


表2 kutacc_af2_opm_act_inputs_t 数据结构表定义

| 参数名 | 类型 | 描述 | 输入/输出 |
| --- | --- | --- | --- |
| n\_seq | int64\_t | 序列数量 | 输入 |
| n\_res | int64\_t | 残基数量 | 输入 |
| input\_act | kutacc\_tensor\_h | 输入激活张量 | 输入 |
| left\_proj | kutacc\_tensor\_h | 左投影 | 输入 |
| right\_proj | kutacc\_tensor\_h | 右投影 | 输入 |
| left\_proj\_ | kutacc\_tensor\_h | 经过掩码处理后的左投影 | 输入 |
| right\_proj\_ | kutacc\_tensor\_h | 经过掩码处理后的右投影 | 输入 |


表3 kutacc_af2_opm_mask_inputs_t 数据结构表定义

| 参数名 | 类型 | 描述 | 输入/输出 |
| --- | --- | --- | --- |
| n\_res\_gather | int64\_t | 聚合后的残基数量 | 输入 |
| mask\_bias | int64\_t | 掩码张量地址偏移量 | 输入 |
| mask | kutacc\_tensor\_h | 掩码张量 | 输入 |
| norm | kutacc\_tensor\_h | 归一化因子张量 | 输入 |


表4 kutacc_af2_opm_weights_t 数据结构表定义

| 参数名 | 类型 | 描述 | 输入/输出 |
| --- | --- | --- | --- |
| c\_m | int64\_t | 输入特征维度 | 输入 |
| c\_i | int64\_t | 投影后的特征维度 | 输入 |
| c\_z | int64\_t | 输出特征维度 | 输入 |
| left\_proj\_w | kutacc\_tensor\_h | 左投影权重 | 输入 |
| left\_proj\_b | kutacc\_tensor\_h | 左投影偏移量 | 输入 |
| right\_proj\_w | kutacc\_tensor\_h | 右投影权重 | 输入 |
| right\_proj\_b | kutacc\_tensor\_h | 右投影偏移量 | 输入 |
| outer\_w | kutacc\_tensor\_h | 输出权重 | 输入 |
| outer\_b | kutacc\_tensor\_h | 输出偏移量 | 输入 |


outer_product_mean整数参数应满足的约束关系：

n_res, n_res_gather, c_i, c_z, n_res, n_res_gather, left_block_size, right_block_size > 0;

n_seq * n_res <INT64_MAX，

left_block_size * right_block_size * c_i * c_i < INT64_MAX，

left_block_size* right_block_size * c_z < INT64_MAX，

单进程时n_res必须等于n_res_gather 多进程下不满足该条件


在构建用例及使用KPEX时应满足以下算子形状约束

表5 KPEX outer_product_mean入参形状约束

| tensor/param | shape/value | 描述 |
| --- | --- | --- |
| input\_ln\_w | [c\_m] | 通过layernorm生成input\_act所需的权重参数 |
| input\_ln\_b | [c\_m] | 通过layernorm生成input\_act所需的偏置参数 |
| left\_proj\_w | [c\_i, c\_m] | 见表4参数 left\_proj\_w |
| left\_proj\_b | [c\_i] | 见表4参数 left\_proj\_b |
| right\_proj\_w | [c\_i, c\_m] | 见表4参数 right\_proj\_w |
| right\_proj\_b | [c\_i] | 见表4参数 right\_proj\_b |
| output\_w | [c\_z, c\_i, c\_i] | 见表4参数 outer\_w |
| output\_b | [c\_z] | 见表4参数 outer\_b |
| act | [n\_seq, n\_res, c\_m] | KPEX输入，经过layernorm生成input\_act |
| mask | [n\_seq, n\_res\_gather] | 见表3参数 mask |
| left\_block\_size | greater than 0 or None | 见表1参数 left\_block\_size |
| right\_block\_size | greater than 0 or None | 见表2参数 right\_block\_size |


#### 示例

C++ interface：

```
// test_outer_product_mean.h
#ifndef KPEX_TPP_ALPHAFOLD_TEST_OPM_H
#define KPEX_TPP_ALPHAFOLD_TEST_OPM_H
#include <ATen/core/Tensor.h>
#include <ATen/ops/empty.h>
#include <ATen/ops/ones.h>
#include <ATen/ops/zeros.h>
#include <ATen/ops/full.h>
#include <ATen/native/cpu/utils.h>
#include <c10/core/ScalarType.h>
namespace alphafold {
at::Tensor test_outer_product_mean(int64_t c_i, int64_t c_m, int64_t c_z, int64_t n_seq, int64_t n_res, int64_t n_res_gather);
}
#endif
// bind.h
#include <torch/extension.h>
#include "test_outer_product_mean.h"
namespace alphafold {
inline void bind(pybind11::module &m)
{
autosubmodule = m.def_submodule("alphafold");
submodule.def("test_outer_product_mean", &test_outer_product_mean, py::arg("c_i"), py::arg("c_m"), py::arg("c_z"), py::arg("n_seq"), py::arg("n_res"), py::arg("n_res_gather"));
}
}
// test.py
import copy
import time
import types
import torch
from torch import nn
import numpy as np
import torch.distributed as dist
import kpex._C as kernel
import kpex
import os
def test_triangle_multiplication(n_res, n_res_gather, c_o, c_i):
out = kernel.alphafold.test_triangle_multiplication(n_res, n_res_gather, c_o, c_i)
return out
// test_outer_product_mean.cpp
#include "test_outer_product_mean.h"
#include "kutacc.h"
#include "outer_product_mean.h"
#include "utils/memory.h"
#include "utils/layernorm.h"
#include <utils/TensorWrapper.h>
namespace alphafold {
at::Tensor test_outer_product_mean(int64_t c_i, int64_t c_m, int64_t c_z, int64_t n_seq, int64_t n_res, int64_t n_res_gather)
{
float a = 0.2f;
float b = 0.5f;
float c = 1.5f;
float d = 2.0f;
at::Tensor act = at::full({n_seq, n_res, c_m}, d, at::TensorOptions().device(kpex::device()).dtype(at::kBFloat16));
at::Tensor mask = at::ones({n_res_gather, n_seq}, at::TensorOptions().device(kpex::device()).dtype(at::kBFloat16));
at::Tensor left_proj = act.new_empty({c_i, n_res, n_seq});
at::Tensor right_proj = act.new_empty({c_i, n_res, n_seq});
at::Tensor left_proj_ = act.new_empty({n_res, c_i, n_seq});
at::Tensor right_proj_ = act.new_empty({n_res, c_i, n_seq});
at::Tensor norm = mask.new_empty({n_res, n_res_gather});
int64_t mask_bias = 0;
at::Tensor input_ln_w = at::full({c_m}, a, at::TensorOptions().device(kpex::device()).dtype(at::kFloat));
at::Tensor input_ln_b = at::full({c_m}, b, at::TensorOptions().device(kpex::device()).dtype(at::kFloat));
at::Tensor left_proj_w = linear_weight_prepack(at::full({c_i, c_m}, c, at::TensorOptions().device(kpex::device()).dtype(at::kBFloat16)));
at::Tensor left_proj_b = at::zeros({c_i}, at::TensorOptions().device(kpex::device()).dtype(at::kFloat));
at::Tensor right_proj_w = linear_weight_prepack(at::ones({c_i, c_m}, at::TensorOptions().device(kpex::device()).dtype(at::kBFloat16)));
at::Tensor right_proj_b = at::zeros({c_i}, at::TensorOptions().device(kpex::device()).dtype(at::kFloat));
at::Tensor output_w = linear_weight_prepack(at::ones({c_z, c_i * c_i}, at::TensorOptions().device(kpex::device()).dtype(at::kBFloat16)));
at::Tensor output_b = at::zeros({c_z}, at::TensorOptions().device(kpex::device()).dtype(at::kFloat));
at::Tensor input_act = layernorm(act.transpose(0, 1), input_ln_w, input_ln_b);
at::Tensor out = act.new_empty({n_res, n_res_gather, c_z});
kutacc::TensorWrapper input_act_tw = convert_to_tensor_wrapper(input_act);
kutacc::TensorWrapper mask_tw = convert_to_tensor_wrapper(mask);
kutacc::TensorWrapper left_proj_w_tw = convert_to_tensor_wrapper(left_proj_w);
kutacc::TensorWrapper left_proj_b_tw = convert_to_tensor_wrapper(left_proj_b);
kutacc::TensorWrapper right_proj_w_tw = convert_to_tensor_wrapper(right_proj_w);
kutacc::TensorWrapper right_proj_b_tw = convert_to_tensor_wrapper(right_proj_b);
kutacc::TensorWrapper left_proj_tw = convert_to_tensor_wrapper(left_proj);
kutacc::TensorWrapper right_proj_tw = convert_to_tensor_wrapper(right_proj);
kutacc::TensorWrapper left_proj_tw_ = convert_to_tensor_wrapper(left_proj_);
kutacc::TensorWrapper right_proj_tw_ = convert_to_tensor_wrapper(right_proj_);
kutacc::TensorWrapper norm_tw = convert_to_tensor_wrapper(norm);
kutacc::TensorWrapper output_w_tw = convert_to_tensor_wrapper(output_w);
kutacc::TensorWrapper output_b_tw = convert_to_tensor_wrapper(output_b);
kutacc::TensorWrapper out_tw = convert_to_tensor_wrapper(out);
int64_t left_block_size = 1024;
int64_t right_block_size = 1024;
kutacc_af2_opm_weights_t_wrapper *opm_weights_ptr = new kutacc_af2_opm_weights_t_wrapper(left_proj_w_tw, left_proj_b_tw, right_proj_w_tw, right_proj_b_tw,
output_w_tw, output_b_tw, c_m, c_i, c_z);
kutacc_af2_opm_act_inputs_t_wrapper *opm_inputs_ptr = new kutacc_af2_opm_act_inputs_t_wrapper(input_act_tw, left_proj_tw, right_proj_tw, left_proj_tw_,
right_proj_tw_, n_seq, n_res);
kutacc_af2_opm_mask_inputs_t_wrapper *opm_mask_ptr = new kutacc_af2_opm_mask_inputs_t_wrapper(mask_tw, norm_tw, n_res_gather, mask_bias);
kutacc_af2_outer_product_mean_calc_left_and_right_mul(opm_inputs_ptr, opm_mask_ptr, opm_weights_ptr);
kutacc_af2_outer_product_mean_chunk(opm_inputs_ptr, opm_mask_ptr, opm_weights_ptr, out_tw.get_tensor(), left_block_size, right_block_size);
delete opm_inputs_ptr;
delete opm_weights_ptr;
delete opm_mask_ptr;
return out;
}
}
```
