MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT / configurePlugin

Method configurePlugin

plugin/embLayerNormPlugin/embLayerNormPlugin.cpp:250–302  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

248}
249
250void 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
304size_t EmbLayerNormPluginDynamic::getWorkspaceSize(
305 PluginTensorDesc const* inputs, int32_t nbInputs, PluginTensorDesc const* outputs, int32_t nbOutputs) const noexcept

Callers

nothing calls this directly

Calls 1

getMHAMaskPackedSizeFunction · 0.85

Tested by

no test coverage detected