test: validate parametric 32x4 layer with parallelism 8
This commit is contained in:
+83
@@ -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
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user