test: validate parametric 32x4 layer with parallelism 8

This commit is contained in:
2026-09-02 09:27:04 +02:00
commit d4ae2417a4
14 changed files with 14571 additions and 0 deletions
+83
View File
@@ -0,0 +1,83 @@
module layer #(
parameter DATA_WIDTH = 16,
parameter FRAC_BITS = 8,
parameter N_INPUTS = 64,
parameter N_NEURONS = 8,
parameter PARALLEL = 8,
parameter ACC_WIDTH = 40
)(
input clk,
input rst,
input start,
input signed [DATA_WIDTH*N_INPUTS-1:0] x_bus,
input signed [DATA_WIDTH*N_INPUTS*N_NEURONS-1:0]
weights_bus,
input signed [DATA_WIDTH*N_NEURONS-1:0]
bias_bus,
output signed [DATA_WIDTH*N_NEURONS-1:0]
y_bus,
output busy,
output done
);
wire [N_NEURONS-1:0] neuron_busy;
wire [N_NEURONS-1:0] neuron_done;
genvar n;
generate
for (n = 0; n < N_NEURONS; n = n + 1) begin : GEN_NEURON
neuron_parallel #(
.DATA_WIDTH(DATA_WIDTH),
.FRAC_BITS(FRAC_BITS),
.N_INPUTS(N_INPUTS),
.PARALLEL(PARALLEL),
.ACC_WIDTH(ACC_WIDTH)
) u_neuron (
.clk(clk),
.rst(rst),
.start(start),
.x_bus(x_bus),
.w_bus(
weights_bus[
n*N_INPUTS*DATA_WIDTH
+: N_INPUTS*DATA_WIDTH
]
),
.bias(
bias_bus[
n*DATA_WIDTH
+: DATA_WIDTH
]
),
.y(
y_bus[
n*DATA_WIDTH
+: DATA_WIDTH
]
),
.busy(neuron_busy[n]),
.done(neuron_done[n])
);
end
endgenerate
assign busy = |neuron_busy;
assign done = &neuron_done;
endmodule
+122
View File
@@ -0,0 +1,122 @@
module mac8 #(
parameter DATA_WIDTH = 16,
parameter ACC_WIDTH = 40
)(
input signed [DATA_WIDTH*8-1:0] x_bus,
input signed [DATA_WIDTH*8-1:0] w_bus,
input signed [ACC_WIDTH-1:0] acc_in,
output signed [ACC_WIDTH-1:0] acc_out
);
wire signed [ACC_WIDTH-1:0] m0;
wire signed [ACC_WIDTH-1:0] m1;
wire signed [ACC_WIDTH-1:0] m2;
wire signed [ACC_WIDTH-1:0] m3;
wire signed [ACC_WIDTH-1:0] m4;
wire signed [ACC_WIDTH-1:0] m5;
wire signed [ACC_WIDTH-1:0] m6;
wire signed [ACC_WIDTH-1:0] m7;
wire signed [ACC_WIDTH-1:0] s0;
wire signed [ACC_WIDTH-1:0] s1;
wire signed [ACC_WIDTH-1:0] s2;
wire signed [ACC_WIDTH-1:0] s3;
wire signed [ACC_WIDTH-1:0] s4;
wire signed [ACC_WIDTH-1:0] s5;
wire signed [ACC_WIDTH-1:0] sum;
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac0 (
.x(x_bus[0*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[0*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m0)
);
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac1 (
.x(x_bus[1*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[1*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m1)
);
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac2 (
.x(x_bus[2*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[2*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m2)
);
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac3 (
.x(x_bus[3*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[3*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m3)
);
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac4 (
.x(x_bus[4*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[4*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m4)
);
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac5 (
.x(x_bus[5*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[5*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m5)
);
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac6 (
.x(x_bus[6*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[6*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m6)
);
mac_unit #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) mac7 (
.x(x_bus[7*DATA_WIDTH +: DATA_WIDTH]),
.w(w_bus[7*DATA_WIDTH +: DATA_WIDTH]),
.acc_in({ACC_WIDTH{1'b0}}),
.acc_out(m7)
);
assign s0 = m0 + m1;
assign s1 = m2 + m3;
assign s2 = m4 + m5;
assign s3 = m6 + m7;
assign s4 = s0 + s1;
assign s5 = s2 + s3;
assign sum = s4 + s5;
assign acc_out = acc_in + sum;
endmodule
+23
View File
@@ -0,0 +1,23 @@
module mac_unit #(
parameter DATA_WIDTH = 16,
parameter ACC_WIDTH = 40
)(
input signed [DATA_WIDTH-1:0] x,
input signed [DATA_WIDTH-1:0] w,
input signed [ACC_WIDTH-1:0] acc_in,
output signed [ACC_WIDTH-1:0] acc_out
);
localparam PROD_WIDTH = 2 * DATA_WIDTH;
wire signed [PROD_WIDTH-1:0] product;
wire signed [ACC_WIDTH-1:0] product_ext;
assign product = x * w;
assign product_ext =
{{(ACC_WIDTH-PROD_WIDTH){product[PROD_WIDTH-1]}}, product};
assign acc_out = acc_in + product_ext;
endmodule
+119
View File
@@ -0,0 +1,119 @@
module neuron_parallel #(
parameter DATA_WIDTH = 16,
parameter FRAC_BITS = 8,
parameter N_INPUTS = 64,
parameter PARALLEL = 8,
parameter ACC_WIDTH = 40
)(
input clk,
input rst,
input start,
input signed [DATA_WIDTH*N_INPUTS-1:0] x_bus,
input signed [DATA_WIDTH*N_INPUTS-1:0] w_bus,
input signed [DATA_WIDTH-1:0] bias,
output reg signed [DATA_WIDTH-1:0] y,
output reg busy,
output reg done
);
localparam GROUPS = N_INPUTS / PARALLEL;
localparam GROUP_INDEX_WIDTH =
(GROUPS <= 1) ? 1 : $clog2(GROUPS);
reg [GROUP_INDEX_WIDTH-1:0] group_index;
reg signed [ACC_WIDTH-1:0] acc;
wire signed [DATA_WIDTH*PARALLEL-1:0] x_group;
wire signed [DATA_WIDTH*PARALLEL-1:0] w_group;
wire signed [ACC_WIDTH-1:0] acc_next;
wire signed [DATA_WIDTH-1:0] bias_ext_small;
wire signed [ACC_WIDTH-1:0] bias_ext;
wire signed [ACC_WIDTH-1:0] final_acc;
wire signed [ACC_WIDTH-1:0] final_value;
assign x_group =
x_bus[group_index*PARALLEL*DATA_WIDTH
+: PARALLEL*DATA_WIDTH];
assign w_group =
w_bus[group_index*PARALLEL*DATA_WIDTH
+: PARALLEL*DATA_WIDTH];
mac8 #(
.DATA_WIDTH(DATA_WIDTH),
.ACC_WIDTH(ACC_WIDTH)
) u_mac8 (
.x_bus(x_group),
.w_bus(w_group),
.acc_in(acc),
.acc_out(acc_next)
);
assign bias_ext_small = bias;
assign bias_ext =
{{(ACC_WIDTH-DATA_WIDTH){bias_ext_small[DATA_WIDTH-1]}},
bias_ext_small};
assign final_acc =
acc_next + (bias_ext <<< FRAC_BITS);
assign final_value =
final_acc >>> FRAC_BITS;
always @(posedge clk) begin
if (rst) begin
group_index <= 0;
acc <= 0;
y <= 0;
busy <= 0;
done <= 0;
end else begin
done <= 0;
if (start && !busy) begin
group_index <= 0;
acc <= 0;
busy <= 1;
end else if (busy) begin
if (group_index == GROUPS-1) begin
acc <= final_acc;
if (final_value <= 0) begin
y <= 0;
end
else if (final_value > 32767) begin
y <= 16'sh7FFF;
end
else begin
y <= final_value[DATA_WIDTH-1:0];
end
busy <= 0;
done <= 1;
end else begin
acc <= acc_next;
group_index <= group_index + 1'b1;
end
end
end
end
endmodule