import macros;macro ImportExpand(s:untyped):untyped = parseStmt($s[2]) # source: src/cplib/tmpl/sheep_old.nim ImportExpand "cplib/tmpl/sheep_old" <=== "when not declared CPLIB_TMPL_SHEEP:\n const CPLIB_TMPL_SHEEP* = 1\n {.warning[UnusedImport]: off.}\n {.hint[XDeclaredButNotUsed]: off.}\n import algorithm\n import sequtils\n import tables\n import macros\n import math\n import sets\n import strutils\n import strformat\n import sugar\n import heapqueue\n import streams\n import deques\n import bitops\n import std/lenientops\n import options\n #入力系\n proc scanf(formatstr: cstring){.header: \"\", varargs.}\n proc getchar(): char {.importc: \"getchar_unlocked\", header: \"\", discardable.}\n proc ii(): int {.inline.} = scanf(\"%lld\\n\", addr result)\n proc lii(N: int): seq[int] {.inline.} = newSeqWith(N, ii())\n proc si(): string {.inline.} =\n result = \"\"\n var c: char\n while true:\n c = getchar()\n if c == ' ' or c == '\\n' or c == '\\255':\n break\n result &= c\n \n # 出力系\n # 1. 実際の処理を行う proc (openArray を受け取る)\n proc print_internal(prop: tuple[f: File, sepc: string, endc: string, flush: bool], args: openArray[string]) =\n for i in 0 ..< args.len:\n prop.f.write(args[i])\n if i != args.len - 1:\n prop.f.write(prop.sepc)\n else:\n prop.f.write(prop.endc)\n if prop.flush:\n prop.f.flushFile()\n\n # 2. ユーザーが呼び出すためのインターフェース (varargs を受け取る)\n proc print*(prop: tuple[f: File, sepc: string, endc: string, flush: bool], args: varargs[string, `$`]) =\n # varargs は内部では openArray として扱えるので、そのまま渡せる\n print_internal(prop, args)\n\n proc print*(args: varargs[string, `$`]) =\n # こちらも内部用の proc を呼ぶ\n print_internal((f: stdout, sepc: \" \", endc: \"\\n\", flush: false), args)\n macro getSymbolName(x: typed): string = x.toStrLit\n macro debug*(args: varargs[untyped]): untyped =\n when defined(debug):\n result = newNimNode(nnkStmtList, args)\n template prop(e: string = \"\"): untyped = (f: stderr, sepc: \"\", endc: e, flush: true)\n for i, arg in args:\n if arg.kind == nnkStrLit:\n result.add(quote do: print(prop(), \"\\\"\", `arg`, \"\\\"\"))\n else:\n result.add(quote do: print(prop(\": \"), getSymbolName(`arg`)))\n result.add(quote do: print(prop(), `arg`))\n if i != args.len - 1: result.add(quote do: print(prop(), \", \"))\n else: result.add(quote do: print(prop(), \"\\n\"))\n else:\n return (quote do: discard)\n #chmin,chmax\n template `max=`(x, y) =\n let yVal = y # yが計算式の場合に評価を1回にするため\n if x < yVal:\n x = yVal\n\n template `min=`(x, y) =\n let yVal = y\n if x > yVal:\n x = yVal\n proc chmin[T](x: var T, y: T):bool=\n if x > y:\n x = y\n return true\n return false\n proc chmax[T](x: var T, y: T):bool=\n if x < y:\n x = y\n return true\n return false\n #bit演算\n proc `%`*(x: int, y: int): int =\n result = x mod y\n if y > 0 and result < 0: result += y\n if y < 0 and result > 0: result += y\n proc `//`*(x: int, y: int): int{.inline.} =\n result = x div y\n if y > 0 and result * y > x: result -= 1\n if y < 0 and result * y < x: result -= 1\n proc `%=`(x: var int, y: int): void = x = x%y\n proc `//=`(x: var int, y: int): void = x = x//y\n proc `**`(x: int, y: int): int = x^y\n proc `**=`(x: var int, y: int): void = x = x^y\n proc `^`(x: int, y: int): int = x xor y\n proc `|`(x: int, y: int): int = x or y\n proc `&`(x: int, y: int): int = x and y\n proc `>>`(x: int, y: int): int = x shr y\n proc `<<`(x: int, y: int): int = x shl y\n proc `~`(x: int): int = not x\n proc `^=`(x: var int, y: int): void = x = x ^ y\n proc `&=`(x: var int, y: int): void = x = x & y\n proc `|=`(x: var int, y: int): void = x = x | y\n proc `>>=`(x: var int, y: int): void = x = x >> y\n proc `<<=`(x: var int, y: int): void = x = x << y\n proc `[]`(x: int, n: int): bool = (x and (1 shl n)) != 0\n #便利な変換\n proc `!`(x: char, a = '0'): int = int(x)-int(a)\n #定数\n when not declared CPLIB_UTILS_CONSTANTS:\n const CPLIB_UTILS_CONSTANTS* = 1\n const INF32*: int32 = 1001000027.int32\n const INF64*: int = int(3300300300300300491)\n \n const INF = INF64\n #converter\n\n #range\n iterator range(start: int, ends: int, step: int): int =\n var i = start\n if step < 0:\n while i > ends:\n yield i\n i += step\n elif step > 0:\n while i < ends:\n yield i\n i += step\n iterator range(ends: int): int = (for i in 0.. r[i]:\n return false\n elif l[i] < r[i]:\n return true\n return len(l) < len(r)\n \n # Yes/No\n proc yes*(b: bool = true): void = print(if b: \"Yes\" else: \"No\")\n proc no*(b: bool = true): void = yes(not b)\n\n proc takahashi(b:bool = true) : void = print(if b: \"Takahashi\" else: \"Aoki\")\n proc aoki(b:bool = true) : void = takahashi(not b)\n\n template dblock(body: untyped) =\n when defined(debug):\n block:\n body\n" when not declared CPLIB_COLLECTIONS_BITSET_AVX2: const CPLIB_COLLECTIONS_BITSET_AVX2* = 1 when not (defined(amd64) and (defined(gcc) or defined(clang))): {.error: "BitSetAvx2 requires amd64 and GCC/Clang".} import bitops type BitSetAvx2* {.byref.} = object bits: seq[uint64] size: int {.emit: """ #include #include #include #define CPLIB_BS_AVX2 __attribute__((target("avx2"))) #define CPLIB_BS_BINARY(name, scalar, vector) \ CPLIB_BS_AVX2 static void name(uint64_t *dst, const uint64_t *x, \ const uint64_t *y, size_t n) { \ /* 256ビットずつ論理演算し、残りを64ビットずつ処理します。 */ \ size_t i = 0; \ for (; i + 4 <= n; i += 4) { \ __m256i a = _mm256_loadu_si256((const __m256i *)(x + i)); \ __m256i b = _mm256_loadu_si256((const __m256i *)(y + i)); \ _mm256_storeu_si256((__m256i *)(dst + i), vector(a, b)); \ } \ for (; i < n; ++i) dst[i] = x[i] scalar y[i]; \ } CPLIB_BS_BINARY(cplib_bs_and, &, _mm256_and_si256) CPLIB_BS_BINARY(cplib_bs_or, |, _mm256_or_si256) CPLIB_BS_BINARY(cplib_bs_xor, ^, _mm256_xor_si256) #undef CPLIB_BS_BINARY CPLIB_BS_AVX2 static void cplib_bs_not(uint64_t *dst, const uint64_t *x, size_t n) { /* 256ビットずつ反転します。 */ const __m256i ones = _mm256_set1_epi64x(-1); size_t i = 0; for (; i + 4 <= n; i += 4) _mm256_storeu_si256((__m256i *)(dst + i), _mm256_xor_si256( _mm256_loadu_si256((const __m256i *)(x + i)), ones)); for (; i < n; ++i) dst[i] = ~x[i]; } CPLIB_BS_AVX2 static void cplib_bs_shl(uint64_t *dst, const uint64_t *x, size_t n, size_t shift) { /* ゼロ初期化済みの別領域へ左シフトし、隣接ワードからの桁上がりも処理します。 */ const size_t offset = shift >> 6; const unsigned bits = shift & 63; const size_t count = n - offset; size_t i = 0; if (bits == 0) { for (; i + 4 <= count; i += 4) _mm256_storeu_si256((__m256i *)(dst + offset + i), _mm256_loadu_si256((const __m256i *)(x + i))); for (; i < count; ++i) dst[offset + i] = x[i]; return; } const __m128i left = _mm_cvtsi32_si128(bits); const __m128i right = _mm_cvtsi32_si128(64 - bits); dst[offset] = x[0] << bits; i = 1; for (; i + 4 <= count; i += 4) { __m256i a = _mm256_loadu_si256((const __m256i *)(x + i)); __m256i b = _mm256_loadu_si256((const __m256i *)(x + i - 1)); _mm256_storeu_si256((__m256i *)(dst + offset + i), _mm256_or_si256( _mm256_sll_epi64(a, left), _mm256_srl_epi64(b, right))); } for (; i < count; ++i) dst[offset + i] = (x[i] << bits) | (x[i - 1] >> (64 - bits)); } CPLIB_BS_AVX2 static void cplib_bs_shr(uint64_t *dst, const uint64_t *x, size_t n, size_t shift) { /* ゼロ初期化済みの別領域へ右シフトし、隣接ワードからの桁下がりも処理します。 */ const size_t offset = shift >> 6; const unsigned bits = shift & 63; const size_t count = n - offset; size_t i = 0; if (bits == 0) { for (; i + 4 <= count; i += 4) _mm256_storeu_si256((__m256i *)(dst + i), _mm256_loadu_si256((const __m256i *)(x + offset + i))); for (; i < count; ++i) dst[i] = x[offset + i]; return; } const __m128i right = _mm_cvtsi32_si128(bits); const __m128i left = _mm_cvtsi32_si128(64 - bits); for (; i + 4 < count; i += 4) { __m256i a = _mm256_loadu_si256((const __m256i *)(x + offset + i)); __m256i b = _mm256_loadu_si256((const __m256i *)(x + offset + i + 1)); _mm256_storeu_si256((__m256i *)(dst + i), _mm256_or_si256( _mm256_srl_epi64(a, right), _mm256_sll_epi64(b, left))); } for (; i + 1 < count; ++i) dst[i] = (x[offset + i] >> bits) | (x[offset + i + 1] << (64 - bits)); dst[count - 1] = x[n - 1] >> bits; } CPLIB_BS_AVX2 static inline __m256i cplib_bs_byte_counts(__m256i x) { /* 4ビットの参照表から各バイトの立っているビット数を求めます。 */ const __m256i table = _mm256_setr_epi8( 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4, 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4); const __m256i mask = _mm256_set1_epi8(15); return _mm256_add_epi8( _mm256_shuffle_epi8(table, _mm256_and_si256(x, mask)), _mm256_shuffle_epi8(table, _mm256_and_si256(_mm256_srli_epi16(x, 4), mask))); } #define CPLIB_BS_COUNT(name, scalar, vector) \ CPLIB_BS_AVX2 static size_t name(const uint64_t *x, const uint64_t *y, size_t n) { \ /* 16ベクトルごとにバイトの和を64ビットへ集約し、桁あふれを防ぎます。 */ \ __m256i total = _mm256_setzero_si256(); \ size_t i = 0; \ while (i + 4 <= n) { \ __m256i local = _mm256_setzero_si256(); \ size_t end = n - i < 64 ? n : i + 64; \ for (; i + 4 <= end; i += 4) { \ __m256i a = _mm256_loadu_si256((const __m256i *)(x + i)); \ __m256i b = _mm256_loadu_si256((const __m256i *)(y + i)); \ local = _mm256_add_epi8(local, cplib_bs_byte_counts(vector)); \ } \ total = _mm256_add_epi64(total, _mm256_sad_epu8(local, _mm256_setzero_si256())); \ } \ uint64_t lanes[4]; \ _mm256_storeu_si256((__m256i *)lanes, total); \ size_t result = lanes[0] + lanes[1] + lanes[2] + lanes[3]; \ for (; i < n; ++i) result += __builtin_popcountll(scalar); \ return result; \ } CPLIB_BS_COUNT(cplib_bs_popcount, x[i], a) CPLIB_BS_COUNT(cplib_bs_andpopcount, x[i] & y[i], _mm256_and_si256(a, b)) CPLIB_BS_COUNT(cplib_bs_orpopcount, x[i] | y[i], _mm256_or_si256(a, b)) CPLIB_BS_COUNT(cplib_bs_xorpopcount, x[i] ^ y[i], _mm256_xor_si256(a, b)) #undef CPLIB_BS_COUNT #undef CPLIB_BS_AVX2 """.} proc avxAnd(dst, x, y: ptr uint64, n: csize_t) {.importc: "cplib_bs_and", nodecl.} proc avxOr(dst, x, y: ptr uint64, n: csize_t) {.importc: "cplib_bs_or", nodecl.} proc avxXor(dst, x, y: ptr uint64, n: csize_t) {.importc: "cplib_bs_xor", nodecl.} proc avxNot(dst, x: ptr uint64, n: csize_t) {.importc: "cplib_bs_not", nodecl.} proc avxShl(dst, x: ptr uint64, n, shift: csize_t) {.importc: "cplib_bs_shl", nodecl.} proc avxShr(dst, x: ptr uint64, n, shift: csize_t) {.importc: "cplib_bs_shr", nodecl.} proc avxPopcount(x, y: ptr uint64, n: csize_t): csize_t {.importc: "cplib_bs_popcount", nodecl.} proc avxAndPopcount(x, y: ptr uint64, n: csize_t): csize_t {.importc: "cplib_bs_andpopcount", nodecl.} proc avxOrPopcount(x, y: ptr uint64, n: csize_t): csize_t {.importc: "cplib_bs_orpopcount", nodecl.} proc avxXorPopcount(x, y: ptr uint64, n: csize_t): csize_t {.importc: "cplib_bs_xorpopcount", nodecl.} proc initBitSet*(N: int): BitSetAvx2 = if N < 0: raise newException(ValueError, "BitSet size must be non-negative") result.size = N result.bits = newSeq[uint64]((N shr 6) + ord((N and 63) != 0)) proc initBitSet*(v: openArray[bool], N: int): BitSetAvx2 = if v.len > N: raise newException(ValueError, "initial value is longer than BitSet size") result = initBitSet(N) for i in 0..= N: raise newException(IndexDefect, "BitSet index out of bounds") result.bits[i shr 6] = result.bits[i shr 6] or (1'u64 shl (i and 63)) proc len*(bitset: BitSetAvx2): int {.inline.} = bitset.size proc checkSameSize(x, y: BitSetAvx2) {.inline.} = if x.size != y.size: raise newException(ValueError, "BitSet sizes must match") proc checkIndex(bitset: BitSetAvx2, idx: Natural) {.inline.} = if idx >= bitset.size: raise newException(IndexDefect, "BitSet index out of bounds") proc trim(bitset: var BitSetAvx2) {.inline.} = let remainder = bitset.size and 63 if remainder != 0: bitset.bits[^1] = bitset.bits[^1] and ((1'u64 shl remainder) - 1) proc `&`*(x, y: BitSetAvx2): BitSetAvx2 = checkSameSize(x, y) result = initBitSet(x.size) if x.bits.len > 0: avxAnd(addr result.bits[0], unsafeAddr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t) proc `&=`*(x: var BitSetAvx2, y: BitSetAvx2) = checkSameSize(x, y) if x.bits.len > 0: avxAnd(addr x.bits[0], addr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t) proc `|`*(x, y: BitSetAvx2): BitSetAvx2 = checkSameSize(x, y) result = initBitSet(x.size) if x.bits.len > 0: avxOr(addr result.bits[0], unsafeAddr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t) proc `|=`*(x: var BitSetAvx2, y: BitSetAvx2) = checkSameSize(x, y) if x.bits.len > 0: avxOr(addr x.bits[0], addr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t) proc `^`*(x, y: BitSetAvx2): BitSetAvx2 = checkSameSize(x, y) result = initBitSet(x.size) if x.bits.len > 0: avxXor(addr result.bits[0], unsafeAddr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t) proc `^=`*(x: var BitSetAvx2, y: BitSetAvx2) = checkSameSize(x, y) if x.bits.len > 0: avxXor(addr x.bits[0], addr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t) proc `<<`*(bitset: BitSetAvx2, x: int): BitSetAvx2 = if x < 0: raise newException(ValueError, "shift count must be non-negative") result = initBitSet(bitset.size) if x < bitset.size: avxShl(addr result.bits[0], unsafeAddr bitset.bits[0], bitset.bits.len.csize_t, x.csize_t) result.trim() proc `>>`*(bitset: BitSetAvx2, x: int): BitSetAvx2 = if x < 0: raise newException(ValueError, "shift count must be non-negative") result = initBitSet(bitset.size) if x < bitset.size: avxShr(addr result.bits[0], unsafeAddr bitset.bits[0], bitset.bits.len.csize_t, x.csize_t) proc `~`*(x: BitSetAvx2): BitSetAvx2 = result = initBitSet(x.size) if x.bits.len > 0: avxNot(addr result.bits[0], unsafeAddr x.bits[0], x.bits.len.csize_t) result.trim() proc popcount*(x: BitSetAvx2): int = if x.bits.len > 0: result = avxPopcount(unsafeAddr x.bits[0], unsafeAddr x.bits[0], x.bits.len.csize_t).int proc andpopcount*(x, y: BitSetAvx2): int = checkSameSize(x, y) if x.bits.len > 0: result = avxAndPopcount(unsafeAddr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t).int proc orpopcount*(x, y: BitSetAvx2): int = checkSameSize(x, y) if x.bits.len > 0: result = avxOrPopcount(unsafeAddr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t).int proc xorpopcount*(x, y: BitSetAvx2): int = checkSameSize(x, y) if x.bits.len > 0: result = avxXorPopcount(unsafeAddr x.bits[0], unsafeAddr y.bits[0], x.bits.len.csize_t).int iterator items*(bitset: BitSetAvx2): int = for wordIndex in 0..