MCPcopy Create free account
hub / github.com/alibaba/MNN / MNNExp

Function MNNExp

source/backend/cpu/arm/CommonOptFunctionNeon.cpp:478–575  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

476
477
478void MNNExp(float* destPtr, const float* srcPtr, float* offset, size_t size) {
479 float32x4_t maxVec = vdupq_n_f32(-offset[2]);
480 float32x4_t sumVec0 = vdupq_n_f32(0);
481 float32x4_t sumVec1 = vdupq_n_f32(0);
482 if (offset[0] == 1.f && offset[1] == 0.f) {
483 while (size >= 8) {
484 float32x4_t srcVec0 = vld1q_f32(srcPtr);
485 float32x4_t srcVec1 = vld1q_f32(srcPtr + 4);
486 auto subVec0 = vsubq_f32(srcVec0, maxVec);
487 auto subVec1 = vsubq_f32(srcVec1, maxVec);
488 auto expVec0 = expApprox(subVec0);
489 auto expVec1 = expApprox(subVec1);
490 vst1q_f32(destPtr, expVec0);
491 vst1q_f32(destPtr + 4, expVec1);
492 sumVec0 = vaddq_f32(sumVec0, expVec0);
493 sumVec1 = vaddq_f32(sumVec1, expVec1);
494 srcPtr += 8;
495 destPtr += 8;
496 size -= 8;
497
498 }
499 while (size >= 4) {
500 float32x4_t srcVec0 = vld1q_f32(srcPtr);
501 auto subVec0 = vsubq_f32(srcVec0, maxVec);
502 auto expVec0 = expApprox(subVec0);
503 sumVec0 = vaddq_f32(sumVec0, expVec0);
504 vst1q_f32(destPtr, expVec0);
505 srcPtr += 4;
506 destPtr += 4;
507 size -= 4;
508 }
509 //merge
510 sumVec0 = vaddq_f32(sumVec0, sumVec1);
511 float32x2_t sumP = vpadd_f32(vget_low_f32(sumVec0), vget_high_f32(sumVec0));
512 sumP = vpadd_f32(sumP, sumP);
513 auto newSum = vget_lane_f32(sumP, 0);
514 if (size > 0) {
515 float tmp[4];
516 memcpy(tmp, srcPtr, size * sizeof(float));
517 float32x4_t srcVec0 = vld1q_f32(tmp);
518 auto subVec0 = vsubq_f32(srcVec0, maxVec);
519 auto expVec0 = expApprox(subVec0);
520 vst1q_f32(tmp, expVec0);
521 for (int i = 0; i < size; ++i) {
522 newSum += tmp[i];
523 destPtr[i] = tmp[i];
524 }
525 }
526 offset[3] += newSum;
527 } else {
528 float32x4_t c0 = vdupq_n_f32(offset[0]);
529 float32x4_t c1 = vdupq_n_f32(offset[1]);
530 while (size >= 8) {
531 float32x4_t srcVec0 = vld1q_f32(srcPtr);
532 float32x4_t srcVec1 = vld1q_f32(srcPtr + 4);
533 auto subVec0 = vsubq_f32(vmulq_f32(srcVec0, c0), maxVec);
534 auto subVec1 = vsubq_f32(vmulq_f32(srcVec1, c0), maxVec);
535 auto expVec0 = vaddq_f32(expApprox(subVec0), c1);

Callers 15

gated_delta_rule_mnnMethod · 0.50
_EXPFunction · 0.50
_EXPM1Function · 0.50
___MNNSoftmaxFunction · 0.50
MNN_CONCURRENCY_BEGINFunction · 0.50
_AVX512_MNNSoftmaxFunction · 0.50
_AVX_MNNSoftmaxFunction · 0.50
_SSE_MNNSoftmaxFunction · 0.50
MNNSiLuFunction · 0.50
MNNSiLuLowpFunction · 0.50
operator()Method · 0.50
operator()Method · 0.50

Calls 1

expApproxFunction · 0.70

Tested by

no test coverage detected