| 58 | } |
| 59 | |
| 60 | void kern_default(const ConvBiasImpl::NCBKernParam& p) { |
| 61 | dt_byte* workspace_ptr = static_cast<dt_byte*>(p.workspace_ptr); |
| 62 | |
| 63 | auto filter_meta_ptr = |
| 64 | reinterpret_cast<const ConvBiasForward::CanonizedFilterMeta*>( |
| 65 | &p.filter_meta); |
| 66 | auto filter_meta = *filter_meta_ptr; |
| 67 | auto layouts = get_layouts(p); |
| 68 | |
| 69 | TensorND src{ |
| 70 | reinterpret_cast<dt_byte*>(const_cast<void*>(p.src_ptr.get_ptr())), |
| 71 | layouts[0]}; |
| 72 | TensorND filter{const_cast<void*>(p.filter_ptr.get_ptr()), layouts[1]}; |
| 73 | auto bias_ptr = reinterpret_cast<dt_byte*>(const_cast<void*>(p.bias_ptr.get_ptr())); |
| 74 | TensorND bias{bias_ptr, layouts[2]}; |
| 75 | TensorND dst{ |
| 76 | reinterpret_cast<dt_byte*>(const_cast<void*>(p.dst_ptr.get_ptr())), |
| 77 | layouts[3]}; |
| 78 | |
| 79 | auto sfb = dst; |
| 80 | if (bias.layout.dtype.enumv() != dst.layout.dtype.enumv()) { |
| 81 | // intermediate result |
| 82 | sfb = TensorND{workspace_ptr, TensorLayout{dst.layout, bias.layout.dtype}}; |
| 83 | } |
| 84 | #define DISPATCH_RAW(in_dt, bias_dt, out_dt, cmode, func) \ |
| 85 | else if ( \ |
| 86 | src.layout.dtype.enumv() == DTypeTrait<dtype::in_dt>::enumv && \ |
| 87 | filter.layout.dtype.enumv() == DTypeTrait<dtype::in_dt>::enumv && \ |
| 88 | (!bias.layout.dtype.valid() || \ |
| 89 | bias.layout.dtype.enumv() == DTypeTrait<dtype::bias_dt>::enumv) && \ |
| 90 | sfb.layout.dtype.enumv() == DTypeTrait<dtype::out_dt>::enumv && \ |
| 91 | p.compute_mode == param::ConvBias::ComputeMode::cmode) { \ |
| 92 | func(src, filter, bias, sfb, workspace_ptr, filter_meta); \ |
| 93 | } |
| 94 | #define DISPATCH(in_dt, out_dt) \ |
| 95 | DISPATCH_RAW( \ |
| 96 | in_dt, out_dt, out_dt, DEFAULT, \ |
| 97 | (megdnn::naive::convolution::forward_bias< \ |
| 98 | DTypeTrait<dtype::in_dt>::ctype, DTypeTrait<dtype::in_dt>::ctype, \ |
| 99 | DTypeTrait<dtype::out_dt>::ctype, \ |
| 100 | DTypeTrait<dtype::out_dt>::ctype>)) |
| 101 | if (0) { |
| 102 | } |
| 103 | DISPATCH(Float32, Float32) |
| 104 | DISPATCH(Int8, Int16) |
| 105 | DISPATCH(Int8, Int32) |
| 106 | DISPATCH(QuantizedS8, QuantizedS32) |
| 107 | DISPATCH(Quantized8Asymm, QuantizedS32) |
| 108 | #if !MEGDNN_DISABLE_FLOAT16 |
| 109 | DISPATCH(Float16, Float16) |
| 110 | DISPATCH_RAW( |
| 111 | Float16, Float16, Float16, FLOAT32, |
| 112 | (megdnn::naive::convolution::forward_bias< |
| 113 | dt_float16, dt_float16, dt_float16, dt_float32>)) |
| 114 | #endif |
| 115 | else { |
| 116 | megdnn_throw(ssprintf( |
| 117 | "unsupported naive ConvBias(%s, %s, %s) -> %s", src.layout.dtype.name(), |
no test coverage detected