Add native self-hosted instance connection to fluxer_desktop
Trimmed monorepo checkout (fluxer_desktop + packages/voice_engine_v2 + tools/ci) with a "Connect to a Different Server" menu item and popout that lets the desktop app switch to any self-hosted Fluxer instance, plus fixes for well-known discovery on single-domain self-hosted deployments and a false-positive ERR_ABORTED on same-origin client redirects during the switch. Defaults to chat.fluxr.chat and uses an isolated userData directory from the official build.
This commit is contained in:
+59
@@ -0,0 +1,59 @@
|
||||
// Tile size: 10x1
|
||||
// Accumulators: 0-9
|
||||
// Col regs: 10-19
|
||||
// Row regs: 20, 21
|
||||
|
||||
vbroadcastss zmm20, dword ptr [rcx]
|
||||
|
||||
vmovaps zmm10, [rax + 0]
|
||||
vmovaps zmm11, [rax + 64]
|
||||
vmovaps zmm12, [rax + 128]
|
||||
vmovaps zmm13, [rax + 192]
|
||||
vmovaps zmm14, [rax + 256]
|
||||
|
||||
vfmadd231ps zmm0, zmm10, zmm20
|
||||
vfmadd231ps zmm1, zmm11, zmm20
|
||||
vfmadd231ps zmm2, zmm12, zmm20
|
||||
vfmadd231ps zmm3, zmm13, zmm20
|
||||
vfmadd231ps zmm4, zmm14, zmm20
|
||||
|
||||
vmovaps zmm15, [rax + 320]
|
||||
vmovaps zmm16, [rax + 384]
|
||||
vmovaps zmm17, [rax + 448]
|
||||
vmovaps zmm18, [rax + 512]
|
||||
vmovaps zmm19, [rax + 576]
|
||||
|
||||
vfmadd231ps zmm5, zmm10, zmm20
|
||||
vfmadd231ps zmm6, zmm11, zmm20
|
||||
vfmadd231ps zmm7, zmm12, zmm20
|
||||
vfmadd231ps zmm8, zmm13, zmm20
|
||||
vfmadd231ps zmm9, zmm14, zmm20
|
||||
|
||||
vbroadcastss zmm21, dword ptr [rcx + 4]
|
||||
|
||||
vmovaps zmm10, [rax + 640]
|
||||
vmovaps zmm11, [rax + 704]
|
||||
vmovaps zmm12, [rax + 768]
|
||||
vmovaps zmm13, [rax + 832]
|
||||
vmovaps zmm14, [rax + 896]
|
||||
|
||||
vfmadd231ps zmm0, zmm10, zmm21
|
||||
vfmadd231ps zmm1, zmm11, zmm21
|
||||
vfmadd231ps zmm2, zmm12, zmm21
|
||||
vfmadd231ps zmm3, zmm13, zmm21
|
||||
vfmadd231ps zmm4, zmm14, zmm21
|
||||
|
||||
vmovaps zmm15, [rax + 960]
|
||||
vmovaps zmm16, [rax + 1024]
|
||||
vmovaps zmm17, [rax + 1088]
|
||||
vmovaps zmm18, [rax + 1152]
|
||||
vmovaps zmm19, [rax + 1216]
|
||||
|
||||
vfmadd231ps zmm5, zmm10, zmm21
|
||||
vfmadd231ps zmm6, zmm11, zmm21
|
||||
vfmadd231ps zmm7, zmm12, zmm21
|
||||
vfmadd231ps zmm8, zmm13, zmm21
|
||||
vfmadd231ps zmm9, zmm14, zmm21
|
||||
|
||||
add rcx, 8
|
||||
add rax, 1280
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
// Tile size: 10x1
|
||||
// Accumulators: 0-9
|
||||
// Col regs: 10-19
|
||||
// Row regs: 20
|
||||
|
||||
vbroadcastss zmm20, dword ptr [rcx]
|
||||
|
||||
vmovaps zmm10, [rax + 0]
|
||||
vmovaps zmm11, [rax + 64]
|
||||
vmovaps zmm12, [rax + 128]
|
||||
vmovaps zmm13, [rax + 192]
|
||||
vmovaps zmm14, [rax + 256]
|
||||
|
||||
vfmadd231ps zmm0, zmm10, zmm20
|
||||
vfmadd231ps zmm1, zmm11, zmm20
|
||||
vfmadd231ps zmm2, zmm12, zmm20
|
||||
vfmadd231ps zmm3, zmm13, zmm20
|
||||
vfmadd231ps zmm4, zmm14, zmm20
|
||||
|
||||
vmovaps zmm15, [rax + 320]
|
||||
vmovaps zmm16, [rax + 384]
|
||||
vmovaps zmm17, [rax + 448]
|
||||
vmovaps zmm18, [rax + 512]
|
||||
vmovaps zmm19, [rax + 576]
|
||||
|
||||
vfmadd231ps zmm5, zmm10, zmm20
|
||||
vfmadd231ps zmm6, zmm11, zmm20
|
||||
vfmadd231ps zmm7, zmm12, zmm20
|
||||
vfmadd231ps zmm8, zmm13, zmm20
|
||||
vfmadd231ps zmm9, zmm14, zmm20
|
||||
|
||||
add rcx, 4
|
||||
add rax, 320
|
||||
+7
@@ -0,0 +1,7 @@
|
||||
vbroadcastss zmm15, dword ptr [rcx]
|
||||
|
||||
vmovups zmm8, [rax]
|
||||
vfmadd231ps zmm0, zmm15, zmm8
|
||||
|
||||
add rcx, 4
|
||||
add rax, 64
|
||||
+68
@@ -0,0 +1,68 @@
|
||||
vmovups zmm31, [rcx]
|
||||
// vbroadcastss zmm17, [rcx + 4 * 0]
|
||||
// vbroadcastss zmm18, [rcx + 4 * 1]
|
||||
// vbroadcastss zmm19, [rcx + 4 * 2]
|
||||
// vbroadcastss zmm20, [rcx + 4 * 3]
|
||||
// vbroadcastss zmm21, [rcx + 4 * 4]
|
||||
// vbroadcastss zmm22, [rcx + 4 * 5]
|
||||
// vbroadcastss zmm23, [rcx + 4 * 6]
|
||||
// vbroadcastss zmm24, [rcx + 4 * 7]
|
||||
// vbroadcastss zmm25, [rcx + 4 * 8]
|
||||
// vbroadcastss zmm26, [rcx + 4 * 9]
|
||||
// vbroadcastss zmm27, [rcx + 4 * 10]
|
||||
// vbroadcastss zmm28, [rcx + 4 * 11]
|
||||
// vbroadcastss zmm29, [rcx + 4 * 12]
|
||||
// vbroadcastss zmm30, [rcx + 4 * 13]
|
||||
// vbroadcastss zmm31, [rcx + 4 * 14]
|
||||
|
||||
vbroadcastss zmm16, xmm31
|
||||
valignd zmm17, zmm31, zmm31, 1
|
||||
vbroadcastss zmm17, xmm17
|
||||
valignd zmm18, zmm31, zmm31, 2
|
||||
vbroadcastss zmm18, xmm18
|
||||
valignd zmm19, zmm31, zmm31, 3
|
||||
vbroadcastss zmm19, xmm19
|
||||
valignd zmm20, zmm31, zmm31, 4
|
||||
vbroadcastss zmm20, xmm20
|
||||
valignd zmm21, zmm31, zmm31, 5
|
||||
vbroadcastss zmm21, xmm21
|
||||
valignd zmm22, zmm31, zmm31, 6
|
||||
vbroadcastss zmm22, xmm22
|
||||
valignd zmm23, zmm31, zmm31, 7
|
||||
vbroadcastss zmm23, xmm23
|
||||
valignd zmm24, zmm31, zmm31, 8
|
||||
vbroadcastss zmm24, xmm24
|
||||
valignd zmm25, zmm31, zmm31, 9
|
||||
vbroadcastss zmm25, xmm25
|
||||
valignd zmm26, zmm31, zmm31, 10
|
||||
vbroadcastss zmm26, xmm26
|
||||
valignd zmm27, zmm31, zmm31, 11
|
||||
vbroadcastss zmm27, xmm27
|
||||
valignd zmm28, zmm31, zmm31, 12
|
||||
vbroadcastss zmm28, xmm28
|
||||
valignd zmm29, zmm31, zmm31, 13
|
||||
vbroadcastss zmm29, xmm29
|
||||
valignd zmm30, zmm31, zmm31, 14
|
||||
vbroadcastss zmm30, xmm30
|
||||
valignd zmm31, zmm31, zmm31, 15
|
||||
vbroadcastss zmm31, xmm31
|
||||
|
||||
vfmadd231ps zmm0, zmm16, [rax + 0]
|
||||
vfmadd231ps zmm1, zmm17, [rax + 64]
|
||||
vfmadd231ps zmm2, zmm18, [rax + 128]
|
||||
vfmadd231ps zmm3, zmm19, [rax + 192]
|
||||
vfmadd231ps zmm4, zmm20, [rax + 256]
|
||||
vfmadd231ps zmm5, zmm21, [rax + 320]
|
||||
vfmadd231ps zmm6, zmm22, [rax + 384]
|
||||
vfmadd231ps zmm7, zmm23, [rax + 448]
|
||||
vfmadd231ps zmm8, zmm24, [rax + 512]
|
||||
vfmadd231ps zmm9, zmm25, [rax + 576]
|
||||
vfmadd231ps zmm10, zmm26, [rax + 640]
|
||||
vfmadd231ps zmm11, zmm27, [rax + 704]
|
||||
vfmadd231ps zmm12, zmm28, [rax + 768]
|
||||
vfmadd231ps zmm13, zmm29, [rax + 832]
|
||||
vfmadd231ps zmm14, zmm30, [rax + 896]
|
||||
vfmadd231ps zmm15, zmm31, [rax + 960]
|
||||
|
||||
add rcx, 64
|
||||
add rax, 1024
|
||||
+24
@@ -0,0 +1,24 @@
|
||||
// slow
|
||||
vbroadcastss xmm16, dword ptr [rcx]
|
||||
vbroadcastss xmm17, dword ptr [rcx + 4]
|
||||
vbroadcastss xmm18, dword ptr [rcx + 8]
|
||||
vbroadcastss xmm19, dword ptr [rcx + 12]
|
||||
|
||||
// fast
|
||||
vmovups xmm31, [rcx]
|
||||
vbroadcastss zmm16, xmm31
|
||||
valignd xmm17, xmm31, xmm31, 1
|
||||
vbroadcastss zmm17, xmm17
|
||||
valignd xmm18, xmm31, xmm31, 2
|
||||
vbroadcastss zmm18, xmm18
|
||||
valignd xmm19, xmm31, xmm31, 3
|
||||
vbroadcastss zmm19, xmm19
|
||||
|
||||
// commmon
|
||||
vfmadd231ps zmm0, zmm16, [rax + 0]
|
||||
vfmadd231ps zmm1, zmm17, [rax + 64]
|
||||
vfmadd231ps zmm2, zmm18, [rax + 128]
|
||||
vfmadd231ps zmm3, zmm19, [rax + 192]
|
||||
|
||||
add rcx, 16
|
||||
add rax, 256
|
||||
+29
@@ -0,0 +1,29 @@
|
||||
vmovups ymm31, [rcx]
|
||||
|
||||
vbroadcastss zmm16, xmm31
|
||||
valignd ymm17, ymm31, ymm31, 1
|
||||
vbroadcastss zmm17, xmm17
|
||||
valignd ymm18, ymm31, ymm31, 2
|
||||
vbroadcastss zmm18, xmm18
|
||||
valignd ymm19, ymm31, ymm31, 3
|
||||
vbroadcastss zmm19, xmm19
|
||||
valignd ymm20, ymm31, ymm31, 4
|
||||
vbroadcastss zmm20, xmm20
|
||||
valignd ymm21, ymm31, ymm31, 5
|
||||
vbroadcastss zmm21, xmm21
|
||||
valignd ymm22, ymm31, ymm31, 6
|
||||
vbroadcastss zmm22, xmm22
|
||||
valignd ymm23, ymm31, ymm31, 7
|
||||
vbroadcastss zmm23, xmm23
|
||||
|
||||
vfmadd231ps zmm0, zmm16, [rax + 0]
|
||||
vfmadd231ps zmm1, zmm17, [rax + 64]
|
||||
vfmadd231ps zmm2, zmm18, [rax + 128]
|
||||
vfmadd231ps zmm3, zmm19, [rax + 192]
|
||||
vfmadd231ps zmm4, zmm20, [rax + 256]
|
||||
vfmadd231ps zmm5, zmm21, [rax + 320]
|
||||
vfmadd231ps zmm6, zmm22, [rax + 384]
|
||||
vfmadd231ps zmm7, zmm23, [rax + 448]
|
||||
|
||||
add rcx, 32
|
||||
add rax, 512
|
||||
+11
@@ -0,0 +1,11 @@
|
||||
vbroadcastss zmm15, dword ptr [rcx]
|
||||
|
||||
vmovaps zmm8, [rax + 0]
|
||||
vfmadd231ps zmm0, zmm15, zmm8
|
||||
|
||||
vbroadcastss zmm16, dword ptr [rcx + 4]
|
||||
vmovaps zmm9, [rax + 64]
|
||||
vfmadd231ps zmm1, zmm16, zmm9
|
||||
|
||||
add rcx, 8
|
||||
add rax, 128
|
||||
+45
@@ -0,0 +1,45 @@
|
||||
// Tile size: 1x12
|
||||
// Accumulators: 0-11
|
||||
// Col regs: zmm14
|
||||
// Row regs: zmm15
|
||||
|
||||
vmovaps zmm15, [rax]
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 0 * 4]
|
||||
vfmadd231ps zmm0, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 1 * 4]
|
||||
vfmadd231ps zmm1, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 2 * 4]
|
||||
vfmadd231ps zmm2, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 3 * 4]
|
||||
vfmadd231ps zmm3, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 4 * 4]
|
||||
vfmadd231ps zmm4, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 5 * 4]
|
||||
vfmadd231ps zmm5, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 6 * 4]
|
||||
vfmadd231ps zmm6, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 7 * 4]
|
||||
vfmadd231ps zmm7, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 8 * 4]
|
||||
vfmadd231ps zmm8, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 9 * 4]
|
||||
vfmadd231ps zmm9, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 10 * 4]
|
||||
vfmadd231ps zmm10, zmm15, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 11 * 4]
|
||||
vfmadd231ps zmm11, zmm15, zmm14
|
||||
|
||||
add rcx, 48
|
||||
add rax, 64
|
||||
+53
@@ -0,0 +1,53 @@
|
||||
// Accumulators: 0-9
|
||||
// Columns: 15-16
|
||||
// Rows: 10-14
|
||||
vbroadcastss zmm10, dword ptr [rcx]
|
||||
vbroadcastss zmm11, dword ptr [rcx + 4]
|
||||
vbroadcastss zmm12, dword ptr [rcx + 8]
|
||||
vbroadcastss zmm13, dword ptr [rcx + 12]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 16]
|
||||
|
||||
vmovaps zmm15, [rax]
|
||||
vmovaps zmm16, [rax + 64]
|
||||
|
||||
vfmadd231ps zmm0, zmm15, zmm10
|
||||
vfmadd231ps zmm1, zmm16, zmm10
|
||||
|
||||
vfmadd231ps zmm2, zmm15, zmm11
|
||||
vfmadd231ps zmm3, zmm16, zmm11
|
||||
|
||||
vfmadd231ps zmm4, zmm15, zmm12
|
||||
vfmadd231ps zmm5, zmm16, zmm12
|
||||
|
||||
vfmadd231ps zmm6, zmm15, zmm13
|
||||
vfmadd231ps zmm7, zmm16, zmm13
|
||||
|
||||
vfmadd231ps zmm8, zmm15, zmm14
|
||||
vfmadd231ps zmm9, zmm16, zmm14
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx + 20]
|
||||
vbroadcastss zmm11, dword ptr [rcx + 24]
|
||||
vbroadcastss zmm12, dword ptr [rcx + 28]
|
||||
vbroadcastss zmm13, dword ptr [rcx + 32]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 36]
|
||||
|
||||
vmovaps zmm15, [rax + 128]
|
||||
vmovaps zmm16, [rax + 192]
|
||||
|
||||
vfmadd231ps zmm0, zmm15, zmm10
|
||||
vfmadd231ps zmm1, zmm16, zmm10
|
||||
|
||||
vfmadd231ps zmm2, zmm15, zmm11
|
||||
vfmadd231ps zmm3, zmm16, zmm11
|
||||
|
||||
vfmadd231ps zmm4, zmm15, zmm12
|
||||
vfmadd231ps zmm5, zmm16, zmm12
|
||||
|
||||
vfmadd231ps zmm6, zmm15, zmm13
|
||||
vfmadd231ps zmm7, zmm16, zmm13
|
||||
|
||||
vfmadd231ps zmm8, zmm15, zmm14
|
||||
vfmadd231ps zmm9, zmm16, zmm14
|
||||
|
||||
add rcx, 40
|
||||
add rax, 256
|
||||
+30
@@ -0,0 +1,30 @@
|
||||
// Accumulators: 0-9
|
||||
// Columns: 15
|
||||
// Rows: 10-14
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx]
|
||||
vbroadcastss zmm11, dword ptr [rcx + 4]
|
||||
vbroadcastss zmm12, dword ptr [rcx + 8]
|
||||
vbroadcastss zmm13, dword ptr [rcx + 12]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 16]
|
||||
|
||||
vmovaps zmm15, [rax]
|
||||
vmovaps zmm16, [rax + 64]
|
||||
|
||||
vfmadd231ps zmm0, zmm15, zmm10
|
||||
vfmadd231ps zmm1, zmm16, zmm10
|
||||
|
||||
vfmadd231ps zmm2, zmm15, zmm11
|
||||
vfmadd231ps zmm3, zmm16, zmm11
|
||||
|
||||
vfmadd231ps zmm4, zmm15, zmm12
|
||||
vfmadd231ps zmm5, zmm16, zmm12
|
||||
|
||||
vfmadd231ps zmm6, zmm15, zmm13
|
||||
vfmadd231ps zmm7, zmm16, zmm13
|
||||
|
||||
vfmadd231ps zmm8, zmm15, zmm14
|
||||
vfmadd231ps zmm9, zmm16, zmm14
|
||||
|
||||
add rcx, 20
|
||||
add rax, 128
|
||||
+71
@@ -0,0 +1,71 @@
|
||||
// Tile size: 2x6
|
||||
// Accumulators: 0-11
|
||||
// Col regs: zmm14-15
|
||||
// Row regs: zmm12-13
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx]
|
||||
vmovaps zmm12, [rax]
|
||||
vmovaps zmm13, [rax + 64]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm14
|
||||
vfmadd231ps zmm1, zmm13, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm15
|
||||
vfmadd231ps zmm3, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm5, zmm13, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 20]
|
||||
|
||||
vfmadd231ps zmm8, zmm12, zmm14
|
||||
vfmadd231ps zmm9, zmm13, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx+24]
|
||||
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
vfmadd231ps zmm11, zmm13, zmm15
|
||||
|
||||
// Iteration two
|
||||
vmovaps zmm12, [rax + 128]
|
||||
vmovaps zmm13, [rax + 192]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 24 + 4]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm14
|
||||
vfmadd231ps zmm1, zmm13, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 24 + 8]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm15
|
||||
vfmadd231ps zmm3, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 24 + 12]
|
||||
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm5, zmm13, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 24 + 16]
|
||||
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 24 + 20]
|
||||
|
||||
vfmadd231ps zmm8, zmm12, zmm14
|
||||
vfmadd231ps zmm9, zmm13, zmm14
|
||||
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
vfmadd231ps zmm11, zmm13, zmm15
|
||||
|
||||
add rax, 256
|
||||
add rcx, 48
|
||||
+39
@@ -0,0 +1,39 @@
|
||||
// Tile size: 2x6
|
||||
// Accumulators: 0-11
|
||||
// Col regs: zmm14-15
|
||||
// Row regs: zmm12-13
|
||||
|
||||
// Load ordered by earliest use for first 2x2 block
|
||||
vbroadcastss zmm14, dword ptr [rcx]
|
||||
vmovaps zmm12, [rax]
|
||||
vmovaps zmm13, [rax + 64]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm14
|
||||
vfmadd231ps zmm1, zmm13, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm15
|
||||
vfmadd231ps zmm3, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm5, zmm13, zmm14
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 20]
|
||||
|
||||
vfmadd231ps zmm8, zmm12, zmm14
|
||||
vfmadd231ps zmm9, zmm13, zmm14
|
||||
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
vfmadd231ps zmm11, zmm13, zmm15
|
||||
|
||||
add rax, 128
|
||||
add rcx, 24
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
// Tile size: 3x4
|
||||
// Accumulators: 0-11
|
||||
// Col regs: zmm12-14
|
||||
// Row regs: zmm15
|
||||
|
||||
vmovaps zmm12, [rax]
|
||||
vmovaps zmm13, [rax+64]
|
||||
vmovaps zmm14, [rax+128]
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 0]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm15
|
||||
vfmadd231ps zmm1, zmm13, zmm15
|
||||
vfmadd231ps zmm2, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps zmm3, zmm12, zmm15
|
||||
vfmadd231ps zmm4, zmm13, zmm15
|
||||
vfmadd231ps zmm5, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
vfmadd231ps zmm8, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps zmm9, zmm12, zmm15
|
||||
vfmadd231ps zmm10, zmm13, zmm15
|
||||
vfmadd231ps zmm11, zmm14, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 192]
|
||||
vmovaps zmm13, [rax + 256]
|
||||
vmovaps zmm14, [rax + 320]
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm15
|
||||
vfmadd231ps zmm1, zmm13, zmm15
|
||||
vfmadd231ps zmm2, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 20]
|
||||
|
||||
vfmadd231ps zmm3, zmm12, zmm15
|
||||
vfmadd231ps zmm4, zmm13, zmm15
|
||||
vfmadd231ps zmm5, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 24]
|
||||
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
vfmadd231ps zmm8, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 28]
|
||||
|
||||
vfmadd231ps zmm9, zmm12, zmm15
|
||||
vfmadd231ps zmm10, zmm13, zmm15
|
||||
vfmadd231ps zmm11, zmm14, zmm15
|
||||
|
||||
add rax, 384
|
||||
add rcx, 32
|
||||
+35
@@ -0,0 +1,35 @@
|
||||
// Tile size: 3x4
|
||||
// Accumulators: 0-11
|
||||
// Col regs: zmm12-14
|
||||
// Row regs: zmm15
|
||||
|
||||
vmovaps zmm12, [rax]
|
||||
vmovaps zmm13, [rax+64]
|
||||
vmovaps zmm14, [rax+128]
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 0]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm15
|
||||
vfmadd231ps zmm1, zmm13, zmm15
|
||||
vfmadd231ps zmm2, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps zmm3, zmm12, zmm15
|
||||
vfmadd231ps zmm4, zmm13, zmm15
|
||||
vfmadd231ps zmm5, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
vfmadd231ps zmm8, zmm14, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps zmm9, zmm12, zmm15
|
||||
vfmadd231ps zmm10, zmm13, zmm15
|
||||
vfmadd231ps zmm11, zmm14, zmm15
|
||||
|
||||
add rax, 192
|
||||
add rcx, 16
|
||||
+69
@@ -0,0 +1,69 @@
|
||||
// Tile size: 4x3
|
||||
// Accumulators: 0-11
|
||||
// Col regs: zmm12
|
||||
// Row regs: zmm13-15
|
||||
|
||||
// Load col of A
|
||||
vmovaps zmm12, [rax]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vbroadcastss zmm13, dword ptr [rcx + 0]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 4]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 8]
|
||||
|
||||
// N.B. Stepping cols in inner loop
|
||||
vfmadd231ps zmm0, zmm12, zmm13
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm8, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax+64]
|
||||
|
||||
vfmadd231ps zmm1, zmm12, zmm13
|
||||
vfmadd231ps zmm5, zmm12, zmm14
|
||||
vfmadd231ps zmm9, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax+128]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm13
|
||||
vfmadd231ps zmm6, zmm12, zmm14
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax+192]
|
||||
|
||||
vfmadd231ps zmm3, zmm12, zmm13
|
||||
vfmadd231ps zmm7, zmm12, zmm14
|
||||
vfmadd231ps zmm11, zmm12, zmm15
|
||||
|
||||
// Load col of A, switching col!
|
||||
vmovaps zmm13, [rax + 256]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vbroadcastss zmm14, dword ptr [rcx + 12]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 16]
|
||||
vbroadcastss zmm12, dword ptr [rcx + 20]
|
||||
|
||||
// N.B. Stepping cols in inner loop
|
||||
vfmadd231ps zmm0, zmm13, zmm14
|
||||
vfmadd231ps zmm4, zmm13, zmm15
|
||||
vfmadd231ps zmm8, zmm13, zmm12
|
||||
|
||||
vmovaps zmm13, [rax + 320]
|
||||
|
||||
vfmadd231ps zmm1, zmm13, zmm14
|
||||
vfmadd231ps zmm5, zmm13, zmm15
|
||||
vfmadd231ps zmm9, zmm13, zmm12
|
||||
|
||||
vmovaps zmm13, [rax + 384]
|
||||
|
||||
vfmadd231ps zmm2, zmm13, zmm14
|
||||
vfmadd231ps zmm6, zmm13, zmm15
|
||||
vfmadd231ps zmm10, zmm13, zmm12
|
||||
|
||||
vmovaps zmm13, [rax + 448]
|
||||
|
||||
vfmadd231ps zmm3, zmm13, zmm14
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
vfmadd231ps zmm11, zmm13, zmm12
|
||||
|
||||
add rcx, 24
|
||||
add rax, 512
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
// Tile size: 4x3
|
||||
// Accumulators: 0-11
|
||||
// Col regs: zmm12
|
||||
// Row regs: zmm13-15
|
||||
|
||||
// Load col of A
|
||||
vmovaps zmm12, [rax]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vbroadcastss zmm13, dword ptr [rcx + 0]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 4]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 8]
|
||||
|
||||
// N.B. Stepping cols in inner loop
|
||||
vfmadd231ps zmm0, zmm12, zmm13
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm8, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax+64]
|
||||
|
||||
vfmadd231ps zmm1, zmm12, zmm13
|
||||
vfmadd231ps zmm5, zmm12, zmm14
|
||||
vfmadd231ps zmm9, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax+128]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm13
|
||||
vfmadd231ps zmm6, zmm12, zmm14
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax+192]
|
||||
|
||||
vfmadd231ps zmm3, zmm12, zmm13
|
||||
vfmadd231ps zmm7, zmm12, zmm14
|
||||
vfmadd231ps zmm11, zmm12, zmm15
|
||||
|
||||
add rcx, 12
|
||||
add rax, 256
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
// Tile size: 5x2
|
||||
// Accumulators: 0-9
|
||||
// Col regs: zmm10-13
|
||||
// Row regs: zmm14-15
|
||||
|
||||
vmovaps zmm10, [rax]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 0]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 4]
|
||||
vmovaps zmm11, [rax + 64]
|
||||
|
||||
// NB stepping column-wise
|
||||
vfmadd231ps zmm0, zmm10, zmm14
|
||||
vfmadd231ps zmm5, zmm10, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 128]
|
||||
|
||||
vfmadd231ps zmm1, zmm11, zmm14
|
||||
vfmadd231ps zmm6, zmm11, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 192]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm14
|
||||
vfmadd231ps zmm7, zmm12, zmm15
|
||||
|
||||
vmovaps zmm10, [rax + 256]
|
||||
|
||||
vfmadd231ps zmm3, zmm13, zmm14
|
||||
vfmadd231ps zmm8, zmm13, zmm15
|
||||
|
||||
vmovaps zmm11, [rax + 320]
|
||||
|
||||
vfmadd231ps zmm4, zmm10, zmm14
|
||||
vfmadd231ps zmm9, zmm10, zmm15
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 8]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 12]
|
||||
|
||||
vmovaps zmm12, [rax + 384]
|
||||
|
||||
// NB stepping column-wise
|
||||
vfmadd231ps zmm0, zmm11, zmm14
|
||||
vfmadd231ps zmm5, zmm11, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 448]
|
||||
|
||||
vfmadd231ps zmm1, zmm12, zmm14
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
|
||||
vmovaps zmm10, [rax + 512]
|
||||
|
||||
vfmadd231ps zmm2, zmm13, zmm14
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
|
||||
vmovaps zmm11, [rax + 576]
|
||||
|
||||
vfmadd231ps zmm3, zmm10, zmm14
|
||||
vfmadd231ps zmm8, zmm10, zmm15
|
||||
|
||||
vfmadd231ps zmm4, zmm11, zmm14
|
||||
vfmadd231ps zmm9, zmm11, zmm15
|
||||
|
||||
add rax, 640
|
||||
add rcx, 16
|
||||
+34
@@ -0,0 +1,34 @@
|
||||
// Tile size: 5x2
|
||||
// Accumulators: 0-9
|
||||
// Col regs: zmm10-14
|
||||
// Row regs: zmm15-16
|
||||
|
||||
vmovaps zmm10, [rax]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 0]
|
||||
vbroadcastss zmm16, dword ptr [rcx + 4]
|
||||
vmovaps zmm11, [rax + 64]
|
||||
|
||||
// NB stepping column-wise
|
||||
vfmadd231ps zmm0, zmm10, zmm15
|
||||
vfmadd231ps zmm5, zmm10, zmm16
|
||||
|
||||
vmovaps zmm12, [rax + 128]
|
||||
|
||||
vfmadd231ps zmm1, zmm11, zmm15
|
||||
vfmadd231ps zmm6, zmm11, zmm16
|
||||
|
||||
vmovaps zmm13, [rax + 192]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm15
|
||||
vfmadd231ps zmm7, zmm12, zmm16
|
||||
|
||||
vmovaps zmm14, [rax + 256]
|
||||
|
||||
vfmadd231ps zmm3, zmm13, zmm15
|
||||
vfmadd231ps zmm8, zmm13, zmm16
|
||||
|
||||
vfmadd231ps zmm4, zmm14, zmm15
|
||||
vfmadd231ps zmm9, zmm14, zmm16
|
||||
|
||||
add rax, 320
|
||||
add rcx, 8
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
// Tile size: 6x1
|
||||
// Accumulators: 0-5
|
||||
// Col regs: 6-11
|
||||
// Row regs: 15
|
||||
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx]
|
||||
vfmadd231ps zmm0, zmm15, [rax]
|
||||
vfmadd231ps zmm1, zmm15, [rax + 64]
|
||||
vfmadd231ps zmm2, zmm15, [rax + 128]
|
||||
vfmadd231ps zmm3, zmm15, [rax + 192]
|
||||
vfmadd231ps zmm4, zmm15, [rax + 256]
|
||||
vfmadd231ps zmm5, zmm15, [rax + 320]
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps zmm0, zmm14, [rax + 384]
|
||||
vfmadd231ps zmm1, zmm14, [rax + 448]
|
||||
vfmadd231ps zmm2, zmm14, [rax + 512]
|
||||
vfmadd231ps zmm3, zmm14, [rax + 576]
|
||||
vfmadd231ps zmm4, zmm14, [rax + 640]
|
||||
vfmadd231ps zmm5, zmm14, [rax + 704]
|
||||
|
||||
add rax, 768
|
||||
add rcx, 8
|
||||
+29
@@ -0,0 +1,29 @@
|
||||
// Tile size: 6x1
|
||||
// Accumulators: 0-5
|
||||
// Col regs: 6-11
|
||||
// Row regs: 15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx]
|
||||
|
||||
vmovups zmm10, [rax]
|
||||
vmulps zmm10, zmm10, zmm15
|
||||
vaddps zmm0, zmm0, zmm10
|
||||
vmovups zmm11, [rax + 64]
|
||||
vmulps zmm11, zmm11, zmm15
|
||||
vaddps zmm1, zmm1, zmm11
|
||||
vmovups zmm12, [rax + 128]
|
||||
vmulps zmm12, zmm12, zmm15
|
||||
vaddps zmm2, zmm2, zmm12
|
||||
vmovups zmm13, [rax + 192]
|
||||
vmulps zmm13, zmm13, zmm15
|
||||
vaddps zmm3, zmm3, zmm13
|
||||
vmovups zmm14, [rax + 256]
|
||||
vmulps zmm14, zmm14, zmm15
|
||||
vaddps zmm4, zmm4, zmm14
|
||||
vmovups zmm15, [rax + 320]
|
||||
vmulps zmm15, zmm15, zmm15
|
||||
vaddps zmm5, zmm5, zmm15
|
||||
|
||||
|
||||
add rcx, 4
|
||||
add rax, 384
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
// Tile size: 6x2
|
||||
// Accumulators: 0-9
|
||||
// Col regs: zmm10-13
|
||||
// Row regs: zmm14-15
|
||||
|
||||
vmovaps zmm12, [rax]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 0]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 4]
|
||||
vmovaps zmm13, [rax + 64]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm14
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 128]
|
||||
|
||||
vfmadd231ps zmm1, zmm13, zmm14
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 192]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm14
|
||||
vfmadd231ps zmm8, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 256]
|
||||
|
||||
vfmadd231ps zmm3, zmm13, zmm14
|
||||
vfmadd231ps zmm9, zmm13, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 320]
|
||||
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 384]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps zmm5, zmm13, zmm14
|
||||
vfmadd231ps zmm11, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx + 12]
|
||||
vmovaps zmm13, [rax + 448]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm14
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 512]
|
||||
|
||||
vfmadd231ps zmm1, zmm13, zmm14
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 576]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm14
|
||||
vfmadd231ps zmm8, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 640]
|
||||
|
||||
vfmadd231ps zmm3, zmm13, zmm14
|
||||
vfmadd231ps zmm9, zmm13, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 704]
|
||||
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
|
||||
vfmadd231ps zmm5, zmm13, zmm14
|
||||
vfmadd231ps zmm11, zmm13, zmm15
|
||||
|
||||
add rax, 768
|
||||
add rcx, 16
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
// Tile size: 6x2
|
||||
// Accumulators: 0-11
|
||||
// Col regs: 12-13
|
||||
// Row regs: 14-15
|
||||
|
||||
vmovaps zmm12, [rax]
|
||||
vbroadcastss zmm14, dword ptr [rcx + 0]
|
||||
vbroadcastss zmm15, dword ptr [rcx + 4]
|
||||
vmovaps zmm13, [rax + 64]
|
||||
|
||||
vfmadd231ps zmm0, zmm12, zmm14
|
||||
vfmadd231ps zmm6, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 128]
|
||||
|
||||
vfmadd231ps zmm1, zmm13, zmm14
|
||||
vfmadd231ps zmm7, zmm13, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 192]
|
||||
|
||||
vfmadd231ps zmm2, zmm12, zmm14
|
||||
vfmadd231ps zmm8, zmm12, zmm15
|
||||
|
||||
vmovaps zmm12, [rax + 256]
|
||||
|
||||
vfmadd231ps zmm3, zmm13, zmm14
|
||||
vfmadd231ps zmm9, zmm13, zmm15
|
||||
|
||||
vmovaps zmm13, [rax + 320]
|
||||
|
||||
vfmadd231ps zmm4, zmm12, zmm14
|
||||
vfmadd231ps zmm10, zmm12, zmm15
|
||||
|
||||
vfmadd231ps zmm5, zmm13, zmm14
|
||||
vfmadd231ps zmm11, zmm13, zmm15
|
||||
|
||||
add rcx, 8
|
||||
add rax, 384
|
||||
+40
@@ -0,0 +1,40 @@
|
||||
// Tile size: 6x1
|
||||
// Accumulators: 0-5
|
||||
// Col regs: 6-11
|
||||
// Row regs: 15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx]
|
||||
|
||||
vmovaps zmm7, [rax + 0]
|
||||
vmovaps zmm8, [rax + 64]
|
||||
vmovaps zmm9, [rax + 128]
|
||||
vmovaps zmm10, [rax + 192]
|
||||
vmovaps zmm11, [rax + 256]
|
||||
vmovaps zmm12, [rax + 320]
|
||||
vmovaps zmm13, [rax + 384]
|
||||
|
||||
vfmadd231ps zmm0, zmm7, zmm15
|
||||
vfmadd231ps zmm1, zmm8, zmm15
|
||||
vfmadd231ps zmm2, zmm9, zmm15
|
||||
vfmadd231ps zmm3, zmm10, zmm15
|
||||
vfmadd231ps zmm4, zmm11, zmm15
|
||||
vfmadd231ps zmm5, zmm12, zmm15
|
||||
vfmadd231ps zmm6, zmm13, zmm15
|
||||
|
||||
vbroadcastss zmm16, dword ptr [rcx + 4]
|
||||
|
||||
vmovaps zmm7, [rax + 448 + 0]
|
||||
vmovaps zmm8, [rax + 448 + 64]
|
||||
vmovaps zmm9, [rax + 448 + 128]
|
||||
vmovaps zmm10, [rax + 448 + 192]
|
||||
vmovaps zmm11, [rax + 448 + 256]
|
||||
vmovaps zmm12, [rax + 448 + 320]
|
||||
vmovaps zmm13, [rax + 448 + 384]
|
||||
|
||||
vfmadd231ps zmm0, zmm7, zmm15
|
||||
vfmadd231ps zmm1, zmm8, zmm15
|
||||
vfmadd231ps zmm2, zmm9, zmm15
|
||||
vfmadd231ps zmm3, zmm10, zmm15
|
||||
vfmadd231ps zmm4, zmm11, zmm15
|
||||
vfmadd231ps zmm5, zmm12, zmm15
|
||||
vfmadd231ps zmm6, zmm13, zmm15
|
||||
+21
@@ -0,0 +1,21 @@
|
||||
// Tile size: 7x1
|
||||
// Accumulators: 0-6
|
||||
// Col regs: 6-13
|
||||
// Row regs: 15
|
||||
vbroadcastss zmm15, dword ptr [rcx]
|
||||
|
||||
vmovaps zmm7, [rax + 0]
|
||||
vmovaps zmm8, [rax + 64]
|
||||
vmovaps zmm9, [rax + 128]
|
||||
vmovaps zmm10, [rax + 192]
|
||||
vmovaps zmm11, [rax + 256]
|
||||
vmovaps zmm12, [rax + 320]
|
||||
vmovaps zmm13, [rax + 384]
|
||||
|
||||
vfmadd231ps zmm0, zmm7, zmm15
|
||||
vfmadd231ps zmm1, zmm8, zmm15
|
||||
vfmadd231ps zmm2, zmm9, zmm15
|
||||
vfmadd231ps zmm3, zmm10, zmm15
|
||||
vfmadd231ps zmm4, zmm11, zmm15
|
||||
vfmadd231ps zmm5, zmm12, zmm15
|
||||
vfmadd231ps zmm6, zmm13, zmm15
|
||||
+30
@@ -0,0 +1,30 @@
|
||||
// Tile size: 8x1
|
||||
// Accumulators: 0-7
|
||||
// Col regs: 8-14
|
||||
// Row regs: 15
|
||||
|
||||
vbroadcastss zmm17, dword ptr [rcx]
|
||||
|
||||
|
||||
vfmadd231ps zmm0, zmm17, [rax + 0]
|
||||
vfmadd231ps zmm1, zmm17, [rax + 64]
|
||||
vfmadd231ps zmm2, zmm17, [rax + 128]
|
||||
vfmadd231ps zmm3, zmm17, [rax + 192]
|
||||
vfmadd231ps zmm4, zmm17, [rax + 256]
|
||||
vfmadd231ps zmm5, zmm17, [rax + 320]
|
||||
vfmadd231ps zmm6, zmm17, [rax + 384]
|
||||
vfmadd231ps zmm7, zmm17, [rax + 448]
|
||||
|
||||
vbroadcastss zmm16, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps zmm0, zmm16, [rax + 0 + 512]
|
||||
vfmadd231ps zmm1, zmm16, [rax + 64 + 512]
|
||||
vfmadd231ps zmm2, zmm16, [rax + 128 + 512]
|
||||
vfmadd231ps zmm3, zmm16, [rax + 192 + 512]
|
||||
vfmadd231ps zmm4, zmm16, [rax + 256 + 512]
|
||||
vfmadd231ps zmm5, zmm16, [rax + 320 + 512]
|
||||
vfmadd231ps zmm6, zmm16, [rax + 384 + 512]
|
||||
vfmadd231ps zmm7, zmm16, [rax + 448 + 512]
|
||||
|
||||
add rcx, 8
|
||||
add rax, 1024
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
// Tile size: 8x1
|
||||
// Accumulators: 0-7
|
||||
// Col regs: 8-14
|
||||
// Row regs: 15
|
||||
|
||||
vbroadcastss zmm15, dword ptr [rcx]
|
||||
|
||||
vmovaps zmm8, [rax + 0]
|
||||
vfmadd231ps zmm0, zmm15, zmm8
|
||||
vmovaps zmm9, [rax + 64]
|
||||
vfmadd231ps zmm1, zmm15, zmm9
|
||||
vmovaps zmm10, [rax + 128]
|
||||
vfmadd231ps zmm2, zmm15, zmm10
|
||||
vmovaps zmm11, [rax + 192]
|
||||
vfmadd231ps zmm3, zmm15, zmm11
|
||||
vmovaps zmm12, [rax + 256]
|
||||
vfmadd231ps zmm4, zmm15, zmm12
|
||||
vmovaps zmm13, [rax + 320]
|
||||
vfmadd231ps zmm5, zmm15, zmm13
|
||||
vmovaps zmm14, [rax + 384]
|
||||
vfmadd231ps zmm6, zmm15, zmm14
|
||||
vmovaps zmm8, [rax + 448]
|
||||
vfmadd231ps zmm7, zmm15, zmm8
|
||||
add rcx, 4
|
||||
add rax, 512
|
||||
+42
@@ -0,0 +1,42 @@
|
||||
// Tile size: 8x2
|
||||
// Accumulators: 0-15
|
||||
// Col regs: 16-23
|
||||
// Row regs: 24-25
|
||||
|
||||
vmovaps zmm16, [rax + 0]
|
||||
vbroadcastss zmm24, dword ptr [rcx + 0]
|
||||
vbroadcastss zmm25, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps zmm0, zmm16, zmm24
|
||||
vfmadd231ps zmm8, zmm16, zmm25
|
||||
|
||||
vmovaps zmm17, [rax + 64]
|
||||
vfmadd231ps zmm1, zmm17, zmm24
|
||||
vfmadd231ps zmm9, zmm17, zmm25
|
||||
|
||||
vmovaps zmm18, [rax + 128]
|
||||
vfmadd231ps zmm2, zmm18, zmm24
|
||||
vfmadd231ps zmm10, zmm18, zmm25
|
||||
|
||||
vmovaps zmm19, [rax + 192]
|
||||
vfmadd231ps zmm3, zmm19, zmm24
|
||||
vfmadd231ps zmm11, zmm19, zmm25
|
||||
|
||||
vmovaps zmm20, [rax + 256]
|
||||
vfmadd231ps zmm4, zmm20, zmm24
|
||||
vfmadd231ps zmm12, zmm20, zmm25
|
||||
|
||||
vmovaps zmm21, [rax + 320]
|
||||
vfmadd231ps zmm5, zmm21, zmm24
|
||||
vfmadd231ps zmm13, zmm21, zmm25
|
||||
|
||||
vmovaps zmm22, [rax + 384]
|
||||
vfmadd231ps zmm6, zmm22, zmm24
|
||||
vfmadd231ps zmm14, zmm22, zmm25
|
||||
|
||||
vmovaps zmm23, [rax + 448]
|
||||
vfmadd231ps zmm7, zmm23, zmm24
|
||||
vfmadd231ps zmm15, zmm23, zmm25
|
||||
|
||||
add rax, 512
|
||||
add rcx, 8
|
||||
+61
@@ -0,0 +1,61 @@
|
||||
// Tile size: 1x8
|
||||
// Accumulators: 0-7
|
||||
// Col regs: 8-14
|
||||
// Row regs: 15
|
||||
|
||||
|
||||
vmovaps zmm15, [rax]
|
||||
|
||||
vbroadcastss zmm8, dword ptr [rcx + 0 * 4]
|
||||
vfmadd231ps zmm0, zmm15, zmm8
|
||||
|
||||
vbroadcastss zmm9, dword ptr [rcx + 1 * 4]
|
||||
vfmadd231ps zmm1, zmm15, zmm9
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx + 2 * 4]
|
||||
vfmadd231ps zmm2, zmm15, zmm10
|
||||
|
||||
vbroadcastss zmm11, dword ptr [rcx + 3 * 4]
|
||||
vfmadd231ps zmm3, zmm15, zmm11
|
||||
|
||||
vbroadcastss zmm12, dword ptr [rcx + 4 * 4]
|
||||
vfmadd231ps zmm4, zmm15, zmm12
|
||||
|
||||
vbroadcastss zmm13, dword ptr [rcx + 5 * 4]
|
||||
vfmadd231ps zmm5, zmm15, zmm13
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx + 6 * 4]
|
||||
vfmadd231ps zmm6, zmm15, zmm10
|
||||
|
||||
vbroadcastss zmm11, dword ptr [rcx + 7 * 4]
|
||||
vfmadd231ps zmm7, zmm15, zmm11
|
||||
|
||||
|
||||
vmovaps zmm15, [rax+64]
|
||||
|
||||
vbroadcastss zmm8, dword ptr [rcx + 8 * 4]
|
||||
vfmadd231ps zmm0, zmm15, zmm8
|
||||
|
||||
vbroadcastss zmm9, dword ptr [rcx + 9 * 4]
|
||||
vfmadd231ps zmm1, zmm15, zmm9
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx + 10 * 4]
|
||||
vfmadd231ps zmm2, zmm15, zmm10
|
||||
|
||||
vbroadcastss zmm11, dword ptr [rcx + 11 * 4]
|
||||
vfmadd231ps zmm3, zmm15, zmm11
|
||||
|
||||
vbroadcastss zmm12, dword ptr [rcx + 12 * 4]
|
||||
vfmadd231ps zmm4, zmm15, zmm12
|
||||
|
||||
vbroadcastss zmm13, dword ptr [rcx + 13 * 4]
|
||||
vfmadd231ps zmm5, zmm15, zmm13
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx + 14 * 4]
|
||||
vfmadd231ps zmm6, zmm15, zmm10
|
||||
|
||||
vbroadcastss zmm11, dword ptr [rcx + 15 * 4]
|
||||
vfmadd231ps zmm7, zmm15, zmm11
|
||||
|
||||
add rcx, 64
|
||||
add rax, 128
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
// Tile size: 1x8
|
||||
// Accumulators: 0-7
|
||||
// Col regs: 8-14
|
||||
// Row regs: 15
|
||||
|
||||
vmovaps zmm15, [rax]
|
||||
|
||||
vbroadcastss zmm8, dword ptr [rcx + 0 * 4]
|
||||
vfmadd231ps zmm0, zmm15, zmm8
|
||||
|
||||
vbroadcastss zmm9, dword ptr [rcx + 1 * 4]
|
||||
vfmadd231ps zmm1, zmm15, zmm9
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx + 2 * 4]
|
||||
vfmadd231ps zmm2, zmm15, zmm10
|
||||
|
||||
vbroadcastss zmm11, dword ptr [rcx + 3 * 4]
|
||||
vfmadd231ps zmm3, zmm15, zmm11
|
||||
|
||||
vbroadcastss zmm12, dword ptr [rcx + 4 * 4]
|
||||
vfmadd231ps zmm4, zmm15, zmm12
|
||||
|
||||
vbroadcastss zmm13, dword ptr [rcx + 5 * 4]
|
||||
vfmadd231ps zmm5, zmm15, zmm13
|
||||
|
||||
vbroadcastss zmm10, dword ptr [rcx + 6 * 4]
|
||||
vfmadd231ps zmm6, zmm15, zmm10
|
||||
|
||||
vbroadcastss zmm11, dword ptr [rcx + 7 * 4]
|
||||
vfmadd231ps zmm7, zmm15, zmm11
|
||||
|
||||
add rcx, 32
|
||||
add rax, 64
|
||||
+151
@@ -0,0 +1,151 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 128 x 1
|
||||
|
||||
zmm0
|
||||
zmm1
|
||||
...
|
||||
zmm7
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of ZMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set size = "128x1" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "8x1/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
sub rbx, 1
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 128 %}{% set from = 0 %}{% set to = 7 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 128 %}{% set from = 0 %}{% set to = 7 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}add_unicast_generic
|
||||
|
||||
{% for row in range(0, 8) %}
|
||||
vaddps zmm{{row}}, zmm{{row}}, [ r10 + {{ row * 64 }} ]
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_unicast_generic:
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
kxnorw k1, k1, k1
|
||||
vgatherdps zmm12{k1}, [r10 + zmm14]
|
||||
vaddps zmm0, zmm0, zmm12
|
||||
|
||||
imul esi, 16
|
||||
vpbroadcastd zmm15, esi
|
||||
|
||||
{% for j in range(1, 8) %}
|
||||
vpaddd zmm14, zmm14, zmm15
|
||||
kxnorw k1, k1, k1
|
||||
vgatherdps zmm12{k1}, [r10 + zmm14]
|
||||
vaddps zmm{{j}}, zmm{{j}}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rbx]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vmovups zmm12, [rax + {{ i * 64 }}]
|
||||
vfmadd231ps zmm{{i}}, zmm12, zmm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}store_noncontiguous
|
||||
|
||||
test r8, 63
|
||||
jnz {{L}}store_unaligned
|
||||
|
||||
{% for row in range(0, 8) %}
|
||||
vmovaps [r8 + {{ row * 64 }}], zmm{{row}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_unaligned:
|
||||
{% for row in range(0, 8) %}
|
||||
vmovups [r8 + {{ row * 64 }}], zmm{{row}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_noncontiguous:
|
||||
{% for r in range(0, 8) %}
|
||||
{% for quarter in range(0, 4) %}
|
||||
vextractf32x4 xmm8, zmm{{r}}, {{quarter}}
|
||||
{% for row in range(0, 4) %}
|
||||
vextractps dword ptr [r8], xmm8, {{row}}
|
||||
add r8, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "128x1" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
+147
@@ -0,0 +1,147 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 16 x 1
|
||||
|
||||
zmm0
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of ZMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
|
||||
{% set size = "16x1" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
cmp rbx, 8
|
||||
jl {{L}}main_loop_packed_packed_tail
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "1x1/packed_packed_loop1/unroll-4.S.raw" %}
|
||||
|
||||
sub rbx, 4
|
||||
cmp rbx, 4
|
||||
jge {{L}}main_loop_packed_packed
|
||||
|
||||
{% for r in range(1, 4) %}
|
||||
vaddps zmm0, zmm0, zmm{{r}}
|
||||
{% endfor %}
|
||||
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed_tail:
|
||||
{% include "1x1/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
sub rbx, 1
|
||||
jnz {{L}}main_loop_packed_packed_tail
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 0 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 0 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 0 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 0 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}add_unicast_generic
|
||||
|
||||
vaddps zmm0, zmm0, [r10]
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_unicast_generic:
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
kxnorw k1, k1, k1
|
||||
vgatherdps zmm12{k1}, [r10 + zmm14]
|
||||
vaddps zmm0, zmm0, zmm12
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rbx]
|
||||
|
||||
{% for i in range(0, 1) %}
|
||||
vmovups zmm12, [rax + {{ i * 64 }}]
|
||||
vfmadd231ps zmm{{i}}, zmm12, zmm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}store_noncontiguous
|
||||
|
||||
test r8, 63
|
||||
jnz {{L}}store_unaligned
|
||||
|
||||
vmovaps [r8], zmm0
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_unaligned:
|
||||
vmovups [r8], zmm0
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_noncontiguous:
|
||||
{% for quarter in range(0, 4) %}
|
||||
vextractf32x4 xmm8, zmm0, {{quarter}}
|
||||
{% for row in range(0, 4) %}
|
||||
vextractps dword ptr [r8], xmm8, {{row}}
|
||||
add r8, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "16x1" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
+165
@@ -0,0 +1,165 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 16 x 12
|
||||
|
||||
zmm0 zmm1 ... zmm11
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of ZMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
|
||||
{% set size = "16x12" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed_tail:
|
||||
{% include "1x12/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
sub rbx, 1
|
||||
jnz {{L}}main_loop_packed_packed_tail
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32 // ymm12 <- xmm12::xmm13
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
{% for i in range(0, 12) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{i}}, zmm{{i}}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups zmm12, zmmword ptr [rax]
|
||||
|
||||
{% for i in range(0, 12) %}
|
||||
vbroadcastss zmm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps zmm{{i}}, zmm12, zmm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 4) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{r}}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
|
||||
// tops of cols
|
||||
lea r8, [ r8 + 4 * rbx ]
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 4) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{ r + 4 }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
|
||||
// tops of cols
|
||||
lea r8, [ r8 + 8 * rbx ]
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 4) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{ r + 8 }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "16x12" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
+143
@@ -0,0 +1,143 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 16 x 8
|
||||
|
||||
zmm0 zmm1 ... zmm8
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of ZMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
|
||||
{% set size = "16x8" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
cmp rbx, 2
|
||||
jl {{L}}main_loop_packed_packed_tail
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "8x8/packed_packed_loop1/avx-512-unroll.S.raw" %}
|
||||
|
||||
sub rbx, 2
|
||||
cmp rbx, 2
|
||||
jge {{L}}main_loop_packed_packed
|
||||
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed_tail:
|
||||
{% include "8x8/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
sub rbx, 1
|
||||
jnz {{L}}main_loop_packed_packed_tail
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 7 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 7 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32 // ymm12 <- xmm12::xmm13
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{i}}, zmm{{i}}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups zmm12, zmmword ptr [rax]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vbroadcastss zmm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps zmm{{i}}, zmm12, zmm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r12, [ r8 + 4 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
lea r13, [ r12 + rbx ]
|
||||
lea r14, [ r12 + 2 * rbx ]
|
||||
lea r15, [ r13 + 2 * rbx ]
|
||||
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 8) %}
|
||||
vextractf32x4 xmm{{ r + 8 }}, zmm{{r}}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 8) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 8 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "16x8" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
+144
@@ -0,0 +1,144 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 32 x 5:
|
||||
|
||||
zmm0 zmm2 zmm4 zmm6 zmm8
|
||||
zmm1 zmm3 zmm5 zmm7 zmm9
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set size = "32x5" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "2x5/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 9 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 9 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 9 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 9 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32 // ymm12 <- xmm12::xmm13
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
{% for i in range(0, 5) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 2 }}, zmm{{ i * 2 }}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
imul esi, 16
|
||||
vpbroadcastd zmm15, esi
|
||||
|
||||
mov r10, [rdi + 8]
|
||||
vpaddd zmm14, zmm14, zmm15
|
||||
|
||||
{% for i in range(0, 5) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 2 + 1 }}, zmm{{ i * 2 + 1 }}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups zmm12, zmmword ptr [rax]
|
||||
vmovups zmm13, zmmword ptr [rax+64]
|
||||
|
||||
{% for i in range(0, 5) %}
|
||||
vbroadcastss zmm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps zmm{{ i * 2 }}, zmm12, zmm14
|
||||
vfmadd231ps zmm{{ i * 2 + 1 }}, zmm13, zmm14
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
lea r12, [ r10 + 2 * rbx ]
|
||||
|
||||
{% for word in range(0, 2) %}
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 5) %}
|
||||
vextractf32x4 xmm{{ r + 11 }}, zmm{{ r * 2 + word }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 5) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 11 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "32x5" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
|
||||
+161
@@ -0,0 +1,161 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 32 x 6:
|
||||
|
||||
zmm0 zmm2 zmm4 zmm6 zmm8 zmm10
|
||||
zmm1 zmm3 zmm5 zmm7 zmm9 zmm11
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set size = "32x6" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "2x6/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32 // ymm12 <- xmm12::xmm13
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
{% for i in range(0, 6) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 2 }}, zmm{{ i * 2 }}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
mov r10, [rdi + 8]
|
||||
imul esi, 16
|
||||
vpbroadcastd zmm15, esi
|
||||
vpaddd zmm14, zmm14, zmm15
|
||||
|
||||
{% for i in range(0, 6) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 2 + 1 }}, zmm{{ i * 2 + 1 }}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups zmm12, zmmword ptr [rax]
|
||||
vmovups zmm13, zmmword ptr [rax+64]
|
||||
|
||||
{% for i in range(0, 6) %}
|
||||
vbroadcastss zmm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps zmm{{ i * 2 }}, zmm12, zmm14
|
||||
vfmadd231ps zmm{{ i * 2 + 1 }}, zmm13, zmm14
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for word in range(0, 2) %}
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 3) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{ r * 2 + word }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 3) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
// tops of cols
|
||||
mov r8, r11
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
|
||||
{% for word in range(0, 2) %}
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 3) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{ (r + 3) * 2 + word }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 3) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "32x6" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
|
||||
+148
@@ -0,0 +1,148 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 48 x 4:
|
||||
|
||||
zmm0 zmm3 zmm6 zmm9
|
||||
zmm1 zmm4 zmm7 zmm10
|
||||
zmm2 zmm5 zmm8 zmm11
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set size = "48x4" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "3x4/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 48 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 48 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32 // ymm12 <- xmm12::xmm13
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 3 }}, zmm{{ i * 3 }}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
imul esi, 16
|
||||
vpbroadcastd zmm15, esi
|
||||
|
||||
{% for j in range(1, 3) %}
|
||||
mov r10, [rdi + 8]
|
||||
vpaddd zmm14, zmm14, zmm15
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 3 + j }}, zmm{{ i * 3 + j }}, zmm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups zmm12, zmmword ptr [rax]
|
||||
vmovups zmm13, zmmword ptr [rax+64]
|
||||
vmovups zmm15, zmmword ptr [rax+128]
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
vbroadcastss zmm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps zmm{{ i * 3 }}, zmm12, zmm14
|
||||
vfmadd231ps zmm{{ i * 3 + 1 }}, zmm13, zmm14
|
||||
vfmadd231ps zmm{{ i * 3 + 2 }}, zmm15, zmm14
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for word in range(0, 3) %}
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 4) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{ r * 3 + word }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "48x4" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
|
||||
+149
@@ -0,0 +1,149 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 64 x 3:
|
||||
|
||||
zmm0 zmm4 zmm8
|
||||
zmm1 zmm5 zmm9
|
||||
zmm2 zmm6 zmm10
|
||||
zmm3 zmm7 zmm11
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set size = "64x3" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "4x3/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 64 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 64 %}{% set from = 0 %}{% set to = 11 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32 // ymm12 <- xmm12::xmm13
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
{% for i in range(0, 3) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 4 }}, zmm{{ i * 4 }}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
imul esi, 16
|
||||
vpbroadcastd zmm15, esi
|
||||
|
||||
{% for j in range(1, 4) %}
|
||||
mov r10, [rdi + 8]
|
||||
vpaddd zmm14, zmm14, zmm15
|
||||
|
||||
{% for i in range(0, 3) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 4 + j }}, zmm{{ i * 4 + j }}, zmm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss zmm13, dword ptr [rbx]
|
||||
vbroadcastss zmm14, dword ptr [rbx+4]
|
||||
vbroadcastss zmm15, dword ptr [rbx+8]
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
vmovups zmm12, zmmword ptr [rax+{{ i * 64 }}]
|
||||
vfmadd231ps zmm{{i}}, zmm12, zmm13
|
||||
vfmadd231ps zmm{{ i + 4 }}, zmm12, zmm14
|
||||
vfmadd231ps zmm{{ i + 8 }}, zmm12, zmm15
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for word in range(0, 4) %}
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 3) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{ r * 4 + word }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 3) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "64x3" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
|
||||
+148
@@ -0,0 +1,148 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 80 x 2:
|
||||
|
||||
zmm0 zmm5
|
||||
zmm1 zmm6
|
||||
zmm2 zmm7
|
||||
zmm3 zmm8
|
||||
zmm4 zmm9
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set size = "80x2" %}{% set suffix = suffix %}{% set G = G %}{% set arch = "avx512" %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "5x2/packed_packed_loop1/avx-512.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 9 %}{% include "f32_scalars.j2" %}
|
||||
{% set mr = 80 %}{% set from = 0 %}{% set to = 9 %}{% include "f32_per_rows.j2" %}
|
||||
{% set mr = 80 %}{% set from = 0 %}{% set to = 9 %}{% include "f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 9 %}{% include "avx512_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm12, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm13, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
vperm2f128 ymm13, ymm12, ymm13, 32 // ymm12 <- xmm12::xmm13
|
||||
vinsertf32x8 zmm14, zmm14, ymm13, 1
|
||||
|
||||
{% for i in range(0, 2) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 5 }}, zmm{{ i * 5 }}, zmm12
|
||||
{% endfor %}
|
||||
|
||||
imul esi, 16
|
||||
vpbroadcastd zmm15, esi
|
||||
|
||||
{% for j in range(1, 5) %}
|
||||
mov r10, [rdi + 8]
|
||||
vpaddd zmm14, zmm14, zmm15
|
||||
|
||||
{% for i in range(0, 2) %}
|
||||
kxnorw k1,k1,k1
|
||||
vgatherdps zmm12{k1}, [ r10 + zmm14 ]
|
||||
add r10, rbx
|
||||
vaddps zmm{{ i * 5 + j }}, zmm{{ i * 5 + j }}, zmm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss zmm14, dword ptr [rbx]
|
||||
vbroadcastss zmm15, dword ptr [rbx+4]
|
||||
|
||||
{% for i in range(0, 5) %}
|
||||
vmovups zmm12, zmmword ptr [rax+{{ i * 64 }}]
|
||||
vfmadd231ps zmm{{i}}, zmm12, zmm14
|
||||
vfmadd231ps zmm{{ i + 5 }}, zmm12, zmm15
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for word in range(0, 5) %}
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% for r in range(0, 2) %}
|
||||
vextractf32x4 xmm{{ r + 12 }}, zmm{{ r * 5 + word }}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 2) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i + 12 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set size = "80x2" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% set arch = "avx512" %}{% include "postamble.j2" %}
|
||||
|
||||
fluxer_desktop/native/webrtc-sender/vendor/tract-linalg-0.23.1/x86_64/avx512/avx512_mmm_load_tile.j2
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
|
||||
{{L}}load_tile:
|
||||
mov r8, [rdi + 8]
|
||||
{% for reg in range(from, to + 1) %}
|
||||
vmovups zmm{{reg}}, zmmword ptr [r8 + {{ (reg - from) * 64 }}]
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
+40
@@ -0,0 +1,40 @@
|
||||
// vim: set syntax=asm :
|
||||
|
||||
{{L}}non_linear:
|
||||
|
||||
{{L}}non_linear_loop_enter:
|
||||
sub rdi, 40
|
||||
{{L}}non_linear_loop:
|
||||
add rdi, 40
|
||||
mov rax, [rdi]
|
||||
|
||||
mov r8, {{ jump_table | length }}
|
||||
cmp rax, 0
|
||||
cmovl rax, r8
|
||||
cmp rax, {{ jump_table | length }}
|
||||
cmovg rax, r8
|
||||
|
||||
{% if msvc %}
|
||||
lea r8, [ offset {{L}}jmp_table ]
|
||||
{% else %}
|
||||
lea r8, [ rip + {{L}}jmp_table ]
|
||||
{% endif %}
|
||||
movsxd r9, dword ptr [ r8 + rax * 4 ]
|
||||
lea r8, [ r8 + r9 ]
|
||||
jmp r8
|
||||
|
||||
{{L}}jmp_table:
|
||||
{% for j in jump_table %}
|
||||
{{long}} {{L}}{{j}}-{{L}}jmp_table
|
||||
{% endfor %}
|
||||
{{long}} {{L}}unsupported-{{L}}jmp_table
|
||||
|
||||
{{L}}unsupported:
|
||||
mov rax, 1
|
||||
jmp {{L}}return
|
||||
|
||||
|
||||
{{L}}done:
|
||||
mov rax, 0
|
||||
jmp {{L}}return
|
||||
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "zmm_ops.j2" import per_col %}
|
||||
|
||||
{{ per_col("per_col_min", "vminps", mr, from, to) }}
|
||||
{{ per_col("per_col_max", "vmaxps", mr, from, to) }}
|
||||
{{ per_col("per_col_add", "vaddps", mr, from, to) }}
|
||||
{{ per_col("per_col_mul", "vmulps", mr, from, to) }}
|
||||
{{ per_col("per_col_sub", "vsubps", mr, from, to) }}
|
||||
{{ per_col("per_col_sub_flipped", "vsubps", mr, from, to, flipped=true) }}
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "zmm_ops.j2" import per_row %}
|
||||
|
||||
{{ per_row("per_row_min", "vminps", mr, from, to) }}
|
||||
{{ per_row("per_row_max", "vmaxps", mr, from, to) }}
|
||||
{{ per_row("per_row_add", "vaddps", mr, from, to) }}
|
||||
{{ per_row("per_row_mul", "vmulps", mr, from, to) }}
|
||||
{{ per_row("per_row_sub", "vsubps", mr, from, to) }}
|
||||
{{ per_row("per_row_sub_flipped", "vsubps", mr, from, to, flipped=true) }}
|
||||
Vendored
+30
@@ -0,0 +1,30 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "zmm_ops.j2" import scalar %}
|
||||
|
||||
{{ scalar("scalar_min", "vminps", from, to) }}
|
||||
{{ scalar("scalar_max", "vmaxps", from, to) }}
|
||||
{{ scalar("scalar_add", "vaddps", from, to) }}
|
||||
{{ scalar("scalar_mul", "vmulps", from, to) }}
|
||||
{{ scalar("scalar_sub", "vsubps", from, to) }}
|
||||
{{ scalar("scalar_sub_flipped", "vsubps", from, to, flipped=true) }}
|
||||
|
||||
{{L}}leaky_relu:
|
||||
// can only use zmm12 to zmm15
|
||||
// ymm15 <- alpha
|
||||
vbroadcastss zmm15, dword ptr [rdi + 8]
|
||||
// ymm14 <- all zero
|
||||
vpxorq zmm14, zmm14, zmm14
|
||||
|
||||
{% for reg in range(from, to + 1) %}
|
||||
vcmpps k1, zmm{{reg}}, zmm14, 1 // 1 means LT
|
||||
// ymm12 <- alpha * x if < 0
|
||||
vmulps zmm{{reg}} {k1}, zmm{{reg}}, zmm15
|
||||
{% endfor %}
|
||||
// select muled of orginal
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale:
|
||||
{{L}}q_shl:
|
||||
{{L}}q_shr:
|
||||
jmp {{L}}unsupported
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "zmm_ops.j2" import per_col %}
|
||||
|
||||
{{ per_col("per_col_min", "vpminsd", mr, from, to) }}
|
||||
{{ per_col("per_col_max", "vpmaxsd", mr, from, to) }}
|
||||
{{ per_col("per_col_add", "vpaddd", mr, from, to) }}
|
||||
{{ per_col("per_col_mul", "vpmulld", mr, from, to) }}
|
||||
{{ per_col("per_col_sub", "vpsubd", mr, from, to) }}
|
||||
{{ per_col("per_col_sub_flipped", "vpsubd", mr, from, to, flipped=true) }}
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "zmm_ops.j2" import per_row %}
|
||||
|
||||
{{ per_row("per_row_min", "vpminsd", mr, from, to) }}
|
||||
{{ per_row("per_row_max", "vpmaxsd", mr, from, to) }}
|
||||
{{ per_row("per_row_add", "vpaddd", mr, from, to) }}
|
||||
{{ per_row("per_row_mul", "vpmulld", mr, from, to) }}
|
||||
{{ per_row("per_row_sub", "vpsubd", mr, from, to) }}
|
||||
{{ per_row("per_row_sub_flipped", "vpsubd", mr, from, to, flipped=true) }}
|
||||
Vendored
+12
@@ -0,0 +1,12 @@
|
||||
// vim: set syntax=asm :
|
||||
{% if not arch %}
|
||||
{% set arch = "ymm" %}
|
||||
{% endif %}
|
||||
{% from "zmm_ops.j2" import scalar %}
|
||||
|
||||
{{ scalar("scalar_min", "vpminsd", from, to) }}
|
||||
{{ scalar("scalar_max", "vpmaxsd", from, to) }}
|
||||
{{ scalar("scalar_mul", "vpmulld", from, to) }}
|
||||
{{ scalar("scalar_add", "vpaddd", from, to) }}
|
||||
{{ scalar("scalar_sub", "vpsubd", from, to) }}
|
||||
{{ scalar("scalar_sub_flipped", "vpsubd", from, to, flipped=true) }}
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
{{L}}return:
|
||||
ldmxcsr [rsp + 4]
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
{% if msvc %}
|
||||
{{arch}}_mmm_f32_{{size}}_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
{{arch}}_mmm_f32_{{size}}_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}{{arch}}_mmm_f32_{{size}}_{{suffix}}
|
||||
{{G}}{{arch}}_mmm_f32_{{size}}_{{suffix}}:
|
||||
.cfi_startproc
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
push rdi
|
||||
push rsi
|
||||
|
||||
mov rdi, rcx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
.cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
|
||||
{% include "dispatcher.j2" %}
|
||||
Vendored
+325
@@ -0,0 +1,325 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
// AVX-512 (zmm, 16-wide) sigmoid. Uses a rational (Padé-style) approximation
|
||||
// of sigmoid clamped to [-18, 18]. Validated against the generic scalar
|
||||
// reference via sigmoid_frame_tests! (see x86_64_fma.rs).
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of ZMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
|
||||
#}
|
||||
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
avx512_sigmoid_f32_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}avx512_sigmoid_f32_{{suffix}}
|
||||
{{G}}avx512_sigmoid_f32_{{suffix}}:
|
||||
.cfi_startproc
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
// move around arguments to mimick SysV rdi,rsi passing
|
||||
push rdi
|
||||
push rsi
|
||||
mov rdi, rcx
|
||||
mov rsi, rdx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
// FIXME
|
||||
// .cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
{% set offset %}{% if msvc %} offset {%else%} rip + {%endif%} {% endset %}
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
cmp rsi, 64
|
||||
jl {{L}}loop_1
|
||||
|
||||
{{L}}loop_4:
|
||||
|
||||
vmovaps zmm4, [rdi]
|
||||
vmovaps zmm5, [rdi + 64]
|
||||
vmovaps zmm6, [rdi + 128]
|
||||
vmovaps zmm7, [rdi + 192]
|
||||
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmaxps zmm4, zmm4, zmm0
|
||||
vmaxps zmm5, zmm5, zmm0
|
||||
vmaxps zmm6, zmm6, zmm0
|
||||
vmaxps zmm7, zmm7, zmm0
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
|
||||
vminps zmm4, zmm4, zmm1
|
||||
vminps zmm5, zmm5, zmm1
|
||||
vminps zmm6, zmm6, zmm1
|
||||
vminps zmm7, zmm7, zmm1 // zmm4..7 <- x
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
|
||||
vmulps zmm8, zmm4, zmm4
|
||||
vmulps zmm9, zmm5, zmm5
|
||||
vmulps zmm10, zmm6, zmm6
|
||||
vmulps zmm11, zmm7, zmm7 // zmm8..11 <- x^2
|
||||
|
||||
vmovaps zmm12, zmm2
|
||||
vmovaps zmm13, zmm2
|
||||
vmovaps zmm14, zmm2
|
||||
vmovaps zmm15, zmm2
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vfmadd132ps zmm13, zmm3, zmm9
|
||||
vfmadd132ps zmm14, zmm3, zmm10
|
||||
vfmadd132ps zmm15, zmm3, zmm11
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_10]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vfmadd132ps zmm13, zmm0, zmm9
|
||||
vfmadd132ps zmm14, zmm0, zmm10
|
||||
vfmadd132ps zmm15, zmm0, zmm11
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_8]
|
||||
vfmadd132ps zmm12, zmm1, zmm8
|
||||
vfmadd132ps zmm13, zmm1, zmm9
|
||||
vfmadd132ps zmm14, zmm1, zmm10
|
||||
vfmadd132ps zmm15, zmm1, zmm11
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vfmadd132ps zmm13, zmm2, zmm9
|
||||
vfmadd132ps zmm14, zmm2, zmm10
|
||||
vfmadd132ps zmm15, zmm2, zmm11
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vmulps zmm4, zmm4, zmm12
|
||||
vmulps zmm5, zmm5, zmm13
|
||||
vmulps zmm6, zmm6, zmm14
|
||||
vmulps zmm7, zmm7, zmm15 // zmm4..7 <- num
|
||||
|
||||
vmovaps zmm12, zmm3
|
||||
vmovaps zmm13, zmm3
|
||||
vmovaps zmm14, zmm3
|
||||
vmovaps zmm15, zmm3
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vfmadd132ps zmm13, zmm0, zmm9
|
||||
vfmadd132ps zmm14, zmm0, zmm10
|
||||
vfmadd132ps zmm15, zmm0, zmm11
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vfmadd132ps zmm12, zmm1, zmm8
|
||||
vfmadd132ps zmm13, zmm1, zmm9
|
||||
vfmadd132ps zmm14, zmm1, zmm10
|
||||
vfmadd132ps zmm15, zmm1, zmm11
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_half]
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vfmadd132ps zmm13, zmm2, zmm9
|
||||
vfmadd132ps zmm14, zmm2, zmm10
|
||||
vfmadd132ps zmm15, zmm2, zmm11
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vfmadd132ps zmm13, zmm3, zmm9
|
||||
vfmadd132ps zmm14, zmm3, zmm10
|
||||
vfmadd132ps zmm15, zmm3, zmm11
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vfmadd132ps zmm13, zmm0, zmm9
|
||||
vfmadd132ps zmm14, zmm0, zmm10
|
||||
vfmadd132ps zmm15, zmm0, zmm11 // zmm12..14 <- denum
|
||||
|
||||
vdivps zmm4, zmm4, zmm12
|
||||
vdivps zmm5, zmm5, zmm13
|
||||
vdivps zmm6, zmm6, zmm14
|
||||
vdivps zmm7, zmm7, zmm15
|
||||
vaddps zmm4, zmm4, zmm1
|
||||
vaddps zmm5, zmm5, zmm1
|
||||
vaddps zmm6, zmm6, zmm1
|
||||
vaddps zmm7, zmm7, zmm1
|
||||
|
||||
vmovaps [rdi], zmm4
|
||||
vmovaps [rdi + 64], zmm5
|
||||
vmovaps [rdi + 128], zmm6
|
||||
vmovaps [rdi + 192], zmm7
|
||||
|
||||
add rdi, 256
|
||||
sub rsi, 64
|
||||
cmp rsi, 64
|
||||
jg {{L}}loop_4
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
{{L}}loop_1:
|
||||
vmovaps zmm4, [rdi]
|
||||
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmaxps zmm4, zmm4, zmm0
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
|
||||
vminps zmm4, zmm4, zmm1 // zmm4 <- x
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
|
||||
vmulps zmm8, zmm4, zmm4 // zmm8 <- x^2
|
||||
|
||||
vmovaps zmm12, zmm2
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_10]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_8]
|
||||
vfmadd132ps zmm12, zmm1, zmm8
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vmulps zmm4, zmm4, zmm12
|
||||
|
||||
vmovaps zmm12, zmm3
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vfmadd132ps zmm12, zmm1, zmm8
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_half]
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
|
||||
vdivps zmm4, zmm4, zmm12
|
||||
vaddps zmm4, zmm4, zmm1
|
||||
|
||||
vmovaps [rdi], zmm4
|
||||
add rdi, 64
|
||||
sub rsi, 16
|
||||
jnz {{L}}loop_1
|
||||
|
||||
{{L}}done:
|
||||
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
ldmxcsr [rsp + 4]
|
||||
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
{% set float %}{% if msvc %} real4 {%else%} .float {%endif%}{% endset %}
|
||||
|
||||
{{L}}coeffs_num_low:
|
||||
{{float}} -18.0 // low
|
||||
{{L}}coeffs_num_high:
|
||||
{{float}} 18.0 // high
|
||||
|
||||
{{L}}coeffs_num_alpha_9:
|
||||
{{float}} 4.37031012579801e-11 // alpha_9
|
||||
{{L}}coeffs_num_alpha_7:
|
||||
{{float}} 1.15627324459942e-07 // alpha_7
|
||||
{{L}}coeffs_num_alpha_5:
|
||||
{{float}} 6.08574864600143e-05 // alpha_5
|
||||
{{L}}coeffs_num_alpha_3:
|
||||
{{float}} 8.51377133304701e-03 // alpha_3
|
||||
{{L}}coeffs_num_alpha_1:
|
||||
{{float}} 2.48287947061529e-01 // alpha_1
|
||||
|
||||
{{L}}coeffs_num_beta_10:
|
||||
{{float}} 6.10247389755681e-13
|
||||
{{L}}coeffs_num_beta_8:
|
||||
{{float}} 5.76102136993427e-09
|
||||
{{L}}coeffs_num_beta_6:
|
||||
{{float}} 6.29106785017040e-06 // beta_6
|
||||
{{L}}coeffs_num_beta_4:
|
||||
{{float}} 1.70198817374094e-03 // beta_4
|
||||
{{L}}coeffs_num_beta_2:
|
||||
{{float}} 1.16817656904453e-01 // beta_2
|
||||
{{L}}coeffs_num_beta_0:
|
||||
{{float}} 9.93151921023180e-01 // beta_0
|
||||
|
||||
{{L}}coeffs_num_half:
|
||||
{{float}} 0.5
|
||||
|
||||
{% if msvc %}
|
||||
avx512_sigmoid_f32_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
+315
@@ -0,0 +1,315 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
// AVX-512 (zmm, 16-wide) tanh. Polynomial-numerator / polynomial-denominator
|
||||
// rational approximation of tanh clamped to [-9, 9]. Validated against the
|
||||
// generic scalar reference via tanh_frame_tests! (see x86_64_fma.rs).
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of ZMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
|
||||
#}
|
||||
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
avx512_tanh_f32_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}avx512_tanh_f32_{{suffix}}
|
||||
{{G}}avx512_tanh_f32_{{suffix}}:
|
||||
.cfi_startproc
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
// move around arguments to mimick SysV rdi,rsi passing
|
||||
push rdi
|
||||
push rsi
|
||||
mov rdi, rcx
|
||||
mov rsi, rdx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
// FIXME
|
||||
// .cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
{% set offset %}{% if msvc %} offset {%else%} rip + {%endif%} {% endset %}
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
cmp rsi, 64
|
||||
jl {{L}}loop_1
|
||||
|
||||
{{L}}loop_4:
|
||||
|
||||
vmovaps zmm4, [rdi]
|
||||
vmovaps zmm5, [rdi + 64]
|
||||
vmovaps zmm6, [rdi + 128]
|
||||
vmovaps zmm7, [rdi + 192]
|
||||
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_13]
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_11]
|
||||
|
||||
vmaxps zmm4, zmm4, zmm0
|
||||
vmaxps zmm5, zmm5, zmm0
|
||||
vmaxps zmm6, zmm6, zmm0
|
||||
vmaxps zmm7, zmm7, zmm0
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
|
||||
vminps zmm4, zmm4, zmm1
|
||||
vminps zmm5, zmm5, zmm1
|
||||
vminps zmm6, zmm6, zmm1
|
||||
vminps zmm7, zmm7, zmm1 // zmm4..7 <- x
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmulps zmm8, zmm4, zmm4
|
||||
vmulps zmm9, zmm5, zmm5
|
||||
vmulps zmm10, zmm6, zmm6
|
||||
vmulps zmm11, zmm7, zmm7 // zmm8..11 <- x^2
|
||||
|
||||
vmovaps zmm12, zmm2
|
||||
vmovaps zmm13, zmm2
|
||||
vmovaps zmm14, zmm2
|
||||
vmovaps zmm15, zmm2
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vfmadd132ps zmm13, zmm3, zmm9
|
||||
vfmadd132ps zmm14, zmm3, zmm10
|
||||
vfmadd132ps zmm15, zmm3, zmm11
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vfmadd132ps zmm13, zmm0, zmm9
|
||||
vfmadd132ps zmm14, zmm0, zmm10
|
||||
vfmadd132ps zmm15, zmm0, zmm11
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps zmm12, zmm1, zmm8
|
||||
vfmadd132ps zmm13, zmm1, zmm9
|
||||
vfmadd132ps zmm14, zmm1, zmm10
|
||||
vfmadd132ps zmm15, zmm1, zmm11
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vfmadd132ps zmm13, zmm2, zmm9
|
||||
vfmadd132ps zmm14, zmm2, zmm10
|
||||
vfmadd132ps zmm15, zmm2, zmm11
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vfmadd132ps zmm13, zmm3, zmm9
|
||||
vfmadd132ps zmm14, zmm3, zmm10
|
||||
vfmadd132ps zmm15, zmm3, zmm11
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vfmadd132ps zmm13, zmm0, zmm9
|
||||
vfmadd132ps zmm14, zmm0, zmm10
|
||||
vfmadd132ps zmm15, zmm0, zmm11
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vmulps zmm4, zmm4, zmm12
|
||||
vmulps zmm5, zmm5, zmm13
|
||||
vmulps zmm6, zmm6, zmm14
|
||||
vmulps zmm7, zmm7, zmm15 // zmm4..7 <- num
|
||||
|
||||
vmovaps zmm12, zmm1
|
||||
vmovaps zmm13, zmm1
|
||||
vmovaps zmm14, zmm1
|
||||
vmovaps zmm15, zmm1
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vfmadd132ps zmm13, zmm2, zmm9
|
||||
vfmadd132ps zmm14, zmm2, zmm10
|
||||
vfmadd132ps zmm15, zmm2, zmm11
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vfmadd132ps zmm13, zmm3, zmm9
|
||||
vfmadd132ps zmm14, zmm3, zmm10
|
||||
vfmadd132ps zmm15, zmm3, zmm11
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vfmadd132ps zmm13, zmm0, zmm9
|
||||
vfmadd132ps zmm14, zmm0, zmm10
|
||||
vfmadd132ps zmm15, zmm0, zmm11 // zmm12..14 <- denum
|
||||
|
||||
vdivps zmm4, zmm4, zmm12
|
||||
vdivps zmm5, zmm5, zmm13
|
||||
vdivps zmm6, zmm6, zmm14
|
||||
vdivps zmm7, zmm7, zmm15
|
||||
|
||||
vmovaps [rdi], zmm4
|
||||
vmovaps [rdi + 64], zmm5
|
||||
vmovaps [rdi + 128], zmm6
|
||||
vmovaps [rdi + 192], zmm7
|
||||
|
||||
add rdi, 256
|
||||
sub rsi, 64
|
||||
cmp rsi, 64
|
||||
jg {{L}}loop_4
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
{{L}}loop_1:
|
||||
vmovaps zmm4, [rdi]
|
||||
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_13]
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_11]
|
||||
|
||||
vmaxps zmm4, zmm4, zmm0
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
|
||||
vminps zmm4, zmm4, zmm1 // zmm4 <- x
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmulps zmm8, zmm4, zmm4 // zmm8 <- x^2
|
||||
|
||||
vmovaps zmm12, zmm2
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps zmm12, zmm1, zmm8
|
||||
vbroadcastss zmm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vbroadcastss zmm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vbroadcastss zmm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
vbroadcastss zmm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vmulps zmm4, zmm4, zmm12
|
||||
|
||||
vmovaps zmm12, zmm1
|
||||
vfmadd132ps zmm12, zmm2, zmm8
|
||||
vfmadd132ps zmm12, zmm3, zmm8
|
||||
vfmadd132ps zmm12, zmm0, zmm8
|
||||
|
||||
vdivps zmm4, zmm4, zmm12
|
||||
|
||||
vmovaps [rdi], zmm4
|
||||
add rdi, 64
|
||||
sub rsi, 16
|
||||
jnz {{L}}loop_1
|
||||
|
||||
{{L}}done:
|
||||
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
ldmxcsr [rsp + 4]
|
||||
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
{% set float %}{% if msvc %} real4 {%else%} .float {%endif%}{% endset %}
|
||||
|
||||
{{L}}coeffs_num_low:
|
||||
{{float}} -9.0 // low
|
||||
{{L}}coeffs_num_high:
|
||||
{{float}} 9.0 // high
|
||||
|
||||
{{L}}coeffs_num_alpha_13:
|
||||
{{float}} -2.76076847742355e-16 // alpha_13
|
||||
{{L}}coeffs_num_alpha_11:
|
||||
{{float}} 2.00018790482477e-13 // alpha_11
|
||||
{{L}}coeffs_num_alpha_9:
|
||||
{{float}} -8.60467152213735e-11 // alpha_9
|
||||
{{L}}coeffs_num_alpha_7:
|
||||
{{float}} 5.12229709037114e-08 // alpha_7
|
||||
{{L}}coeffs_num_alpha_5:
|
||||
{{float}} 1.48572235717979e-05 // alpha_5
|
||||
{{L}}coeffs_num_alpha_3:
|
||||
{{float}} 6.37261928875436e-04 // alpha_3
|
||||
{{L}}coeffs_num_alpha_1:
|
||||
{{float}} 4.89352455891786e-03 // alpha_1
|
||||
|
||||
{{L}}coeffs_num_beta_6:
|
||||
{{float}} 1.19825839466702e-06 // beta_6
|
||||
{{L}}coeffs_num_beta_4:
|
||||
{{float}} 1.18534705686654e-04 // beta_4
|
||||
{{L}}coeffs_num_beta_2:
|
||||
{{float}} 2.26843463243900e-03 // beta_2
|
||||
{{L}}coeffs_num_beta_0:
|
||||
{{float}} 4.89352518554385e-03 // beta_0
|
||||
|
||||
{% if msvc %}
|
||||
avx512_tanh_f32_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
|
||||
{% macro scalar(label, op, from, to, flipped=false) %}
|
||||
{{L}}{{label}}:
|
||||
vbroadcastss zmm12, dword ptr [rdi + 8]
|
||||
{% if flipped %}
|
||||
{% for reg in range(from, to + 1) %}
|
||||
{{op}} zmm{{reg}}, zmm{{reg}}, zmm12
|
||||
{% endfor %}
|
||||
{% else %}
|
||||
{% for reg in range(from, to + 1) %}
|
||||
{{op}} zmm{{reg}}, zmm12, zmm{{reg}}
|
||||
{% endfor %}
|
||||
{% endif %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
{% endmacro %}
|
||||
|
||||
{% macro per_row(label, op, mr, from, to, flipped=false) %}
|
||||
{{L}}{{label}}:
|
||||
mov rax, [ rdi + 8 ]
|
||||
|
||||
{% set mr_over_16 = mr // 16 %}
|
||||
{% set mr_over_16_min_1 = mr // 16 - 1 %}
|
||||
|
||||
{% for ix in range(0, mr_over_16_min_1 + 1) %}
|
||||
vmovups zmm{{ to + 1 + ix }}, [rax + {{ ix * 64 }}]
|
||||
{% endfor %}
|
||||
|
||||
{% if flipped %}
|
||||
{% for acc in range(from, to + 1) %}
|
||||
{{op}} zmm{{acc}}, zmm{{acc}}, zmm{{ acc % mr_over_16 + to + 1 }}
|
||||
{% endfor %}
|
||||
{% else %}
|
||||
{% for acc in range(from, to + 1) %}
|
||||
{{op}} zmm{{acc}}, zmm{{ acc % mr_over_16 + to + 1 }}, zmm{{acc}}
|
||||
{% endfor %}
|
||||
{% endif %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
{% endmacro %}
|
||||
|
||||
{% macro per_col(label, op, mr, from, to, flipped=false) %}
|
||||
{{L}}{{label}}:
|
||||
mov rax, [ rdi + 8 ]
|
||||
|
||||
{% set mr_over_16 = mr // 16 %}
|
||||
{% set mr_over_16_min_1 = mr // 16 - 1 %}
|
||||
|
||||
{% set tmp = to + 1 %}
|
||||
|
||||
{% set cols = (to + 1 - from) // mr_over_16 %}
|
||||
{% set cols_min_1 = (to + 1 - from) // mr_over_16 - 1 %}
|
||||
// {{ to - from + 1 }} cols:{{cols}}
|
||||
|
||||
{% for right in range(0, cols_min_1 + 1) %}
|
||||
vbroadcastss zmm{{tmp}}, dword ptr [ rax ]
|
||||
add rax, 4
|
||||
|
||||
{% for down in range(0, mr_over_16_min_1 + 1) %}
|
||||
{% set acc = mr_over_16 * right + from + down %}
|
||||
{% if flipped %}
|
||||
{{op}} zmm{{acc}}, zmm{{acc}}, zmm{{tmp}}
|
||||
{% else %}
|
||||
{{op}} zmm{{acc}}, zmm{{tmp}}, zmm{{acc}}
|
||||
{% endif %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
{% endmacro %}
|
||||
Vendored
+13
@@ -0,0 +1,13 @@
|
||||
// Build-time capability probe for the assembler, used by build.rs
|
||||
// (assembler_supports_avx512vnni). Older binutils — notably the Debian stretch
|
||||
// x86_64 toolchain in CI — predate AVX-512 VNNI (added in binutils ~2.30) and
|
||||
// cannot assemble `vpdpbusd ymm` even when targeting a VNNI-capable CPU. If
|
||||
// this file fails to assemble, build.rs skips the VNNI kernel and the
|
||||
// `tract_avx512vnni` cfg, and the runtime falls back to the AVX2 i32 path.
|
||||
// Not linked into anything.
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.globl tract_avx512vnni_probe
|
||||
tract_avx512vnni_probe:
|
||||
vpdpbusd ymm0, ymm1, ymm2
|
||||
ret
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
// Accumulators: 0-7
|
||||
// Columns: 14-15
|
||||
// Rows: 8-13
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovaps ymm10, [rax + 0]
|
||||
vmovaps ymm11, [rax + 32]
|
||||
vmovaps ymm12, [rax + 64]
|
||||
vmovaps ymm13, [rax + 96]
|
||||
vmovaps ymm14, [rax + 128]
|
||||
|
||||
vfmadd231ps ymm0, ymm10, ymm15
|
||||
vfmadd231ps ymm1, ymm11, ymm15
|
||||
vfmadd231ps ymm2, ymm12, ymm15
|
||||
vfmadd231ps ymm3, ymm13, ymm15
|
||||
vfmadd231ps ymm4, ymm14, ymm15
|
||||
|
||||
vmovaps ymm10, [rax + 160]
|
||||
vmovaps ymm11, [rax + 192]
|
||||
vmovaps ymm12, [rax + 224]
|
||||
vmovaps ymm13, [rax + 256]
|
||||
vmovaps ymm14, [rax + 288]
|
||||
|
||||
vfmadd231ps ymm5, ymm10, ymm15
|
||||
vfmadd231ps ymm6, ymm11, ymm15
|
||||
vfmadd231ps ymm7, ymm12, ymm15
|
||||
vfmadd231ps ymm8, ymm13, ymm15
|
||||
vfmadd231ps ymm9, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
|
||||
vmovaps ymm10, [rax + 320]
|
||||
vmovaps ymm11, [rax + 352]
|
||||
vmovaps ymm12, [rax + 384]
|
||||
vmovaps ymm13, [rax + 416]
|
||||
vmovaps ymm14, [rax + 448]
|
||||
|
||||
vfmadd231ps ymm0, ymm10, ymm15
|
||||
vfmadd231ps ymm1, ymm11, ymm15
|
||||
vfmadd231ps ymm2, ymm12, ymm15
|
||||
vfmadd231ps ymm3, ymm13, ymm15
|
||||
vfmadd231ps ymm4, ymm14, ymm15
|
||||
|
||||
vmovaps ymm10, [rax + 480]
|
||||
vmovaps ymm11, [rax + 512]
|
||||
vmovaps ymm12, [rax + 544]
|
||||
vmovaps ymm13, [rax + 576]
|
||||
vmovaps ymm14, [rax + 608]
|
||||
|
||||
vfmadd231ps ymm5, ymm10, ymm15
|
||||
vfmadd231ps ymm6, ymm11, ymm15
|
||||
vfmadd231ps ymm7, ymm12, ymm15
|
||||
vfmadd231ps ymm8, ymm13, ymm15
|
||||
vfmadd231ps ymm9, ymm14, ymm15
|
||||
|
||||
add rcx, 8
|
||||
add rax, 640
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
// Tile size: 10x1
|
||||
// Accumulators: 0-9
|
||||
// Col regs: 10-14
|
||||
// Row regs: 15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovaps ymm10, [rax + 0]
|
||||
vmovaps ymm11, [rax + 32]
|
||||
vmovaps ymm12, [rax + 64]
|
||||
vmovaps ymm13, [rax + 96]
|
||||
vmovaps ymm14, [rax + 128]
|
||||
|
||||
vfmadd231ps ymm0, ymm10, ymm15
|
||||
vfmadd231ps ymm1, ymm11, ymm15
|
||||
vfmadd231ps ymm2, ymm12, ymm15
|
||||
vfmadd231ps ymm3, ymm13, ymm15
|
||||
vfmadd231ps ymm4, ymm14, ymm15
|
||||
|
||||
vmovaps ymm10, [rax + 160]
|
||||
vmovaps ymm11, [rax + 192]
|
||||
vmovaps ymm12, [rax + 224]
|
||||
vmovaps ymm13, [rax + 256]
|
||||
vmovaps ymm14, [rax + 288]
|
||||
|
||||
vfmadd231ps ymm5, ymm10, ymm15
|
||||
vfmadd231ps ymm6, ymm11, ymm15
|
||||
vfmadd231ps ymm7, ymm12, ymm15
|
||||
vfmadd231ps ymm8, ymm13, ymm15
|
||||
vfmadd231ps ymm9, ymm14, ymm15
|
||||
|
||||
add rcx, 4
|
||||
add rax, 320
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
// Accumulators: 0-9
|
||||
// Columns: 14-15
|
||||
// Rows: 10-13
|
||||
vbroadcastss ymm10, dword ptr [rcx]
|
||||
vbroadcastss ymm11, dword ptr [rcx + 4]
|
||||
vbroadcastss ymm12, dword ptr [rcx + 8]
|
||||
vbroadcastss ymm13, dword ptr [rcx + 12]
|
||||
|
||||
vmovaps ymm14, [rax]
|
||||
vmovaps ymm15, [rax + 32]
|
||||
|
||||
vfmadd231ps ymm0, ymm14, ymm10
|
||||
vfmadd231ps ymm1, ymm15, ymm10
|
||||
|
||||
vfmadd231ps ymm2, ymm14, ymm11
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps ymm4, ymm14, ymm12
|
||||
vfmadd231ps ymm5, ymm15, ymm12
|
||||
|
||||
vfmadd231ps ymm6, ymm14, ymm13
|
||||
vfmadd231ps ymm7, ymm15, ymm13
|
||||
|
||||
vfmadd231ps ymm8, ymm14, ymm11
|
||||
vfmadd231ps ymm9, ymm15, ymm11
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rcx + 20]
|
||||
vbroadcastss ymm11, dword ptr [rcx + 24]
|
||||
vbroadcastss ymm12, dword ptr [rcx + 28]
|
||||
vbroadcastss ymm13, dword ptr [rcx + 32]
|
||||
|
||||
vmovaps ymm14, [rax + 64]
|
||||
vmovaps ymm15, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm0, ymm14, ymm10
|
||||
vfmadd231ps ymm1, ymm15, ymm10
|
||||
|
||||
vfmadd231ps ymm2, ymm14, ymm11
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 36]
|
||||
|
||||
vfmadd231ps ymm4, ymm14, ymm12
|
||||
vfmadd231ps ymm5, ymm15, ymm12
|
||||
|
||||
vfmadd231ps ymm6, ymm14, ymm13
|
||||
vfmadd231ps ymm7, ymm15, ymm13
|
||||
|
||||
vfmadd231ps ymm8, ymm14, ymm11
|
||||
vfmadd231ps ymm9, ymm15, ymm11
|
||||
+30
@@ -0,0 +1,30 @@
|
||||
// Accumulators: 0-9
|
||||
// Columns: 14-15
|
||||
// Rows: 10-13
|
||||
vbroadcastss ymm10, dword ptr [rcx]
|
||||
vbroadcastss ymm11, dword ptr [rcx + 4]
|
||||
vbroadcastss ymm12, dword ptr [rcx + 8]
|
||||
vbroadcastss ymm13, dword ptr [rcx + 12]
|
||||
|
||||
vmovaps ymm14, [rax]
|
||||
vmovaps ymm15, [rax + 32]
|
||||
|
||||
vfmadd231ps ymm0, ymm14, ymm10
|
||||
vfmadd231ps ymm1, ymm15, ymm10
|
||||
|
||||
vfmadd231ps ymm2, ymm14, ymm11
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
|
||||
// Use register 11 as it's "middle" use, leading to a decent
|
||||
// trade-off between required use next iteration and when it has
|
||||
// to be used this iteration.
|
||||
vbroadcastss ymm11, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps ymm4, ymm14, ymm12
|
||||
vfmadd231ps ymm5, ymm15, ymm12
|
||||
|
||||
vfmadd231ps ymm6, ymm14, ymm13
|
||||
vfmadd231ps ymm7, ymm15, ymm13
|
||||
|
||||
vfmadd231ps ymm8, ymm14, ymm11
|
||||
vfmadd231ps ymm9, ymm15, ymm11
|
||||
+71
@@ -0,0 +1,71 @@
|
||||
// Tile size: 2x6
|
||||
// Accumulators: 0-11
|
||||
// Col regs: ymm14-15
|
||||
// Row regs: ymm12-13
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx]
|
||||
vmovaps ymm12, [rax]
|
||||
vmovaps ymm13, [rax + 32]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm14
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm15
|
||||
vfmadd231ps ymm3, ymm13, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm5, ymm13, ymm14
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 20]
|
||||
|
||||
vfmadd231ps ymm8, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm13, ymm14
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx+24]
|
||||
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
vfmadd231ps ymm11, ymm13, ymm15
|
||||
|
||||
// Iteration two
|
||||
vmovaps ymm12, [rax + 64]
|
||||
vmovaps ymm13, [rax + 96]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 24 + 4]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm14
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 24 + 8]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm15
|
||||
vfmadd231ps ymm3, ymm13, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 24 + 12]
|
||||
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm5, ymm13, ymm14
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 24 + 16]
|
||||
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 24 + 20]
|
||||
|
||||
vfmadd231ps ymm8, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm13, ymm14
|
||||
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
vfmadd231ps ymm11, ymm13, ymm15
|
||||
|
||||
add rax, 128
|
||||
add rcx, 48
|
||||
+39
@@ -0,0 +1,39 @@
|
||||
// Tile size: 2x6
|
||||
// Accumulators: 0-11
|
||||
// Col regs: ymm14-15
|
||||
// Row regs: ymm12-13
|
||||
|
||||
// Load ordered by earliest use for first 2x2 block
|
||||
vbroadcastss ymm14, dword ptr [rcx]
|
||||
vmovaps ymm12, [rax]
|
||||
vmovaps ymm13, [rax + 32]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm14
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm15
|
||||
vfmadd231ps ymm3, ymm13, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm5, ymm13, ymm14
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 20]
|
||||
|
||||
vfmadd231ps ymm8, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm13, ymm14
|
||||
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
vfmadd231ps ymm11, ymm13, ymm15
|
||||
|
||||
add rax, 64
|
||||
add rcx, 24
|
||||
+60
@@ -0,0 +1,60 @@
|
||||
// Tile size: 3x4
|
||||
// Accumulators: 0-11
|
||||
// Col regs: ymm12-14
|
||||
// Row regs: ymm15
|
||||
|
||||
vmovaps ymm12, [rax]
|
||||
vmovaps ymm13, [rax+32]
|
||||
vmovaps ymm14, [rax+64]
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 0]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm15
|
||||
vfmadd231ps ymm1, ymm13, ymm15
|
||||
vfmadd231ps ymm2, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm15
|
||||
vfmadd231ps ymm4, ymm13, ymm15
|
||||
vfmadd231ps ymm5, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
vfmadd231ps ymm8, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
vfmadd231ps ymm10, ymm13, ymm15
|
||||
vfmadd231ps ymm11, ymm14, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 96]
|
||||
vmovaps ymm13, [rax + 128]
|
||||
vmovaps ymm14, [rax + 160]
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 16]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm15
|
||||
vfmadd231ps ymm1, ymm13, ymm15
|
||||
vfmadd231ps ymm2, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 20]
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm15
|
||||
vfmadd231ps ymm4, ymm13, ymm15
|
||||
vfmadd231ps ymm5, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 24]
|
||||
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
vfmadd231ps ymm8, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 28]
|
||||
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
vfmadd231ps ymm10, ymm13, ymm15
|
||||
vfmadd231ps ymm11, ymm14, ymm15
|
||||
+32
@@ -0,0 +1,32 @@
|
||||
// Tile size: 3x4
|
||||
// Accumulators: 0-11
|
||||
// Col regs: ymm12-14
|
||||
// Row regs: ymm15
|
||||
|
||||
vmovaps ymm12, [rax]
|
||||
vmovaps ymm13, [rax+32]
|
||||
vmovaps ymm14, [rax+64]
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 0]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm15
|
||||
vfmadd231ps ymm1, ymm13, ymm15
|
||||
vfmadd231ps ymm2, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm15
|
||||
vfmadd231ps ymm4, ymm13, ymm15
|
||||
vfmadd231ps ymm5, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
vfmadd231ps ymm8, ymm14, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 12]
|
||||
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
vfmadd231ps ymm10, ymm13, ymm15
|
||||
vfmadd231ps ymm11, ymm14, ymm15
|
||||
+69
@@ -0,0 +1,69 @@
|
||||
// Tile size: 4x3
|
||||
// Accumulators: 0-11
|
||||
// Col regs: ymm12
|
||||
// Row regs: ymm13-15
|
||||
|
||||
// Load col of A
|
||||
vmovaps ymm12, [rax]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vbroadcastss ymm13, dword ptr [rcx + 0]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 4]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 8]
|
||||
|
||||
// N.B. Stepping cols in inner loop
|
||||
vfmadd231ps ymm0, ymm12, ymm13
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+32]
|
||||
|
||||
vfmadd231ps ymm1, ymm12, ymm13
|
||||
vfmadd231ps ymm5, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+64]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm13
|
||||
vfmadd231ps ymm6, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+96]
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm13
|
||||
vfmadd231ps ymm7, ymm12, ymm14
|
||||
vfmadd231ps ymm11, ymm12, ymm15
|
||||
|
||||
// Load col of A, switching col!
|
||||
vmovaps ymm13, [rax + 128]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vbroadcastss ymm14, dword ptr [rcx + 12]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 16]
|
||||
vbroadcastss ymm12, dword ptr [rcx + 20]
|
||||
|
||||
// N.B. Stepping cols in inner loop
|
||||
vfmadd231ps ymm0, ymm13, ymm14
|
||||
vfmadd231ps ymm4, ymm13, ymm15
|
||||
vfmadd231ps ymm8, ymm13, ymm12
|
||||
|
||||
vmovaps ymm13, [rax + 160]
|
||||
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
vfmadd231ps ymm5, ymm13, ymm15
|
||||
vfmadd231ps ymm9, ymm13, ymm12
|
||||
|
||||
vmovaps ymm13, [rax + 192]
|
||||
|
||||
vfmadd231ps ymm2, ymm13, ymm14
|
||||
vfmadd231ps ymm6, ymm13, ymm15
|
||||
vfmadd231ps ymm10, ymm13, ymm12
|
||||
|
||||
vmovaps ymm13, [rax + 224]
|
||||
|
||||
vfmadd231ps ymm3, ymm13, ymm14
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
vfmadd231ps ymm11, ymm13, ymm12
|
||||
|
||||
add rcx, 24
|
||||
add rax, 256
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
// Tile size: 4x3
|
||||
// Accumulators: 0-11
|
||||
// Col regs: ymm12
|
||||
// Row regs: ymm13-15
|
||||
|
||||
// Load col of A
|
||||
vmovaps ymm12, [rax]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vbroadcastss ymm13, dword ptr [rcx + 0]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 4]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 8]
|
||||
|
||||
// N.B. Stepping cols in inner loop
|
||||
vfmadd231ps ymm0, ymm12, ymm13
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+32]
|
||||
|
||||
vfmadd231ps ymm1, ymm12, ymm13
|
||||
vfmadd231ps ymm5, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+64]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm13
|
||||
vfmadd231ps ymm6, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+96]
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm13
|
||||
vfmadd231ps ymm7, ymm12, ymm14
|
||||
vfmadd231ps ymm11, ymm12, ymm15
|
||||
|
||||
add rcx, 12
|
||||
add rax, 128
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
// Tile size: 5x2
|
||||
// Accumulators: 0-9
|
||||
// Col regs: ymm10-13
|
||||
// Row regs: ymm14-15
|
||||
|
||||
vmovaps ymm10, [rax]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 0]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
vmovaps ymm11, [rax + 32]
|
||||
|
||||
// NB stepping column-wise
|
||||
vfmadd231ps ymm0, ymm10, ymm14
|
||||
vfmadd231ps ymm5, ymm10, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 64]
|
||||
|
||||
vfmadd231ps ymm1, ymm11, ymm14
|
||||
vfmadd231ps ymm6, ymm11, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm14
|
||||
vfmadd231ps ymm7, ymm12, ymm15
|
||||
|
||||
vmovaps ymm10, [rax + 128]
|
||||
|
||||
vfmadd231ps ymm3, ymm13, ymm14
|
||||
vfmadd231ps ymm8, ymm13, ymm15
|
||||
|
||||
vmovaps ymm11, [rax + 160]
|
||||
|
||||
vfmadd231ps ymm4, ymm10, ymm14
|
||||
vfmadd231ps ymm9, ymm10, ymm15
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 8]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 12]
|
||||
|
||||
vmovaps ymm12, [rax + 192]
|
||||
|
||||
// NB stepping column-wise
|
||||
vfmadd231ps ymm0, ymm11, ymm14
|
||||
vfmadd231ps ymm5, ymm11, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 224]
|
||||
|
||||
vfmadd231ps ymm1, ymm12, ymm14
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
|
||||
vmovaps ymm10, [rax + 256]
|
||||
|
||||
vfmadd231ps ymm2, ymm13, ymm14
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
|
||||
vmovaps ymm11, [rax + 288]
|
||||
|
||||
vfmadd231ps ymm3, ymm10, ymm14
|
||||
vfmadd231ps ymm8, ymm10, ymm15
|
||||
|
||||
vfmadd231ps ymm4, ymm11, ymm14
|
||||
vfmadd231ps ymm9, ymm11, ymm15
|
||||
|
||||
add rax, 320
|
||||
add rcx, 16
|
||||
+34
@@ -0,0 +1,34 @@
|
||||
// Tile size: 5x2
|
||||
// Accumulators: 0-9
|
||||
// Col regs: ymm10-13
|
||||
// Row regs: ymm14-15
|
||||
|
||||
vmovaps ymm10, [rax]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 0]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
vmovaps ymm11, [rax + 32]
|
||||
|
||||
// NB stepping column-wise
|
||||
vfmadd231ps ymm0, ymm10, ymm14
|
||||
vfmadd231ps ymm5, ymm10, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 64]
|
||||
|
||||
vfmadd231ps ymm1, ymm11, ymm14
|
||||
vfmadd231ps ymm6, ymm11, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm14
|
||||
vfmadd231ps ymm7, ymm12, ymm15
|
||||
|
||||
vmovaps ymm11, [rax + 128]
|
||||
|
||||
vfmadd231ps ymm3, ymm13, ymm14
|
||||
vfmadd231ps ymm8, ymm13, ymm15
|
||||
|
||||
vfmadd231ps ymm4, ymm11, ymm14
|
||||
vfmadd231ps ymm9, ymm11, ymm15
|
||||
|
||||
add rax, 160
|
||||
add rcx, 8
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
// Tile size: 6x1
|
||||
// Accumulators: 0-5
|
||||
// Col regs: 6-11
|
||||
// Row regs: 15
|
||||
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
vfmadd231ps ymm0, ymm15, [rax]
|
||||
vfmadd231ps ymm1, ymm15, [rax + 32]
|
||||
vfmadd231ps ymm2, ymm15, [rax + 64]
|
||||
vfmadd231ps ymm3, ymm15, [rax + 96]
|
||||
vfmadd231ps ymm4, ymm15, [rax + 128]
|
||||
vfmadd231ps ymm5, ymm15, [rax + 160]
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rcx + 4]
|
||||
|
||||
vfmadd231ps ymm0, ymm14, [rax + 192]
|
||||
vfmadd231ps ymm1, ymm14, [rax + 224]
|
||||
vfmadd231ps ymm2, ymm14, [rax + 256]
|
||||
vfmadd231ps ymm3, ymm14, [rax + 288]
|
||||
vfmadd231ps ymm4, ymm14, [rax + 320]
|
||||
vfmadd231ps ymm5, ymm14, [rax + 352]
|
||||
|
||||
add rax, 384
|
||||
add rcx, 8
|
||||
+29
@@ -0,0 +1,29 @@
|
||||
// Tile size: 6x1
|
||||
// Accumulators: 0-5
|
||||
// Col regs: 6-11
|
||||
// Row regs: 15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovups ymm10, [rax]
|
||||
vmulps ymm10, ymm10, ymm15
|
||||
vaddps ymm0, ymm0, ymm10
|
||||
vmovups ymm11, [rax + 32]
|
||||
vmulps ymm11, ymm11, ymm15
|
||||
vaddps ymm1, ymm1, ymm11
|
||||
vmovups ymm12, [rax + 64]
|
||||
vmulps ymm12, ymm12, ymm15
|
||||
vaddps ymm2, ymm2, ymm12
|
||||
vmovups ymm13, [rax + 96]
|
||||
vmulps ymm13, ymm13, ymm15
|
||||
vaddps ymm3, ymm3, ymm13
|
||||
vmovups ymm14, [rax + 128]
|
||||
vmulps ymm14, ymm14, ymm15
|
||||
vaddps ymm4, ymm4, ymm14
|
||||
vmovups ymm15, [rax + 160]
|
||||
vmulps ymm15, ymm15, ymm15
|
||||
vaddps ymm5, ymm5, ymm15
|
||||
|
||||
|
||||
add rcx, 4
|
||||
add rax, 192
|
||||
+70
@@ -0,0 +1,70 @@
|
||||
// Tile size: 6x2
|
||||
// Accumulators: 0-9
|
||||
// Col regs: ymm10-13
|
||||
// Row regs: ymm14-15
|
||||
|
||||
vmovaps ymm12, [rax]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 0]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
vmovaps ymm13, [rax + 32]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm14
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 64]
|
||||
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 128]
|
||||
|
||||
vfmadd231ps ymm3, ymm13, ymm14
|
||||
vfmadd231ps ymm9, ymm13, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 160]
|
||||
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 192]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 8]
|
||||
|
||||
vfmadd231ps ymm5, ymm13, ymm14
|
||||
vfmadd231ps ymm11, ymm13, ymm15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx + 12]
|
||||
vmovaps ymm13, [rax + 224]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm14
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 256]
|
||||
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 288]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 320]
|
||||
|
||||
vfmadd231ps ymm3, ymm13, ymm14
|
||||
vfmadd231ps ymm9, ymm13, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 352]
|
||||
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vfmadd231ps ymm5, ymm13, ymm14
|
||||
vfmadd231ps ymm11, ymm13, ymm15
|
||||
|
||||
add rax, 384
|
||||
add rcx, 16
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
// Tile size: 6x2
|
||||
// Accumulators: 0-11
|
||||
// Col regs: 12-13
|
||||
// Row regs: 14-15
|
||||
|
||||
vmovaps ymm12, [rax]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 0]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 4]
|
||||
vmovaps ymm13, [rax + 32]
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm14
|
||||
vfmadd231ps ymm6, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 64]
|
||||
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
vfmadd231ps ymm7, ymm13, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 128]
|
||||
|
||||
vfmadd231ps ymm3, ymm13, ymm14
|
||||
vfmadd231ps ymm9, ymm13, ymm15
|
||||
|
||||
vmovaps ymm13, [rax + 160]
|
||||
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vfmadd231ps ymm5, ymm13, ymm14
|
||||
vfmadd231ps ymm11, ymm13, ymm15
|
||||
|
||||
add rcx, 8
|
||||
add rax, 192
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
// Tile size: 6x1
|
||||
// Accumulators: 0-5
|
||||
// Col regs: 6-11
|
||||
// Row regs: 15
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovaps ymm6, [rax + 0]
|
||||
vmovaps ymm7, [rax + 32]
|
||||
vmovaps ymm8, [rax + 64]
|
||||
vmovaps ymm9, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm0, ymm6, ymm15
|
||||
vmovaps ymm10, [rax + 128]
|
||||
|
||||
vfmadd231ps ymm1, ymm7, ymm15
|
||||
vmovaps ymm11, [rax + 160]
|
||||
vfmadd231ps ymm2, ymm8, ymm15
|
||||
vbroadcastss ymm14, dword ptr [rcx+4]
|
||||
vfmadd231ps ymm3, ymm9, ymm15
|
||||
vmovaps ymm12, [rax + 192]
|
||||
vfmadd231ps ymm4, ymm10, ymm15
|
||||
vmovaps ymm13, [rax + 224]
|
||||
vfmadd231ps ymm5, ymm11, ymm15
|
||||
|
||||
vmovaps ymm6, [rax + 256]
|
||||
vfmadd231ps ymm0, ymm12, ymm14
|
||||
vmovaps ymm7, [rax + 288]
|
||||
vfmadd231ps ymm1, ymm13, ymm14
|
||||
|
||||
vmovaps ymm8, [rax + 128]
|
||||
vfmadd231ps ymm2, ymm6, ymm14
|
||||
|
||||
vmovaps ymm9, [rax + 160]
|
||||
vfmadd231ps ymm3, ymm7, ymm14
|
||||
|
||||
vfmadd231ps ymm4, ymm8, ymm14
|
||||
vfmadd231ps ymm5, ymm9, ymm14
|
||||
+22
@@ -0,0 +1,22 @@
|
||||
// Tile size: 6x1
|
||||
// Accumulators: 0-5
|
||||
// Col regs: 6-11
|
||||
// Row regs: 15
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovaps ymm6, [rax + 0]
|
||||
vmovaps ymm7, [rax + 32]
|
||||
vmovaps ymm8, [rax + 64]
|
||||
vmovaps ymm9, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm0, ymm6, ymm15
|
||||
vfmadd231ps ymm1, ymm7, ymm15
|
||||
|
||||
vmovaps ymm10, [rax + 128]
|
||||
vfmadd231ps ymm2, ymm8, ymm15
|
||||
|
||||
vmovaps ymm11, [rax + 160]
|
||||
vfmadd231ps ymm3, ymm9, ymm15
|
||||
|
||||
vfmadd231ps ymm4, ymm10, ymm15
|
||||
vfmadd231ps ymm5, ymm11, ymm15
|
||||
+48
@@ -0,0 +1,48 @@
|
||||
// Accumulators: 0-7
|
||||
// Columns: 14-15
|
||||
// Rows: 8-13
|
||||
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 4]
|
||||
|
||||
vmovaps ymm8, [rax]
|
||||
vmovaps ymm9, [rax + 32]
|
||||
vmovaps ymm10, [rax + 64]
|
||||
vmovaps ymm11, [rax + 96]
|
||||
vmovaps ymm12, [rax + 128]
|
||||
vmovaps ymm13, [rax + 160]
|
||||
|
||||
vfmadd231ps ymm0, ymm15, ymm8
|
||||
vfmadd231ps ymm1, ymm15, ymm9
|
||||
vfmadd231ps ymm2, ymm15, ymm10
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
vfmadd231ps ymm4, ymm15, ymm12
|
||||
vfmadd231ps ymm5, ymm15, ymm13
|
||||
|
||||
vmovaps ymm8, [rax + 192]
|
||||
vmovaps ymm9, [rax + 224]
|
||||
vmovaps ymm10, [rax + 256]
|
||||
vmovaps ymm11, [rax + 288]
|
||||
vmovaps ymm12, [rax + 320]
|
||||
vmovaps ymm13, [rax + 352]
|
||||
|
||||
vfmadd231ps ymm6, ymm15, ymm8
|
||||
vfmadd231ps ymm7, ymm15, ymm9
|
||||
vfmadd231ps ymm0, ymm14, ymm10
|
||||
vfmadd231ps ymm1, ymm14, ymm11
|
||||
vfmadd231ps ymm2, ymm14, ymm12
|
||||
vfmadd231ps ymm3, ymm14, ymm13
|
||||
|
||||
vmovaps ymm8, [rax + 384]
|
||||
vmovaps ymm9, [rax + 416]
|
||||
vmovaps ymm10, [rax + 448]
|
||||
vmovaps ymm11, [rax + 480]
|
||||
|
||||
vfmadd231ps ymm4, ymm14, ymm8
|
||||
vfmadd231ps ymm5, ymm14, ymm9
|
||||
vfmadd231ps ymm6, ymm14, ymm10
|
||||
vfmadd231ps ymm7, ymm14, ymm11
|
||||
|
||||
add rcx, 8
|
||||
add rax, 512
|
||||
+33
@@ -0,0 +1,33 @@
|
||||
// Tile size: 8x1
|
||||
// Accumulators: 0-7
|
||||
// Col regs: 8-14
|
||||
// Row regs: 15
|
||||
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovaps ymm8, [rax + 0]
|
||||
vmovaps ymm9, [rax + 32]
|
||||
vmovaps ymm10, [rax + 64]
|
||||
vmovaps ymm11, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm0, ymm8, ymm15
|
||||
vfmadd231ps ymm1, ymm9, ymm15
|
||||
|
||||
vmovaps ymm12, [rax + 128]
|
||||
vmovaps ymm13, [rax + 160]
|
||||
|
||||
vfmadd231ps ymm2, ymm10, ymm15
|
||||
vfmadd231ps ymm3, ymm11, ymm15
|
||||
|
||||
vmovaps ymm14, [rax + 192]
|
||||
vmovaps ymm11, [rax + 224]
|
||||
|
||||
vfmadd231ps ymm4, ymm12, ymm15
|
||||
vfmadd231ps ymm5, ymm13, ymm15
|
||||
|
||||
|
||||
vfmadd231ps ymm6, ymm14, ymm15
|
||||
vfmadd231ps ymm7, ymm11, ymm15
|
||||
|
||||
add rcx, 4
|
||||
add rax, 256
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
// Tile size: 1x8
|
||||
// Accumulators: 0-7
|
||||
// Col regs: 8-14
|
||||
// Row regs: 15
|
||||
|
||||
|
||||
vmovaps ymm15, [rax]
|
||||
|
||||
vbroadcastss ymm8, dword ptr [rcx + 0 * 4]
|
||||
vfmadd231ps ymm0, ymm15, ymm8
|
||||
|
||||
vbroadcastss ymm9, dword ptr [rcx + 1 * 4]
|
||||
vfmadd231ps ymm1, ymm15, ymm9
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rcx + 2 * 4]
|
||||
vfmadd231ps ymm2, ymm15, ymm10
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 3 * 4]
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
|
||||
vbroadcastss ymm12, dword ptr [rcx + 4 * 4]
|
||||
vfmadd231ps ymm4, ymm15, ymm12
|
||||
|
||||
vbroadcastss ymm13, dword ptr [rcx + 5 * 4]
|
||||
vfmadd231ps ymm5, ymm15, ymm13
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rcx + 6 * 4]
|
||||
vfmadd231ps ymm6, ymm15, ymm10
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 7 * 4]
|
||||
vfmadd231ps ymm7, ymm15, ymm11
|
||||
|
||||
|
||||
vmovaps ymm15, [rax]
|
||||
|
||||
vbroadcastss ymm8, dword ptr [rcx + 0 * 4]
|
||||
vfmadd231ps ymm0, ymm15, ymm8
|
||||
|
||||
vbroadcastss ymm9, dword ptr [rcx + 1 * 4]
|
||||
vfmadd231ps ymm1, ymm15, ymm9
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rcx + 2 * 4]
|
||||
vfmadd231ps ymm2, ymm15, ymm10
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 3 * 4]
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
|
||||
vbroadcastss ymm12, dword ptr [rcx + 4 * 4]
|
||||
vfmadd231ps ymm4, ymm15, ymm12
|
||||
|
||||
vbroadcastss ymm13, dword ptr [rcx + 5 * 4]
|
||||
vfmadd231ps ymm5, ymm15, ymm13
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rcx + 6 * 4]
|
||||
vfmadd231ps ymm6, ymm15, ymm10
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 7 * 4]
|
||||
vfmadd231ps ymm7, ymm15, ymm11
|
||||
+30
@@ -0,0 +1,30 @@
|
||||
// Tile size: 1x8
|
||||
// Accumulators: 0-7
|
||||
// Col regs: 8-14
|
||||
// Row regs: 15
|
||||
|
||||
vmovaps ymm15, [rax]
|
||||
|
||||
vbroadcastss ymm8, dword ptr [rcx + 0 * 4]
|
||||
vfmadd231ps ymm0, ymm15, ymm8
|
||||
|
||||
vbroadcastss ymm9, dword ptr [rcx + 1 * 4]
|
||||
vfmadd231ps ymm1, ymm15, ymm9
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rcx + 2 * 4]
|
||||
vfmadd231ps ymm2, ymm15, ymm10
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 3 * 4]
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
|
||||
vbroadcastss ymm12, dword ptr [rcx + 4 * 4]
|
||||
vfmadd231ps ymm4, ymm15, ymm12
|
||||
|
||||
vbroadcastss ymm13, dword ptr [rcx + 5 * 4]
|
||||
vfmadd231ps ymm5, ymm15, ymm13
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rcx + 6 * 4]
|
||||
vfmadd231ps ymm6, ymm15, ymm10
|
||||
|
||||
vbroadcastss ymm11, dword ptr [rcx + 7 * 4]
|
||||
vfmadd231ps ymm7, ymm15, ymm11
|
||||
Vendored
+682
@@ -0,0 +1,682 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 8x8:
|
||||
|
||||
ymm0 ymm1 ymm2 ymm3 ymm4 ymm5 ymm6 ymm7
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
avx2_mmm_i32_8x8_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}avx2_mmm_i32_8x8_{{suffix}}
|
||||
{{G}}avx2_mmm_i32_8x8_{{suffix}}:
|
||||
.cfi_startproc
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
push rdi
|
||||
push rsi
|
||||
|
||||
mov rdi, rcx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
.cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
|
||||
{% include "dispatcher.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov r12, [rdi + 32] // packing
|
||||
mov rbx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rcx, [rdi + 8] // k
|
||||
test rcx, rcx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
cmp r12, 1
|
||||
je {{L}}main_loop_packed_packed_i8i8
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
vmovaps ymm12, [rax]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{i}} * 4]
|
||||
vpmulld ymm13, ymm12, ymm14
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm13
|
||||
{% endfor %}
|
||||
|
||||
add rax, 32
|
||||
add rbx, 32
|
||||
dec rcx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed_i8i8:
|
||||
movq xmm8, qword ptr [rax] // read 8 bytes
|
||||
vpmovsxbw ymm8, xmm8 // promote byte to i32x8
|
||||
|
||||
vpbroadcastb ymm9, byte ptr [rbx] // broadcast 1 byte from B
|
||||
vpbroadcastb ymm10, byte ptr [rbx + 1] // broadcast 1 byte from B
|
||||
vpbroadcastb ymm11, byte ptr [rbx + 2] // broadcast 1 byte from B
|
||||
vpbroadcastb ymm12, byte ptr [rbx + 3] // broadcast 1 byte from B
|
||||
vpmovsxbw ymm9, xmm9 // promote byte to i32x8
|
||||
vpmovsxbw ymm10, xmm10 // promote byte to i32x8
|
||||
vpmovsxbw ymm11, xmm11 // promote byte to i32x8
|
||||
vpmovsxbw ymm12, xmm12 // promote byte to i32x8
|
||||
|
||||
vpmullw ymm9, ymm9, ymm8
|
||||
vpmullw ymm10, ymm10, ymm8
|
||||
vpmullw ymm11, ymm11, ymm8
|
||||
vpmullw ymm12, ymm12, ymm8
|
||||
vpmovsxwd ymm9, xmm9 // promote byte to i32x8
|
||||
vpmovsxwd ymm10, xmm10 // promote byte to i32x8
|
||||
vpmovsxwd ymm11, xmm11 // promote byte to i32x8
|
||||
vpmovsxwd ymm12, xmm12 // promote byte to i32x8
|
||||
vpaddd ymm0, ymm0, ymm9
|
||||
vpaddd ymm1, ymm1, ymm10
|
||||
vpaddd ymm2, ymm2, ymm11
|
||||
vpaddd ymm3, ymm3, ymm12
|
||||
|
||||
vpbroadcastb ymm9, byte ptr [rbx + 4]
|
||||
vpbroadcastb ymm10, byte ptr [rbx + 5]
|
||||
vpbroadcastb ymm11, byte ptr [rbx + 6]
|
||||
vpbroadcastb ymm12, byte ptr [rbx + 7]
|
||||
vpmovsxbw ymm9, xmm9
|
||||
vpmovsxbw ymm10, xmm10
|
||||
vpmovsxbw ymm11, xmm11
|
||||
vpmovsxbw ymm12, xmm12
|
||||
|
||||
vpmullw ymm9, ymm9, ymm8
|
||||
vpmullw ymm10, ymm10, ymm8
|
||||
vpmullw ymm11, ymm11, ymm8
|
||||
vpmullw ymm12, ymm12, ymm8
|
||||
vpmovsxwd ymm9, xmm9 // promote byte to i32x8
|
||||
vpmovsxwd ymm10, xmm10 // promote byte to i32x8
|
||||
vpmovsxwd ymm11, xmm11 // promote byte to i32x8
|
||||
vpmovsxwd ymm12, xmm12 // promote byte to i32x8
|
||||
vpaddd ymm4, ymm4, ymm9
|
||||
vpaddd ymm5, ymm5, ymm10
|
||||
vpaddd ymm6, ymm6, ymm11
|
||||
vpaddd ymm7, ymm7, ymm12
|
||||
|
||||
add rbx, 8
|
||||
add rax, 8
|
||||
dec rcx
|
||||
jnz {{L}}main_loop_packed_packed_i8i8
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_i32_scalars.j2" %}
|
||||
{% set mr = 8 %}{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_i32_per_rows.j2" %}
|
||||
{% set mr = 8 %}{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_i32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
mov r8, [rdi + 32] // item size
|
||||
|
||||
cmp r8, 4
|
||||
je {{L}}non_linear_addc_i32
|
||||
|
||||
{#
|
||||
// This is not great as vgatherdps reads 32-bits values and goes beyond our buffer. Probably harmless though.
|
||||
// Commented and replaced with the "mov al" loop beyond to pacify valgrind.
|
||||
// ymm14 and ymm15 are the same as in the non_linear_addc_i32 case (compute them before the test right above here.
|
||||
// {% for i in range(0, 8) %}
|
||||
// vpcmpeqd ymm15, ymm15, ymm15
|
||||
// vgatherdps ymm12, [ r10 + ymm14 ], ymm15 // 0xxx 1xxx 2xxx 3xxx 4xxx 5xxx 6xxx 7xxx
|
||||
//
|
||||
// // we need to go through vpmovsxbd, shuffling naively erases signs
|
||||
// vpshufb ymm12, ymm12, ymm10 // 0123 0123 0123 0123 4567 4567 4567 4567
|
||||
//
|
||||
// vpermd ymm12, ymm11, ymm12 // 0123 4567
|
||||
// vpmovsxbd ymm12, xmm12 // sign extend
|
||||
//
|
||||
// vpaddd ymm{{i}}, ymm{{i}}, ymm12
|
||||
// add r10, rbx
|
||||
// {% endfor %}
|
||||
#}
|
||||
|
||||
{% for col in range(0, 8) %}
|
||||
mov r8, r10
|
||||
{% for half in range(0, 2) %}
|
||||
{% for lane in range(0, 4) %}
|
||||
mov al, [ r8 ]
|
||||
add r8, rsi
|
||||
movsx eax, al
|
||||
pinsrd xmm10, eax, {{lane}}
|
||||
{% endfor %}
|
||||
vperm2f128 ymm10, ymm10, ymm10, 1
|
||||
{% endfor %}
|
||||
vpaddd ymm{{col}}, ymm{{col}}, ymm10
|
||||
add r10, rbx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}non_linear_addc_i32:
|
||||
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
vpermq ymm14, ymm14, 78 // 0b01001110
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
vpermq ymm14, ymm14, 78 // 0b01001110
|
||||
|
||||
|
||||
{% if msvc %}
|
||||
vpbroadcastd ymm10, dword ptr [ offset byte_shuffle ]
|
||||
vmovups ymm11, dword ptr [ offset i128_shuffle ]
|
||||
{% else %}
|
||||
vpbroadcastd ymm10, [ rip + {{L}}byte_shuffle ]
|
||||
vmovups ymm11, [ rip + {{L}}i128_shuffle ]
|
||||
{% endif %}
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r10 + ymm14 ], ymm15
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm12
|
||||
add r10, rbx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% if msvc %}
|
||||
.data
|
||||
byte_shuffle dd 201851904 // 0x0c080400
|
||||
i128_shuffle dd 0, 4
|
||||
.code
|
||||
{% else %}
|
||||
{{L}}byte_shuffle: .int 201851904 // 0x0c080400
|
||||
{{L}}i128_shuffle: .int 0, 4
|
||||
{% endif %}
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups ymm12, [rax]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vpmulld ymm15, ymm12, ymm14
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm15
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale:
|
||||
mov r8, [ rdi + 16 ] // policy
|
||||
vbroadcastss ymm8, dword ptr [rdi + 24] // multi
|
||||
|
||||
mov rax, 1
|
||||
movq xmm9, rax
|
||||
vpbroadcastq ymm9, xmm9 // ymm9 <- 1
|
||||
|
||||
mov rax, [ rdi + 8 ] // xmm10 <- shift + 31
|
||||
add rax, 31
|
||||
movq xmm10, rax
|
||||
vpbroadcastq ymm10, xmm10
|
||||
|
||||
mov rax, 1
|
||||
movq xmm11, rax
|
||||
vpsubq ymm12, ymm10, ymm9 // shift+31 - 1
|
||||
vpsllq ymm11, ymm9, xmm12 // ymm11 <- 1 << (shift + 31 - 1)
|
||||
|
||||
cmp r8, 1
|
||||
je {{L}}q_scale_rounding_zero
|
||||
cmp r8, 2
|
||||
je {{L}}q_scale_rounding_away
|
||||
cmp r8, 3
|
||||
je {{L}}q_scale_rounding_minus_inf
|
||||
cmp r8, 4
|
||||
je {{L}}q_scale_rounding_plus_inf
|
||||
cmp r8, 5
|
||||
je {{L}}q_scale_rounding_even
|
||||
cmp r8, 6
|
||||
je {{L}}q_scale_rounding_odd
|
||||
|
||||
jmp {{L}}unsupported
|
||||
|
||||
{{L}}q_scale_rounding_zero: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsubq ymm14, ymm14, ymm9
|
||||
vpsubq ymm15, ymm15, ymm9
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_away: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_minus_inf: // signum * ( (abs << 32 + 1<<30+shift) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
// sign extract for nudging in the right direction
|
||||
vpxor ymm13, ymm13, ymm13
|
||||
vpcmpgtd ymm13, ymm{{i}}, ymm13 // ymm13 <- s0, s1, ..s8 (signums, as all ones or all zeros)
|
||||
vpsrld ymm13, ymm13, 31 // then just 0 or 1
|
||||
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
// reinterpret ymm13=s0i32..s7 as i64 and blend with zero to pick the even ones as i64
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm14, ymm14, ymm12
|
||||
|
||||
vpsrldq ymm13, ymm13, 4 // ymm13 <- s1, s2, .., s7, 0
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm15, ymm15, ymm12
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_plus_inf: // signum * ( (abs << 32 + 1<<30+shift) >> shift )
|
||||
|
||||
vpbroadcastd ymm9, xmm9
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpxor ymm13, ymm13, ymm13
|
||||
|
||||
// sign extract for nudging in the right direction
|
||||
vpcmpgtd ymm13, ymm{{i}}, ymm13 // ymm13 <- s0, s1, ..s8 (signums, as all ones or all zeros)
|
||||
vpaddd ymm13, ymm13, ymm9 // if val >= 0 { 0i32 } else { 1i32 }
|
||||
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
// reinterpret ymm13=s0i32..s7 as i64 and blend with zero to pick the even ones as i64
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm14, ymm14, ymm12
|
||||
|
||||
vpsrldq ymm13, ymm13, 4 // ymm13 <- s1, s2, .., s7, 0
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm15, ymm15, ymm12
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_even: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpsrlq ymm12, ymm14, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpaddq ymm14, ymm14, ymm12
|
||||
vpsubq ymm14, ymm14, ymm9
|
||||
|
||||
vpsrlq ymm12, ymm15, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpaddq ymm15, ymm15, ymm12
|
||||
vpsubq ymm15, ymm15, ymm9
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_odd: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpsrlq ymm12, ymm14, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpsubq ymm14, ymm14, ymm12
|
||||
|
||||
vpsrlq ymm12, ymm15, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpsubq ymm15, ymm15, ymm12
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shl:
|
||||
mov eax, [ rdi + 8 ] // xmm10 <- -shift (8 times)
|
||||
movd xmm10, eax
|
||||
vpbroadcastd ymm10, xmm10
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpsllvd ymm{{i}}, ymm{{i}}, ymm10
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr:
|
||||
mov r8, [ rdi + 16 ] // policy
|
||||
|
||||
mov eax, 1
|
||||
movd xmm9, eax
|
||||
vpbroadcastd ymm9, xmm9 // ymm9 <- 1u32 (8 times)
|
||||
|
||||
mov eax, [ rdi + 8 ] // xmm10 <- shift (8 times)
|
||||
movd xmm10, eax
|
||||
vpbroadcastd ymm10, xmm10
|
||||
|
||||
mov ebx, 1
|
||||
mov cl, al
|
||||
sub cl, 1 // rcx <- shift -1
|
||||
sal ebx, cl // rbx <- (1 << (shift - 1))
|
||||
movd xmm11, ebx
|
||||
vpbroadcastd ymm11, xmm11 // ymm11 <- "half"
|
||||
|
||||
vpxor ymm12, ymm12, ymm12 // ymm12 <- zeroes
|
||||
|
||||
cmp r8, 1
|
||||
je {{L}}q_shr_rounding_zero
|
||||
cmp r8, 2
|
||||
je {{L}}q_shr_rounding_away
|
||||
cmp r8, 3
|
||||
je {{L}}q_shr_rounding_minus_inf
|
||||
cmp r8, 4
|
||||
je {{L}}q_shr_rounding_plus_inf
|
||||
cmp r8, 5
|
||||
je {{L}}q_shr_rounding_even
|
||||
cmp r8, 6
|
||||
je {{L}}q_shr_rounding_odd
|
||||
|
||||
jmp {{L}}unsupported
|
||||
|
||||
{{L}}q_shr_rounding_zero:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsubd ymm14, ymm14, ymm9
|
||||
vpaddd ymm14, ymm14, ymm11
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_away:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpaddd ymm14, ymm14, ymm11
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_minus_inf:
|
||||
{% for i in range(0, 8) %}
|
||||
vpsubd ymm{{i}}, ymm{{i}}, ymm9
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm11
|
||||
vpsravd ymm{{i}}, ymm{{i}}, ymm10
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_plus_inf:
|
||||
{% for i in range(0, 8) %}
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm11
|
||||
vpsravd ymm{{i}}, ymm{{i}}, ymm10
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_even:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsravd ymm13, ymm14, ymm10
|
||||
vpand ymm13, ymm13, ymm9
|
||||
vpsubd ymm13, ymm13, ymm9 // nudge = ((abs >>l shift) & 0x01) - 1
|
||||
vpaddd ymm14, ymm14, ymm13 // add nudge
|
||||
vpaddd ymm14, ymm14, ymm11 // add half
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_odd:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsravd ymm13, ymm14, ymm10
|
||||
vpand ymm13, ymm13, ymm9
|
||||
vpsubd ymm13, ymm12, ymm13 // nudge = - ((abs >>l shift) & 0x01)
|
||||
vpaddd ymm14, ymm14, ymm13 // add nudge
|
||||
vpaddd ymm14, ymm14, ymm11 // add half
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rdx, [rdi + 24] // col stride
|
||||
mov rcx, [rdi + 32] // item size
|
||||
|
||||
cmp rcx, 4
|
||||
je {{L}}store_strides_i32
|
||||
|
||||
{% for col in range(0, 8) %}
|
||||
mov r10, r8
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov byte ptr [r10], bl
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm{{col}}, ymm{{col}}, ymm{{col}}, 1
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov byte ptr [r10], bl
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
add r8, rdx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_strides_i32:
|
||||
{% for col in range(0, 8) %}
|
||||
mov r10, r8
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov dword ptr [r10], ebx
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm{{col}}, ymm{{col}}, ymm{{col}}, 1
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov dword ptr [r10], ebx
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
add r8, rdx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}return:
|
||||
ldmxcsr [rsp + 4]
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
|
||||
{{L}}one_32bit:
|
||||
{% if msvc %}
|
||||
dd 1
|
||||
{% else %}
|
||||
.int 1
|
||||
{% endif %}
|
||||
|
||||
{% if msvc %}
|
||||
avx2_mmm_i32_8x8_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
+676
@@ -0,0 +1,676 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 8x8:
|
||||
|
||||
ymm0 ymm1 ymm2 ymm3 ymm4 ymm5 ymm6 ymm7
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
avx512vnni_mmm_i32_8x8_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}avx512vnni_mmm_i32_8x8_{{suffix}}
|
||||
{{G}}avx512vnni_mmm_i32_8x8_{{suffix}}:
|
||||
.cfi_startproc
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
push rdi
|
||||
push rsi
|
||||
|
||||
mov rdi, rcx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
.cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
|
||||
{% include "dispatcher.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov r12, [rdi + 32] // packing
|
||||
mov rbx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rcx, [rdi + 8] // k
|
||||
test rcx, rcx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
cmp r12, 1
|
||||
je {{L}}main_loop_packed_packed_i8i8
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
vmovaps ymm12, [rax]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{i}} * 4]
|
||||
vpmulld ymm13, ymm12, ymm14
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm13
|
||||
{% endfor %}
|
||||
|
||||
add rax, 32
|
||||
add rbx, 32
|
||||
dec rcx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed_i8i8:
|
||||
// PackedI8K4 layout: per K=4 block, the A panel is 8 rows x 4 K-bytes (32
|
||||
// bytes, lane m = A[m, 4kb..4kb+3]) and the B panel is 8 cols x 4 K-bytes
|
||||
// (lane n = B[n, 4kb..4kb+3]). VPDPBUSD is u8 x s8, so A is offset by +128
|
||||
// (-> u8) and the resulting 128*sum_k(B[n]) bias is removed per column after
|
||||
// the loop, leaving the i32 accumulators identical to the AVX2 path.
|
||||
|
||||
add rcx, 3
|
||||
shr rcx, 2 // rcx <- ceil(k/4) K=4 blocks
|
||||
|
||||
mov r8d, 0x01010101
|
||||
movd xmm11, r8d
|
||||
vpbroadcastd ymm11, xmm11 // ymm11 <- u8 ones (sum of B)
|
||||
|
||||
mov r8d, 0x80808080
|
||||
movd xmm12, r8d
|
||||
vpbroadcastd ymm12, xmm12 // ymm12 <- byte 0x80 (A + 128)
|
||||
|
||||
vpxor ymm10, ymm10, ymm10 // ymm10 <- per-col sum_k B[n]
|
||||
|
||||
{{L}}loop_4k_i8i8:
|
||||
vmovdqu ymm8, [rax] // A block: lane m = A[m,4kb..]
|
||||
vpaddb ymm8, ymm8, ymm12 // s8 -> u8 (+128, modular)
|
||||
|
||||
vmovdqu ymm9, [rbx] // B block: lane n = B[n,4kb..]
|
||||
vpdpbusd ymm10, ymm11, ymm9 // sum_k B[n] += sum_t B[n,4kb+t]
|
||||
|
||||
{% for n in range(0, 8) %}
|
||||
vpbroadcastd ymm13, dword ptr [rbx + {{n}} * 4]
|
||||
vpdpbusd ymm{{n}}, ymm8, ymm13 // acc[n][m] += sum_t (A[m]+128)*B[n]
|
||||
{% endfor %}
|
||||
|
||||
add rax, 32
|
||||
add rbx, 32
|
||||
dec rcx
|
||||
jnz {{L}}loop_4k_i8i8
|
||||
|
||||
// remove the +128 bias added on A: acc[n] -= 128 * sum_k B[n]
|
||||
vpslld ymm10, ymm10, 7 // lane n <- 128 * sum_k B[n]
|
||||
{% for n in range(0, 8) %}
|
||||
mov r8d, {{n}}
|
||||
movd xmm14, r8d
|
||||
vpbroadcastd ymm14, xmm14 // index = n in every lane
|
||||
vpermd ymm15, ymm14, ymm10 // splat 128*sum_k B[n]
|
||||
vpsubd ymm{{n}}, ymm{{n}}, ymm15
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_i32_scalars.j2" %}
|
||||
{% set mr = 8 %}{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_i32_per_rows.j2" %}
|
||||
{% set mr = 8 %}{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_i32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
mov r8, [rdi + 32] // item size
|
||||
|
||||
cmp r8, 4
|
||||
je {{L}}non_linear_addc_i32
|
||||
|
||||
{#
|
||||
// This is not great as vgatherdps reads 32-bits values and goes beyond our buffer. Probably harmless though.
|
||||
// Commented and replaced with the "mov al" loop beyond to pacify valgrind.
|
||||
// ymm14 and ymm15 are the same as in the non_linear_addc_i32 case (compute them before the test right above here.
|
||||
// {% for i in range(0, 8) %}
|
||||
// vpcmpeqd ymm15, ymm15, ymm15
|
||||
// vgatherdps ymm12, [ r10 + ymm14 ], ymm15 // 0xxx 1xxx 2xxx 3xxx 4xxx 5xxx 6xxx 7xxx
|
||||
//
|
||||
// // we need to go through vpmovsxbd, shuffling naively erases signs
|
||||
// vpshufb ymm12, ymm12, ymm10 // 0123 0123 0123 0123 4567 4567 4567 4567
|
||||
//
|
||||
// vpermd ymm12, ymm11, ymm12 // 0123 4567
|
||||
// vpmovsxbd ymm12, xmm12 // sign extend
|
||||
//
|
||||
// vpaddd ymm{{i}}, ymm{{i}}, ymm12
|
||||
// add r10, rbx
|
||||
// {% endfor %}
|
||||
#}
|
||||
|
||||
{% for col in range(0, 8) %}
|
||||
mov r8, r10
|
||||
{% for half in range(0, 2) %}
|
||||
{% for lane in range(0, 4) %}
|
||||
mov al, [ r8 ]
|
||||
add r8, rsi
|
||||
movsx eax, al
|
||||
pinsrd xmm10, eax, {{lane}}
|
||||
{% endfor %}
|
||||
vperm2f128 ymm10, ymm10, ymm10, 1
|
||||
{% endfor %}
|
||||
vpaddd ymm{{col}}, ymm{{col}}, ymm10
|
||||
add r10, rbx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}non_linear_addc_i32:
|
||||
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
vpermq ymm14, ymm14, 78 // 0b01001110
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
vpermq ymm14, ymm14, 78 // 0b01001110
|
||||
|
||||
|
||||
{% if msvc %}
|
||||
vpbroadcastd ymm10, dword ptr [ offset byte_shuffle ]
|
||||
vmovups ymm11, dword ptr [ offset i128_shuffle ]
|
||||
{% else %}
|
||||
vpbroadcastd ymm10, [ rip + {{L}}byte_shuffle ]
|
||||
vmovups ymm11, [ rip + {{L}}i128_shuffle ]
|
||||
{% endif %}
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r10 + ymm14 ], ymm15
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm12
|
||||
add r10, rbx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% if msvc %}
|
||||
.data
|
||||
byte_shuffle dd 201851904 // 0x0c080400
|
||||
i128_shuffle dd 0, 4
|
||||
.code
|
||||
{% else %}
|
||||
{{L}}byte_shuffle: .int 201851904 // 0x0c080400
|
||||
{{L}}i128_shuffle: .int 0, 4
|
||||
{% endif %}
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups ymm12, [rax]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vpmulld ymm15, ymm12, ymm14
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm15
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale:
|
||||
mov r8, [ rdi + 16 ] // policy
|
||||
vbroadcastss ymm8, dword ptr [rdi + 24] // multi
|
||||
|
||||
mov rax, 1
|
||||
movq xmm9, rax
|
||||
vpbroadcastq ymm9, xmm9 // ymm9 <- 1
|
||||
|
||||
mov rax, [ rdi + 8 ] // xmm10 <- shift + 31
|
||||
add rax, 31
|
||||
movq xmm10, rax
|
||||
vpbroadcastq ymm10, xmm10
|
||||
|
||||
mov rax, 1
|
||||
movq xmm11, rax
|
||||
vpsubq ymm12, ymm10, ymm9 // shift+31 - 1
|
||||
vpsllq ymm11, ymm9, xmm12 // ymm11 <- 1 << (shift + 31 - 1)
|
||||
|
||||
cmp r8, 1
|
||||
je {{L}}q_scale_rounding_zero
|
||||
cmp r8, 2
|
||||
je {{L}}q_scale_rounding_away
|
||||
cmp r8, 3
|
||||
je {{L}}q_scale_rounding_minus_inf
|
||||
cmp r8, 4
|
||||
je {{L}}q_scale_rounding_plus_inf
|
||||
cmp r8, 5
|
||||
je {{L}}q_scale_rounding_even
|
||||
cmp r8, 6
|
||||
je {{L}}q_scale_rounding_odd
|
||||
|
||||
jmp {{L}}unsupported
|
||||
|
||||
{{L}}q_scale_rounding_zero: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsubq ymm14, ymm14, ymm9
|
||||
vpsubq ymm15, ymm15, ymm9
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_away: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_minus_inf: // signum * ( (abs << 32 + 1<<30+shift) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
// sign extract for nudging in the right direction
|
||||
vpxor ymm13, ymm13, ymm13
|
||||
vpcmpgtd ymm13, ymm{{i}}, ymm13 // ymm13 <- s0, s1, ..s8 (signums, as all ones or all zeros)
|
||||
vpsrld ymm13, ymm13, 31 // then just 0 or 1
|
||||
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
// reinterpret ymm13=s0i32..s7 as i64 and blend with zero to pick the even ones as i64
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm14, ymm14, ymm12
|
||||
|
||||
vpsrldq ymm13, ymm13, 4 // ymm13 <- s1, s2, .., s7, 0
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm15, ymm15, ymm12
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_plus_inf: // signum * ( (abs << 32 + 1<<30+shift) >> shift )
|
||||
|
||||
vpbroadcastd ymm9, xmm9
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpxor ymm13, ymm13, ymm13
|
||||
|
||||
// sign extract for nudging in the right direction
|
||||
vpcmpgtd ymm13, ymm{{i}}, ymm13 // ymm13 <- s0, s1, ..s8 (signums, as all ones or all zeros)
|
||||
vpaddd ymm13, ymm13, ymm9 // if val >= 0 { 0i32 } else { 1i32 }
|
||||
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
// reinterpret ymm13=s0i32..s7 as i64 and blend with zero to pick the even ones as i64
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm14, ymm14, ymm12
|
||||
|
||||
vpsrldq ymm13, ymm13, 4 // ymm13 <- s1, s2, .., s7, 0
|
||||
vpxor ymm12, ymm12, ymm12
|
||||
vpblendd ymm12, ymm12, ymm13, 85 // 0x55
|
||||
vpsubq ymm15, ymm15, ymm12
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_even: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpsrlq ymm12, ymm14, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpaddq ymm14, ymm14, ymm12
|
||||
vpsubq ymm14, ymm14, ymm9
|
||||
|
||||
vpsrlq ymm12, ymm15, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpaddq ymm15, ymm15, ymm12
|
||||
vpsubq ymm15, ymm15, ymm9
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale_rounding_odd: // signum * ( (abs + nudge) >> shift )
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsrldq ymm15, ymm14, 4 // ymm15 <- a1, a2, a3, a4, a5, a6, a7, 0
|
||||
vpmuldq ymm14, ymm14, ymm8 // ymm14 <- a0*c, a2*c, a4*c, a6*c
|
||||
vpmuldq ymm15, ymm15, ymm8 // ymm15 <- a1*c, a3*c, a5*c, a7*c
|
||||
|
||||
vpsrlq ymm12, ymm14, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpsubq ymm14, ymm14, ymm12
|
||||
|
||||
vpsrlq ymm12, ymm15, xmm10
|
||||
vpand ymm12, ymm12, ymm9
|
||||
vpsubq ymm15, ymm15, ymm12
|
||||
|
||||
vpaddq ymm14, ymm14, ymm11
|
||||
vpaddq ymm15, ymm15, ymm11
|
||||
|
||||
vpsrlq ymm14, ymm14, xmm10
|
||||
vpsrlq ymm15, ymm15, xmm10
|
||||
|
||||
vpslldq ymm15, ymm15, 4
|
||||
vpblendd ymm14, ymm15, ymm14, 85 // 0x55
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shl:
|
||||
mov eax, [ rdi + 8 ] // xmm10 <- -shift (8 times)
|
||||
movd xmm10, eax
|
||||
vpbroadcastd ymm10, xmm10
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpsllvd ymm{{i}}, ymm{{i}}, ymm10
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr:
|
||||
mov r8, [ rdi + 16 ] // policy
|
||||
|
||||
mov eax, 1
|
||||
movd xmm9, eax
|
||||
vpbroadcastd ymm9, xmm9 // ymm9 <- 1u32 (8 times)
|
||||
|
||||
mov eax, [ rdi + 8 ] // xmm10 <- shift (8 times)
|
||||
movd xmm10, eax
|
||||
vpbroadcastd ymm10, xmm10
|
||||
|
||||
mov ebx, 1
|
||||
mov cl, al
|
||||
sub cl, 1 // rcx <- shift -1
|
||||
sal ebx, cl // rbx <- (1 << (shift - 1))
|
||||
movd xmm11, ebx
|
||||
vpbroadcastd ymm11, xmm11 // ymm11 <- "half"
|
||||
|
||||
vpxor ymm12, ymm12, ymm12 // ymm12 <- zeroes
|
||||
|
||||
cmp r8, 1
|
||||
je {{L}}q_shr_rounding_zero
|
||||
cmp r8, 2
|
||||
je {{L}}q_shr_rounding_away
|
||||
cmp r8, 3
|
||||
je {{L}}q_shr_rounding_minus_inf
|
||||
cmp r8, 4
|
||||
je {{L}}q_shr_rounding_plus_inf
|
||||
cmp r8, 5
|
||||
je {{L}}q_shr_rounding_even
|
||||
cmp r8, 6
|
||||
je {{L}}q_shr_rounding_odd
|
||||
|
||||
jmp {{L}}unsupported
|
||||
|
||||
{{L}}q_shr_rounding_zero:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsubd ymm14, ymm14, ymm9
|
||||
vpaddd ymm14, ymm14, ymm11
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_away:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpaddd ymm14, ymm14, ymm11
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_minus_inf:
|
||||
{% for i in range(0, 8) %}
|
||||
vpsubd ymm{{i}}, ymm{{i}}, ymm9
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm11
|
||||
vpsravd ymm{{i}}, ymm{{i}}, ymm10
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_plus_inf:
|
||||
{% for i in range(0, 8) %}
|
||||
vpaddd ymm{{i}}, ymm{{i}}, ymm11
|
||||
vpsravd ymm{{i}}, ymm{{i}}, ymm10
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_even:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsravd ymm13, ymm14, ymm10
|
||||
vpand ymm13, ymm13, ymm9
|
||||
vpsubd ymm13, ymm13, ymm9 // nudge = ((abs >>l shift) & 0x01) - 1
|
||||
vpaddd ymm14, ymm14, ymm13 // add nudge
|
||||
vpaddd ymm14, ymm14, ymm11 // add half
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_shr_rounding_odd:
|
||||
{% for i in range(0, 8) %}
|
||||
vpabsd ymm14, ymm{{i}}
|
||||
vpsravd ymm13, ymm14, ymm10
|
||||
vpand ymm13, ymm13, ymm9
|
||||
vpsubd ymm13, ymm12, ymm13 // nudge = - ((abs >>l shift) & 0x01)
|
||||
vpaddd ymm14, ymm14, ymm13 // add nudge
|
||||
vpaddd ymm14, ymm14, ymm11 // add half
|
||||
vpsravd ymm14, ymm14, ymm10
|
||||
vpsignd ymm{{i}}, ymm14, ymm{{i}}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rdx, [rdi + 24] // col stride
|
||||
mov rcx, [rdi + 32] // item size
|
||||
|
||||
cmp rcx, 4
|
||||
je {{L}}store_strides_i32
|
||||
|
||||
{% for col in range(0, 8) %}
|
||||
mov r10, r8
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov byte ptr [r10], bl
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm{{col}}, ymm{{col}}, ymm{{col}}, 1
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov byte ptr [r10], bl
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
add r8, rdx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_strides_i32:
|
||||
{% for col in range(0, 8) %}
|
||||
mov r10, r8
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov dword ptr [r10], ebx
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm{{col}}, ymm{{col}}, ymm{{col}}, 1
|
||||
{% for row in range(0, 4) %}
|
||||
extractps ebx, xmm{{col}}, {{row}}
|
||||
mov dword ptr [r10], ebx
|
||||
add r10, rsi
|
||||
{% endfor %}
|
||||
add r8, rdx
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}return:
|
||||
ldmxcsr [rsp + 4]
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
|
||||
{{L}}one_32bit:
|
||||
{% if msvc %}
|
||||
dd 1
|
||||
{% else %}
|
||||
.int 1
|
||||
{% endif %}
|
||||
|
||||
{% if msvc %}
|
||||
avx512vnni_mmm_i32_8x8_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
+40
@@ -0,0 +1,40 @@
|
||||
// vim: set syntax=asm :
|
||||
|
||||
{{L}}non_linear:
|
||||
|
||||
{{L}}non_linear_loop_enter:
|
||||
sub rdi, 40
|
||||
{{L}}non_linear_loop:
|
||||
add rdi, 40
|
||||
mov rax, [rdi]
|
||||
|
||||
mov r8, {{ jump_table | length }}
|
||||
cmp rax, 0
|
||||
cmovl rax, r8
|
||||
cmp rax, {{ jump_table | length }}
|
||||
cmovg rax, r8
|
||||
|
||||
{% if msvc %}
|
||||
lea r8, [ offset {{L}}jmp_table ]
|
||||
{% else %}
|
||||
lea r8, [ rip + {{L}}jmp_table ]
|
||||
{% endif %}
|
||||
movsxd r9, dword ptr [ r8 + rax * 4 ]
|
||||
lea r8, [ r8 + r9 ]
|
||||
jmp r8
|
||||
|
||||
{{L}}jmp_table:
|
||||
{% for j in jump_table %}
|
||||
{{long}} {{L}}{{j}}-{{L}}jmp_table
|
||||
{% endfor %}
|
||||
{{long}} {{L}}unsupported-{{L}}jmp_table
|
||||
|
||||
{{L}}unsupported:
|
||||
mov rax, 1
|
||||
jmp {{L}}return
|
||||
|
||||
|
||||
{{L}}done:
|
||||
mov rax, 0
|
||||
jmp {{L}}return
|
||||
|
||||
Vendored
+143
@@ -0,0 +1,143 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
/* mmm 16 x 5:
|
||||
|
||||
ymm0 ymm2 ymm4 ymm6 ymm8
|
||||
ymm1 ymm3 ymm5 ymm7 ymm9
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "16x5" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "2x5/packed_packed_loop1/avx.S.raw" %}
|
||||
|
||||
add rcx, 20
|
||||
add rax, 64
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
// NON LINEAR / ADDC
|
||||
|
||||
{% set from = 0 %}{% set to = 9 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 9 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 9 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 9 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
lea r8, [ r10 + rsi * 8 ]
|
||||
|
||||
{% for i in range(0, 5) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r10 + ymm14 ], ymm15
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm13, [ r8 + ymm14 ], ymm15
|
||||
add r10, rbx
|
||||
add r8, rbx
|
||||
vaddps ymm{{ i * 2 }}, ymm{{ i * 2 }}, ymm12
|
||||
vaddps ymm{{ i * 2 + 1 }}, ymm{{ i * 2 + 1 }}, ymm13
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups ymm12, [rax]
|
||||
vmovups ymm13, [rax + 32]
|
||||
|
||||
{% for i in range(0, 5) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps ymm{{ i * 2 }}, ymm12, ymm14
|
||||
vfmadd231ps ymm{{ i * 2 + 1 }}, ymm13, ymm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r12, [ r8 + 4 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
cmp rbx, 64
|
||||
jne {{L}}store_strides_generic
|
||||
|
||||
{% for row in range(0, 2) %}
|
||||
{% for col in range(0, 5) %}
|
||||
vmovups ymmword ptr [r{{ col + 8 }}], ymm{{ col * 2 + row }}
|
||||
add r{{ col + 8 }}, 32
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_strides_generic:
|
||||
// tops of cols
|
||||
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% if quarter != 0 %}
|
||||
// move next four rows at top (xmm0,2,..10)
|
||||
vperm2f128 ymm0, ymm0, ymm1, {{quarter}}
|
||||
vperm2f128 ymm2, ymm2, ymm3, {{quarter}}
|
||||
vperm2f128 ymm4, ymm4, ymm5, {{quarter}}
|
||||
vperm2f128 ymm6, ymm6, ymm7, {{quarter}}
|
||||
vperm2f128 ymm8, ymm8, ymm9, {{quarter}}
|
||||
{% endif %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 5) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i * 2 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set type = "f32" %}{% set size = "16x5" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+131
@@ -0,0 +1,131 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 16 x 6:
|
||||
|
||||
ymm0 ymm2 ymm4 ymm6 ymm8 ymm10
|
||||
ymm1 ymm3 ymm5 ymm7 ymm9 ymm11
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "16x6" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "2x6/packed_packed_loop1/original.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
// NON LINEAR / ADDC
|
||||
|
||||
{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 16 %}{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
lea r8, [ r10 + rsi * 8 ]
|
||||
|
||||
{% for i in range(0, 6) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r10 + ymm14 ], ymm15
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm13, [ r8 + ymm14 ], ymm15
|
||||
add r10, rbx
|
||||
add r8, rbx
|
||||
vaddps ymm{{ i * 2 }}, ymm{{ i * 2 }}, ymm12
|
||||
vaddps ymm{{ i * 2 + 1 }}, ymm{{ i * 2 + 1 }}, ymm13
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups ymm12, [rax]
|
||||
vmovups ymm13, [rax + 32]
|
||||
|
||||
{% for i in range(0, 6) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps ymm{{ i * 2 }}, ymm12, ymm14
|
||||
vfmadd231ps ymm{{ i * 2 + 1 }}, ymm13, ymm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r12, [ r8 + 4 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
lea r13, [ r12 + rbx ]
|
||||
|
||||
{% for quarter in range(0, 4) %}
|
||||
{% if quarter != 0 %}
|
||||
// move next four rows at top (xmm0,2,..10)
|
||||
vperm2f128 ymm0, ymm0, ymm1, {{quarter}}
|
||||
vperm2f128 ymm2, ymm2, ymm3, {{quarter}}
|
||||
vperm2f128 ymm4, ymm4, ymm5, {{quarter}}
|
||||
vperm2f128 ymm6, ymm6, ymm7, {{quarter}}
|
||||
vperm2f128 ymm8, ymm8, ymm9, {{quarter}}
|
||||
vperm2f128 ymm10, ymm10, ymm11, {{quarter}}
|
||||
{% endif %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 6) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{ i * 2 }}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set type = "f32" %}{% set size = "16x6" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+158
@@ -0,0 +1,158 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
/* mmm 24 x 4:
|
||||
|
||||
ymm0 ymm3 ymm6 ymm10
|
||||
ymm1 ymm4 ymm7 ymm11
|
||||
ymm2 ymm5 ymm8 ymm12
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "24x4" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "3x4/packed_packed_loop1/avx.S.raw" %}
|
||||
|
||||
add rcx, 16
|
||||
add rax, 96
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
// NON LINEAR / ADDC
|
||||
|
||||
{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 24 %}{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 24 %}{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}unicast_generic
|
||||
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r9 + rbx]
|
||||
lea r11, [ r10 + rbx ]
|
||||
lea r12, [ r11 + rbx ]
|
||||
|
||||
{% for col in range(0, 4) %}
|
||||
{% for row in range(0, 3) %}
|
||||
vmovups ymm12, [ r{{ col + 8 }} ]
|
||||
add r{{ col + 8 }}, 32
|
||||
vaddps ymm{{ col * 3 + row }}, ymm{{ col * 3 + row }}, ymm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}unicast_generic:
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
// mov r12, [0]
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
lea r9, [ r8 + rsi * 8 ]
|
||||
lea r10, [ r9 + rsi * 8 ]
|
||||
|
||||
{% for col in range(0, 4) %}
|
||||
{% for row in range(0, 3) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r{{ row + 8 }} + ymm14 ], ymm15
|
||||
add r{{ row + 8 }}, rbx
|
||||
vaddps ymm{{ col * 3 + row }}, ymm{{ col * 3 + row }}, ymm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups ymm12, [rax]
|
||||
vmovups ymm13, [rax + 32]
|
||||
vmovups ymm15, [rax + 64]
|
||||
{% for i in range(0, 4) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps ymm{{ i * 3 }}, ymm12, ymm14
|
||||
vfmadd231ps ymm{{ i * 3 + 1 }}, ymm13, ymm14
|
||||
vfmadd231ps ymm{{ i * 3 + 2 }}, ymm15, ymm14
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}store_strides_generic
|
||||
|
||||
{% for col in range(0, 4) %}
|
||||
{% for row in range(0, 3) %}
|
||||
vmovups ymmword ptr [r{{ col + 8 }}], ymm{{ col * 3 + row }}
|
||||
add r{{ col + 8 }}, 32
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_strides_generic:
|
||||
{% for col in range(0, 4) %}
|
||||
{% for row in range(0, 3) %}
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ col + 8 }}], xmm{{ col * 3 + row }}, {{i}}
|
||||
add r{{ col + 8 }}, rsi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm{{ col * 3 + row }}, ymm{{ col * 3 + row }}, ymm{{ col * 3 + row }}, 1
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ col + 8 }}], xmm{{ col * 3 + row }}, {{i}}
|
||||
add r{{ col + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set type = "f32" %}{% set size = "24x4" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+424
@@ -0,0 +1,424 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 64 x 1
|
||||
|
||||
ymm0
|
||||
ymm1
|
||||
ymm2
|
||||
ymm3
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "32x1" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
mov r8, [rdi + 32] // packing
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
cmp r8, 1
|
||||
jz {{L}}q40f32
|
||||
|
||||
cmp r8, 2
|
||||
jz {{L}}q40f16
|
||||
|
||||
cmp r8, 3
|
||||
jz {{L}}f16f16
|
||||
|
||||
cmp r8, 4
|
||||
jz {{L}}f16f32
|
||||
|
||||
cmp r8, 5
|
||||
jz {{L}}f32f16
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed:
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovaps ymm8, [rax]
|
||||
vmovaps ymm9, [rax + 32]
|
||||
vmovaps ymm10, [rax + 64]
|
||||
vmovaps ymm11, [rax + 96]
|
||||
|
||||
vfmadd231ps ymm0, ymm15, ymm8
|
||||
vfmadd231ps ymm1, ymm15, ymm9
|
||||
vfmadd231ps ymm2, ymm15, ymm10
|
||||
vfmadd231ps ymm3, ymm15, ymm11
|
||||
|
||||
add rcx, 4
|
||||
add rax, 128
|
||||
sub rbx, 1
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% if msvc %}
|
||||
{{L}}q40f32_mask:
|
||||
{{long}} 0F0F0F0Fh
|
||||
{{L}}q40f32_eight:
|
||||
{{long}} 08h
|
||||
{% else %}
|
||||
{{L}}q40f32_mask:
|
||||
{{long}} 0x0F0F0F0F
|
||||
{{L}}q40f32_eight:
|
||||
{{long}} 8
|
||||
{% endif %}
|
||||
|
||||
{{L}}q40f32:
|
||||
// ymm0-3: acc
|
||||
// ymm4-7: scales
|
||||
// ymm13: 8
|
||||
// ymm14: mask
|
||||
// ymm15: b value
|
||||
vbroadcastss ymm14, dword ptr [{{offset}} {{L}}q40f32_mask]
|
||||
vbroadcastss ymm13, dword ptr [{{offset}} {{L}}q40f32_eight]
|
||||
|
||||
{{L}}q40f32_outerloop:
|
||||
// scales
|
||||
vmovaps xmm4, [rax]
|
||||
vmovaps xmm5, [rax + 16]
|
||||
vmovaps xmm6, [rax + 32]
|
||||
vmovaps xmm7, [rax + 48]
|
||||
vcvtph2ps ymm4, xmm4
|
||||
vcvtph2ps ymm5, xmm5
|
||||
vcvtph2ps ymm6, xmm6
|
||||
vcvtph2ps ymm7, xmm7
|
||||
add rax, 64
|
||||
|
||||
mov rdx, 32
|
||||
|
||||
{{L}}q40f32_innerloop:
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
vmovaps xmm8, [rax] // 32 nibbles
|
||||
|
||||
vpand xmm10, xmm8, xmm14 // 16 bytes
|
||||
|
||||
vpmovzxbd ymm9, xmm10 // 8 u32
|
||||
|
||||
vpermilpd xmm10, xmm10, 1 // swap 64bit halves
|
||||
vpmovzxbd ymm10, xmm10 // 8 u32
|
||||
|
||||
vpsrlw xmm8, xmm8, 4
|
||||
vpand xmm12, xmm8, xmm14 // 16 bytes
|
||||
vpmovzxbd ymm11, xmm12 // 8 u32
|
||||
vpermilpd xmm12, xmm12, 1 // swap 64bit halves
|
||||
vpmovzxbd ymm12, xmm12 // 8 u32
|
||||
|
||||
vpsubd ymm9, ymm9, ymm13
|
||||
vpsubd ymm10, ymm10, ymm13
|
||||
vpsubd ymm11, ymm11, ymm13
|
||||
vpsubd ymm12, ymm12, ymm13
|
||||
|
||||
vcvtdq2ps ymm9, ymm9
|
||||
vcvtdq2ps ymm10, ymm10
|
||||
vcvtdq2ps ymm11, ymm11
|
||||
vcvtdq2ps ymm12, ymm12
|
||||
|
||||
vmulps ymm9, ymm9, ymm4
|
||||
vmulps ymm10, ymm10, ymm5
|
||||
vmulps ymm11, ymm11, ymm6
|
||||
vmulps ymm12, ymm12, ymm7
|
||||
|
||||
vfmadd231ps ymm0, ymm15, ymm9
|
||||
vfmadd231ps ymm1, ymm15, ymm10
|
||||
vfmadd231ps ymm2, ymm15, ymm11
|
||||
vfmadd231ps ymm3, ymm15, ymm12
|
||||
|
||||
add rax, 16
|
||||
add rcx, 4
|
||||
sub rdx, 1
|
||||
jnz {{L}}q40f32_innerloop
|
||||
|
||||
sub rbx, 32
|
||||
jnz {{L}}q40f32_outerloop
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q40f16:
|
||||
// ymm0-3: acc
|
||||
// ymm4-7: scales
|
||||
// ymm13: 8
|
||||
// ymm14: mask
|
||||
// ymm15: b value
|
||||
vbroadcastss ymm14, dword ptr [{{offset}} {{L}}q40f32_mask]
|
||||
vbroadcastss ymm13, dword ptr [{{offset}} {{L}}q40f32_eight]
|
||||
|
||||
{{L}}q40f16_outerloop:
|
||||
// scales
|
||||
vmovaps xmm4, [rax]
|
||||
vmovaps xmm5, [rax + 16]
|
||||
vmovaps xmm6, [rax + 32]
|
||||
vmovaps xmm7, [rax + 48]
|
||||
vcvtph2ps ymm4, xmm4
|
||||
vcvtph2ps ymm5, xmm5
|
||||
vcvtph2ps ymm6, xmm6
|
||||
vcvtph2ps ymm7, xmm7
|
||||
add rax, 64
|
||||
|
||||
mov rdx, 32
|
||||
|
||||
{{L}}q40f16_innerloop:
|
||||
vpbroadcastw ymm15, word ptr [rcx]
|
||||
vcvtph2ps ymm15, xmm15
|
||||
|
||||
vmovaps xmm8, [rax] // 32 nibbles
|
||||
|
||||
vpand xmm10, xmm8, xmm14 // 16 bytes
|
||||
|
||||
vpmovzxbd ymm9, xmm10 // 8 u32
|
||||
|
||||
vpermilpd xmm10, xmm10, 1 // swap 64bit halves
|
||||
vpmovzxbd ymm10, xmm10 // 8 u32
|
||||
|
||||
vpsrlw xmm8, xmm8, 4
|
||||
vpand xmm12, xmm8, xmm14 // 16 bytes
|
||||
vpmovzxbd ymm11, xmm12 // 8 u32
|
||||
vpermilpd xmm12, xmm12, 1 // swap 64bit halves
|
||||
vpmovzxbd ymm12, xmm12 // 8 u32
|
||||
|
||||
vpsubd ymm9, ymm9, ymm13
|
||||
vpsubd ymm10, ymm10, ymm13
|
||||
vpsubd ymm11, ymm11, ymm13
|
||||
vpsubd ymm12, ymm12, ymm13
|
||||
|
||||
vcvtdq2ps ymm9, ymm9
|
||||
vcvtdq2ps ymm10, ymm10
|
||||
vcvtdq2ps ymm11, ymm11
|
||||
vcvtdq2ps ymm12, ymm12
|
||||
|
||||
vmulps ymm9, ymm9, ymm4
|
||||
vmulps ymm10, ymm10, ymm5
|
||||
vmulps ymm11, ymm11, ymm6
|
||||
vmulps ymm12, ymm12, ymm7
|
||||
|
||||
vfmadd231ps ymm0, ymm15, ymm9
|
||||
vfmadd231ps ymm1, ymm15, ymm10
|
||||
vfmadd231ps ymm2, ymm15, ymm11
|
||||
vfmadd231ps ymm3, ymm15, ymm12
|
||||
|
||||
add rax, 16
|
||||
add rcx, 2
|
||||
sub rdx, 1
|
||||
jnz {{L}}q40f16_innerloop
|
||||
|
||||
sub rbx, 32
|
||||
jnz {{L}}q40f16_outerloop
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}f16f16:
|
||||
{{align}} 16
|
||||
vpbroadcastw ymm15, word ptr [rcx]
|
||||
|
||||
vmovaps xmm4, [rax]
|
||||
vmovaps xmm5, [rax + 16]
|
||||
vmovaps xmm6, [rax + 32]
|
||||
vmovaps xmm7, [rax + 48]
|
||||
|
||||
vcvtph2ps ymm15, xmm15
|
||||
vcvtph2ps ymm4, xmm4
|
||||
vcvtph2ps ymm5, xmm5
|
||||
vcvtph2ps ymm6, xmm6
|
||||
vcvtph2ps ymm7, xmm7
|
||||
|
||||
vfmadd231ps ymm0, ymm15, ymm4
|
||||
vfmadd231ps ymm1, ymm15, ymm5
|
||||
vfmadd231ps ymm2, ymm15, ymm6
|
||||
vfmadd231ps ymm3, ymm15, ymm7
|
||||
|
||||
add rcx, 2
|
||||
add rax, 64
|
||||
sub rbx, 1
|
||||
jnz {{L}}f16f16
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}f32f16:
|
||||
{{align}} 16
|
||||
vpbroadcastw ymm15, word ptr [rcx]
|
||||
|
||||
vmovaps ymm4, [rax]
|
||||
vmovaps ymm5, [rax + 32]
|
||||
vmovaps ymm6, [rax + 64]
|
||||
vmovaps ymm7, [rax + 96]
|
||||
|
||||
vcvtph2ps ymm15, xmm15
|
||||
|
||||
vfmadd231ps ymm0, ymm15, ymm4
|
||||
vfmadd231ps ymm1, ymm15, ymm5
|
||||
vfmadd231ps ymm2, ymm15, ymm6
|
||||
vfmadd231ps ymm3, ymm15, ymm7
|
||||
|
||||
add rcx, 2
|
||||
add rax, 128
|
||||
sub rbx, 1
|
||||
jnz {{L}}f32f16
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}f16f32:
|
||||
{{align}} 16
|
||||
vbroadcastss ymm15, dword ptr [rcx]
|
||||
|
||||
vmovaps xmm4, [rax]
|
||||
vmovaps xmm5, [rax + 16]
|
||||
vmovaps xmm6, [rax + 32]
|
||||
vmovaps xmm7, [rax + 48]
|
||||
|
||||
vcvtph2ps ymm4, xmm4
|
||||
vcvtph2ps ymm5, xmm5
|
||||
vcvtph2ps ymm6, xmm6
|
||||
vcvtph2ps ymm7, xmm7
|
||||
|
||||
vfmadd231ps ymm0, ymm15, ymm4
|
||||
vfmadd231ps ymm1, ymm15, ymm5
|
||||
vfmadd231ps ymm2, ymm15, ymm6
|
||||
vfmadd231ps ymm3, ymm15, ymm7
|
||||
|
||||
add rcx, 4
|
||||
add rax, 64
|
||||
sub rbx, 1
|
||||
jnz {{L}}f16f32
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
|
||||
{% set from = 0 %}{% set to = 3 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 3 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 3 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 3 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}add_unicast_generic
|
||||
|
||||
{% for row in range(0, 4) %}
|
||||
vaddps ymm{{row}}, ymm{{row}}, [ r10 + {{ row * 32 }} ]
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_unicast_generic:
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r10 + ymm14 ], ymm15
|
||||
|
||||
vaddps ymm{{i}}, ymm{{i}}, ymm12
|
||||
lea r10, [ r10 + rsi * 8 ]
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rbx]
|
||||
|
||||
{% for i in range(0, 4) %}
|
||||
vmovups ymm12, [rax + {{ i * 32 }}]
|
||||
vfmadd231ps ymm{{i}}, ymm12, ymm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov r11, [rdi + 32] // item size
|
||||
|
||||
cmp r11, 2
|
||||
je {{L}}store_f16
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}store_generic
|
||||
|
||||
{% for row in range(0, 4) %}
|
||||
vmovups [r8 + {{ row * 32 }}], ymm{{row}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_generic:
|
||||
|
||||
{% for vec in range(0, 4) %}
|
||||
{% for half in range(0, 2) %}
|
||||
{% if half == 0 %}
|
||||
movaps xmm9, xmm{{vec}}
|
||||
{% else %}
|
||||
vperm2f128 ymm9, ymm{{vec}}, ymm{{vec}}, 1
|
||||
{% endif %}
|
||||
{% for row in range(0, 4) %}
|
||||
vextractps dword ptr [r8], xmm9, {{row}}
|
||||
add r8, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_f16:
|
||||
|
||||
vcvtps2ph xmm0, ymm0, 0
|
||||
vcvtps2ph xmm1, ymm1, 0
|
||||
vcvtps2ph xmm2, ymm2, 0
|
||||
vcvtps2ph xmm3, ymm3, 0
|
||||
|
||||
cmp rsi, 2
|
||||
jne {{L}}store_generic_f16
|
||||
|
||||
{% for row in range(0, 4) %}
|
||||
vmovups [r8 + {{ row * 16 }}], xmm{{row}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_generic_f16:
|
||||
|
||||
{% for vec in range(0, 4) %}
|
||||
{% for row in range(0, 8) %}
|
||||
pextrw word ptr [r8], xmm{{vec}}, {{row}}
|
||||
add r8, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set type = "f32" %}{% set size = "32x1" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+336
@@ -0,0 +1,336 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
/* mmm 16 x 5:
|
||||
|
||||
ymm0 ymm4 ymm8
|
||||
ymm1 ymm5 ymm9
|
||||
ymm2 ymm6 ymm10
|
||||
ymm3 ymm7 ymm11
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "32x3" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rbx, [rdi + 8] // k
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov r8, [rdi + 32] // packing
|
||||
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
cmp r8, 1
|
||||
jz {{L}}main_loop_packed_packed_f32_f16
|
||||
|
||||
cmp r8, 2
|
||||
jz {{L}}main_loop_packed_packed_f16_f32
|
||||
|
||||
cmp r8, 3
|
||||
jz {{L}}main_loop_packed_packed_f16_f16
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "4x3/packed_packed_loop1/avx.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed_f32_f16:
|
||||
// Load col of A
|
||||
vmovaps ymm12, [rax]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vpbroadcastw xmm13, word ptr [rcx + 0]
|
||||
vpbroadcastw xmm14, word ptr [rcx + 2]
|
||||
vpbroadcastw xmm15, word ptr [rcx + 4]
|
||||
|
||||
vcvtph2ps ymm13, xmm13
|
||||
vcvtph2ps ymm14, xmm14
|
||||
vcvtph2ps ymm15, xmm15
|
||||
|
||||
// N.B. Stepping cols in inner loop
|
||||
vfmadd231ps ymm0, ymm12, ymm13
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+32]
|
||||
|
||||
vfmadd231ps ymm1, ymm12, ymm13
|
||||
vfmadd231ps ymm5, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+64]
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm13
|
||||
vfmadd231ps ymm6, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vmovaps ymm12, [rax+96]
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm13
|
||||
vfmadd231ps ymm7, ymm12, ymm14
|
||||
vfmadd231ps ymm11, ymm12, ymm15
|
||||
|
||||
add rcx, 6
|
||||
add rax, 128
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed_f32_f16
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed_f16_f32:
|
||||
// Load col of A
|
||||
vmovaps xmm12, [rax]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vbroadcastss ymm13, dword ptr [rcx + 0]
|
||||
vbroadcastss ymm14, dword ptr [rcx + 4]
|
||||
vbroadcastss ymm15, dword ptr [rcx + 8]
|
||||
|
||||
vcvtph2ps ymm12, xmm12
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm13
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps xmm12, [rax+16]
|
||||
vcvtph2ps ymm12, xmm12
|
||||
|
||||
vfmadd231ps ymm1, ymm12, ymm13
|
||||
vfmadd231ps ymm5, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
|
||||
vmovaps xmm12, [rax+32]
|
||||
vcvtph2ps ymm12, xmm12
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm13
|
||||
vfmadd231ps ymm6, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vmovaps xmm12, [rax+48]
|
||||
vcvtph2ps ymm12, xmm12
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm13
|
||||
vfmadd231ps ymm7, ymm12, ymm14
|
||||
vfmadd231ps ymm11, ymm12, ymm15
|
||||
|
||||
add rcx, 12
|
||||
add rax, 64
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed_f16_f32
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed_f16_f16:
|
||||
// Load col of A
|
||||
vmovaps xmm12, [rax]
|
||||
|
||||
// Fill 3 cols of B
|
||||
vpbroadcastw xmm13, word ptr [rcx + 0]
|
||||
vpbroadcastw xmm14, word ptr [rcx + 2]
|
||||
vpbroadcastw xmm15, word ptr [rcx + 4]
|
||||
|
||||
vcvtph2ps ymm12, xmm12
|
||||
vcvtph2ps ymm13, xmm13
|
||||
vcvtph2ps ymm14, xmm14
|
||||
vcvtph2ps ymm15, xmm15
|
||||
|
||||
vfmadd231ps ymm0, ymm12, ymm13
|
||||
vfmadd231ps ymm4, ymm12, ymm14
|
||||
vfmadd231ps ymm8, ymm12, ymm15
|
||||
|
||||
vmovaps xmm12, [rax+16]
|
||||
vcvtph2ps ymm12, xmm12
|
||||
|
||||
vfmadd231ps ymm1, ymm12, ymm13
|
||||
vfmadd231ps ymm5, ymm12, ymm14
|
||||
vfmadd231ps ymm9, ymm12, ymm15
|
||||
|
||||
vmovaps xmm12, [rax+32]
|
||||
vcvtph2ps ymm12, xmm12
|
||||
|
||||
vfmadd231ps ymm2, ymm12, ymm13
|
||||
vfmadd231ps ymm6, ymm12, ymm14
|
||||
vfmadd231ps ymm10, ymm12, ymm15
|
||||
|
||||
vmovaps xmm12, [rax+48]
|
||||
vcvtph2ps ymm12, xmm12
|
||||
|
||||
vfmadd231ps ymm3, ymm12, ymm13
|
||||
vfmadd231ps ymm7, ymm12, ymm14
|
||||
vfmadd231ps ymm11, ymm12, ymm15
|
||||
|
||||
add rcx, 6
|
||||
add rax, 64
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed_f16_f16
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
// NON LINEAR / ADDC
|
||||
|
||||
{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 32 %}{% set from = 0 %}{% set to = 11 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 11 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}unicast_generic
|
||||
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r9 + rbx]
|
||||
lea r11, [ r10 + rbx ]
|
||||
|
||||
{% for col in range(0, 3) %}
|
||||
{% for row in range(0, 4) %}
|
||||
vmovups ymm12, [ r{{ col + 8 }} ]
|
||||
add r{{ col + 8 }}, 32
|
||||
vaddps ymm{{ col * 4 + row }}, ymm{{ col * 4 + row }}, ymm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}unicast_generic:
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
// mov r12, [0]
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
lea r9, [ r8 + rsi * 8 ]
|
||||
lea r10, [ r9 + rsi * 8 ]
|
||||
lea r11, [ r10 + rsi * 8 ]
|
||||
|
||||
{% for col in range(0, 3) %}
|
||||
{% for row in range(0, 4) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r{{ row + 8 }} + ymm14 ], ymm15
|
||||
add r{{ row + 8 }}, rbx
|
||||
vaddps ymm{{ col * 4 + row }}, ymm{{ col * 4 + row }}, ymm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss ymm13, dword ptr [rbx]
|
||||
vbroadcastss ymm14, dword ptr [rbx + 4]
|
||||
vbroadcastss ymm15, dword ptr [rbx + 8]
|
||||
{% for i in range(0, 4) %}
|
||||
vmovups ymm12, [rax + {{ i * 32 }}]
|
||||
vfmadd231ps ymm{{ 0 + i }}, ymm12, ymm13
|
||||
vfmadd231ps ymm{{ 4 + i }}, ymm12, ymm14
|
||||
vfmadd231ps ymm{{ 8 + i }}, ymm12, ymm15
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
mov r11, [rdi + 32] // item size
|
||||
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
|
||||
cmp r11, 2
|
||||
je {{L}}store_f16
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}store_strides_generic
|
||||
|
||||
{% for col in range(0, 3) %}
|
||||
{% for row in range(0, 4) %}
|
||||
vmovups ymmword ptr [r{{ col + 8 }}], ymm{{ col * 4 + row }}
|
||||
add r{{ col + 8 }}, 32
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_strides_generic:
|
||||
|
||||
{% for col in range(0, 3) %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ col + 8 }}], xmm{{ col * 4 + row }}, {{i}}
|
||||
add r{{ col + 8 }}, rsi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm{{ col * 4 + row }}, ymm{{ col * 4 + row }}, ymm{{ col * 4 + row }}, 1
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ col + 8 }}], xmm{{ col * 4 + row }}, {{i}}
|
||||
add r{{ col + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_f16:
|
||||
|
||||
{% for reg in range(0, 12) %}
|
||||
vcvtps2ph xmm{{reg}}, ymm{{reg}}, 0
|
||||
{% endfor %}
|
||||
|
||||
cmp rsi, 2
|
||||
jne {{L}}store_generic_f16
|
||||
|
||||
{% for col in range(0, 3) %}
|
||||
{% for row in range(0, 4) %}
|
||||
vmovups [r{{ col + 8 }} + {{ row * 16 }}], xmm{{ col * 4 + row }}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_generic_f16:
|
||||
{% for col in range(0, 3) %}
|
||||
{% for vec in range(0, 4) %}
|
||||
{% for row in range(0, 8) %}
|
||||
pextrw word ptr [r{{ col + 8 }}], xmm{{ col * 4 + vec }}, {{row}}
|
||||
add r{{ col + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set type = "f32" %}{% set size = "32x3" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+158
@@ -0,0 +1,158 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
/* mmm 40 x 5:
|
||||
|
||||
ymm0 ymm5
|
||||
ymm1 ymm6
|
||||
ymm2 ymm7
|
||||
ymm3 ymm8
|
||||
ymm4 ymm9
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "40x2" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "5x2/packed_packed_loop1/avx.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
// NON LINEAR / ADDC
|
||||
|
||||
{% set from = 0 %}{% set to = 9 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 40 %}{% set from = 0 %}{% set to = 9 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 40 %}{% set from = 0 %}{% set to = 9 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 9 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}unicast_generic
|
||||
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r9 + rbx]
|
||||
lea r11, [ r10 + rbx ]
|
||||
lea r12, [ r11 + rbx ]
|
||||
|
||||
|
||||
{% for col in range(0, 2) %}
|
||||
{% for row in range(0, 5) %}
|
||||
vmovups ymm12, [ r{{ col + 8 }} ]
|
||||
add r{{ col + 8 }}, 32
|
||||
vaddps ymm{{ col * 5 + row }}, ymm{{ col * 5 + row }}, ymm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}unicast_generic:
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
lea r9, [ r8 + rsi * 8]
|
||||
lea r10, [ r9 + rsi * 8]
|
||||
lea r11, [ r10 + rsi * 8]
|
||||
lea r12, [ r11 + rsi * 8]
|
||||
|
||||
{% for col in range(0, 2) %}
|
||||
{% for row in range(0, 5) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r{{ row + 8 }} + ymm14 ], ymm15
|
||||
add r{{ row + 8 }}, rbx
|
||||
vaddps ymm{{ col * 5 + row }}, ymm{{ col * 5 + row }}, ymm12
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss ymm10, dword ptr [rbx]
|
||||
vbroadcastss ymm11, dword ptr [rbx + 4]
|
||||
{% for i in range(0, 5) %}
|
||||
vmovups ymm12, [rax + {{ i * 32 }}]
|
||||
vfmadd231ps ymm{{ 0 + i }}, ymm12, ymm10
|
||||
vfmadd231ps ymm{{ 5 + i }}, ymm12, ymm11
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
lea r12, [ r10 + 2 * rbx ]
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}store_strides_generic
|
||||
|
||||
{% for col in range(0, 2) %}
|
||||
{% for row in range(0, 5) %}
|
||||
vmovups ymmword ptr [r{{ col + 8 }}], ymm{{ col * 5 + row }}
|
||||
add r{{ col + 8 }}, 32
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_strides_generic:
|
||||
{% for col in range(0, 2) %}
|
||||
{% for row in range(0, 5) %}
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ col + 8 }}], xmm{{ col * 5 + row }}, {{i}}
|
||||
add r{{ col + 8 }}, rsi
|
||||
{% endfor %}
|
||||
vperm2f128 ymm{{ col * 5 + row }}, ymm{{ col * 5 + row }}, ymm{{ col * 5 + row }}, 1
|
||||
{% for i in range(0, 4) %}
|
||||
vextractps dword ptr [r{{ col + 8 }}], xmm{{ col * 5 + row }}, {{i}}
|
||||
add r{{ col + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set type = "f32" %}{% set size = "40x2" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+142
@@ -0,0 +1,142 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 64 x 1
|
||||
|
||||
ymm0
|
||||
ymm1
|
||||
...
|
||||
ymm8
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "64x1" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rcx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rbx, [rdi + 8] // k
|
||||
test rbx, rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
test rbx, 1
|
||||
jz {{L}}main_loop_packed_packed
|
||||
{% include "8x1/packed_packed_loop1/avx.S.raw" %}
|
||||
|
||||
dec rbx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{align}} 16
|
||||
{{L}}main_loop_packed_packed:
|
||||
{% include "8x1/packed_packed_loop1/avx-unroll.S.raw" %}
|
||||
|
||||
sub rbx, 2
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{% set from = 0 %}{% set to = 7 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 64 %}{% set from = 0 %}{% set to = 7 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 64 %}{% set from = 0 %}{% set to = 7 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}add_unicast_generic
|
||||
|
||||
{% for row in range(0, 8) %}
|
||||
vaddps ymm{{row}}, ymm{{row}}, [ r10 + {{ row * 32 }} ]
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_unicast_generic:
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r10 + ymm14 ], ymm15
|
||||
|
||||
vaddps ymm{{i}}, ymm{{i}}, ymm12
|
||||
lea r10, [ r10 + rsi * 8 ]
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vbroadcastss ymm14, dword ptr [rbx]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vmovups ymm12, [rax + {{ i * 32 }}]
|
||||
vfmadd231ps ymm{{i}}, ymm12, ymm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
|
||||
cmp rsi, 4
|
||||
jne {{L}}store_generic
|
||||
|
||||
{% for row in range(0, 8) %}
|
||||
vmovups [r8 + {{ row * 32 }}], ymm{{row}}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store_generic:
|
||||
|
||||
{% for vec in range(0, 8) %}
|
||||
{% for half in range(0, 2) %}
|
||||
{% if half == 0 %}
|
||||
movaps xmm9, xmm{{vec}}
|
||||
{% else %}
|
||||
vperm2f128 ymm9, ymm{{vec}}, ymm{{vec}}, 1
|
||||
{% endif %}
|
||||
{% for row in range(0, 4) %}
|
||||
vextractps dword ptr [r8], xmm9, {{row}}
|
||||
add r8, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
|
||||
{% set type = "f32" %}{% set size = "64x1" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+129
@@ -0,0 +1,129 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
/* mmm 16 x 6:
|
||||
|
||||
ymm0 ymm2 ymm4 ymm6 ymm8 ymm10
|
||||
ymm1 ymm3 ymm5 ymm7 ymm9 ymm11
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
*/
|
||||
#}
|
||||
|
||||
{% set type = "f32" %}{% set size = "8x8" %}{% set suffix = suffix %}{% set G = G %}{% include "preamble.j2" %}
|
||||
|
||||
{{L}}clear:
|
||||
vzeroall
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_mat_mul:
|
||||
mov rbx, [rdi + 24] // B
|
||||
mov rax, [rdi + 16] // A
|
||||
|
||||
mov rcx, [rdi + 8] // k
|
||||
test rcx, rcx
|
||||
jz {{L}}non_linear_loop
|
||||
|
||||
{{L}}main_loop_packed_packed:
|
||||
vmovaps ymm12, [rax]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{i}} * 4]
|
||||
vfmadd231ps ymm{{i}}, ymm12, ymm14
|
||||
{% endfor %}
|
||||
|
||||
add rax, 32
|
||||
add rbx, 32
|
||||
dec rcx
|
||||
jnz {{L}}main_loop_packed_packed
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
// NON LINEAR / ADDC
|
||||
|
||||
{% set from = 0 %}{% set to = 7 %}{% set type = "f32" %}{% include "fma_mmm_f32_scalars.j2" %}
|
||||
{% set mr = 8 %}{% set from = 0 %}{% set to = 7 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_rows.j2" %}
|
||||
{% set mr = 8 %}{% set from = 0 %}{% set to = 7 %}{% set type = "f32" %}{% include "fma_mmm_f32_per_cols.j2" %}
|
||||
{% set from = 0 %}{% set to = 7 %}{% include "fma_mmm_load_tile.j2" %}
|
||||
|
||||
{{L}}add_unicast:
|
||||
|
||||
mov r10, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
mov eax, 0
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm14, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
{% for i in range(0, 4) %}
|
||||
pinsrd xmm15, eax, {{i}}
|
||||
add eax, esi
|
||||
{% endfor %}
|
||||
|
||||
vperm2f128 ymm14, ymm14, ymm15, 32 // ymm14 <- xmm14::xmm15
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vpcmpeqd ymm15, ymm15, ymm15
|
||||
vgatherdps ymm12, [ r10 + ymm14 ], ymm15
|
||||
add r10, rbx
|
||||
vaddps ymm{{i}}, ymm{{i}}, ymm12
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}add_row_col_products:
|
||||
mov rax, [ rdi + 8 ]
|
||||
mov rbx, [ rdi + 16 ]
|
||||
|
||||
vmovups ymm12, [rax]
|
||||
|
||||
{% for i in range(0, 8) %}
|
||||
vbroadcastss ymm14, dword ptr [rbx + {{ i * 4 }} ]
|
||||
vfmadd231ps ymm{{i}}, ymm12, ymm14
|
||||
{% endfor %}
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}store:
|
||||
mov r8, [rdi + 8] // c ptr
|
||||
mov rsi, [rdi + 16] // row stride
|
||||
mov rbx, [rdi + 24] // col stride
|
||||
|
||||
// tops of cols
|
||||
lea r9, [ r8 + rbx ]
|
||||
lea r10, [ r8 + 2 * rbx ]
|
||||
lea r12, [ r8 + 4 * rbx ]
|
||||
lea r11, [ r10 + rbx ]
|
||||
lea r13, [ r12 + rbx ]
|
||||
lea r14, [ r12 + 2 * rbx ]
|
||||
lea r15, [ r13 + 2 * rbx ]
|
||||
|
||||
{% for quarter in range(0, 2) %}
|
||||
{% if quarter != 0 %}
|
||||
// move next four rows at top (xmm0,2,..10)
|
||||
{% for r in range(0, 8) %}
|
||||
vperm2f128 ymm{{r}}, ymm{{r}}, ymm{{r}}, {{quarter}}
|
||||
{% endfor %}
|
||||
{% endif %}
|
||||
{% for row in range(0, 4) %}
|
||||
{% for i in range(0, 8) %}
|
||||
vextractps dword ptr [r{{ i + 8 }}], xmm{{i}}, {{row}}
|
||||
add r{{ i + 8 }}, rsi
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
|
||||
{% set type = "f32" %}{% set size = "8x8" %}{% set suffix = suffix %}{% set G = G %}{% set L = L %}{% include "postamble.j2" %}
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "fma_mmm_ymm_ops.j2" import per_col %}
|
||||
|
||||
{{ per_col("per_col_min", "vminps", mr, from, to, type=type) }}
|
||||
{{ per_col("per_col_max", "vmaxps", mr, from, to, type=type) }}
|
||||
{{ per_col("per_col_add", "vaddps", mr, from, to, type=type) }}
|
||||
{{ per_col("per_col_mul", "vmulps", mr, from, to, type=type) }}
|
||||
{{ per_col("per_col_sub", "vsubps", mr, from, to, type=type) }}
|
||||
{{ per_col("per_col_sub_flipped", "vsubps", mr, from, to, type=type, flipped=true) }}
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "fma_mmm_ymm_ops.j2" import per_row %}
|
||||
|
||||
{{ per_row("per_row_min", "vminps", mr, from, to, type=type) }}
|
||||
{{ per_row("per_row_max", "vmaxps", mr, from, to, type=type) }}
|
||||
{{ per_row("per_row_add", "vaddps", mr, from, to, type=type) }}
|
||||
{{ per_row("per_row_mul", "vmulps", mr, from, to, type=type) }}
|
||||
{{ per_row("per_row_sub", "vsubps", mr, from, to, type=type) }}
|
||||
{{ per_row("per_row_sub_flipped", "vsubps", mr, from, to, type=type, flipped=true) }}
|
||||
Vendored
+38
@@ -0,0 +1,38 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "fma_mmm_ymm_ops.j2" import scalar %}
|
||||
|
||||
{{ scalar("scalar_min", "vminps", from, to, type=type) }}
|
||||
{{ scalar("scalar_max", "vmaxps", from, to, type=type) }}
|
||||
{{ scalar("scalar_add", "vaddps", from, to, type=type) }}
|
||||
{{ scalar("scalar_mul", "vmulps", from, to, type=type) }}
|
||||
{{ scalar("scalar_sub", "vsubps", from, to, type=type) }}
|
||||
{{ scalar("scalar_sub_flipped", "vsubps", from, to, type=type, flipped=true) }}
|
||||
|
||||
{{L}}leaky_relu:
|
||||
// can only use ymm12 to ymm15
|
||||
// ymm15 <- alpha
|
||||
{% if type == "f32" %}
|
||||
vbroadcastss ymm15, dword ptr [rdi + 8]
|
||||
{% else %}
|
||||
pinsrw xmm15, word ptr [rdi + 8], 0
|
||||
vcvtph2ps ymm15, xmm15
|
||||
vbroadcastss ymm15, xmm15
|
||||
{% endif %}
|
||||
|
||||
// ymm14 <- all zero
|
||||
vpxor ymm14, ymm14, ymm14
|
||||
|
||||
{% for reg in range(from, to + 1) %}
|
||||
// ymm12 <- alpha * x
|
||||
vmulps ymm12, ymm{{reg}}, ymm15
|
||||
vcmpps ymm13, ymm14, ymm{{reg}}, 1 // 1 means LT
|
||||
vblendvps ymm{{reg}}, ymm12, ymm{{reg}}, ymm13
|
||||
{% endfor %}
|
||||
// select muled of orginal
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
|
||||
{{L}}q_scale:
|
||||
{{L}}q_shl:
|
||||
{{L}}q_shr:
|
||||
jmp {{L}}unsupported
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "fma_mmm_ymm_ops.j2" import per_col %}
|
||||
|
||||
{{ per_col("per_col_min", "vpminsd", mr, from, to, type="i32") }}
|
||||
{{ per_col("per_col_max", "vpmaxsd", mr, from, to, type="i32") }}
|
||||
{{ per_col("per_col_add", "vpaddd", mr, from, to, type="i32") }}
|
||||
{{ per_col("per_col_mul", "vpmulld", mr, from, to, type="i32") }}
|
||||
{{ per_col("per_col_sub", "vpsubd", mr, from, to, type="i32") }}
|
||||
{{ per_col("per_col_sub_flipped", "vpsubd", mr, from, to, type="i32", flipped=true) }}
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "fma_mmm_ymm_ops.j2" import per_row %}
|
||||
|
||||
{{ per_row("per_row_min", "vpminsd", mr, from, to, type="i32") }}
|
||||
{{ per_row("per_row_max", "vpmaxsd", mr, from, to, type="i32") }}
|
||||
{{ per_row("per_row_add", "vpaddd", mr, from, to, type="i32") }}
|
||||
{{ per_row("per_row_mul", "vpmulld", mr, from, to, type="i32") }}
|
||||
{{ per_row("per_row_sub", "vpsubd", mr, from, to, type="i32") }}
|
||||
{{ per_row("per_row_sub_flipped", "vpsubd", mr, from, to, type="i32", flipped=true) }}
|
||||
Vendored
+24
@@ -0,0 +1,24 @@
|
||||
// vim: set syntax=asm :
|
||||
{% from "fma_mmm_ymm_ops.j2" import scalar %}
|
||||
|
||||
{{ scalar("scalar_min", "vpminsd", from, to, type="i32") }}
|
||||
{{ scalar("scalar_max", "vpmaxsd", from, to, type="i32") }}
|
||||
{{ scalar("scalar_mul", "vpmulld", from, to, type="i32") }}
|
||||
{{ scalar("scalar_add", "vpaddd", from, to, type="i32") }}
|
||||
{{ scalar("scalar_sub", "vpsubd", from, to, type="i32") }}
|
||||
{{ scalar("scalar_sub_flipped", "vpsubd", from, to, type="i32", flipped=true) }}
|
||||
|
||||
{{L}}leaky_relu:
|
||||
// can only use ymm12 to ymm15
|
||||
// ymm15 <- alpha
|
||||
vbroadcastss ymm15, dword ptr [rdi + 8]
|
||||
// ymm14 <- all zero
|
||||
vpxor ymm14, ymm14, ymm14
|
||||
|
||||
{% for reg in range(from, to + 1) %}
|
||||
vpmulld ymm12, ymm{{reg}}, ymm15
|
||||
vpcmpgtd ymm13, ymm14, ymm{{reg}}
|
||||
vblendvps ymm{{reg}}, ymm{{reg}}, ymm12, ymm13
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
Vendored
+9
@@ -0,0 +1,9 @@
|
||||
// vim: set syntax=asm :
|
||||
|
||||
{{L}}load_tile:
|
||||
mov r8, [rdi + 8]
|
||||
{% for reg in range(from, to + 1) %}
|
||||
vmovups ymm{{reg}}, ymmword ptr [r8 + {{ (reg - from) * 32 }}]
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
Vendored
+92
@@ -0,0 +1,92 @@
|
||||
|
||||
{% macro scalar(label, op, from, to, type="f32", flipped=false) %}
|
||||
{{L}}{{label}}:
|
||||
{% if type == "f16" %}
|
||||
pinsrw xmm12, word ptr [rdi + 8], 0
|
||||
vcvtph2ps ymm12, xmm12
|
||||
vbroadcastss ymm12, xmm12
|
||||
{% else %}
|
||||
vbroadcastss ymm12, dword ptr [rdi + 8]
|
||||
{% endif %}
|
||||
|
||||
{% if flipped %}
|
||||
{% for reg in range(from, to + 1) %}
|
||||
{{op}} ymm{{reg}}, ymm{{reg}}, ymm12
|
||||
{% endfor %}
|
||||
{% else %}
|
||||
{% for reg in range(from, to + 1) %}
|
||||
{{op}} ymm{{reg}}, ymm12, ymm{{reg}}
|
||||
{% endfor %}
|
||||
{% endif %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
{% endmacro %}
|
||||
|
||||
{% macro per_row(label, op, mr, from, to, type="f32", flipped=false) %}
|
||||
{{L}}{{label}}:
|
||||
mov rax, [ rdi + 8 ]
|
||||
|
||||
{% set mr_over_8 = mr // 8 %}
|
||||
{% set mr_over_8_min_1 = mr // 8 - 1 %}
|
||||
|
||||
{% if type == "f16" %}
|
||||
{% for ix in range(0, mr_over_8_min_1 + 1) %}
|
||||
vmovups xmm{{ to + 1 + ix }}, [rax + {{ ix * 16 }}]
|
||||
{% endfor %}
|
||||
{% for ix in range(0, mr_over_8_min_1 + 1) %}
|
||||
vcvtph2ps ymm{{ to + 1 + ix }}, xmm{{ to + 1 + ix }}
|
||||
{% endfor %}
|
||||
{% else %}
|
||||
{% for ix in range(0, mr_over_8_min_1 + 1) %}
|
||||
vmovups ymm{{ to + 1 + ix }}, [rax + {{ ix * 32 }}]
|
||||
{% endfor %}
|
||||
{% endif %}
|
||||
|
||||
{% if flipped %}
|
||||
{% for acc in range(from, to + 1) %}
|
||||
{{op}} ymm{{acc}}, ymm{{acc}}, ymm{{ acc % mr_over_8 + to + 1 }}
|
||||
{% endfor %}
|
||||
{% else %}
|
||||
{% for acc in range(from, to + 1) %}
|
||||
{{op}} ymm{{acc}}, ymm{{ acc % mr_over_8 + to + 1 }}, ymm{{acc}}
|
||||
{% endfor %}
|
||||
{% endif %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
{% endmacro %}
|
||||
|
||||
{% macro per_col(label, op, mr, from, to, type="f32", flipped=false) %}
|
||||
{{L}}{{label}}:
|
||||
mov rax, [ rdi + 8 ]
|
||||
|
||||
{% set mr_over_8 = mr // 8 %}
|
||||
{% set mr_over_8_min_1 = mr // 8 - 1 %}
|
||||
|
||||
{% set tmp = to + 1 %}
|
||||
|
||||
{% set cols = (to + 1 - from) // mr_over_8 %}
|
||||
{% set cols_min_1 = (to + 1 - from) // mr_over_8 - 1 %}
|
||||
|
||||
|
||||
{% for right in range(0, cols_min_1 + 1) %}
|
||||
{% if type == "f16" %}
|
||||
pinsrw xmm{{tmp}}, word ptr [ rax ], 0
|
||||
add rax, 2
|
||||
vcvtph2ps ymm{{tmp}}, xmm{{tmp}}
|
||||
vbroadcastss ymm{{tmp}}, xmm{{tmp}}
|
||||
{% else %}
|
||||
vbroadcastss ymm{{tmp}}, dword ptr [ rax ]
|
||||
add rax, 4
|
||||
{% endif %}
|
||||
{% for down in range(0, mr_over_8_min_1 + 1) %}
|
||||
{% set acc = mr_over_8 * right + from + down %}
|
||||
{% if flipped %}
|
||||
{{op}} ymm{{acc}}, ymm{{acc}}, ymm{{tmp}}
|
||||
{% else %}
|
||||
{{op}} ymm{{acc}}, ymm{{tmp}}, ymm{{acc}}
|
||||
{% endif %}
|
||||
{% endfor %}
|
||||
{% endfor %}
|
||||
|
||||
jmp {{L}}non_linear_loop
|
||||
{% endmacro %}
|
||||
Vendored
+319
@@ -0,0 +1,319 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
|
||||
#}
|
||||
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
fma_sigmoid_f32_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}fma_sigmoid_f32_{{suffix}}
|
||||
{{G}}fma_sigmoid_f32_{{suffix}}:
|
||||
.cfi_startproc
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
// move around arguments to mimick SysV rdi,rsi passing
|
||||
push rdi
|
||||
push rsi
|
||||
mov rdi, rcx
|
||||
mov rsi, rdx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
// FIXME
|
||||
// .cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
cmp rsi, 32
|
||||
jl {{L}}loop_1
|
||||
|
||||
{{L}}loop_4:
|
||||
|
||||
vmovaps ymm4, [rdi]
|
||||
vmovaps ymm5, [rdi + 32]
|
||||
vmovaps ymm6, [rdi + 64]
|
||||
vmovaps ymm7, [rdi + 96]
|
||||
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_13]
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_11]
|
||||
|
||||
vmaxps ymm4, ymm4, ymm0
|
||||
vmaxps ymm5, ymm5, ymm0
|
||||
vmaxps ymm6, ymm6, ymm0
|
||||
vmaxps ymm7, ymm7, ymm0
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
|
||||
vminps ymm4, ymm4, ymm1
|
||||
vminps ymm5, ymm5, ymm1
|
||||
vminps ymm6, ymm6, ymm1
|
||||
vminps ymm7, ymm7, ymm1 // ymm4..7 <- x
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmulps ymm8, ymm4, ymm4
|
||||
vmulps ymm9, ymm5, ymm5
|
||||
vmulps ymm10, ymm6, ymm6
|
||||
vmulps ymm11, ymm7, ymm7 // ymm8..11 <- x^2
|
||||
|
||||
vmovaps ymm12, ymm2
|
||||
vmovaps ymm13, ymm2
|
||||
vmovaps ymm14, ymm2
|
||||
vmovaps ymm15, ymm2
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm13, ymm3, ymm9
|
||||
vfmadd132ps ymm14, ymm3, ymm10
|
||||
vfmadd132ps ymm15, ymm3, ymm11
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vfmadd132ps ymm13, ymm0, ymm9
|
||||
vfmadd132ps ymm14, ymm0, ymm10
|
||||
vfmadd132ps ymm15, ymm0, ymm11
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps ymm12, ymm1, ymm8
|
||||
vfmadd132ps ymm13, ymm1, ymm9
|
||||
vfmadd132ps ymm14, ymm1, ymm10
|
||||
vfmadd132ps ymm15, ymm1, ymm11
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vfmadd132ps ymm13, ymm2, ymm9
|
||||
vfmadd132ps ymm14, ymm2, ymm10
|
||||
vfmadd132ps ymm15, ymm2, ymm11
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm13, ymm3, ymm9
|
||||
vfmadd132ps ymm14, ymm3, ymm10
|
||||
vfmadd132ps ymm15, ymm3, ymm11
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vfmadd132ps ymm13, ymm0, ymm9
|
||||
vfmadd132ps ymm14, ymm0, ymm10
|
||||
vfmadd132ps ymm15, ymm0, ymm11
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vmulps ymm4, ymm4, ymm12
|
||||
vmulps ymm5, ymm5, ymm13
|
||||
vmulps ymm6, ymm6, ymm14
|
||||
vmulps ymm7, ymm7, ymm15 // ymm4..7 <- num
|
||||
|
||||
vmovaps ymm12, ymm1
|
||||
vmovaps ymm13, ymm1
|
||||
vmovaps ymm14, ymm1
|
||||
vmovaps ymm15, ymm1
|
||||
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_half]
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vfmadd132ps ymm13, ymm2, ymm9
|
||||
vfmadd132ps ymm14, ymm2, ymm10
|
||||
vfmadd132ps ymm15, ymm2, ymm11
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm13, ymm3, ymm9
|
||||
vfmadd132ps ymm14, ymm3, ymm10
|
||||
vfmadd132ps ymm15, ymm3, ymm11
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vfmadd132ps ymm13, ymm0, ymm9
|
||||
vfmadd132ps ymm14, ymm0, ymm10
|
||||
vfmadd132ps ymm15, ymm0, ymm11 // ymm12..14 <- denum
|
||||
|
||||
vdivps ymm4, ymm4, ymm12
|
||||
vdivps ymm5, ymm5, ymm13
|
||||
vdivps ymm6, ymm6, ymm14
|
||||
vdivps ymm7, ymm7, ymm15
|
||||
vaddps ymm4, ymm4, ymm1
|
||||
vaddps ymm5, ymm5, ymm1
|
||||
vaddps ymm6, ymm6, ymm1
|
||||
vaddps ymm7, ymm7, ymm1
|
||||
|
||||
vmovaps [rdi], ymm4
|
||||
vmovaps [rdi + 32], ymm5
|
||||
vmovaps [rdi + 64], ymm6
|
||||
vmovaps [rdi + 96], ymm7
|
||||
|
||||
add rdi, 128
|
||||
sub rsi, 32
|
||||
cmp rsi, 32
|
||||
jg {{L}}loop_4
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
{{L}}loop_1:
|
||||
vmovaps ymm4, [rdi]
|
||||
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_13]
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_11]
|
||||
|
||||
vmaxps ymm4, ymm4, ymm0
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
|
||||
vminps ymm4, ymm4, ymm1 // ymm4 <- x
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmulps ymm8, ymm4, ymm4 // ymm8 <- x^2
|
||||
|
||||
vmovaps ymm12, ymm2
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps ymm12, ymm1, ymm8
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vmulps ymm4, ymm4, ymm12
|
||||
|
||||
vmovaps ymm12, ymm1
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_half]
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
|
||||
vdivps ymm4, ymm4, ymm12
|
||||
vaddps ymm4, ymm4, ymm1
|
||||
|
||||
vmovaps [rdi], ymm4
|
||||
add rdi, 32
|
||||
sub rsi, 8
|
||||
jnz {{L}}loop_1
|
||||
{{L}}done:
|
||||
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
ldmxcsr [rsp + 4]
|
||||
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
{% set float %}{% if msvc %} real4 {%else%} .float {%endif%}{% endset %}
|
||||
|
||||
{{L}}coeffs_num_low:
|
||||
{{float}} -18.6 // low
|
||||
{{L}}coeffs_num_high:
|
||||
{{float}} 18.6 // high
|
||||
|
||||
{{L}}coeffs_num_alpha_13:
|
||||
{{float}} -4.433153405e-18
|
||||
{{L}}coeffs_num_alpha_11:
|
||||
{{float}} 1.169974371e-14
|
||||
{{L}}coeffs_num_alpha_9:
|
||||
{{float}} -1.875289645e-11
|
||||
{{L}}coeffs_num_alpha_7:
|
||||
{{float}} 4.257889523e-8
|
||||
{{L}}coeffs_num_alpha_5:
|
||||
{{float}} 0.00004811817576
|
||||
{{L}}coeffs_num_alpha_3:
|
||||
{{float}} 0.008163842030
|
||||
{{L}}coeffs_num_alpha_1:
|
||||
{{float}} 0.2499999971
|
||||
|
||||
{{L}}coeffs_num_beta_6:
|
||||
{{float}} 3.922935744e-6
|
||||
{{L}}coeffs_num_beta_4:
|
||||
{{float}} 0.001524872358
|
||||
{{L}}coeffs_num_beta_2:
|
||||
{{float}} 0.1159886749
|
||||
{{L}}coeffs_num_beta_0:
|
||||
{{float}} 1.0;
|
||||
|
||||
{{L}}coeffs_num_half:
|
||||
{{float}} 0.5
|
||||
|
||||
{% if msvc %}
|
||||
fma_sigmoid_f32_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
Vendored
+313
@@ -0,0 +1,313 @@
|
||||
{#
|
||||
// vim: set syntax=asm :
|
||||
|
||||
System V ABI:
|
||||
args: rdi, rsi, rdx, rcx, r8, r9
|
||||
preserve: rbx, rsp, rbp, r12, r13, r14, r15
|
||||
scratch: rax, rdi, rsi, rdx, rcx, r8, r9, r10, r11
|
||||
return: rax (+rdx)
|
||||
|
||||
Windows ABI:
|
||||
args: RCX, RDX, R8, R9
|
||||
preserve: RBX, RBP, RDI, RSI, RSP, R12, R13, R14, R15, and XMM6-15
|
||||
scratch: RAX, RCX, RDX, R8, R9, R10, R11, XMM0-5, and the upper portions of YMM0-15 and ZMM0-15
|
||||
return: rax (+rdx)
|
||||
|
||||
#}
|
||||
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
fma_tanh_f32_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}fma_tanh_f32_{{suffix}}
|
||||
{{G}}fma_tanh_f32_{{suffix}}:
|
||||
.cfi_startproc
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
// move around arguments to mimick SysV rdi,rsi passing
|
||||
push rdi
|
||||
push rsi
|
||||
mov rdi, rcx
|
||||
mov rsi, rdx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
// FIXME
|
||||
// .cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
{% set offset %}{% if msvc %} offset {%else%} rip + {%endif%} {% endset %}
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
cmp rsi, 32
|
||||
jl {{L}}loop_1
|
||||
|
||||
{{L}}loop_4:
|
||||
|
||||
vmovaps ymm4, [rdi]
|
||||
vmovaps ymm5, [rdi + 32]
|
||||
vmovaps ymm6, [rdi + 64]
|
||||
vmovaps ymm7, [rdi + 96]
|
||||
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_13]
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_11]
|
||||
|
||||
vmaxps ymm4, ymm4, ymm0
|
||||
vmaxps ymm5, ymm5, ymm0
|
||||
vmaxps ymm6, ymm6, ymm0
|
||||
vmaxps ymm7, ymm7, ymm0
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
|
||||
vminps ymm4, ymm4, ymm1
|
||||
vminps ymm5, ymm5, ymm1
|
||||
vminps ymm6, ymm6, ymm1
|
||||
vminps ymm7, ymm7, ymm1 // ymm4..7 <- x
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmulps ymm8, ymm4, ymm4
|
||||
vmulps ymm9, ymm5, ymm5
|
||||
vmulps ymm10, ymm6, ymm6
|
||||
vmulps ymm11, ymm7, ymm7 // ymm8..11 <- x^2
|
||||
|
||||
vmovaps ymm12, ymm2
|
||||
vmovaps ymm13, ymm2
|
||||
vmovaps ymm14, ymm2
|
||||
vmovaps ymm15, ymm2
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm13, ymm3, ymm9
|
||||
vfmadd132ps ymm14, ymm3, ymm10
|
||||
vfmadd132ps ymm15, ymm3, ymm11
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vfmadd132ps ymm13, ymm0, ymm9
|
||||
vfmadd132ps ymm14, ymm0, ymm10
|
||||
vfmadd132ps ymm15, ymm0, ymm11
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps ymm12, ymm1, ymm8
|
||||
vfmadd132ps ymm13, ymm1, ymm9
|
||||
vfmadd132ps ymm14, ymm1, ymm10
|
||||
vfmadd132ps ymm15, ymm1, ymm11
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vfmadd132ps ymm13, ymm2, ymm9
|
||||
vfmadd132ps ymm14, ymm2, ymm10
|
||||
vfmadd132ps ymm15, ymm2, ymm11
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm13, ymm3, ymm9
|
||||
vfmadd132ps ymm14, ymm3, ymm10
|
||||
vfmadd132ps ymm15, ymm3, ymm11
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vfmadd132ps ymm13, ymm0, ymm9
|
||||
vfmadd132ps ymm14, ymm0, ymm10
|
||||
vfmadd132ps ymm15, ymm0, ymm11
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vmulps ymm4, ymm4, ymm12
|
||||
vmulps ymm5, ymm5, ymm13
|
||||
vmulps ymm6, ymm6, ymm14
|
||||
vmulps ymm7, ymm7, ymm15 // ymm4..7 <- num
|
||||
|
||||
vmovaps ymm12, ymm1
|
||||
vmovaps ymm13, ymm1
|
||||
vmovaps ymm14, ymm1
|
||||
vmovaps ymm15, ymm1
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vfmadd132ps ymm13, ymm2, ymm9
|
||||
vfmadd132ps ymm14, ymm2, ymm10
|
||||
vfmadd132ps ymm15, ymm2, ymm11
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm13, ymm3, ymm9
|
||||
vfmadd132ps ymm14, ymm3, ymm10
|
||||
vfmadd132ps ymm15, ymm3, ymm11
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vfmadd132ps ymm13, ymm0, ymm9
|
||||
vfmadd132ps ymm14, ymm0, ymm10
|
||||
vfmadd132ps ymm15, ymm0, ymm11 // ymm12..14 <- denum
|
||||
|
||||
vdivps ymm4, ymm4, ymm12
|
||||
vdivps ymm5, ymm5, ymm13
|
||||
vdivps ymm6, ymm6, ymm14
|
||||
vdivps ymm7, ymm7, ymm15
|
||||
|
||||
vmovaps [rdi], ymm4
|
||||
vmovaps [rdi + 32], ymm5
|
||||
vmovaps [rdi + 64], ymm6
|
||||
vmovaps [rdi + 96], ymm7
|
||||
|
||||
add rdi, 128
|
||||
sub rsi, 32
|
||||
cmp rsi, 32
|
||||
jg {{L}}loop_4
|
||||
|
||||
cmp rsi, 0
|
||||
je {{L}}done
|
||||
|
||||
{{L}}loop_1:
|
||||
vmovaps ymm4, [rdi]
|
||||
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_low]
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_high]
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_13]
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_11]
|
||||
|
||||
vmaxps ymm4, ymm4, ymm0
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_9]
|
||||
|
||||
vminps ymm4, ymm4, ymm1 // ymm4 <- x
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_alpha_7]
|
||||
|
||||
vmulps ymm8, ymm4, ymm4 // ymm8 <- x^2
|
||||
|
||||
vmovaps ymm12, ymm2
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_alpha_5]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_alpha_3]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_alpha_1]
|
||||
vfmadd132ps ymm12, ymm1, ymm8
|
||||
vbroadcastss ymm1, dword ptr [{{offset}} {{L}}coeffs_num_beta_6]
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vbroadcastss ymm2, dword ptr [{{offset}} {{L}}coeffs_num_beta_4]
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vbroadcastss ymm3, dword ptr [{{offset}} {{L}}coeffs_num_beta_2]
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
vbroadcastss ymm0, dword ptr [{{offset}} {{L}}coeffs_num_beta_0]
|
||||
vmulps ymm4, ymm4, ymm12
|
||||
|
||||
vmovaps ymm12, ymm1
|
||||
vfmadd132ps ymm12, ymm2, ymm8
|
||||
vfmadd132ps ymm12, ymm3, ymm8
|
||||
vfmadd132ps ymm12, ymm0, ymm8
|
||||
|
||||
vdivps ymm4, ymm4, ymm12
|
||||
|
||||
vmovaps [rdi], ymm4
|
||||
add rdi, 32
|
||||
sub rsi, 8
|
||||
jnz {{L}}loop_1
|
||||
|
||||
{{L}}done:
|
||||
|
||||
// ----------------------------------------------------------------------
|
||||
|
||||
ldmxcsr [rsp + 4]
|
||||
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
{% set float %}{% if msvc %} real4 {%else%} .float {%endif%}{% endset %}
|
||||
|
||||
{{L}}coeffs_num_low:
|
||||
{{float}} -8.9
|
||||
{{L}}coeffs_num_high:
|
||||
{{float}} 8.9
|
||||
|
||||
{{L}}coeffs_num_alpha_13:
|
||||
{{float}} -8.488492677e-14
|
||||
{{L}}coeffs_num_alpha_11:
|
||||
{{float}} 5.277853000e-11
|
||||
{{L}}coeffs_num_alpha_9:
|
||||
{{float}} -2.022500419e-8
|
||||
{{L}}coeffs_num_alpha_7:
|
||||
{{float}} 0.00001115424833
|
||||
{{L}}coeffs_num_alpha_5:
|
||||
{{float}} 0.003103950131
|
||||
{{L}}coeffs_num_alpha_3:
|
||||
{{float}} 0.1308400453
|
||||
{{L}}coeffs_num_alpha_1:
|
||||
{{float}} 0.9999999934
|
||||
|
||||
{{L}}coeffs_num_beta_6:
|
||||
{{float}} 0.0002546136580
|
||||
{{L}}coeffs_num_beta_4:
|
||||
{{float}} 0.02449515379
|
||||
{{L}}coeffs_num_beta_2:
|
||||
{{float}} 0.4641733162
|
||||
{{L}}coeffs_num_beta_0:
|
||||
{{float}} 1.0
|
||||
|
||||
|
||||
|
||||
{% if msvc %}
|
||||
fma_tanh_f32_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
+38
@@ -0,0 +1,38 @@
|
||||
{{L}}return:
|
||||
ldmxcsr [rsp + 4]
|
||||
add rsp, 8
|
||||
|
||||
pop r15
|
||||
pop r14
|
||||
pop r13
|
||||
pop r12
|
||||
pop rbx
|
||||
|
||||
{% if family == "windows" %}
|
||||
pop rsi
|
||||
pop rdi
|
||||
|
||||
vmovaps xmm15, [rsp+16*9]
|
||||
vmovaps xmm14, [rsp+16*8]
|
||||
vmovaps xmm13, [rsp+16*7]
|
||||
vmovaps xmm12, [rsp+16*6]
|
||||
vmovaps xmm11, [rsp+16*5]
|
||||
vmovaps xmm10, [rsp+16*4]
|
||||
vmovaps xmm9, [rsp+16*3]
|
||||
vmovaps xmm8, [rsp+16*2]
|
||||
vmovaps xmm7, [rsp+16*1]
|
||||
vmovaps xmm6, [rsp]
|
||||
{% endif %}
|
||||
|
||||
mov rsp, rbp
|
||||
pop rbp
|
||||
ret
|
||||
|
||||
{% if msvc %}
|
||||
fma_mmm_{{type}}_{{size}}_{{suffix}} endp
|
||||
_text ends
|
||||
end
|
||||
|
||||
{% else %}
|
||||
.cfi_endproc
|
||||
{% endif %}
|
||||
+64
@@ -0,0 +1,64 @@
|
||||
|
||||
{% if msvc %}
|
||||
|
||||
_text segment
|
||||
fma_mmm_{{type}}_{{size}}_{{suffix}} proc
|
||||
|
||||
{% else %}
|
||||
|
||||
.intel_syntax noprefix
|
||||
.text
|
||||
.p2align 5
|
||||
.globl {{G}}fma_mmm_{{type}}_{{size}}_{{suffix}}
|
||||
{{G}}fma_mmm_{{type}}_{{size}}_{{suffix}}:
|
||||
.cfi_startproc
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbp
|
||||
mov rbp, rsp
|
||||
|
||||
{% if family == "windows" %}
|
||||
// https://www.agner.org/optimize/calling_conventions.pdf xmm6-15 are not scratch
|
||||
// https://stackoverflow.com/questions/43358429/save-value-of-xmm-registers
|
||||
and rsp,-16
|
||||
lea rsp,[rsp-160]
|
||||
vmovaps [rsp], xmm6
|
||||
vmovaps [rsp+16*1],xmm7
|
||||
vmovaps [rsp+16*2],xmm8
|
||||
vmovaps [rsp+16*3],xmm9
|
||||
vmovaps [rsp+16*4],xmm10
|
||||
vmovaps [rsp+16*5],xmm11
|
||||
vmovaps [rsp+16*6],xmm12
|
||||
vmovaps [rsp+16*7],xmm13
|
||||
vmovaps [rsp+16*8],xmm14
|
||||
vmovaps [rsp+16*9],xmm15
|
||||
|
||||
push rdi
|
||||
push rsi
|
||||
|
||||
mov rdi, rcx
|
||||
|
||||
{% endif %}
|
||||
|
||||
push rbx
|
||||
push r12
|
||||
push r13
|
||||
push r14
|
||||
push r15
|
||||
|
||||
sub rsp, 8
|
||||
|
||||
{% if family == "unix" %}
|
||||
.cfi_def_cfa_offset 64
|
||||
{% endif %}
|
||||
stmxcsr [rsp + 4]
|
||||
{% if msvc %}
|
||||
mov rax, 1FC0h
|
||||
{% else %}
|
||||
mov rax, 0x1FC0
|
||||
{% endif %}
|
||||
mov [rsp], eax
|
||||
ldmxcsr [rsp]
|
||||
|
||||
{% include "dispatcher.j2" %}
|
||||
Reference in New Issue
Block a user