Skip to content
Closed
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
23 changes: 19 additions & 4 deletions Deeploy/Targets/PULPOpen/Templates/FloatConvGradTemplate.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,15 +142,30 @@ def hoistTransientBuffers(self, ctxt: NetworkContext,
return ctxt, operatorRepresentation, [im2col_name, bt_name]


class _ConvGradXTemplate(NodeTemplate):
"""Captures the FULL (untiled) output spatial dims before TilingVariableReplacement
rewrites dim_im_out_* into tile refs, so the gather kernel can numerically detect
spatial tiling and fall back to scatter."""

def alignToContext(self, ctxt: NetworkContext,
operatorRepresentation: OperatorRepresentation) -> Tuple[NetworkContext, Dict, List[str]]:
if not _is_tiled_expr(operatorRepresentation.get('dim_im_out_x')):
operatorRepresentation['full_dim_im_out_x'] = operatorRepresentation['dim_im_out_x']
operatorRepresentation['full_dim_im_out_y'] = operatorRepresentation['dim_im_out_y']
operatorRepresentation.setdefault('full_dim_im_out_x', operatorRepresentation['dim_im_out_x'])
operatorRepresentation.setdefault('full_dim_im_out_y', operatorRepresentation['dim_im_out_y'])
return ctxt, operatorRepresentation, []


# Templates for ConvGradX operations
referenceConvGradX2DTemplate = NodeTemplate("""
referenceConvGradX2DTemplate = _ConvGradXTemplate("""
// 2D FP ConvGradX (dX) NCHW trainlib naive (Name: ${nodeName}, Op: ${nodeOp})
${grad_out_type.typeName} ref_${grad_out} = ${grad_out}; // dY
${weight_type.typeName} ref_${weight} = ${weight}; // W
${grad_in_type.typeName} ref_${grad_in} = ${grad_in}; // dX

for (uint32_t n=0; n<${batch}; ++n) {
PULP_ConvGradX2d_fp${grad_out_type.referencedType.typeWidth}_fp${weight_type.referencedType.typeWidth}_fp${grad_in_type.referencedType.typeWidth}_CHW_scatter_tiled(
PULP_ConvGradX2d_fp${grad_out_type.referencedType.typeWidth}_fp${weight_type.referencedType.typeWidth}_fp${grad_in_type.referencedType.typeWidth}_CHW_gather_tiled(
ref_${grad_out},
${dim_im_out_x}, ${dim_im_out_y}, ${ch_im_out},
ref_${weight},
Expand All @@ -161,8 +176,8 @@ def hoistTransientBuffers(self, ctxt: NetworkContext,
${dim_im_in_x}, ${dim_im_in_y},
${padding_y_top}, ${padding_y_bottom}, ${padding_x_left}, ${padding_x_right},
${offset_grad_in_h}, ${offset_grad_in_w},
${offset_grad_out_h}, ${offset_grad_out_w}

${offset_grad_out_h}, ${offset_grad_out_w},
${full_dim_im_out_x}, ${full_dim_im_out_y}
);

ref_${grad_out} += ${ch_im_out} * ${dim_im_out_y} * ${dim_im_out_x};
Expand Down
11 changes: 11 additions & 0 deletions TargetLibraries/GAP9/inc/kernel/GAP9Kernels.h
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,17 @@ void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_scatter_tiled(
uint32_t padding_y_bottom, uint16_t offset_grad_in_h,
uint16_t offset_grad_in_w, uint16_t offset_grad_out_h,
uint16_t offset_grad_out_w);
void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_gather_tiled(
const float *__restrict__ pGradOut, uint32_t dim_im_out_x,
uint32_t dim_im_out_y, uint32_t ch_im_out,
const float *__restrict__ pWeight, uint32_t ch_im_in, uint32_t dim_kernel_x,
uint32_t dim_kernel_y, uint32_t stride_h, uint32_t stride_w,
float *__restrict__ pGradIn, uint32_t dim_im_in_x, uint32_t dim_im_in_y,
uint32_t padding_x_left, uint32_t padding_x_right, uint32_t padding_y_top,
uint32_t padding_y_bottom, uint16_t offset_grad_in_h,
uint16_t offset_grad_in_w, uint16_t offset_grad_out_h,
uint16_t offset_grad_out_w, uint32_t full_dim_im_out_x,
uint32_t full_dim_im_out_y);

void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_Im2Col_tiled(
const float *__restrict__ pGradOut, uint32_t dim_im_out_x,
Expand Down
11 changes: 11 additions & 0 deletions TargetLibraries/PULPOpen/inc/kernel/ConvGrad.h
Original file line number Diff line number Diff line change
Expand Up @@ -135,6 +135,17 @@ void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_scatter_tiled(
uint32_t padding_y_bottom, uint16_t offset_grad_in_h,
uint16_t offset_grad_in_w, uint16_t offset_grad_out_h,
uint16_t offset_grad_out_w);
void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_gather_tiled(
const float *__restrict__ pGradOut, uint32_t dim_im_out_x,
uint32_t dim_im_out_y, uint32_t ch_im_out,
const float *__restrict__ pWeight, uint32_t ch_im_in, uint32_t dim_kernel_x,
uint32_t dim_kernel_y, uint32_t stride_h, uint32_t stride_w,
float *__restrict__ pGradIn, uint32_t dim_im_in_x, uint32_t dim_im_in_y,
uint32_t padding_x_left, uint32_t padding_x_right, uint32_t padding_y_top,
uint32_t padding_y_bottom, uint16_t offset_grad_in_h,
uint16_t offset_grad_in_w, uint16_t offset_grad_out_h,
uint16_t offset_grad_out_w, uint32_t full_dim_im_out_x,
uint32_t full_dim_im_out_y);

// Tiled im2col+GEMM with co_block (ForkTransformer)
void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_Im2Col_tiled(
Expand Down
100 changes: 100 additions & 0 deletions TargetLibraries/PULPOpen/src/ConvGrad.c
Original file line number Diff line number Diff line change
Expand Up @@ -240,6 +240,106 @@ void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_scatter_tiled(
}
}

// Gather variant of CHW ConvGradX: output-stationary. For each dX[ci,ih,iw] it
// accumulates the full (co,ky,kx) reduction in a REGISTER and writes dX once,
// instead of the scatter's Cout*K read-modify-writes per output (a dependent
// FP-accumulate chain through L1 that stalls the FPU ~44 cyc/MAC). Same args /
// tile semantics as the scatter kernel; drop-in replacement.
void PULP_ConvGradX2d_fp32_fp32_fp32_CHW_gather_tiled(
const float *__restrict__ pGradOut, uint32_t dim_im_out_x,
uint32_t dim_im_out_y, uint32_t ch_im_out,
const float *__restrict__ pWeight, uint32_t ch_im_in, uint32_t dim_kernel_x,
uint32_t dim_kernel_y, uint32_t stride_h, uint32_t stride_w,
float *__restrict__ pGradIn, uint32_t dim_im_in_x, uint32_t dim_im_in_y,
uint32_t padding_x_left, uint32_t padding_x_right, uint32_t padding_y_top,
uint32_t padding_y_bottom, uint16_t offset_grad_in_h,
uint16_t offset_grad_in_w, uint16_t offset_grad_out_h,
uint16_t offset_grad_out_w, uint32_t full_dim_im_out_x,
uint32_t full_dim_im_out_y) {
(void)padding_x_right;
(void)padding_y_bottom;

// The output-stationary gather needs dY fully readable. When the conv is
// SPATIALLY tiled (this tile's dY spatial < the full dY spatial) the needed
// halo spans neighbouring tiles, so fall back to the halo-safe scatter
// kernel. Numeric tile-vs-full check (not a tiled/untiled flag) so a single
// full tile that is merely L3-managed still takes the fast gather path.
if (dim_im_out_x != full_dim_im_out_x || dim_im_out_y != full_dim_im_out_y) {
PULP_ConvGradX2d_fp32_fp32_fp32_CHW_scatter_tiled(
pGradOut, dim_im_out_x, dim_im_out_y, ch_im_out, pWeight, ch_im_in,
dim_kernel_x, dim_kernel_y, stride_h, stride_w, pGradIn, dim_im_in_x,
dim_im_in_y, padding_x_left, padding_x_right, padding_y_top,
padding_y_bottom, offset_grad_in_h, offset_grad_in_w, offset_grad_out_h,
offset_grad_out_w);
return;
}

const uint32_t Hout_t = dim_im_out_x;
const uint32_t Wout_t = dim_im_out_y;
const uint32_t Hin_t = dim_im_in_x;
const uint32_t Win_t = dim_im_in_y;
const uint32_t Cout = ch_im_out;
const uint32_t Cin = ch_im_in;
const uint32_t P = dim_kernel_x;
const uint32_t Q = dim_kernel_y;
const int32_t pad_top = (int32_t)padding_x_left;
const int32_t pad_left = (int32_t)padding_y_top;
const int32_t sh = (int32_t)stride_h;
const int32_t sw = (int32_t)stride_w;
const int32_t hx0 = (int32_t)offset_grad_in_h;
const int32_t wx0 = (int32_t)offset_grad_in_w;
const int32_t oy0 = (int32_t)offset_grad_out_h;
const int32_t ox0 = (int32_t)offset_grad_out_w;

// Parallel over Cin — each core owns exclusive dX[ci_start..ci_stop]
const int core_id = pi_core_id();
const uint32_t ci_chunk = (Cin + NUM_CORES - 1u) / NUM_CORES;
const uint32_t ci_start = (uint32_t)core_id * ci_chunk;
uint32_t ci_stop = ci_start + ci_chunk;
if (ci_stop > Cin)
ci_stop = Cin;
if (ci_start >= ci_stop)
return;

const size_t dyStrideCo = (size_t)Hout_t * Wout_t; // dY[co] stride
const size_t wStrideCo = (size_t)Cin * P * Q; // W[co] stride

for (uint32_t ci = ci_start; ci < ci_stop; ++ci) {
float *dx_ci = pGradIn + (size_t)ci * Hin_t * Win_t;
for (uint32_t ih = 0; ih < Hin_t; ++ih) {
const int32_t gih = hx0 + (int32_t)ih; // global input row
for (uint32_t iw = 0; iw < Win_t; ++iw) {
const int32_t giw = wx0 + (int32_t)iw;
float acc = 0.0f;
for (uint32_t ky = 0; ky < P; ++ky) {
// gih = ly*sh - pad_top + ky => ly = (gih + pad_top - ky) / sh
const int32_t ly_num = gih + pad_top - (int32_t)ky;
if (ly_num < 0 || (ly_num % sh) != 0)
continue;
const int32_t lyt = (ly_num / sh) - oy0; // tile-local output row
if (lyt < 0 || lyt >= (int32_t)Hout_t)
continue;
for (uint32_t kx = 0; kx < Q; ++kx) {
const int32_t lx_num = giw + pad_left - (int32_t)kx;
if (lx_num < 0 || (lx_num % sw) != 0)
continue;
const int32_t lxt = (lx_num / sw) - ox0;
if (lxt < 0 || lxt >= (int32_t)Wout_t)
continue;
// accumulate the Cout reduction in the register `acc`
const float *dy_p = pGradOut + (size_t)lyt * Wout_t + (uint32_t)lxt;
const float *w_p =
pWeight + ((size_t)ci * P * Q) + (uint32_t)ky * Q + kx;
for (uint32_t co = 0; co < Cout; ++co)
acc += dy_p[co * dyStrideCo] * w_p[co * wStrideCo];
}
}
dx_ci[ih * Win_t + iw] = acc;
}
}
}
}

// ============================================================================
// Regular Conv — Im2Col+GEMM tiled ConvGradX (ForkTransformer)
// ============================================================================
Expand Down
Loading