0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

Shrike-LiteでAtCoder問題を解く(16):ABC467C - Adjacent Sums (easy)(4組パック+Prefix XOR編)

0
Last updated at Posted at 2026-07-23

はじめに

前回は、RP2040側のSPI転送方法を見直しました。

結果として、FPGA側はNaive実装のままでも、18bit通信形式で扱える最大値であるN=262,143を約0.65秒で処理できました。

今回は、ABC467Cの仕上げに向けて、さらに通信パフォーマンスの向上を狙います。

ABC467Cでは、A_iB_iが0または1です。
そこで、

  • A_(i+1)B_iの1組を2bitで表す
  • 1byteに4組を詰める

という構成へ変更します。

FPGA側では一度に4組のデータを受け取ることになりますので、Naive実装から処理回路を変更します。

Naive実装では1組の受信ごとに処理を行っていましたが、1byteで受信する4組のPrefix XORを計算することにより、4組分の更新をまとめて処理できるようにします。

今回も記事とコードの草稿作成にはAIを使用しています。


前回までの構成

前回のNaive実装では、1組の入力を次の1byteで送信していました。

bit7:5 = SEND_PAIRコマンド
bit4:2 = 未使用
bit1   = A_(i+1)
bit0   = B_i

つまり、1byteを送っても、問題の入力として使っているのは2bitだけです。

また、FPGA側では1組を受信するたびに、現在値を更新し、不一致なら回答となる操作回数を1増やしていました。

RP2040側の転送呼び出しは前回かなり改善できましたが、SPI線上では依然として1組につき1byteを送っています。

そこで今回は、通信データを4分の1に圧縮します。


1byteに4組の入力を詰める

A_(i+1)B_iは、それぞれ1bitです。

1組は次の2bitで表せます。

bit1 = A_(i+1)
bit0 = B_i

これを1byteへ4組詰めます。

bit7:6 = {A_(i+1), B_i}
bit5:4 = {A_(i+2), B_(i+1)}
bit3:2 = {A_(i+3), B_(i+2)}
bit1:0 = {A_(i+4), B_(i+3)}

上位側から順に処理します。

たとえば、4組の入力が次の場合を考えます。

(A_(i+1), B_i)     = (1, 0)
(A_(i+2), B_(i+1)) = (0, 1)
(A_(i+3), B_(i+2)) = (1, 1)
(A_(i+4), B_(i+3)) = (0, 0)

送信byteは次のようになります。

10 01 11 00

16進数では0x9Cです。

ヘッダとデータ部を分ける

最初の5byteは、これまでと同様にコマンド付きのヘッダとして送信します。

1byte目:N[17:15]
2byte目:N[14:10]
3byte目:N[9:5]
4byte目:N[4:0]
5byte目:A_1

A_1を受信した後はデータ受信モードへ入り、それ以降のbyteはすべて4組パックされたデータとして扱います。

そのため、データ部にはコマンドbitを持たせません。

送信するデータbyte数は次の通りです。

ceil((N - 1) / 4)

入力stream全体では、

5 + ceil((N - 1) / 4) byte

になります。

Naive版ではN+4 byteだったので、大きなNではほぼ4分の1です。


Prefix XORで4組をまとめて処理する

4組を1byteに詰めても、FPGA内部で1組ずつ4クロックかけて処理すると、回路制御が少し複雑になります。

そこで、今回は4組分を組合せ回路でまとめて計算できるようにします。

ABC467Cの解法は、B_1からB_iまでの累積XORを計算し、A_(i+1)と比較するというものでした。

ここで、この4組を処理する直前までのBの累積XOR値をx、今回受信した4個のBb0からb3、対応する4個のAa0からa3とします。

今回受信した4個のBについて、各位置までのPrefix XORは次のようになります。

p0 = b0
p1 = b0 XOR b1
p2 = b0 XOR b1 XOR b2
p3 = b0 XOR b1 XOR b2 XOR b3

直前までの累積XOR値xを反映すると、4組を順番に処理したときの各位置での累積XOR値は、

x0 = x XOR p0
x1 = x XOR p1
x2 = x XOR p2
x3 = x XOR p3

となります。

このp0からp3が、今回受信した4個のBに対するPrefix XORです。

x0からx3を、同じ位置のAと比較します。

m0 = (x0 != a0)
m1 = (x1 != a1)
m2 = (x2 != a2)
m3 = (x3 != a3)

最後に、4個の不一致bitを加算します。

package_cost = m0 + m1 + m2 + m3

package_costをこれまでの操作回数に加算し、最後の累積XOR値であるx3を、次の4組を処理するときのxとして引き継ぎます。

これにより、1byte分のデータを受信するたびに、最大4組分の不一致判定と操作回数の加算をまとめて行えます。

2候補のうち片方だけ数える

ABC467CではM=2なので、求める列の先頭値を0と仮定した候補と、1と仮定した候補では、各位置の値が常に反転関係になります。

ある位置で片方がA_iと一致すれば、もう片方は不一致です。

したがって、候補0の不一致数をcost0とすると、候補1の不一致数は次のように求められます。

cost1 = N - cost0

最終的な答えは、

min(cost0, N - cost0)

です。

Naive実装では、候補0と候補1の回路を並列にもたせて、同時に計算していました。

今回の実装では少し回路が大きくなりそうなので、2候補をそれぞれ数えるのではなく、片方の不一致数だけを保持することにします。


最後の1byteに含まれる端数を処理する

N-1が4の倍数とは限りません。

最後の1byteには、1組から3組しか有効データがない場合があります。

RP2040側では、存在しない組を0で埋めて送信します。

FPGA側では、処理済み組数とN-1を比較して、最後のbyteで有効な組数を求めます。

残り1組:先頭の1組だけ有効
残り2組:先頭の2組だけ有効
残り3組:先頭の3組だけ有効
残り4組以上:4組すべて有効

加算時には、無効な位置の不一致bitを0として扱います。

これにより、末尾を0埋めしても答えへ影響しません。


処理結果を返信する形式

答えは従来と同じ3byteで返信します。

1byte目には、答えが有効であることを示すVALIDに加えて、処理が想定通りに完了したことを示すCOUNT_OKを持たせます。

1byte目
bit7   = VALID
bit6   = COUNT_OK
bit5:2 = 予約
bit1:0 = ANSWER[17:16]

2byte目
bit7:0 = ANSWER[15:8]

3byte目
bit7:0 = ANSWER[7:0]

COUNT_OKは、最終データbyteを処理した時点で、次の条件を満たしたことを表します。

処理した組数       = N - 1
受信データbyte数   = ceil((N - 1) / 4)

FPGA内部で異常を検出した場合も、COUNT_OK=0として返信します。

これにより、答えに加えて、想定した組数とデータbyte数で処理を完了できたことも確認できます。


AIへVerilog実装を依頼する

今回は、前回のNaive版を元にして、次の内容でAIへ実装を依頼しました。

Shrike-LiteでABC467C - Adjacent Sums (easy)を高速化するため、
既存のNaive版main.vを元に、4組パック+Prefix XOR版を実装してください。

プロジェクト名とbitstream名はabc467c_prefix_xor_burstです。

spi_target.vは変更しないでください。
SystemVerilog固有構文は使用せず、Verilogで実装してください。
既存の日本語コメントは、変更に直接必要な箇所以外は書き換えないでください。

NとA_1を送る最初の5byteは、Naive版と同じコマンド形式を使用してください。

候補0では先頭の値を0と仮定します。
A_1を受信した時点で、候補0の不一致数をA_1の値で初期化してください。
Bの累積XOR値は0で初期化してください。

A_1受信後はデータ受信モードへ入り、
以降の各byteを4組の入力データとして処理してください。

各byteは次の形式です。

- bit7:6 = {A_(i+1), B_i}
- bit5:4 = {A_(i+2), B_(i+1)}
- bit3:2 = {A_(i+3), B_(i+2)}
- bit1:0 = {A_(i+4), B_(i+3)}

現在のBの累積XOR値と、受信した4個のBから、
4位置分のPrefix XORを組合せ回路で生成してください。

各位置の値を対応するAと比較し、
1byte分の不一致数を0から4で求めてください。

最後のデータbyteでは、
残っている有効組数に対応する不一致だけを加算してください。
存在しない組はRP2040側で0埋めしますが、答えには加算しないでください。

候補0の不一致数だけを保持し、最終回答は

min(cost0, N - cost0)

で求めてください。

受信組数は18bitで数えてください。
受信データbyte数は17bitで数えてください。

最終データbyteを処理するときは、
今回の不一致数、受信組数、受信データbyte数を加えた
更新後の値を使って、最終回答とCOUNT_OKを計算してください。

MISOは3byteで返信してください。

1byte目:
- bit7 = VALID
- bit6 = COUNT_OK
- bit5:2 = 0
- bit1:0 = ANSWER[17:16]

2byte目:
- bit7:0 = ANSWER[15:8]

3byte目:
- bit7:0 = ANSWER[7:0]

COUNT_OKは、最終データbyteを含めた処理後に、
受信組数がN-1、
受信データbyte数がceil((N-1)/4)と一致し、
内部処理で異常を検出していない場合に1としてください。

COUNT_OKと既存の異常検出処理は、
4組パック対応に必要な範囲以外では変更しないでください。

RESETでは、問題固有レジスタ、データ受信モード、
各カウンタ、回答レジスタ、tx_dataを初期化してください。

Verilogコード(問題固有部分)

次のようなコードが提示されました。

    // ===== 問題ごとに変更する部分 =====
    // ABC467C - Adjacent Sums (easy)
    // M=2なので、求める列の先頭値を0と仮定した候補をPrefix XORで計算する。
    // もう一方の候補は全位置の一致・不一致が反転することを利用して求める。
    localparam NOP          = 3'b000;
    localparam SEND_N_17_15 = 3'b001;
    localparam SEND_N_14_10 = 3'b010;
    localparam SEND_N_9_5   = 3'b011;
    localparam SEND_N_4_0   = 3'b100;
    localparam SEND_A1      = 3'b101;
    localparam RESET        = 3'b111;

    reg [17:0] n_value;
    reg [17:0] pair_count;
    reg [16:0] package_count;

    reg        value_reg;
    reg [17:0] answer_count;

    reg [17:0] answer_reg;
    reg [1:0]  reply_index;
    reg        answer_ready;
    reg        count_ok;
    reg        protocol_error;
    reg        stream_active;

    // 受信byteの上位側から順に、4組のA/Bを取り出す。
    wire a0;
    wire b0;
    wire a1;
    wire b1;
    wire a2;
    wire b2;
    wire a3;
    wire b3;

    wire value_next0;
    wire value_next1;
    wire value_next2;
    wire value_next3;

    wire diff0_raw;
    wire diff1_raw;
    wire diff2_raw;
    wire diff3_raw;
    wire diff0;
    wire diff1;
    wire diff2;
    wire diff3;
    wire [1:0] sum01;
    wire [1:0] sum23;
    wire [2:0] package_cost;

    wire [17:0] total_pair_count;
    wire [17:0] remaining_pairs;
    wire [2:0] valid_count;
    wire [17:0] pair_count_next;
    wire [16:0] package_count_next;
    wire [18:0] expected_package_count_wide;
    wire [16:0] expected_package_count;
    wire        last_value_next;
    wire [17:0] answer0_next;
    wire [17:0] answer1;
    wire [17:0] next_answer;
    wire        final_package;
    wire        count_match;

    assign a0 = rx_data[7];
    assign b0 = rx_data[6];
    assign a1 = rx_data[5];
    assign b1 = rx_data[4];
    assign a2 = rx_data[3];
    assign b2 = rx_data[2];
    assign a3 = rx_data[1];
    assign b3 = rx_data[0];

    // M=2では加算剰余はXORと同じ。
    assign value_next0 = value_reg ^ b0;
    assign value_next1 = value_next0 ^ b1;
    assign value_next2 = value_next1 ^ b2;
    assign value_next3 = value_next2 ^ b3;

    assign diff0_raw = a0 ^ value_next0;
    assign diff1_raw = a1 ^ value_next1;
    assign diff2_raw = a2 ^ value_next2;
    assign diff3_raw = a3 ^ value_next3;

    assign total_pair_count =
        (n_value >= 18'd1) ? (n_value - 18'd1) : 18'd0;
    assign remaining_pairs =
        (pair_count < total_pair_count) ?
        (total_pair_count - pair_count) : 18'd0;

    assign valid_count =
        (remaining_pairs >= 18'd4) ? 3'd4 :
        (remaining_pairs == 18'd3) ? 3'd3 :
        (remaining_pairs == 18'd2) ? 3'd2 :
        (remaining_pairs == 18'd1) ? 3'd1 : 3'd0;

    // 最終パッケージの無効位置は加算しない。
    assign diff0 = (valid_count >= 3'd1) ? diff0_raw : 1'b0;
    assign diff1 = (valid_count >= 3'd2) ? diff1_raw : 1'b0;
    assign diff2 = (valid_count >= 3'd3) ? diff2_raw : 1'b0;
    assign diff3 = (valid_count >= 3'd4) ? diff3_raw : 1'b0;

    // 4個の不一致を2段の加算木で集計する。
    assign sum01 = {1'b0, diff0} + {1'b0, diff1};
    assign sum23 = {1'b0, diff2} + {1'b0, diff3};
    assign package_cost = {1'b0, sum01} + {1'b0, sum23};

    assign pair_count_next =
        pair_count + {15'd0, valid_count};
    assign package_count_next = package_count + 17'd1;

    // 最大Nでも途中の加算があふれないよう19bitで計算する。
    assign expected_package_count_wide =
        (n_value >= 18'd2) ?
        (({1'b0, n_value} - 19'd1 + 19'd3) >> 2) : 19'd0;
    assign expected_package_count =
        expected_package_count_wide[16:0];

    assign last_value_next =
        (valid_count == 3'd1) ? value_next0 :
        (valid_count == 3'd2) ? value_next1 :
        (valid_count == 3'd3) ? value_next2 :
        (valid_count == 3'd4) ? value_next3 : value_reg;

    assign answer0_next =
        answer_count + {15'd0, package_cost};
    assign answer1 = n_value - answer0_next;
    assign next_answer =
        (answer0_next <= answer1) ? answer0_next : answer1;

    // 現在のパッケージを反映したカウンタで最終判定する。
    assign final_package =
        (pair_count_next == total_pair_count);
    assign count_match =
        (pair_count_next == total_pair_count) &&
        (package_count_next == expected_package_count);

    always @(posedge clk or negedge rst_n) begin
        if (!rst_n) begin
            n_value        <= 18'd0;
            pair_count     <= 18'd0;
            package_count  <= 17'd0;
            value_reg      <= 1'b0;
            answer_count   <= 18'd0;
            answer_reg     <= 18'd0;
            reply_index    <= 2'd0;
            answer_ready   <= 1'b0;
            count_ok       <= 1'b0;
            protocol_error <= 1'b0;
            stream_active  <= 1'b0;
            tx_data        <= 8'h00;
        end else if (rx_data_strobe) begin
            if (stream_active) begin
                if (pair_count > total_pair_count) begin
                    count_ok       <= 1'b0;
                    protocol_error <= 1'b1;
                    stream_active  <= 1'b0;
                    tx_data        <= 8'h00;
                end else if (valid_count == 3'd0) begin
                    count_ok       <= 1'b0;
                    protocol_error <= 1'b1;
                    stream_active  <= 1'b0;
                    tx_data        <= 8'h00;
                end else if (pair_count_next > total_pair_count) begin
                    count_ok       <= 1'b0;
                    protocol_error <= 1'b1;
                    stream_active  <= 1'b0;
                    tx_data        <= 8'h00;
                end else begin
                    answer_count  <= answer0_next;
                    pair_count    <= pair_count_next;
                    package_count <= package_count_next;
                    value_reg     <= last_value_next;

                    if (final_package) begin
                        answer_reg     <= next_answer;
                        reply_index    <= 2'd0;
                        answer_ready   <= 1'b1;
                        count_ok       <= count_match;
                        protocol_error <= protocol_error | ~count_match;
                        stream_active  <= 1'b0;
                        tx_data        <= {
                            1'b1,
                            count_match & ~protocol_error,
                            4'b0000,
                            next_answer[17:16]
                        };
                    end
                end
            end else begin
                case (rx_data[7:5])
                    NOP: begin
                        // 答えは3byteで返す。
                        // 1byte目: VALID、COUNT_OK、予約、ANSWER[17:16]
                        // 2byte目: ANSWER[15:8]
                        // 3byte目: ANSWER[7:0]
                        if (answer_ready) begin
                            case (reply_index)
                                2'd0: begin
                                    tx_data     <= answer_reg[15:8];
                                    reply_index <= 2'd1;
                                end
                                2'd1: begin
                                    tx_data     <= answer_reg[7:0];
                                    reply_index <= 2'd2;
                                end
                                default: begin
                                    tx_data      <= 8'h00;
                                    reply_index  <= 2'd0;
                                    answer_ready <= 1'b0;
                                end
                            endcase
                        end else begin
                            tx_data <= 8'h00;
                        end
                    end

                    SEND_N_17_15: begin
                        n_value[17:15] <= rx_data[2:0];
                    end

                    SEND_N_14_10: begin
                        n_value[14:10] <= rx_data[4:0];
                    end

                    SEND_N_9_5: begin
                        n_value[9:5] <= rx_data[4:0];
                    end

                    SEND_N_4_0: begin
                        n_value[4:0] <= rx_data[4:0];
                    end

                    SEND_A1: begin
                        // 求める列の先頭値を0と仮定する候補を開始する。
                        value_reg      <= 1'b0;
                        answer_count   <= rx_data[0] ? 18'd1 : 18'd0;
                        pair_count     <= 18'd0;
                        package_count  <= 17'd0;
                        answer_reg     <= 18'd0;
                        reply_index    <= 2'd0;
                        answer_ready   <= 1'b0;
                        count_ok       <= 1'b0;
                        protocol_error <= (n_value < 18'd2);
                        stream_active  <= (n_value >= 18'd2);
                        tx_data        <= 8'h00;
                    end

                    RESET: begin
                        n_value        <= 18'd0;
                        pair_count     <= 18'd0;
                        package_count  <= 17'd0;
                        value_reg      <= 1'b0;
                        answer_count   <= 18'd0;
                        answer_reg     <= 18'd0;
                        reply_index    <= 2'd0;
                        answer_ready   <= 1'b0;
                        count_ok       <= 1'b0;
                        protocol_error <= 1'b0;
                        stream_active  <= 1'b0;
                        tx_data        <= 8'h00;
                    end

                    default: begin
                        // 3'b110は通常コマンドとして使用しない。
                    end
                endcase
            end
        end
    end

Verilogコードの合成結果

提示されたVerilogコードを合成し、bitstreamを生成します。

リソース使用量は次の通りでした。

Shrike-LiteでAtCoder問題を解く(16)_001.png

LUT使用率は4割程度ですが、CLB使用率は7割を超えました。

4組分のPrefix XOR、比較、加算回路を並列に置いたため、Naive版より回路規模は大きくなっています。

それでもShrike-Liteへ収まり、bitstream生成まで完了しました。


MicroPython側も4組パックへ変更する

RP2040側では、4組分の入力を1byteへ詰めます。

主要部分は次のようになります。

pair_index = 0

while pair_index < n - 1:
    packed_data = 0

    # 上位側から順に、各2bitをA、Bの順で格納する。
    for position in range(4):
        if pair_index < n - 1:
            pair_data = (
                ((a_values[pair_index + 1] & 0x01) << 1)
                | (b_values[pair_index] & 0x01)
            )
            shift = 6 - (position * 2)
            packed_data |= pair_data << shift
            pair_index += 1

    tx_stream[index] = packed_data
    index += 1

作成したstreamは、前回と同じく最大256byte単位へ分割します。

def send_packages(packages):
    for tx_package, rx_package in packages:
        spi_transfer(tx_package, rx_package)

SPI設定は前回の試験結果を引き継ぎます。

RP2040 CPUクロック  125MHz
SPIクロック         4MHz
SPI Mode           Mode 0
転送単位            最大256byte

AIへMicroPython実装を依頼する

前節で整理した通信形式に基づき、RP2040側のMicroPythonコードもAIへ作成を依頼しました。

Shrike-LiteでABC467C - Adjacent Sums (easy)の
4組パック+Prefix XOR版を実機テストするため、
既存のNaive版MicroPythonコードを元に、
abc467c_prefix_xor_burst_test_estimate.pyを作成してください。

開発環境はThonnyとMicroPythonです。

bitstream名はabc467c_prefix_xor_burst.binです。

RP2040とShrike-Liteの接続、FPGAへのbitstream書き込み、
SPI初期化、RESET処理などの共通部分は、
既存コードの構成を引き継いでください。

SPI設定は次の通りです。

- RP2040 CPUクロック:125MHz
- SPIクロック:4MHz
- SPI Mode 0
- 1回の転送サイズ:最大256byte
- CSはwrite_readinto()の呼び出し中だけLowにする

NとA_1は、最初の5byteで送信してください。

1byte目:N[17:15]
2byte目:N[14:10]
3byte目:N[9:5]
4byte目:N[4:0]
5byte目:A_1

A_1より後のデータは、
A_(i+1)とB_iの1組を2bitで表し、
1byteへ4組ずつ格納してください。

各byteの形式は次の通りです。

- bit7:6 = {A_(i+1), B_i}
- bit5:4 = {A_(i+2), B_(i+1)}
- bit3:2 = {A_(i+3), B_(i+2)}
- bit1:0 = {A_(i+4), B_(i+3)}

最後のbyteに4組未満しかデータがない場合は、
存在しない組を0で埋めてください。

作成した送信streamは最大256byte単位へ分割し、
write_readinto()でまとめて送信してください。

SPI転送に使用するバッファは可能な範囲で再利用し、
測定中には不要なbytearray生成やGCを行わないでください。

FPGAからは3byteで結果を受信してください。

1byte目:
- bit7 = VALID
- bit6 = COUNT_OK
- bit1:0 = ANSWER[17:16]

2byte目:
- bit7:0 = ANSWER[15:8]

3byte目:
- bit7:0 = ANSWER[7:0]

MicroPython側にもABC467Cの参照計算を実装し、
FPGAから受信した答えと比較してください。

機能テストには、次の内容を含めてください。

- 公式サンプル
- すべて0の入力
- 最後のbyteの有効組数が1~4組となるケース
- 送信streamが256byte境界の前後となるケース
- 固定seedで生成する再現可能な疑似乱数ケース

各テストについて、テスト名、N、送信byte数、
SPI転送回数、受信データ、VALID、COUNT_OK、
期待値、受信した答え、PASSまたはFAIL、処理時間を
1行で表示してください。

すべての機能テスト終了後に、
PASS数、FAIL数、テスト総数を表示してください。

さらに、AとBをすべて0としたベンチマークを実装してください。

現在の18bit通信形式で扱えるNの最大値は、
262,143とします。

指数探索と二分探索によって、
2秒以内に処理できるNを調べてください。

18bit通信上限まで2秒以内だった場合は、
通信形式の上限へ到達したことを表示してください。

MicroPythonの実行時間にはばらつきがあるため、
探索結果は厳密な最大値ではなく、
おおよその目安として扱ってください。

コードには処理内容が分かる日本語コメントを入れてください。

MicroPythonコード(全文)

次のコードが提示されました。

from machine import Pin, SPI
import gc
import time
import shrike


# ===== 共通部分:bitstream名とShrike-Liteのピン設定 =====
BITSTREAM = "abc467c_prefix_xor_burst.bin"

SCK = 2
CS = 1
MOSI = 3
MISO = 0
FPGA_RESET = 14

# ===== SPI転送設定 =====
SPI_BAUDRATE = 4_000_000
PACKAGE_SIZE = 256

# ===== ベンチマーク設定 =====
TIME_LIMIT_US = 2_000_000
MIN_N = 2

# 現在の通信形式では、Nを18bitで送信する。
# そのため、FPGA側を変更せずに扱える最大値は2^18-1。
PROTOCOL_MAX_N = (1 << 18) - 1
MAX_N = PROTOCOL_MAX_N

INITIAL_N = 1_024
ESTIMATE_UNIT = 100


# ===== 共通部分:FPGAへのbitstream書き込みとリセット =====
shrike.reset()
shrike.flash(BITSTREAM)

reset_pin = Pin(FPGA_RESET, Pin.OUT, value=1)
reset_pin.value(0)
time.sleep_ms(100)
reset_pin.value(1)
time.sleep_ms(100)


# ===== 共通部分:SPI Masterの初期化 =====
cs = Pin(CS, Pin.OUT, value=1)

spi = SPI(
    0,
    baudrate=SPI_BAUDRATE,
    polarity=0,
    phase=0,
    bits=8,
    firstbit=SPI.MSB,
    sck=Pin(SCK),
    mosi=Pin(MOSI),
    miso=Pin(MISO)
)


# ===== ABC467C固有処理 =====
NOP = 0b000
SEND_N_17_15 = 0b001
SEND_N_14_10 = 0b010
SEND_N_9_5 = 0b011
SEND_N_4_0 = 0b100
SEND_A1 = 0b101
RESET = 0b111


def make_command(command, data=0):
    return (command << 5) | (data & 0x1F)


NOP_BYTE = make_command(NOP)
RESET_BYTE = make_command(RESET)
PACKED_ZERO = 0x00


# ===== 再利用するSPI送受信バッファ =====

# RESETと答え受信に使用する1byteバッファ
single_tx = bytearray(1)
single_rx = bytearray(1)

# NとA_1を送る5byteヘッダ
header_tx = bytearray(5)
header_rx = bytearray(5)

# ベンチマーク用の4組ゼロデータパッケージ
packed_zero_tx = bytearray(PACKAGE_SIZE)
packed_zero_rx = bytearray(PACKAGE_SIZE)

for i in range(PACKAGE_SIZE):
    packed_zero_tx[i] = PACKED_ZERO


# ===== SPI送受信 =====

def spi_transfer(tx_buffer, rx_buffer):
    # 複数byteを一度のwrite_readinto()で転送する。
    cs.value(0)
    spi.write_readinto(tx_buffer, rx_buffer)
    cs.value(1)


def spi_exchange_1byte(value):
    # RESETと答え受信では、再利用する1byteバッファを使用する。
    single_tx[0] = value
    spi_transfer(single_tx, single_rx)
    return single_rx[0]


# ===== ABC467C固有の通信処理 =====

def set_header(n, a_first):
    if n < MIN_N or n > PROTOCOL_MAX_N:
        raise ValueError("N is outside the 18bit protocol range")

    header_tx[0] = make_command(
        SEND_N_17_15,
        (n >> 15) & 0x07
    )
    header_tx[1] = make_command(
        SEND_N_14_10,
        (n >> 10) & 0x1F
    )
    header_tx[2] = make_command(
        SEND_N_9_5,
        (n >> 5) & 0x1F
    )
    header_tx[3] = make_command(
        SEND_N_4_0,
        n & 0x1F
    )
    header_tx[4] = make_command(
        SEND_A1,
        a_first
    )


def reset_problem():
    # RESETの処理結果を次のNOPでSPI送信側へ反映させる。
    spi_exchange_1byte(RESET_BYTE)
    spi_exchange_1byte(NOP_BYTE)


def receive_answer():
    rx_hi = spi_exchange_1byte(NOP_BYTE)
    rx_mid = spi_exchange_1byte(NOP_BYTE)
    rx_lo = spi_exchange_1byte(NOP_BYTE)

    valid = (rx_hi >> 7) & 0x01
    count_ok = (rx_hi >> 6) & 0x01
    answer = (
        ((rx_hi & 0x03) << 16)
        | (rx_mid << 8)
        | rx_lo
    )

    return (
        valid,
        count_ok,
        answer,
        (rx_hi, rx_mid, rx_lo)
    )


# ===== 機能テスト用の入力stream作成 =====

def build_input_stream(a_values, b_values):
    n = len(a_values)

    if n < MIN_N:
        raise ValueError("N must be at least 2")

    if n > PROTOCOL_MAX_N:
        raise ValueError("N exceeds the 18bit protocol range")

    if len(b_values) != n - 1:
        raise ValueError("len(B) must be N - 1")

    # N送信4byte、A_1送信1byte、4組データ送信
    data_byte_count = (n - 1 + 3) // 4
    tx_stream = bytearray(5 + data_byte_count)
    index = 0

    tx_stream[index] = make_command(
        SEND_N_17_15,
        (n >> 15) & 0x07
    )
    index += 1

    tx_stream[index] = make_command(
        SEND_N_14_10,
        (n >> 10) & 0x1F
    )
    index += 1

    tx_stream[index] = make_command(
        SEND_N_9_5,
        (n >> 5) & 0x1F
    )
    index += 1

    tx_stream[index] = make_command(
        SEND_N_4_0,
        n & 0x1F
    )
    index += 1

    tx_stream[index] = make_command(
        SEND_A1,
        a_values[0]
    )
    index += 1

    pair_index = 0

    while pair_index < n - 1:
        packed_data = 0

        # 上位側から順に、各2bitをA、Bの順で格納する。
        for position in range(4):
            if pair_index < n - 1:
                pair_data = (
                    ((a_values[pair_index + 1] & 0x01) << 1)
                    | (b_values[pair_index] & 0x01)
                )
                shift = 6 - (position * 2)
                packed_data |= pair_data << shift
                pair_index += 1

        tx_stream[index] = packed_data
        index += 1

    return tx_stream


def make_packages(tx_stream):
    # パッケージ生成は測定前に行う。
    packages = []
    start = 0
    stream_length = len(tx_stream)

    while start < stream_length:
        end = start + PACKAGE_SIZE

        if end > stream_length:
            end = stream_length

        tx_package = tx_stream[start:end]
        rx_package = bytearray(len(tx_package))
        packages.append((tx_package, rx_package))

        start = end

    return packages


def send_packages(packages):
    for tx_package, rx_package in packages:
        spi_transfer(tx_package, rx_package)


def run_test_case(name, a_values, b_values):
    n = len(a_values)
    expected = solve_reference(a_values, b_values)

    # 入力streamとパッケージは測定前に作る。
    tx_stream = build_input_stream(
        a_values,
        b_values
    )
    packages = make_packages(tx_stream)

    reset_problem()

    # GCの実行時間はTIME_USに含めない。
    gc.collect()

    start_us = time.ticks_us()

    send_packages(packages)
    (
        valid,
        count_ok,
        result,
        rx_bytes
    ) = receive_answer()

    elapsed_us = time.ticks_diff(
        time.ticks_us(),
        start_us
    )

    passed = (
        valid == 1
        and count_ok == 1
        and result == expected
    )
    status = "PASS" if passed else "FAIL"

    print(
        "NAME={} N={} STREAM_BYTES={} PACKAGES={} "
        "RX=[0x{:02X},0x{:02X},0x{:02X}] "
        "VALID={} COUNT_OK={} "
        "EXPECT={} RESULT={} {} TIME_US={}".format(
            name,
            n,
            len(tx_stream),
            len(packages),
            rx_bytes[0],
            rx_bytes[1],
            rx_bytes[2],
            valid,
            count_ok,
            expected,
            result,
            status,
            elapsed_us
        )
    )

    return passed


# ===== 参照用のABC467C計算 =====

def solve_reference(a_values, b_values):
    value0 = 0
    value1 = 1

    cost0 = 1 if value0 != a_values[0] else 0
    cost1 = 1 if value1 != a_values[0] else 0

    for i in range(len(b_values)):
        value0 ^= b_values[i]
        value1 ^= b_values[i]

        if value0 != a_values[i + 1]:
            cost0 += 1

        if value1 != a_values[i + 1]:
            cost1 += 1

    return cost0 if cost0 < cost1 else cost1


def make_boundary_test_case(n):
    # 256byte送信境界を確認する機能テストを作る。
    a_values = [0] * n
    b_values = [0] * (n - 1)

    for i in range(n):
        a_values[i] = (
            i
            ^ (i >> 2)
            ^ (i >> 5)
        ) & 0x01

    for i in range(n - 1):
        b_values[i] = (
            (i * 3)
            ^ (i >> 1)
            ^ 1
        ) & 0x01

    return (
        "package_boundary_n{}".format(n),
        a_values,
        b_values
    )


def make_pseudo_random_test_case(n, seed):
    # 固定seedの疑似乱数で再現可能な機能テストを作る。
    a_values = [0] * n
    b_values = [0] * (n - 1)
    state = seed & 0x7FFFFFFF

    for i in range(n):
        state = (
            (state * 1103515245 + 12345)
            & 0x7FFFFFFF
        )
        a_values[i] = (state >> 15) & 0x01

    for i in range(n - 1):
        state = (
            (state * 1103515245 + 12345)
            & 0x7FFFFFFF
        )
        b_values[i] = (state >> 15) & 0x01

    return (
        "pseudo_random_n{}".format(n),
        a_values,
        b_values
    )


# ===== ベンチマーク用のゼロ入力転送 =====

def prepare_zero_case(n):
    # 測定中にバッファを生成しないよう、
    # ヘッダと最終パッケージを測定前に準備する。
    set_header(n, 0)

    data_byte_count = (n - 1 + 3) // 4
    full_package_count = data_byte_count // PACKAGE_SIZE
    tail_count = data_byte_count % PACKAGE_SIZE

    if tail_count == 0:
        tail_tx = None
        tail_rx = None
    else:
        tail_tx = bytearray(tail_count)
        tail_rx = bytearray(tail_count)

        for i in range(tail_count):
            tail_tx[i] = PACKED_ZERO

    return (
        full_package_count,
        tail_count,
        tail_tx,
        tail_rx
    )


def send_zero_case(
    full_package_count,
    tail_count,
    tail_tx,
    tail_rx
):
    # NとA_1を5byteまとめて送信する。
    spi_transfer(header_tx, header_rx)

    # 4組ゼロデータを256byte単位で繰り返し送信する。
    for _ in range(full_package_count):
        spi_transfer(packed_zero_tx, packed_zero_rx)

    # 最後の端数だけ短いパッケージで送信する。
    if tail_count != 0:
        spi_transfer(tail_tx, tail_rx)


def measure_zero_case(n, label="SEARCH"):
    (
        full_package_count,
        tail_count,
        tail_tx,
        tail_rx
    ) = prepare_zero_case(n)

    reset_problem()

    # 各測定の直前にGCを実行する。
    # バッファ生成時間とGC時間はTIME_USに含めない。
    gc.collect()

    start_us = time.ticks_us()

    send_zero_case(
        full_package_count,
        tail_count,
        tail_tx,
        tail_rx
    )

    (
        valid,
        count_ok,
        answer,
        _
    ) = receive_answer()

    elapsed_us = time.ticks_diff(
        time.ticks_us(),
        start_us
    )

    correct = (
        valid == 1
        and count_ok == 1
        and answer == 0
    )
    within_limit = (
        correct
        and elapsed_us <= TIME_LIMIT_US
    )

    package_count = (
        1
        + full_package_count
        + (1 if tail_count != 0 else 0)
    )

    print(
        "{} N={} PACKAGES={} TIME_US={} VALID={} COUNT_OK={} "
        "ANSWER={} {}".format(
            label,
            n,
            package_count,
            elapsed_us,
            valid,
            count_ok,
            answer,
            "PASS" if within_limit else "FAIL"
        )
    )

    return within_limit, elapsed_us, correct


# ===== 指数探索と二分探索 =====

def find_upper_bound():
    # まず指数探索で2秒を超える最初のNを探す。
    # 現在の18bit通信形式で扱える最大値まで探索する。
    n = INITIAL_N
    last_pass = MIN_N

    while True:
        if n > MAX_N:
            n = MAX_N

        passed, _, correct = measure_zero_case(
            n,
            "EXPAND"
        )

        if not correct:
            raise RuntimeError(
                "FPGA reply error during benchmark"
            )

        if not passed:
            return last_pass, n

        last_pass = n

        if n == MAX_N:
            # 2秒境界へ到達する前に18bit上限へ到達した。
            return MAX_N, MAX_N

        n *= 2


def binary_search_limit(low_pass, high_fail):
    if low_pass == MAX_N:
        return MAX_N, None

    low = low_pass
    high = high_fail

    # 実測値にはばらつきがあるため、厳密な最大値ではなく
    # おおよその2秒境界を得る目的で二分探索する。
    while high - low > 1:
        mid = (low + high) // 2

        passed, _, correct = measure_zero_case(
            mid,
            "BINARY"
        )

        if not correct:
            raise RuntimeError(
                "FPGA reply error during benchmark"
            )

        if passed:
            low = mid
        else:
            high = mid

    return low, high


def round_to_estimate_unit(value):
    estimated = (
        (value + (ESTIMATE_UNIT // 2))
        // ESTIMATE_UNIT
    ) * ESTIMATE_UNIT

    if estimated < MIN_N:
        return MIN_N

    if estimated > MAX_N:
        return MAX_N

    return estimated


# ===== 機能テスト =====

TEST_CASES = [
    (
        "official_sample_1",
        [1, 1, 1],
        [1, 1]
    ),
    (
        "official_sample_2",
        [1, 1],
        [0]
    ),
    (
        "official_sample_3",
        [0, 0, 0, 1, 1, 0, 1, 0, 1, 0],
        [0, 1, 0, 1, 0, 1, 0, 1, 0]
    ),
    (
        "all_zero",
        [0, 0, 0, 0, 0],
        [0, 0, 0, 0]
    ),
    (
        "two_elements_mismatch",
        [0, 1],
        [0]
    ),
    (
        "tail_valid_2",
        [0, 1, 0],
        [1, 0]
    ),
    (
        "tail_valid_3",
        [0, 1, 1, 1],
        [1, 0, 1]
    ),
    (
        "tail_valid_4",
        [0, 1, 1, 0, 0],
        [1, 0, 1, 1]
    ),
    (
        "second_package_tail_valid_1",
        [0, 1, 1, 0, 1, 0],
        [1, 0, 1, 1, 0]
    ),
    make_boundary_test_case(1005),
    make_boundary_test_case(1006),
    make_pseudo_random_test_case(73, 0x467C),
]


print(
    "CONFIG "
    "SPI_BAUDRATE={} "
    "PACKAGE_SIZE={} "
    "TIME_LIMIT_US={} "
    "MAX_N={} "
    "PROTOCOL_BITS=18".format(
        SPI_BAUDRATE,
        PACKAGE_SIZE,
        TIME_LIMIT_US,
        MAX_N
    )
)

reset_problem()

pass_count = 0

for name, a_values, b_values in TEST_CASES:
    if run_test_case(
        name,
        a_values,
        b_values
    ):
        pass_count += 1

fail_count = len(TEST_CASES) - pass_count

print(
    "FUNCTION_SUMMARY PASS={} FAIL={} TOTAL={}".format(
        pass_count,
        fail_count,
        len(TEST_CASES)
    )
)

if fail_count != 0:
    raise RuntimeError("Functional test failed")


# ===== 2秒前後となるNの推定 =====

# MicroPythonの実行時間にはばらつきがあるため、
# 厳密な最大値ではなく、おおよその目安として扱う。
low_pass, high_fail = find_upper_bound()

raw_pass_n, raw_fail_n = binary_search_limit(
    low_pass,
    high_fail
)

if raw_fail_n is None:
    print(
        "SEARCH_BOUNDARY "
        "PASS_N={} "
        "FAIL_N=NONE "
        "LIMIT_REASON=18BIT_PROTOCOL_MAX".format(
            raw_pass_n
        )
    )

    print(
        "BENCHMARK_ESTIMATE "
        "TIME_LIMIT_US={} "
        "AT_LEAST_N={} "
        "PROTOCOL_MAX_REACHED=1".format(
            TIME_LIMIT_US,
            raw_pass_n
        )
    )
else:
    estimated_n = round_to_estimate_unit(
        raw_pass_n
    )

    print(
        "SEARCH_BOUNDARY PASS_N={} FAIL_N={}".format(
            raw_pass_n,
            raw_fail_n
        )
    )

    print(
        "BENCHMARK_ESTIMATE "
        "ESTIMATED_N_AROUND_2S={} "
        "ROUND_UNIT={} "
        "PROTOCOL_MAX_REACHED=0".format(
            estimated_n,
            ESTIMATE_UNIT
        )
    )

実機で機能テストする

公式サンプルに加えて、次のケースを確認しました。

  • 最後の1byteに含まれる有効組数が1~4組の場合
  • 送信streamが256byte境界の前後になる場合
  • 固定seedで生成した疑似乱数入力

すべてPASSしました。

FUNCTION_SUMMARY PASS=12 FAIL=0 TOTAL=12

18bit通信上限まで処理時間を確認する

前回と同様に、AとBをすべて0としたベンチマークも行いました。

最大値は、18bit通信形式で扱える次の値です。

N = 2^18 - 1 = 262,143

結果は次の通りでした。

EXPAND N=262143 PACKAGES=257 TIME_US=162744
VALID=1 COUNT_OK=1 ANSWER=0 PASS

約0.163秒です。

前回のNaive版は、同じN=262,143を約0.650秒で処理していました。

Naive版                約650ms
4組パック+Prefix XOR版 約163ms

1byteへ4組を詰めた効果が、そのまま処理時間へ表れました。

SPIは4MHzなので、1byteの物理転送時間は2µsです。

最大Nでは、4組パックされたデータbyte数は次のようになります。

ceil((262,143 - 1) / 4) = 65,536 byte

ヘッダと答え受信を加えても、SPI線上の転送時間は約131msです。

実測は約163msなので、今回も物理転送時間に比較的近いところで動作しています。

第14回の素朴な1byte転送版では、2秒で処理できるNは約13,500でした。

処理可能量は大幅に増えていますね。


長い組合せ回路が気になる

高速化と実機試験は成功しました。

ただし、今回の回路では、4組パックされた1byteの受信が完了すると、通常のデータbyteについて次の処理をまとめて行っています。

4組分のPrefix XORを計算
    ↓
各位置のAと比較
    ↓
4組分の不一致数を加算
    ↓
これまでの不一致数へ加算

最終データbyteでは、さらに次の処理が続きます。

最終回答を計算
    ↓
次のクロックで回答レジスタへ取り込む値を決める

特に最終データbyteを処理する経路は、1クロックの間に通過する組合せ回路がかなり長く見えます。

高速化には成功しましたが、これはすごく気になります。


なぜ気になるのか

FPGAの同期回路では、前段レジスタの出力を元に組合せ回路が次の値を計算し、その結果を次段レジスタの入力へ伝えます。

次のクロックが立ち上がると、次段レジスタがその値を取り込みます。

つまり、クロックの1サイクルの間に、前段レジスタの出力を処理し、次段レジスタの入力まで結果を届ける必要があります。

Shrike-LiteのFPGA内部クロックは50MHzです。

したがって、クロックの立ち上がりから次の立ち上がりまでの20nsの間に、各レジスタ間の計算結果が次段レジスタの入力へ到着し、安定している必要があります。

前段レジスタの出力と次段レジスタの入力をつなぐ組合せ回路には、クロックを待って処理を始めるという動作はありません。

入力が変化すると、それに応じて出力も変化します。

しかし、実際のFPGA内部では、入力が変化した瞬間に出力が切り替わるわけではありません。

XOR、比較、加算などの論理回路を信号が通過するには、それぞれ少しずつ時間がかかります。

回路同士を結ぶ配線を信号が伝わる時間も必要です。

Shrike-LiteでAtCoder問題を解く(16)_002.png

1段ごとの遅延は短くても、複数の回路を直列に通過すると、その時間は積み重なります。

今回の回路では、Prefix XOR、比較、4組分の不一致数の集計、これまでの不一致数への加算を、途中にレジスタを挟まずにまとめて処理しています。

最終データbyteでは、さらに最終回答の計算と、回答レジスタへ取り込む値の選択まで続きます。

これだけ組合せ回路の経路が長くても、レジスタ間の信号伝搬は20ns以内に収まるのでしょうか?

一連の実機試験にはPASSしましたが、この回路が50MHzで安定して動作できるだけの時間的な余裕があるかは分かりません。

手元の個体では動作していても、FPGAの個体差や温度、電圧などの条件によっては、期待通りに動かない設計になっている可能性もあります。

そのような疑問を確認するための機能が、ForgeFPGA Workshopには備わっています。

それがTiming Analysisです。


Timing Analysisを開いて概要を確認する

ForgeFPGA WorkshopでTiming Analysisを開くには、上部ツールバーのTiming Analysisボタンをクリックします。

Shrike-LiteでAtCoder問題を解く(16)_003.png

するとTiming Analysisウィンドウが表示されます。

Shrike-LiteでAtCoder問題を解く(16)_004.png

ウィンドウ上には、bitstream生成時に計算されたFPGA内部回路の動作時間情報が表示されています。

とりあえず真っ先に確認するのは、Clocksセクションの右のほう、Achievable Period(ps)Achievable Frequency(MHz)です。

今回は32888ps30.406MHzの表示が見えます。

乱暴に言えば、これは、32.888ns(=32888ps)の時間がかかる組合せ回路の経路があって、この回路設計のままでは30.406MHzまでクロックを下げないと動作しないかもしれませんよ、という内容です。

心配が現実になりました。


今回のまとめ

今回は、ABC467Cの入力形式とFPGA側の処理を変更し、前回のNaive版の約0.650秒から約0.163秒へ、ほぼ4倍に高速化しました。

一方、Timing Analysisを開いて概要を確認すると、次の値が表示されました。

Achievable Period       32888ps
Achievable Frequency    30.406MHz

実機試験にはPASSしましたが、現在の回路は50MHzで安定して動作できる設計になっていない可能性があります。

次回

次回は、Timing Analysisの結果を詳しく確認し、長い組合せ回路の処理を複数クロックへ分割して、50MHzのタイミング違反を解消してみます。お楽しみに。


前回:
Shrike-LiteでAtCoder問題を解く(15):RP2040側でSPI転送を高速化する

次回:
Shrike-LiteでAtCoder問題を解く(17):ABC467C - Adjacent Sums (easy)(タイミング違反解消編・前編)

0
0
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?