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:
2026-07-01 18:22:43 -04:00
commit 682afacd30
1763 changed files with 613720 additions and 0 deletions
@@ -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
@@ -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
@@ -0,0 +1,7 @@
vbroadcastss zmm15, dword ptr [rcx]
vmovups zmm8, [rax]
vfmadd231ps zmm0, zmm15, zmm8
add rcx, 4
add rax, 64
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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
@@ -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
@@ -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) }}
@@ -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) }}
@@ -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
@@ -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) }}
@@ -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) }}
@@ -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) }}
@@ -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 %}
@@ -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" %}
@@ -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 %}
@@ -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 %}
@@ -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 %}
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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 %}
@@ -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 %}
@@ -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
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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" %}
@@ -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) }}
@@ -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) }}
@@ -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
@@ -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) }}
@@ -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) }}
@@ -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
@@ -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
@@ -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 %}
@@ -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 %}
@@ -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 %}
@@ -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 %}
@@ -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" %}