-
Notifications
You must be signed in to change notification settings - Fork 86
Expand file tree
/
Copy pathpersistent_rnn.h
More file actions
261 lines (221 loc) · 10.6 KB
/
Copy pathpersistent_rnn.h
File metadata and controls
261 lines (221 loc) · 10.6 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
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
/*! \file persistent_rnn.h
\date May 10, 2016
\brief C language interface to persistent RNN kernels, modeled after the CUDNNv5 RNN interface
for maximum compatibility.
*/
#if !defined(PRNN_H_)
#define PRNN_H_
#define PRNN_MAJOR 0
#define PRNN_MINOR 2
#define PRNN_PATCHLEVEL 0
#define PRNN_VERSION (PRNN_MAJOR * 1000 + PRNN_MINOR * 100 + PRNN_PATCHLEVEL)
// Standard Library Includes
#include <stddef.h>
#if defined (__cplusplus)
extern "C" {
#endif
struct prnnContext;
typedef struct prnnContext* prnnHandle_t;
size_t prnnGetVersion(void);
/*
* PRNN return codes
*/
typedef enum
{
PRNN_STATUS_SUCCESS = 0,
PRNN_STATUS_NOT_INITIALIZED = 1,
PRNN_STATUS_ALLOC_FAILED = 2,
PRNN_STATUS_BAD_PARAM = 3,
PRNN_STATUS_INTERNAL_ERROR = 4,
PRNN_STATUS_INVALID_VALUE = 5,
PRNN_STATUS_ARCH_MISMATCH = 6,
PRNN_STATUS_MAPPING_ERROR = 7,
PRNN_STATUS_EXECUTION_FAILED = 8,
PRNN_STATUS_NOT_SUPPORTED = 9
} prnnStatus_t;
// human-readable error messages
const char* prnnGetErrorString(prnnStatus_t status);
prnnStatus_t prnnCreate (prnnHandle_t* handle);
prnnStatus_t prnnDestroy (prnnHandle_t handle);
prnnStatus_t prnnSetStream (prnnHandle_t handle, void* streamId);
prnnStatus_t prnnGetStream (prnnHandle_t handle, void** streamId);
/* Data structures to represent input data and the Neural Network Layer */
typedef struct prnnTensorStruct* prnnTensorDescriptor_t;
typedef struct prnnTensorStruct* prnnFilterDescriptor_t;
typedef struct prnnDropoutStruct* prnnDropoutDescriptor_t;
/*
* PRNN data type
*/
typedef enum
{
PRNN_DATA_FLOAT = 0,
PRNN_DATA_DOUBLE = 1,
PRNN_DATA_HALF = 2,
PRNN_INVALID_DATA = 3,
} prnnDataType_t;
/* Maximum supported number of tensor dimensions */
#define PRNN_DIM_MAX 8
/* Create an instance of a generic Tensor descriptor */
prnnStatus_t prnnCreateTensorDescriptor(prnnTensorDescriptor_t* tensorDesc);
typedef enum
{
PRNN_TENSOR_NCHW = 0, /* row major (wStride = 1, hStride = w) */
PRNN_TENSOR_NHWC = 1 /* feature maps interleaved ( cStride = 1 )*/
} prnnTensorFormat_t;
prnnStatus_t prnnSetTensorNdDescriptor(prnnTensorDescriptor_t tensorDesc,
prnnDataType_t dataType,
int nbDims,
const int* dimA,
const int* strideA);
prnnStatus_t prnnGetTensorNdDescriptor(const prnnTensorDescriptor_t tensorDesc,
int nbDimsRequested,
prnnDataType_t* dataType,
int* nbDims,
int* dimA,
int* strideA);
/* Destroy an instance of Tensor4d descriptor */
prnnStatus_t prnnDestroyTensorDescriptor(prnnTensorDescriptor_t tensorDesc);
/* RNN API */
typedef enum
{
PRNN_RNN_RELU = 0, // Stock RNN with ReLu activation
PRNN_RNN_TANH = 1, // Stock RNN with tanh activation
PRNN_LSTM = 2, // LSTM with no peephole connections
PRNN_GRU = 3 // Using h' = tanh(r * Uh(t-1) + Wx) and h = (1 - z) * h' + z * h(t-1);
} prnnRNNMode_t;
typedef enum
{
PRNN_UNIDIRECTIONAL = 0,
PRNN_BIDIRECTIONAL = 1, // Using output concatination at each step. Do we also want to support output sum?
PRNN_REVERSE = 2
} prnnDirectionMode_t;
typedef enum
{
PRNN_LINEAR_INPUT = 0,
PRNN_SKIP_INPUT = 1
} prnnRNNInputMode_t;
typedef enum
{
PRNN_PERSISTENT_BACKEND = 0,
PRNN_CUDNN_BACKEND = 1,
PRNN_BEST_BACKEND = 2
} prnnBackend_t;
struct prnnRNNStruct;
typedef struct prnnRNNStruct* prnnRNNDescriptor_t;
prnnStatus_t prnnCreateRNNDescriptor(prnnRNNDescriptor_t* rnnDesc);
prnnStatus_t prnnDestroyRNNDescriptor(prnnRNNDescriptor_t rnnDesc);
prnnStatus_t prnnSetRNNDescriptor(prnnRNNDescriptor_t rnnDesc,
int hiddenSize,
int numLayers,
prnnDropoutDescriptor_t dropoutDesc, // Between layers, not between recurrent steps.
prnnRNNInputMode_t inputMode,
prnnDirectionMode_t direction,
prnnRNNMode_t mode,
prnnDataType_t dataType,
prnnBackend_t backend);
// dataType in the RNN descriptor is used to determine math precision
// dataType in weight descriptors and input descriptors is used to describe storage
prnnStatus_t prnnGetRNNWorkspaceSize(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const int seqLength,
const prnnTensorDescriptor_t* xDesc,
size_t* sizeInBytes
);
prnnStatus_t prnnGetRNNTrainingReserveSize(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const int seqLength,
const prnnTensorDescriptor_t* xDesc,
size_t* sizeInBytes
);
prnnStatus_t prnnGetRNNParamsSize(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const prnnTensorDescriptor_t* xDesc,
size_t* sizeInBytes
);
prnnStatus_t prnnGetRNNLinLayerMatrixParams(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const int layer,
const prnnTensorDescriptor_t* xDesc,
const prnnFilterDescriptor_t wDesc,
const void* w,
const int linLayerID,
prnnFilterDescriptor_t linLayerMatDesc,
void** linLayerMat
);
prnnStatus_t prnnGetRNNLinLayerBiasParams(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const int layer,
const prnnTensorDescriptor_t* xDesc,
const prnnFilterDescriptor_t wDesc,
const void* w,
const int linLayerID,
prnnFilterDescriptor_t linLayerBiasDesc,
void** linLayerBias
);
prnnStatus_t prnnRNNForward(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const int seqLength,
const prnnTensorDescriptor_t* xDesc,
const void* x,
const prnnTensorDescriptor_t hxDesc,
const void* hx,
const prnnTensorDescriptor_t cxDesc,
const void* cx,
const prnnFilterDescriptor_t wDesc,
const void* w,
const prnnTensorDescriptor_t* yDesc,
void* y,
const prnnTensorDescriptor_t hyDesc,
void* hy,
const prnnTensorDescriptor_t cyDesc,
void* cy,
void* workspace,
size_t workSpaceSizeInBytes,
void* reserveSpace,
size_t reserveSpaceSizeInBytes);
prnnStatus_t prnnRNNBackwardData(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const int seqLength,
const prnnTensorDescriptor_t* yDesc,
const void* y,
const prnnTensorDescriptor_t* dyDesc,
const void* dy,
const prnnTensorDescriptor_t dhyDesc,
const void* dhy,
const prnnTensorDescriptor_t dcyDesc,
const void* dcy,
const prnnFilterDescriptor_t wDesc,
const void* w,
const prnnTensorDescriptor_t hxDesc,
const void* hx,
const prnnTensorDescriptor_t cxDesc,
const void* cx,
const prnnTensorDescriptor_t* dxDesc,
void* dx,
const prnnTensorDescriptor_t dhxDesc,
void* dhx,
const prnnTensorDescriptor_t dcxDesc,
void* dcx,
void* workspace,
size_t workSpaceSizeInBytes,
void* reserveSpace,
size_t reserveSpaceSizeInBytes);
prnnStatus_t prnnRNNBackwardWeights(prnnHandle_t handle,
const prnnRNNDescriptor_t rnnDesc,
const int seqLength,
const prnnTensorDescriptor_t* xDesc,
const void* x,
const prnnTensorDescriptor_t hxDesc,
const void* hx,
const prnnTensorDescriptor_t* yDesc,
const void* y,
const void* workspace,
size_t workSpaceSizeInBytes,
const prnnFilterDescriptor_t dwDesc,
void* dw,
const void* reserveSpace,
size_t reserveSpaceSizeInBytes);
#if defined (__cplusplus)
}
#endif
#endif