| 248 | } |
| 249 | |
| 250 | void EmbLayerNormPluginDynamic::configurePlugin(DynamicPluginTensorDesc const* inputs, int32_t nbInputs, |
| 251 | DynamicPluginTensorDesc const* outputs, int32_t nbOutputs) noexcept |
| 252 | { |
| 253 | BERT_DEBUG_MSG("EmbLayerNormPluginDynamic configurePlugin."); |
| 254 | |
| 255 | // Validate input arguments |
| 256 | PLUGIN_ASSERT(nbOutputs == 2); |
| 257 | PLUGIN_ASSERT(nbInputs == 3); |
| 258 | |
| 259 | PLUGIN_ASSERT(inputs[0].desc.dims.nbDims == 2); |
| 260 | int32_t const S = inputs[0].desc.dims.d[SDIM]; |
| 261 | mS = S; |
| 262 | int32_t const B = inputs[0].desc.dims.d[BDIM]; |
| 263 | TRT_UNUSED B; |
| 264 | PLUGIN_ASSERT(mS == static_cast<size_t>(inputs[1].desc.dims.d[SDIM])); |
| 265 | PLUGIN_ASSERT(B == inputs[1].desc.dims.d[BDIM]); |
| 266 | PLUGIN_ASSERT(mS == static_cast<size_t>(inputs[2].desc.dims.d[SDIM])); |
| 267 | PLUGIN_ASSERT(B == inputs[2].desc.dims.d[BDIM]); |
| 268 | |
| 269 | PLUGIN_ASSERT(outputs[0].desc.dims.nbDims == 5); |
| 270 | PLUGIN_ASSERT(static_cast<size_t>(outputs[0].desc.dims.d[SDIM]) == mS); |
| 271 | PLUGIN_ASSERT(outputs[0].desc.dims.d[BDIM] == B); |
| 272 | PLUGIN_ASSERT(static_cast<size_t>(outputs[0].desc.dims.d[2]) == mLd); |
| 273 | PLUGIN_ASSERT(outputs[0].desc.dims.d[3] == 1); |
| 274 | PLUGIN_ASSERT(outputs[0].desc.dims.d[4] == 1); |
| 275 | |
| 276 | if (mUseFullMask) |
| 277 | { |
| 278 | // user force full_mask |
| 279 | PLUGIN_ASSERT(outputs[1].desc.dims.nbDims == 2); |
| 280 | PLUGIN_ASSERT(outputs[1].desc.dims.d[0] == B); |
| 281 | PLUGIN_ASSERT((outputs[1].desc.dims.d[1] == -1) || (outputs[1].desc.dims.d[1] == packedMaskSize384) |
| 282 | || (outputs[1].desc.dims.d[1] == packedMaskSize128)); |
| 283 | } |
| 284 | else |
| 285 | { |
| 286 | // auto detect using mhatype |
| 287 | if (S != -1 && B != -1) |
| 288 | { |
| 289 | PLUGIN_ASSERT(outputs[1].desc.dims.nbDims == 2); |
| 290 | PLUGIN_ASSERT(outputs[1].desc.dims.d[0] == B); |
| 291 | int32_t packedSize = getMHAMaskPackedSize(mSM, mMhaType, S); |
| 292 | TRT_UNUSED packedSize; |
| 293 | PLUGIN_ASSERT(outputs[1].desc.dims.d[1] == -1 || outputs[1].desc.dims.d[1] == packedSize); |
| 294 | } |
| 295 | } |
| 296 | |
| 297 | PLUGIN_ASSERT(inputs[0].desc.type == DataType::kINT32); |
| 298 | PLUGIN_ASSERT(inputs[1].desc.type == DataType::kINT32); |
| 299 | PLUGIN_ASSERT(inputs[2].desc.type == DataType::kINT32); |
| 300 | PLUGIN_ASSERT(outputs[0].desc.type == mType); |
| 301 | PLUGIN_ASSERT(outputs[1].desc.type == DataType::kINT32); |
| 302 | } |
| 303 | |
| 304 | size_t EmbLayerNormPluginDynamic::getWorkspaceSize( |
| 305 | PluginTensorDesc const* inputs, int32_t nbInputs, PluginTensorDesc const* outputs, int32_t nbOutputs) const noexcept |
nothing calls this directly
no test coverage detected