Skip to content
Draft
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
72 changes: 39 additions & 33 deletions lib/kernels/include/kernels/batch_norm_kernels.h
Original file line number Diff line number Diff line change
@@ -1,47 +1,53 @@
#ifndef _FLEXFLOW_KERNELS_BATCH_NORM_KERNELS_H
#define _FLEXFLOW_KERNELS_BATCH_NORM_KERNELS_H

#include "kernels/accessor.h"
#include "kernels/allocation.h"
#include "kernels/batch_norm_per_device_state.dtg.h"
#include "kernels/device_handle_t.dtg.h"
#include "kernels/device_stream_t.dtg.h"
#include "kernels/ff_handle.h"
#include "op-attrs/ops/batch_norm_attrs.dtg.h"
#include "op-attrs/tensor_shape.dtg.h"
#include "pcg/device_type.dtg.h"

namespace FlexFlow::Kernels::BatchNorm {
namespace FlexFlow {

std::optional<BatchNormPerDeviceState>
init_kernel(DeviceType device_type,
device_handle_t const &handle,
Allocator &allocator,
float *runningMean,
int output_n,
int output_c,
int output_h,
int output_w,
bool relu);

void forward_kernel(device_stream_t const &stream,
BatchNormPerDeviceState const &per_device_state,
float const *input_ptr,
float *output_ptr,
float const *scale_ptr,
float const *bias_ptr);

void backward_kernel(device_stream_t const &stream,
BatchNormPerDeviceState const &per_device_state,
float const *output_ptr,
float *output_grad_ptr,
float const *input_ptr,
float *input_grad_ptr,
float const *scale_ptr,
float *scale_grad_ptr,
float *bias_grad_ptr,
size_t numElements);

void cleanup_kernel(
batch_norm_init_kernel(DeviceType device_type,
Allocator &allocator,
BatchNormAttrs const &attrs,
TensorShape const &input_shape,
TensorShape const &output_shape);

void batch_norm_forward_kernel(
device_stream_t const &stream,
device_handle_t const &handle,
std::optional<BatchNormPerDeviceState> const &per_device_state,
BatchNormAttrs const &attrs,
GenericTensorAccessorR const &input,
GenericTensorAccessorR const &gamma,
GenericTensorAccessorR const &beta,
GenericTensorAccessorW const &output);

void batch_norm_backward_kernel(
device_stream_t const &stream,
device_handle_t const &handle,
std::optional<BatchNormPerDeviceState> const &per_device_state,
BatchNormAttrs const &attrs,
GenericTensorAccessorR const &output,
GenericTensorAccessorR const &output_grad,
GenericTensorAccessorR const &input,
GenericTensorAccessorW const &input_grad,
GenericTensorAccessorR const &gamma,
GenericTensorAccessorR const &beta,
GenericTensorAccessorW const &gamma_grad,
GenericTensorAccessorW const &beta_grad);

void batch_norm_cleanup_kernel(
DeviceType device_type,
Allocator &allocator,
std::optional<BatchNormPerDeviceState> const &per_device_state);
std::optional<BatchNormPerDeviceState> &per_device_state);

} // namespace FlexFlow

} // namespace FlexFlow::Kernels::BatchNorm
#endif
37 changes: 18 additions & 19 deletions lib/kernels/include/kernels/batch_norm_kernels_cpu.h
Original file line number Diff line number Diff line change
@@ -1,28 +1,27 @@
#ifndef _FLEXFLOW_LIB_KERNELS_INCLUDE_KERNELS_BATCH_NORM_KERNELS_CPU_H
#define _FLEXFLOW_LIB_KERNELS_INCLUDE_KERNELS_BATCH_NORM_KERNELS_CPU_H

#include "kernels/allocation.h"
#include "kernels/batch_norm_per_device_state.dtg.h"
#include "kernels/device_stream_t.dtg.h"
#include "kernels/accessor.h"
#include "op-attrs/ops/batch_norm_attrs.dtg.h"

namespace FlexFlow::Kernels::BatchNorm {
namespace FlexFlow {

void cpu_forward_kernel(BatchNormPerDeviceState const &per_device_state,
float const *input_ptr,
float *output_ptr,
float const *scale_ptr,
float const *bias_ptr);
void batch_norm_cpu_forward_kernel(BatchNormAttrs const &attrs,
GenericTensorAccessorR const &input,
GenericTensorAccessorR const &gamma,
GenericTensorAccessorR const &beta,
GenericTensorAccessorW const &output);

void cpu_backward_kernel(BatchNormPerDeviceState const &per_device_state,
float const *output_ptr,
float *output_grad_ptr,
float const *input_ptr,
float *input_grad_ptr,
float const *scale_ptr,
float *scale_grad_ptr,
float *bias_grad_ptr,
size_t numElements);
void batch_norm_cpu_backward_kernel(BatchNormAttrs const &attrs,
GenericTensorAccessorR const &output,
GenericTensorAccessorR const &output_grad,
GenericTensorAccessorR const &input,
GenericTensorAccessorW const &input_grad,
GenericTensorAccessorR const &gamma,
GenericTensorAccessorR const &beta,
GenericTensorAccessorW const &gamma_grad,
GenericTensorAccessorW const &beta_grad);

} // namespace FlexFlow::Kernels::BatchNorm
} // namespace FlexFlow

#endif
74 changes: 39 additions & 35 deletions lib/kernels/include/kernels/batch_norm_kernels_gpu.h
Original file line number Diff line number Diff line change
@@ -1,43 +1,47 @@
#ifndef _FLEXFLOW_LIB_KERNELS_INCLUDE_KERNELS_BATCH_NORM_KERNELS_GPU_H
#define _FLEXFLOW_LIB_KERNELS_INCLUDE_KERNELS_BATCH_NORM_KERNELS_GPU_H

#include "kernels/accessor.h"
#include "kernels/allocation.h"
#include "kernels/batch_norm_per_device_state.dtg.h"
#include "kernels/device.h"
#include "kernels/ff_handle.h"

namespace FlexFlow::Kernels::BatchNorm {

BatchNormPerDeviceState gpu_init_kernel(PerDeviceFFHandle const &handle,
Allocator &allocator,
float *runningMean,
int output_n,
int output_c,
int output_h,
int output_w,
bool relu);

void gpu_forward_kernel(ffStream_t stream,
BatchNormPerDeviceState const &per_device_statem,
float const *input_ptr,
float *output_ptr,
float const *scale_ptr,
float const *bias_ptr);

void gpu_backward_kernel(ffStream_t stream,
BatchNormPerDeviceState const &per_device_state,
float const *output_ptr,
float *output_grad_ptr,
float const *input_ptr,
float *input_grad_ptr,
float const *scale_ptr,
float *scale_grad_ptr,
float *bias_grad_ptr,
size_t numElements);

void gpu_cleanup_kernel(Allocator &allocator,
BatchNormPerDeviceState &per_device_state);

} // namespace FlexFlow::Kernels::BatchNorm
#include "op-attrs/ops/batch_norm_attrs.dtg.h"

namespace FlexFlow {

BatchNormPerDeviceState
batch_norm_gpu_init_kernel(Allocator &allocator,
BatchNormAttrs const &attrs,
TensorShape const &input_shape,
TensorShape const &output_shape);

void batch_norm_gpu_forward_kernel(
ffStream_t stream,
PerDeviceFFHandle const &handle,
BatchNormPerDeviceState const &per_device_state,
BatchNormAttrs const &attrs,
GenericTensorAccessorR const &input,
GenericTensorAccessorR const &gamma,
GenericTensorAccessorR const &beta,
GenericTensorAccessorW const &output);

void batch_norm_gpu_backward_kernel(
ffStream_t stream,
PerDeviceFFHandle const &handle,
BatchNormPerDeviceState const &per_device_state,
BatchNormAttrs const &attrs,
GenericTensorAccessorR const &output,
GenericTensorAccessorR const &output_grad,
GenericTensorAccessorR const &input,
GenericTensorAccessorW const &input_grad,
GenericTensorAccessorR const &gamma,
GenericTensorAccessorR const &beta,
GenericTensorAccessorW const &gamma_grad,
GenericTensorAccessorW const &beta_grad);

void batch_norm_gpu_cleanup_kernel(Allocator &allocator,
BatchNormPerDeviceState &per_device_state);

} // namespace FlexFlow

#endif
34 changes: 7 additions & 27 deletions lib/kernels/include/kernels/batch_norm_per_device_state.dtg.toml
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,8 @@ features = []

includes = [
"kernels/device.h",
"kernels/ff_handle.h",
]

[[fields]]
name = "handle"
type = "::FlexFlow::PerDeviceFFHandle"

[[fields]]
name = "inputTensor"
type = "ffTensorDescriptor_t"
Expand All @@ -24,10 +19,6 @@ type = "ffTensorDescriptor_t"
name = "biasTensor"
type = "ffTensorDescriptor_t"

[[fields]]
name = "actiDesc"
type = "ffActivationDescriptor_t"

[[fields]]
name = "mode"
type = "ffBatchNormMode_t"
Expand All @@ -49,21 +40,10 @@ name = "saveVar"
type = "float *"

[[fields]]
name = "output_n"
type = "int"

[[fields]]
name = "output_c"
type = "int"

[[fields]]
name = "output_h"
type = "int"

[[fields]]
name = "output_w"
type = "int"

[[fields]]
name = "relu"
type = "bool"
name = "gradSums"
type = "float *"
docstring = '''
Scratch for the fused backward pass: the per-channel sums of the gradient and
of the gradient times the normalized input, handed from the kernel that
reduces them to the kernel that applies them.
'''
11 changes: 8 additions & 3 deletions lib/kernels/include/kernels/create_accessor_with_contents.h
Original file line number Diff line number Diff line change
Expand Up @@ -155,7 +155,8 @@ GenericTensorAccessorW create_4d_accessor_w_with_contents(
type_to_data_type_enum_v<T>,
};

GenericTensorAccessorW accessor = allocator.allocate_tensor(shape);
Allocator cpu_allocator = create_local_cpu_memory_allocator();
GenericTensorAccessorW cpu_accessor = cpu_allocator.allocate_tensor(shape);

for (nonnegative_int dim0_idx :
nonnegative_range(dim0_size.nonnegative_int_from_positive_int())) {
Expand All @@ -165,7 +166,7 @@ GenericTensorAccessorW create_4d_accessor_w_with_contents(
nonnegative_range(dim2_size.nonnegative_int_from_positive_int())) {
for (nonnegative_int dim3_idx :
nonnegative_range(dim3_size.nonnegative_int_from_positive_int())) {
accessor.at<type_to_data_type_enum_v<T>>(TensorDimsCoord{
cpu_accessor.at<type_to_data_type_enum_v<T>>(TensorDimsCoord{
FFOrdered{dim0_idx, dim1_idx, dim2_idx, dim3_idx}}) =
contents.at(dim0_idx.unwrap_nonnegative())
.at(dim1_idx.unwrap_nonnegative())
Expand All @@ -176,7 +177,11 @@ GenericTensorAccessorW create_4d_accessor_w_with_contents(
}
}

return accessor;
GenericTensorAccessorW result = allocator.allocate_tensor(shape);
copy_accessor_data_to_l_from_r(
result, read_only_accessor_from_write_accessor(cpu_accessor));

return result;
}

template <typename T>
Expand Down
Loading
Loading