Skip to content

Commit f18a9cb

Browse files
tinebpclaude
andcommitted
dxa: size smem address datapath to LMEM byte-address width
The DXA shared-memory (LMEM) byte address was carried as VX_CFG_MEM_ADDR_WIDTH (48b on XLEN=64) — the global-memory width — even though it only ever indexes one core's LMEM. Introduce DXA_SMEM_ADDR_W = VX_CFG_LMEM_LOG_SIZE (the LMEM byte-address width, same as VX_local_mem's ADDR_WIDTH) and use it across the whole smem datapath: dxa_req_data_t.smem_addr, dxa_setup_params_t.initial_smem_base, setup's r_/s_initial_smem_base, addr_gen's smem_byte_addr_r/km_row_base_r and out_smem_byte_addr, worker's ag_/sw_smem_byte_addr and the gmem_req SMEM_ADDR_W param, and smem_wr's pend_/defer_/fb_byte_addr_r and fb_load_smem_byte_addr. The unit producer slices the LMEM-relative address to DXA_SMEM_ADDR_W; UNUSED_VAR covers the LMEM-bounded upper bits of km_step_in_row and per_lane_stride_bytes. GMEM addresses and the already-tight LMEM word-address internals are unchanged. Elaborates clean on build64. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
1 parent 949ee87 commit f18a9cb

7 files changed

Lines changed: 39 additions & 30 deletions

File tree

hw/rtl/dxa/VX_dxa_addr_gen.sv

Lines changed: 11 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -36,7 +36,7 @@ module VX_dxa_addr_gen import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
3636
output wire out_valid,
3737
input wire out_ready,
3838
output wire [GMEM_ADDR_WIDTH-1:0] out_cl_addr,
39-
output wire [`VX_CFG_MEM_ADDR_WIDTH-1:0] out_smem_byte_addr,
39+
output wire [DXA_SMEM_ADDR_W-1:0] out_smem_byte_addr,
4040
output wire [CL_OFF_BITS-1:0] out_byte_offset,
4141
output wire [CL_OFF_BITS:0] out_valid_length,
4242
output wire out_oob,
@@ -65,15 +65,15 @@ module VX_dxa_addr_gen import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
6565
reg [DXA_MAX_OUTER_DIMS-1:0][31:0] oob_limit_r; // OOB limit per dim
6666

6767
// SMEM byte address tracking.
68-
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] smem_byte_addr_r;
68+
reg [DXA_SMEM_ADDR_W-1:0] smem_byte_addr_r;
6969
// K-major scatter-mode state.
7070
// km_row_base_r: SMEM base for the current outer-dim row (= initial +
7171
// dim_count[0] * elem_bytes). Updates only at row wrap.
7272
// km_dest_kmajor_r / km_per_lane_stride / km_elem_bytes: stable params.
7373
reg km_dest_kmajor_r;
7474
reg [15:0] km_per_lane_stride_r;
7575
reg [3:0] km_elem_bytes_r;
76-
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] km_row_base_r;
76+
reg [DXA_SMEM_ADDR_W-1:0] km_row_base_r;
7777

7878
// Pass-through latched params.
7979
reg [31:0] cfill_r;
@@ -181,11 +181,11 @@ module VX_dxa_addr_gen import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
181181
row_len_r <= setup_params.row_len_bytes;
182182
line_idx_r <= '0;
183183
cfill_r <= setup_params.cfill;
184-
smem_byte_addr_r <= `VX_CFG_MEM_ADDR_WIDTH'(setup_params.initial_smem_base);
184+
smem_byte_addr_r <= setup_params.initial_smem_base;
185185
km_dest_kmajor_r <= setup_params.dest_kmajor;
186186
km_per_lane_stride_r <= setup_params.per_lane_stride_bytes;
187187
km_elem_bytes_r <= setup_params.elem_bytes;
188-
km_row_base_r <= `VX_CFG_MEM_ADDR_WIDTH'(setup_params.initial_smem_base);
188+
km_row_base_r <= setup_params.initial_smem_base;
189189
for (int d = 0; d < DXA_MAX_OUTER_DIMS; d++) begin
190190
dim_count_r[d] <= '0;
191191
dim_tile_r[d] <= setup_params.dim_tiles[d];
@@ -200,9 +200,9 @@ module VX_dxa_addr_gen import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
200200
// layout). On row wrap, the address is overridden
201201
// below to the next column's row base.
202202
if (km_dest_kmajor_r) begin
203-
smem_byte_addr_r <= smem_byte_addr_r + `VX_CFG_MEM_ADDR_WIDTH'(km_step_in_row);
203+
smem_byte_addr_r <= smem_byte_addr_r + DXA_SMEM_ADDR_W'(km_step_in_row);
204204
end else begin
205-
smem_byte_addr_r <= smem_byte_addr_r + `VX_CFG_MEM_ADDR_WIDTH'(cur_valid_length);
205+
smem_byte_addr_r <= smem_byte_addr_r + DXA_SMEM_ADDR_W'(cur_valid_length);
206206
end
207207

208208
if (is_last_line && is_last_outer) begin
@@ -218,8 +218,8 @@ module VX_dxa_addr_gen import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
218218
// cursor — the override above is unconditional, so it would
219219
// step PAST the row base; we restore it here for K-major.
220220
if (km_dest_kmajor_r) begin
221-
smem_byte_addr_r <= km_row_base_r + `VX_CFG_MEM_ADDR_WIDTH'(km_elem_bytes_r);
222-
km_row_base_r <= km_row_base_r + `VX_CFG_MEM_ADDR_WIDTH'(km_elem_bytes_r);
221+
smem_byte_addr_r <= km_row_base_r + DXA_SMEM_ADDR_W'(km_elem_bytes_r);
222+
km_row_base_r <= km_row_base_r + DXA_SMEM_ADDR_W'(km_elem_bytes_r);
223223
end
224224
if (dim0_steps) begin
225225
dim_count_r[0] <= dim_count_r[0] + 1;
@@ -245,7 +245,8 @@ module VX_dxa_addr_gen import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
245245

246246
`UNUSED_VAR (cur_cl_byte_addr[CL_OFF_BITS-1:0])
247247
`UNUSED_VAR (total_end[31:CL_OFF_BITS])
248-
`UNUSED_VAR (setup_params.initial_smem_base)
248+
// km_step_in_row is LMEM-bounded; only the low DXA_SMEM_ADDR_W bits feed smem_byte_addr_r.
249+
`UNUSED_VAR (km_step_in_row[31:DXA_SMEM_ADDR_W])
249250

250251
`ifdef DBG_TRACE_DXA
251252
always @(posedge clk) begin

hw/rtl/dxa/VX_dxa_gmem_req.sv

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ module VX_dxa_gmem_req import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
2929
parameter GMEM_ADDR_WIDTH = `VX_CFG_MEM_ADDR_WIDTH - `CLOG2(`VX_CFG_L1_LINE_SIZE),
3030
parameter GMEM_TAG_WIDTH = L1_MEM_ARB_TAG_WIDTH,
3131
parameter CL_OFF_BITS = `CLOG2(`VX_CFG_L1_LINE_SIZE),
32-
parameter SMEM_ADDR_W = `VX_CFG_MEM_ADDR_WIDTH
32+
parameter SMEM_ADDR_W = DXA_SMEM_ADDR_W
3333
) (
3434
input wire clk,
3535
input wire reset,

hw/rtl/dxa/VX_dxa_pkg.sv

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,10 @@ package VX_dxa_pkg;
2424
localparam DXA_LMEM_WORD_SIZE = `VX_CFG_LMEM_NUM_BANKS * (`VX_CFG_XLEN / 8);
2525
localparam DXA_LMEM_ADDR_W = LMEM_DMA_ADDR_WIDTH;
2626

27+
// LMEM byte-address width: an LMEM-relative byte offset, same as
28+
// VX_local_mem's ADDR_WIDTH = `CLOG2(SIZE) = `VX_CFG_LMEM_LOG_SIZE.
29+
localparam DXA_SMEM_ADDR_W = `VX_CFG_LMEM_LOG_SIZE;
30+
2731
localparam DXA_DESC_SLOT_BITS = `CLOG2(`VX_DCR_DXA_DESC_COUNT);
2832
localparam DXA_DESC_SLOT_W = `UP(DXA_DESC_SLOT_BITS);
2933

@@ -32,7 +36,7 @@ package VX_dxa_pkg;
3236
logic [NC_WIDTH-1:0] core_id;
3337
logic [UUID_WIDTH-1:0] uuid;
3438
logic [NW_WIDTH-1:0] wid;
35-
logic [`VX_CFG_XLEN-1:0] smem_addr; // from lane 0 rs1
39+
logic [DXA_SMEM_ADDR_W-1:0] smem_addr; // from lane 0 rs1; LMEM byte address
3640
logic [31:0] meta; // from lane 1 rs1 (desc[3:0], bar[30:4], 1[31]); 32-bit ABI word
3741
logic [4:0][31:0] coords; // [0]=lane2.rs1,[1]=lane3.rs1,[2]=lane0.rs2,[3]=lane1.rs2,[4]=lane2.rs2; element indices, 32-bit ABI
3842
logic [`VX_CFG_NUM_WARPS-1:0] cta_mask; // from rs2 lane 3
@@ -67,7 +71,7 @@ package VX_dxa_pkg;
6771
// All multiplies happen during setup; fast path uses additions only.
6872
typedef struct packed {
6973
logic [`VX_CFG_MEM_ADDR_WIDTH-1:0] initial_gmem_base;
70-
logic [`VX_CFG_XLEN-1:0] initial_smem_base;
74+
logic [DXA_SMEM_ADDR_W-1:0] initial_smem_base;
7175
logic [31:0] row_len_bytes;
7276
// Rolling-cursor deltas applied at each outer-dim step:
7377
// delta[0]: dim-0 step = stride[0]

hw/rtl/dxa/VX_dxa_setup.sv

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -150,10 +150,10 @@ module VX_dxa_setup import VX_gpu_pkg::*, VX_dxa_pkg::*; (
150150
reg [BAR_ADDR_W-1:0] r_bar_addr;
151151
reg r_notify_smem_done;
152152
reg r_is_multicast;
153-
reg [`VX_CFG_NUM_WARPS-1:0] r_cta_mask;
153+
reg [`VX_CFG_NUM_WARPS-1:0] r_cta_mask;
154154
reg [31:0] r_smem_stride;
155-
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] r_initial_gmem_base;
156-
reg [`VX_CFG_XLEN-1:0] r_initial_smem_base;
155+
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] r_initial_gmem_base;
156+
reg [DXA_SMEM_ADDR_W-1:0] r_initial_smem_base;
157157
reg [31:0] r_row_len_bytes;
158158
reg [DXA_MAX_OUTER_DIMS-1:0][31:0] r_delta;
159159
reg [DXA_MAX_OUTER_DIMS-1:0][31:0] r_dim_tiles;
@@ -174,10 +174,10 @@ module VX_dxa_setup import VX_gpu_pkg::*, VX_dxa_pkg::*; (
174174
reg [BAR_ADDR_W-1:0] s_bar_addr;
175175
reg s_notify_smem_done;
176176
reg s_is_multicast;
177-
reg [`VX_CFG_NUM_WARPS-1:0] s_cta_mask;
177+
reg [`VX_CFG_NUM_WARPS-1:0] s_cta_mask;
178178
reg [31:0] s_smem_stride;
179-
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] s_initial_gmem_base;
180-
reg [`VX_CFG_XLEN-1:0] s_initial_smem_base;
179+
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] s_initial_gmem_base;
180+
reg [DXA_SMEM_ADDR_W-1:0] s_initial_smem_base;
181181
reg [31:0] s_row_len_bytes;
182182
reg [DXA_MAX_OUTER_DIMS-1:0][31:0] s_delta;
183183
reg [DXA_MAX_OUTER_DIMS-1:0][31:0] s_dim_tiles;

hw/rtl/dxa/VX_dxa_smem_wr.sv

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,7 @@ module VX_dxa_smem_wr import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
5151
output wire sw_ready,
5252
input wire [TAG_W-1:0] sw_tag,
5353
input wire [GMEM_DATAW-1:0] sw_data,
54-
input wire [`VX_CFG_MEM_ADDR_WIDTH-1:0] sw_smem_byte_addr,
54+
input wire [DXA_SMEM_ADDR_W-1:0] sw_smem_byte_addr,
5555
input wire [CL_OFF_BITS-1:0] sw_byte_offset,
5656
input wire [CL_OFF_BITS:0] sw_valid_length,
5757
input wire sw_oob,
@@ -112,7 +112,7 @@ module VX_dxa_smem_wr import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
112112
// ════════════════════════════════════════════════════════════════════
113113
reg pend_valid_r;
114114
reg [TAG_W-1:0] pend_tag_r;
115-
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] pend_smem_byte_addr_r;
115+
reg [DXA_SMEM_ADDR_W-1:0] pend_smem_byte_addr_r;
116116
reg [CL_OFF_BITS-1:0] pend_byte_offset_r;
117117
reg [CL_OFF_BITS:0] pend_valid_length_r;
118118
reg pend_last_r;
@@ -124,7 +124,7 @@ module VX_dxa_smem_wr import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
124124
// ════════════════════════════════════════════════════════════════════
125125
reg defer_valid_r;
126126
reg [TAG_W-1:0] defer_tag_r;
127-
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] defer_smem_byte_addr_r;
127+
reg [DXA_SMEM_ADDR_W-1:0] defer_smem_byte_addr_r;
128128
reg [CL_OFF_BITS-1:0] defer_byte_offset_r;
129129
reg [CL_OFF_BITS:0] defer_valid_length_r;
130130
reg [GMEM_DATAW-1:0] defer_data_r;
@@ -142,7 +142,7 @@ module VX_dxa_smem_wr import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
142142
reg [SMEM_ADDR_WIDTH-1:0] fb_start_word_r;
143143
// K-major scatter state: per-beat target byte address (this CL's
144144
// element-0 destination plus N*per_lane_stride per beat).
145-
reg [`VX_CFG_MEM_ADDR_WIDTH-1:0] fb_byte_addr_r;
145+
reg [DXA_SMEM_ADDR_W-1:0] fb_byte_addr_r;
146146

147147
// K-major drain quantum (in bytes) = 1 element per beat in scatter mode,
148148
// SMEM_WORD_SIZE bytes per beat in row-major streaming mode.
@@ -257,7 +257,7 @@ module VX_dxa_smem_wr import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
257257
wire fb_load_now = use_sw_for_fb || use_pend_for_fb || use_defer_for_fb;
258258

259259
wire [GMEM_DATAW-1:0] fb_load_data;
260-
wire [`VX_CFG_MEM_ADDR_WIDTH-1:0] fb_load_smem_byte_addr;
260+
wire [DXA_SMEM_ADDR_W-1:0] fb_load_smem_byte_addr;
261261
wire [CL_OFF_BITS-1:0] fb_load_byte_offset;
262262
wire [CL_OFF_BITS:0] fb_load_valid_length;
263263
wire [TAG_W-1:0] fb_load_tag;
@@ -323,7 +323,7 @@ module VX_dxa_smem_wr import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
323323
if (dest_kmajor) begin
324324
fb_data_r <= fb_data_r >> km_shift_bits;
325325
fb_level_r <= fb_level_r - drain_q_bytes;
326-
fb_byte_addr_r <= fb_byte_addr_r + `VX_CFG_MEM_ADDR_WIDTH'(per_lane_stride_bytes);
326+
fb_byte_addr_r <= fb_byte_addr_r + DXA_SMEM_ADDR_W'(per_lane_stride_bytes);
327327
end else begin
328328
fb_data_r <= fb_data_r >> SMEM_DATAW;
329329
fb_level_r <= fb_level_r - FILL_W'(SMEM_WORD_SIZE);
@@ -458,6 +458,8 @@ module VX_dxa_smem_wr import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
458458
wire [SMEM_ADDR_WIDTH-1:0] replay_addr = base_word_addr + beat_offset;
459459
// Only the low SMEM_ADDR_WIDTH+SMEM_OFF_W bits of smem_stride are used.
460460
`UNUSED_VAR (smem_stride)
461+
// per_lane_stride_bytes is LMEM-bounded; only the low DXA_SMEM_ADDR_W bits feed fb_byte_addr_r.
462+
`UNUSED_VAR (per_lane_stride_bytes[15:DXA_SMEM_ADDR_W])
461463

462464
wire replay_is_last = replay_has_remaining
463465
&& (replay_remaining_use == (`VX_CFG_NUM_WARPS'(1) << replay_next_idx));

hw/rtl/dxa/VX_dxa_unit.sv

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,15 +61,17 @@ module VX_dxa_unit import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
6161
assign dxa_req_data_in.core_id = NC_WIDTH'(CORE_ID);
6262
assign dxa_req_data_in.uuid = execute_if.data.header.uuid;
6363
assign dxa_req_data_in.wid = execute_if.data.header.wid;
64-
assign dxa_req_data_in.smem_addr = lmem_rel_byte_addr;
64+
assign dxa_req_data_in.smem_addr = lmem_rel_byte_addr[DXA_SMEM_ADDR_W-1:0];
6565
assign dxa_req_data_in.meta = lane1_rs1[31:0];
6666
assign dxa_req_data_in.coords[0] = lane2_rs1[31:0];
6767
assign dxa_req_data_in.coords[1] = lane3_rs1[31:0];
6868
assign dxa_req_data_in.coords[2] = lane0_rs2[31:0];
6969
assign dxa_req_data_in.coords[3] = lane1_rs2[31:0];
7070
assign dxa_req_data_in.coords[4] = lane2_rs2[31:0];
7171
assign dxa_req_data_in.cta_mask = lane3_rs2[`VX_CFG_NUM_WARPS-1:0];
72-
// meta/coords carry 32-bit ABI values; high bits unused on XLEN>32 builds
72+
// smem_addr (LMEM byte width) and meta/coords (32-bit ABI) take low bits;
73+
// high bits unused when XLEN exceeds the field width
74+
`UNUSED_VAR (lmem_rel_byte_addr)
7375
`UNUSED_VAR (lane1_rs1)
7476
`UNUSED_VAR (lane2_rs1)
7577
`UNUSED_VAR (lane3_rs1)

hw/rtl/dxa/VX_dxa_worker.sv

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -67,7 +67,7 @@ module VX_dxa_worker import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
6767
wire ag_valid;
6868
wire ag_ready;
6969
wire [GMEM_ADDR_WIDTH-1:0] ag_cl_addr;
70-
wire [`VX_CFG_MEM_ADDR_WIDTH-1:0] ag_smem_byte_addr;
70+
wire [DXA_SMEM_ADDR_W-1:0] ag_smem_byte_addr;
7171
wire [GMEM_OFF_BITS-1:0] ag_byte_offset;
7272
wire [GMEM_OFF_BITS:0] ag_valid_length;
7373
wire ag_oob;
@@ -82,7 +82,7 @@ module VX_dxa_worker import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
8282
wire sw_ready;
8383
wire [TAG_W-1:0] sw_tag;
8484
wire [GMEM_DATAW-1:0] sw_data;
85-
wire [`VX_CFG_MEM_ADDR_WIDTH-1:0] sw_smem_byte_addr;
85+
wire [DXA_SMEM_ADDR_W-1:0] sw_smem_byte_addr;
8686
wire [GMEM_OFF_BITS-1:0] sw_byte_offset;
8787
wire [GMEM_OFF_BITS:0] sw_valid_length;
8888
wire sw_oob;
@@ -168,7 +168,7 @@ module VX_dxa_worker import VX_gpu_pkg::*, VX_dxa_pkg::*; #(
168168
.GMEM_ADDR_WIDTH (GMEM_ADDR_WIDTH),
169169
.GMEM_TAG_WIDTH (GMEM_TAG_WIDTH),
170170
.CL_OFF_BITS (GMEM_OFF_BITS),
171-
.SMEM_ADDR_W (`VX_CFG_MEM_ADDR_WIDTH)
171+
.SMEM_ADDR_W (DXA_SMEM_ADDR_W)
172172
) gmem_req (
173173
.clk (clk),
174174
.reset (reset),

0 commit comments

Comments
 (0)