3
0
Fork 0
mirror of https://github.com/YosysHQ/yosys synced 2026-07-24 08:02:32 +00:00

opt_argmax pass

This commit is contained in:
Akash Levy 2026-06-02 04:11:17 -07:00
parent c7b2c16405
commit b3ea5770cd
4 changed files with 1434 additions and 1 deletions

View file

@ -0,0 +1,324 @@
module opt_argmax_basic (
input wire [15:0] sig,
input wire [15:0][3:0] sig3,
input wire [15:0][7:0] sig2,
output reg [3:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 16; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_w8 (
input wire [7:0] sig,
input wire [7:0][2:0] sig3,
input wire [7:0][4:0] sig2,
output reg [2:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 8; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_w32 (
input wire [31:0] sig,
input wire [31:0][4:0] sig3,
input wire [31:0][5:0] sig2,
output reg [4:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 32; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_flat (
input wire [7:0] sig,
input wire [23:0] sig3,
input wire [39:0] sig2,
output reg [2:0] se_target_idx
);
function automatic [2:0] idx_at(input [2:0] pos);
idx_at = sig3[pos * 3 +: 3];
endfunction
function automatic [4:0] val_at(input [2:0] pos);
val_at = sig2[idx_at(pos) * 5 +: 5];
endfunction
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 8; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(val_at(se_target_idx) < val_at(k[2:0]))) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_value_w1 (
input wire [7:0] sig,
input wire [7:0][2:0] sig3,
input wire [7:0] sig2,
output reg [2:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 8; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_value_w16 (
input wire [7:0] sig,
input wire [7:0][2:0] sig3,
input wire [7:0][15:0] sig2,
output reg [2:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 8; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_two_regions (
input wire [7:0] sig_a,
input wire [7:0][2:0] sig3_a,
input wire [7:0][7:0] sig2_a,
input wire [7:0] sig_b,
input wire [7:0][2:0] sig3_b,
input wire [7:0][5:0] sig2_b,
output reg [2:0] idx_a,
output reg [2:0] idx_b
);
always_comb begin
idx_a = '0;
for (int k = 1; k < 8; k++) begin
if (!sig_a[idx_a] && sig_a[k]) begin
idx_a = k;
end else if (sig_a[idx_a] && sig_a[k] &&
(sig2_a[sig3_a[idx_a]] < sig2_a[sig3_a[k]])) begin
idx_a = k;
end
end
idx_b = '0;
for (int k = 1; k < 8; k++) begin
if (!sig_b[idx_b] && sig_b[k]) begin
idx_b = k;
end else if (sig_b[idx_b] && sig_b[k] &&
(sig2_b[sig3_b[idx_b]] < sig2_b[sig3_b[k]])) begin
idx_b = k;
end
end
end
endmodule
module opt_argmax_shared_consumer (
input wire [7:0] sig,
input wire [7:0][2:0] sig3,
input wire [7:0][7:0] sig2,
input wire [2:0] salt,
output reg [2:0] se_target_idx,
output wire [2:0] also_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 8; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
assign also_idx = se_target_idx ^ salt;
endmodule
module opt_argmax_tie_high (
input wire [15:0] sig,
input wire [15:0][3:0] sig3,
input wire [15:0][7:0] sig2,
output reg [3:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 16; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] <= sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_nonzero_default (
input wire [15:0] sig,
input wire [15:0][3:0] sig3,
input wire [15:0][7:0] sig2,
output reg [3:0] se_target_idx
);
always_comb begin
se_target_idx = 4'd1;
for (int k = 1; k < 16; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_min (
input wire [15:0] sig,
input wire [15:0][3:0] sig3,
input wire [15:0][7:0] sig2,
output reg [3:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 16; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] > sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_w12 (
input wire [11:0] sig,
input wire [11:0][3:0] sig3,
input wire [11:0][7:0] sig2,
output reg [3:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 12; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_bad_index_width (
input wire [15:0] sig,
input wire [15:0][4:0] sig3,
input wire [15:0][7:0] sig2,
output reg [3:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 16; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx][3:0]] < sig2[sig3[k][3:0]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_stress_noop (
input wire [63:0] sel,
input wire [63:0] a,
input wire [63:0] b,
output wire [63:0] y
);
wire [63:0] mux0 = sel[0] ? a : b;
wire [63:0] mux1 = sel[1] ? mux0 : {mux0[31:0], mux0[63:32]};
wire [63:0] mux2 = sel[2] ? mux1 : (mux1 ^ a);
wire [63:0] mux3 = sel[3] ? mux2 : (mux2 & b);
wire [63:0] mux4 = sel[4] ? mux3 : (mux3 | a);
wire [63:0] mux5 = sel[5] ? mux4 : {mux4[47:0], mux4[63:48]};
assign y = sel[6] ? mux5 : ~mux5;
endmodule
module opt_argmax_unrelated (
input wire [3:0] a,
input wire [3:0] b,
input wire sel,
output wire [3:0] y
);
assign y = sel ? a : b;
endmodule
module opt_argmax_multi_match (
input wire [15:0] sig,
input wire [15:0][3:0] sig3,
input wire [15:0][7:0] sig2,
output reg [3:0] se_target_idx
);
always_comb begin
se_target_idx = '0;
for (int k = 1; k < 16; k++) begin
if (!sig[se_target_idx] && sig[k]) begin
se_target_idx = k;
end else if (sig[se_target_idx] && sig[k] &&
(sig2[sig3[se_target_idx]] < sig2[sig3[k]])) begin
se_target_idx = k;
end
end
end
endmodule
module opt_argmax_multi_keep (
input wire [3:0] a,
input wire [3:0] b,
input wire sel,
output wire [3:0] y
);
assign y = sel ? a : b;
endmodule

View file

@ -0,0 +1,332 @@
# Tests for opt_argmax.
log -header "Small masked argmax self-equivalence"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_w8
proc; opt_clean
rename opt_argmax_w8 gold
read -sv opt_argmax.sv
verific -import opt_argmax_w8
proc; opt_clean
select -module opt_argmax_w8
opt_argmax
select -clear
opt_clean
rename opt_argmax_w8 gate
miter -equiv -flatten -make_assert gold gate miter
hierarchy -top miter
proc; opt; memory; opt
sat -prove-asserts -verify
design -reset
log -pop
log -header "Basic masked argmax structural rewrite"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_basic
proc; opt_clean
opt_argmax
opt_clean
select -assert-min 1 w:*argmax*
select -assert-count 16 t:$bmux
select -assert-count 15 t:$lt
select -assert-count 29 t:$mux
select -assert-count 30 t:$and
select -assert-count 29 t:$or
select -assert-count 15 t:$not
select -assert-none c:LessThan_*
design -reset
log -pop
log -header "Flat-bus masked argmax self-equivalence"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_flat
proc; opt_clean
rename opt_argmax_flat gold
read -sv opt_argmax.sv
verific -import opt_argmax_flat
proc; opt_clean
select -module opt_argmax_flat
opt_argmax
select -clear
opt_clean
select -assert-min 1 w:*argmax*
rename opt_argmax_flat gate
miter -equiv -flatten -make_assert gold gate miter
hierarchy -top miter
proc; opt; memory; opt
sat -prove-asserts -verify
design -reset
log -pop
log -header "Scaled masked argmax: 8 entries structural"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_w8
proc; opt_clean
opt_argmax
opt_clean
select -assert-min 1 w:*argmax*
design -reset
log -pop
log -header "Scaled masked argmax: 32 entries structural"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_w32
proc; opt_clean
opt_argmax
opt_clean
select -assert-min 1 w:*argmax*
design -reset
log -pop
log -header "Value width edge: 1-bit values"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_value_w1
proc; opt_clean
rename opt_argmax_value_w1 gold
read -sv opt_argmax.sv
verific -import opt_argmax_value_w1
proc; opt_clean
select -module opt_argmax_value_w1
opt_argmax
select -clear
opt_clean
select -assert-min 1 w:*argmax*
rename opt_argmax_value_w1 gate
miter -equiv -flatten -make_assert gold gate miter
hierarchy -top miter
proc; opt; memory; opt
sat -prove-asserts -verify
design -reset
log -pop
log -header "Value width edge: 16-bit values"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_value_w16
proc; opt_clean
rename opt_argmax_value_w16 gold
read -sv opt_argmax.sv
verific -import opt_argmax_value_w16
proc; opt_clean
select -module opt_argmax_value_w16
opt_argmax
select -clear
opt_clean
select -assert-min 1 w:*argmax*
rename opt_argmax_value_w16 gate
miter -equiv -flatten -make_assert gold gate miter
hierarchy -top miter
proc; opt; memory; opt
sat -prove-asserts -verify
design -reset
log -pop
log -header "Same module: two independent argmax regions"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_two_regions
proc; opt_clean
rename opt_argmax_two_regions gold
read -sv opt_argmax.sv
verific -import opt_argmax_two_regions
proc; opt_clean
select -module opt_argmax_two_regions
opt_argmax
select -clear
opt_clean
select -assert-min 2 w:*argmax*
rename opt_argmax_two_regions gate
miter -equiv -flatten -make_assert gold gate miter
hierarchy -top miter
proc; opt; memory; opt
sat -prove-asserts -verify
design -reset
log -pop
log -header "Shared consumer of argmax output remains equivalent"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_shared_consumer
proc; opt_clean
rename opt_argmax_shared_consumer gold
read -sv opt_argmax.sv
verific -import opt_argmax_shared_consumer
proc; opt_clean
select -module opt_argmax_shared_consumer
opt_argmax
select -clear
opt_clean
select -assert-min 1 w:*argmax*
rename opt_argmax_shared_consumer gate
miter -equiv -flatten -make_assert gold gate miter
hierarchy -top miter
proc; opt; memory; opt
sat -prove-asserts -verify
design -reset
log -pop
log -header "Max width leaves argmax unchanged"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_basic
proc; opt_clean
opt_argmax -max_width 8
select -assert-none w:*argmax*
design -reset
log -pop
log -header "Negative: non-power-of-two candidate count"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_w12
proc; opt_clean
opt_argmax
select -assert-none w:*argmax*
design -reset
log -pop
log -header "Negative: mismatched index-map width"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_bad_index_width
proc; opt_clean
opt_argmax
select -assert-none w:*argmax*
design -reset
log -pop
log -header "Negative: strict tie behavior changed"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_tie_high
proc; opt_clean
opt_argmax
select -assert-none w:*argmax*
design -reset
log -pop
log -header "Negative: nonzero all-invalid default"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_nonzero_default
proc; opt_clean
opt_argmax
select -assert-none w:*argmax*
design -reset
log -pop
log -header "Negative: min-selection comparator"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_min
proc; opt_clean
opt_argmax
select -assert-none w:*argmax*
design -reset
log -pop
log -header "Negative: unrelated mux logic unchanged"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_unrelated
proc; opt_clean
select -assert-count 1 t:$mux
opt_argmax
select -assert-none w:*argmax*
select -assert-count 1 t:$mux
design -reset
log -pop
log -header "Negative: bmux-heavy unrelated stress module"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_stress_noop
proc; opt_clean
opt_argmax
select -assert-none w:*argmax*
design -reset
log -pop
log -header "Multi-module: only matching module rewrites"
log -push
design -reset
verific -cfg veri_optimize_wide_selector 1
verific -cfg db_infer_wide_muxes_post_elaboration 0
read -sv opt_argmax.sv
verific -import opt_argmax_multi_match opt_argmax_multi_keep
proc; opt_clean
opt_argmax
opt_clean
select -assert-min 1 opt_argmax_multi_match/w:*argmax*
select -assert-none opt_argmax_multi_keep/w:*argmax*
design -reset
log -pop