-
Notifications
You must be signed in to change notification settings - Fork 251
Expand file tree
/
Copy pathnorm_samples.h
More file actions
119 lines (109 loc) · 4.83 KB
/
Copy pathnorm_samples.h
File metadata and controls
119 lines (109 loc) · 4.83 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
#pragma once
/*
* SPDX-FileCopyrightText: Copyright (c) 2020 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: Apache-2.0
*/
#pragma once
#include <iostream>
#include <inttypes.h>
#include <stdlib.h>
#include <string.h>
#include <ctype.h>
#include <assert.h>
#include <tuple>
#include <functional>
#include <cudnn_frontend.h>
/**
* @brief Run a Group BN forward sample with 2 peer stat tensors.
*
* @param tensorDims an array with shape (N, C, H, W) for input tensor dims. Stride in NHWC or NCHW will take care of
memory format
* @param perChannelSum an array with shape (1, C, 1, 1) to denote the sum values for each channel in the input tensor
* @param epsilon a scalar array with shape (1, 1, 1, 1) to represent the epsilon value for the BN
* @param peerDims an array with shape (num GPUs, 2 * C, 1, 1) to denote the tensor dimensions for peer stat tensor in
GBN
*
*/
cudnn_frontend::ExecutionPlan
run_batch_norm_forward(cudnnHandle_t &handle_,
int64_t *tensorDims,
int64_t *perChannelSum,
int64_t *epsilon,
int64_t *peerDims,
cudnnDataType_t in_out_data_type);
/**
* @param xDevPtr input tensor device pointer
* @param yDevPtr output tensor device pointer
* @param scaledevPtr input scale device pointer for BN scaling
* @param biasdevPtr input scale device pointer for BN bias
* @param in_meandevPtr Input mean device pointer
* @param in_vardevPtr Input variance device pointer
* @param out_meandevPtr output mean device pointer
* @param out_vardevPtr output variance device pointer
* @param saved_meandevPtr saved mean device pointer for BN backward
* @param saved_inv_vardevPtr saved inverse variance device pointer for BN backward
* @param peer_devPtr1 peer stat tensor 1 device pointer
* @param peer_devPtr2 peer stat tensor 2 device pointer
* @param epsilon_val episilon value as a double
* @param exponential_decay_factor exponential_decay_factor as a value
*
**/
void
execute_batch_norm_forward(cudnnHandle_t &handle_,
cudnn_frontend::ExecutionPlan plan,
void *xDevPtr,
void *yDevPtr,
void *scaledevPtr,
void *biasdevPtr,
void *in_meandevPtr,
void *in_vardevPtr,
void *out_meandevPtr,
void *out_vardevPtr,
void *saved_meandevPtr,
void *saved_inv_vardevPtr,
void *peer_devPtr1,
void *peer_devPtr2,
double epsilon_val,
double exponential_decay_factor);
/**
* @brief Run a Group BN backward sample with 2 peer stat tensors.
*
* @param tensorDims an array with shape (N, C, H, W) for input tensor dims. Stride in NHWC or NCHW will take care of
* memory format
* @param perChannelSum an array with shape (1, C, 1, 1) to denote the sum values for each channel in the input tensor
* @param epsilon a scalar array with shape (1, 1, 1, 1) to represent the epsilon value for the BN
* @param peerDims an array with shape (num GPUs, 2 * C, 1, 1) to denote the tensor dimensions for peer stat tensor in
* GBN
* @param xDevPtr input tensor device pointer
* @param yDevPtr output tensor device pointer
* @param scaledevPtr input scale device pointer for BN scaling
* @param biasdevPtr input scale device pointer for BN bias
* @param in_meandevPtr Input mean device pointer
* @param in_vardevPtr Input variance device pointer
* @param out_meandevPtr output mean device pointer
* @param out_vardevPtr output variance device pointer
* @param saved_meandevPtr saved mean device pointer for BN backward
* @param saved_inv_vardevPtr saved inverse variance device pointer for BN backward
* @param peer_devPtr1 peer stat tensor 1 device pointer
* @param peer_devPtr2 peer stat tensor 2 device pointer
* @param epsilon_val episilon value as a double
* @param exponential_decay_factor exponential_decay_factor as a value
*
*/
void
run_batch_norm_backward(int64_t *tensorDims,
int64_t *perChannelSum,
int64_t *epsilon,
int64_t *peerDims,
void *xDevPtr,
void *dyDevPtr,
void *scaledevPtr,
void *saved_meandevPtr,
void *saved_inv_vardevPtr,
void *peer_devPtr1,
void *peer_devPtr2,
void *dscaledevPtr,
void *dbiasdevPtr,
void *dxDevPtr,
double epsilon_val,
cudnnDataType_t in_out_data_type);