2 ポイント 投稿者 GN⁺ 2024-03-11 | 1件のコメント | WhatsAppで共有
  • 1BRCのボトルネックは、CSVの温度値10億件を極端な速度でパースすることだった。Quân Anh Maiのmerykitty SWARコードは、ifなしで固定ALU演算により温度を整数化する点で注目を集めた
  • このコードは、long1つに入った8バイトを一度に扱う**SWAR(SIMD Within A Register)**方式で、一般的なCPUレジスタ上で複数の文字を並列的に処理する
  • 処理の流れは、マイナス記号の検出、符号の除去、小数点位置の検出、XY.Zへの整列、ASCII数字の変換、マジック乗算、符号の適用という順で進む
  • 入力形式は -XX.X-X.XX.XXX.X の4種類で、小数点位置を基準にバイトを移動して、異なる長さを同じビット配置に揃える
  • 分岐や反復を減らす代わりに、ASCIIコードの性質、2の補数、ビットマスク、乗算のシフト加算的性質を緻密に活用して高性能なパースを実現している

1BRCでボトルネックになった温度パース

  • One Billion Row Challenge(1BRC)では、CSVファイルの温度値を非常に高速にパースすることが主要なボトルネックとして浮上した
  • 以前の最適化だけでも、慣用的な並列Javaコードは71秒から1.7秒まで高速化されていた
  • 温度形式自体は単純だが、10億件を1秒未満でパースしようとすると、小さなコストでも大きく積み上がる
    • 可能な形式は -XX.X-X.XX.XXX.X
  • 初期の参加者たちは Double.parseDouble() を使っていたが、その後はループのないカスタムパーサが登場した
  • Quân Anh Maiの@merykittyソリューションの一部は、ifなしで単一ファイル読み取りを処理し、1BRC上位ソリューションの標準要素のように広まった
  • 優勝者のThomas Wuerthingerは、自身のソリューションに貢献したチームの一員としてQuân Anhを明記している

merykittyコードがしていること

  • このコードは、8バイトのCSV入力が入ったlongを受け取り、実際の温度の10倍にあたる整数温度値を返す
  • 入力はmmapされたCSVファイルから直接ネイティブメモリ読み取りで取得され、この部分は別の関心事として切り分けられている
  • 演算は固定順序の18個のALU操作で構成される
    • ビットシフト、AND、NOT、XOR
    • 加算、減算、乗算
    • Long.numberOfTrailingZeros()
  • numberOfTrailingZeros() はJDKコンパイラのintrinsicによって特殊なCPU命令を使用する
  • 専用のSIMD命令ではなく、一般的なCPUレジスタと命令で複数バイトを扱うため、SWAR方式に該当する
  • 例示コードは原文を読みやすくするために少し改変したもので、原文は CalculateAverage_merykitty.java にある

全体の処理手順

  • コードは次の順序で温度をパースする
    • 先頭文字が - かどうかを確認して負数かを検出する
    • 符号文字があれば、そのバイトを0にする
    • 小数点 . の位置を見つける
    • 数字が XY.Z テンプレートに合うように long 内のビットを移動する
    • ASCII文字を実際の数値に変換する
    • 各桁に 1x10x100x の重みを掛けて合算する
    • 最後に符号を適用する
  • 見た目は高水準のパース問題だが、各段階はALU演算だけで実装されている

1段階: マイナス記号の検出

  • 符号検出は次のコードから始まる
long negatedInput = ~inputData;
long broadcastSign = (negatedInput << 59) >> 63;
  • 説明上は順序を入れ替えて ( ~(inputData << 59) ) >> 63 のように見ることもできる
  • ASCIIではマイナス -ビット4が0で、数字文字はそのビットが1であるという性質を利用する
  • 入力を左に59ビットシフトすると、先頭文字の判別ビットが最上位ビットへ移動する
  • NOTでビットを反転し、その後に算術右シフトを63ビット行うと、最上位ビットがlong全体に広がる
  • 結果のbroadcastSignは、マイナスがあれば全ビットが1、なければ全ビットが0になる

2段階: 符号文字の除去

  • 負数かどうかはbroadcastSignに保存されたので、入力データからは符号文字を除去する
long maskToRemoveSign = ~(broadcastSign & 0xFF);
long withSignRemoved = inputData & maskToRemoveSign;
  • broadcastSign がすべて1なら、broadcastSign & 0xFF で最下位8ビットだけが1になる
  • これをNOTすると、最下位8ビットだけが0のマスクができる
  • inputData とANDを取ると、最下位バイトの - が除去される
  • マイナスがなければ broadcastSign は0なので、マスクは全ビット1となり、数字バイトはそのまま保持される

3段階: 小数点位置の特定

  • 小数点位置は次のコードで計算される
int dotPos = Long.numberOfTrailingZeros(negatedInput & DOT_DETECTOR);
  • . 文字もマイナスと同様にビット4が0という特性を持つ
  • 可能な小数点位置のビット4だけを確認するため、DOT_DETECTOR = 0x10101000 というマスクを使う
  • 元の入力を反転したnegatedInputでは、小数点位置の該当ビットが1になる
  • Long.numberOfTrailingZeros() はこの1ビットの位置を返す
  • 例の -10.8 では小数点がビット位置28にあり、dotPos = 28 となる

4段階: 固定テンプレートへの整列

  • 小数点位置を基準に入力を左へシフトし、常に同じテンプレートに合わせる
long alignedToTemplate = withSignRemoved << (28 - dotPos);
  • 目標のテンプレートは次の通り
0 0 0 Z . Y X 0
  • ここで X は十の位、Y は一の位、Z は小数第1位を表す
  • 0 はASCIIの "0" ではなく、値0のバイトを意味する
  • 符号除去後の入力は4つの配置のいずれかになり得る
    • 0 0 0 Z . Y X 0
    • 0 0 0 0 Z . Y 0
    • 0 0 0 0 Z . Y X
    • 0 0 0 0 0 Z . Y
  • -10.8 はすでに dotPos = 28 なのでシフト量は0である
  • -7.7 は小数点位置がビット20なので、8ビット、つまり1バイト左にシフトして X の位置に0が置かれる

5段階: ASCII数字を値に変換

  • 整列後はASCII文字から数字値だけを残す
long digits = alignedToTemplate & ASCII_TO_DIGIT_MASK;
  • ASCII数字の 0 から 9 は16進数で 0x30 から 0x39 である
  • 下位4ビットだけを残せば、文字コードがそのまま実際の数値になる
  • テンプレート内の数字位置にだけ F があるマスクを適用する
0 0 0 Z . Y X 0
000000F000F0F00
  • 例の -10.8 では、マスク適用後に Z=8Y=0X=1 を表す値だけが残る

6段階: マジック乗算で桁の重みを合算

  • 最終的な絶対値は 100 * X + 10 * Y + Z で計算する必要がある
  • 乗算がシフトと加算の組み合わせである性質を使い、複数桁の重み計算を1回の乗算で処理する
  • まず X + Y + Z を考えると、digits を0、16、24ビット位置にシフトした値を足すことで、特定のビット範囲に合計を集められる
  • このシフト加算の組み合わせは次のような乗算で表せる
0x1 + 0x10000 + 0x1000000
  • 実際には各桁の重みが異なるため、MAGIC_MULTIPLIER は次のように構成される
MAGIC_MULTIPLIER = 0x1 + 10 * 0x10000 + 100 * 0x1000000;
  • 計算式は次の通り
absValue = ((digits * MAGIC_MULTIPLIER) >>> 32) & 0x3FF;
  • 0x3FF は10ビット幅の結果だけを取り出すためのマスクである
  • 100 * X は10ビットまで大きくなって隣接ビットと重なる可能性があるが、Y * 100 の右2ビットが0になる性質により、必要なビット空間が確保される
  • merykittyはこの部分に // That was close :) というコメントを残している

7段階: 分岐なしで符号を適用

  • この時点で絶対値 absValue と符号情報 broadcastSign がある
  • broadcastSign は正数なら0、負数なら-1として動作する
  • 2の補数では負数は次の式で表される
-n = NOT(n) + 1
  • XORは条件付きNOTのように使える
    • n XOR -1NOT(n)
    • n XOR 0n
  • 条件付きの +1-broadcastSign で処理する
temperature = (absValue ^ broadcastSign) - broadcastSign;
  • 結果として、if なしで正数はそのまま、負数は2の補数の負値へ変換される

おまけ: 次のCSV行の開始位置の計算

  • 1BRC全体のソリューションでは、次のCSV行の開始位置も低コストで計算する必要がある
  • 小数点の後には常に小数1桁と改行が続くため、小数点位置を基準に次の行の開始位置を求める
  • dotPos はビット単位の位置なので、8で割るために3ビット右シフトを使う
nextLineStart = (dotPos >>> 3) + 3;
  • +3 は、小数点、小数1桁、改行の後の最初のバイトを指すための値である

結論

  • merykittyのSWARコードは、固定されたビット演算だけで4種類の温度文字列形式を統一してパースする
  • 核心は、ASCIIコードのビット特性、小数点位置ベースの整列、マスクによる数字抽出、乗算による桁重みの合算、2の補数に基づく符号適用にある
  • 段階ごとに分ければ動作を追えるが、これをオンラインチャレンジの数日以内に組み上げた点が印象的である

1件のコメント

 
GN⁺ 2024-03-11
Hacker News の意見
  • 段階的な説明が本当に素晴らしい
    2年以上前に、byte array view var handle が Java/Scala で効率的な SWAR ルーチンを作るのにかなり適していることに気づいた
    Base16/64 文字列のパース、java.time.*、数値をバイト配列から直接パースするなど、SWAR の活用例はここにも多い: https://github.com/plokhotnyuk/jsoniter-scala/blob/master/js...
  • 記事も良く、コードの文脈では優れた解法だが、この方式はデータが正しい形式であることを前提としている
    実戦で鍛えられたパーサーの大きな価値は、効率的なエラーチェックと復旧にある
    • 不正な入力が出力にどのように影響し得るのかを分解して見ると面白そうだ
      そして、現在のコードスタイルのように何らかの番兵エラー値を返すよう検出するには、どれほどの作業が必要なのかも気になる
      自分でやってみるほど興味があるわけではないが ;-)
  • 数値ビットフィールドに各桁の 10 の累乗を掛け、MUL でシフト/加算するテクニックはかなり知られた方法だ
    Lemire の記事を参照: https://lemire.me/blog/2023/11/28/parsing-8-bit-integers-qui...
  • 記事によると、SWAR は SIMD Within A Register のこと
  • こういう内容が好きなら、simdjson の論文も似たテクニックを使っていて、文章が非常によく書かれており、例も良い
    論文: https://arxiv.org/abs/1902.08318
    Github: https://github.com/simdjson/simdjson
    • これは SWAR ではないが、なぜ面白いかは分かる
  • BRC がなぜ入出力ボトルネックに引っかからないのか説明してもらえる? CPU がボトルネックだという点が理解できない
    • 最近のシステムでは、ローカルディスク I/O はもはやボトルネックではない: https://benhoyt.com/writings/io-is-no-longer-the-bottleneck/
      さらに公式 1BRC は、I/O 速度を完全に除外するために結果を RAM ディスク上で評価すると明記している: https://github.com/gunnarmorling/1brc?tab=readme-ov-file#eva...
      “Programs are run from a RAM disk (i.o. the IO overhead for loading the file from disk is not relevant)”
    • 背景として、Daniel Lemire へのインタビューがある。彼は入出力が常にボトルネックとは限らないという観察を土台にキャリア全体を築いてきた人だ: https://corecursive.com/frontiers-of-performance-with-daniel...
    • この問題を詳しく見たわけではないが、逆から考え始めることはできる。なぜメモリ I/Oがボトルネックだと思うのか?
      限られた理解では、大きなテキストファイルを順次 L1 に取り込み、各値ごとに一度読む。多くのプロセッサではこうした読み取りを 1 サイクルに 2 回できる。遅い部分は RAM から L1 に持ってくるところだろうが、シーケンシャル読み取りはかなり速い
      その後、各読み取りに対して処理を行う。ざっと見ると、最適化版ではおよそ 4 サイクル程度のように思える。その後、結果をどこかに書く必要があり、おそらくその前にランダム読み取りが 1 回か 2 回必要になる。この部分を I/O ボトルネックと見ているのか?
      CPU 制約が明白だという意味ではないが、そうでないことも明白には見えない
      追記: 「ディスク I/O」のことを指していた可能性を考慮していなかった。他の人たちが言っているように、ここでは実質的に要因ではない
    • テストは memfs で実行される。ファイルもすべて開始時点から RAM 上にある
    • データセットは Linux カーネルのページキャッシュに収まるほど小さく、ベンチマークは 5 回連続で繰り返されるため、1 回目の反復はディスク I/O がボトルネックになり得るが、残り 4 回はそうではない
      つまり、すべてのデータは RAM、より正確にはページキャッシュ上にある
  • 68000 では SWAR をかなり効果的に使っていた。1 命令で4バイトを並列処理していた
    記憶が正しければ、オーバーフロー処理が厄介だった。この記事は本当に気に入った
  • 「一人で働く人が、T シャツとコーヒーマグが報酬のオンラインチャレンジを数日軽くやりながら、これだけのものを作り上げたのが本当にミステリーだ」とあるが、なぜそれがミステリーなのか?
    いまだに CPU を実際にプログラムでき、自分が何をしているのか理解している人たちはいる
    本当のミステリーは、自分をプログラマーと呼ぶ大多数に深い理解が不足しており、しかも自分が深刻に不足していることにすら気づいていないように見える点にある
  • C# ではこうした SWAR の小技を使う必要はない。代わりに第一級のクロスプラットフォーム SIMD APIを提供している
    実際にうまく動作することは、これまで公開された 1BRC の中で最速に見える C# の解法から確認できる: https://hotforknowledge.com/2024/01/13/1brc-in-dotnet-among-...
  • これを SSE でベクトル化できるのか? 中核処理の大部分は 32 ビット整数 4 個のベクトルで可能に見える
    問題は、初期ベクトルを構成し、結果を取り出すコストが過大ではないかという点だ
    • 可能で、他の複数の 1BRC 実装もそうしていた
      ただし HotSpot が自力でやれるかは疑わしく、1BRC の提出物の多くが起動オーバーヘッドを減らすために Graal で実行されていた点も別にある
      基本の SSE2 には 32 ビットや 64 ビットの乗算がないため、32×32→64 ビット乗算が問題になるが、SSE4.1 にはまさに必要な pmuldq が追加されている。ただし結果が 64 ビットなので、32 ビット整数のベクトル全体を処理するにはこの演算が 2 回必要になる
    • 温度フィールドが名前フィールドと混在しているため、SSE で追加の利得を得るのは難しそうだ
      また温度フィールドは長さが可変なので、列単位で保存されていたとしても利得は出ない可能性が高い
      ただし、名前と温度の間の区切り文字探しには SSE がうまく適用されていた
    • こういうコードは、最初からであれ HotSpot がホットスポットを検出した後であれ、自動ベクトル化されそうだ