Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 7 additions & 3 deletions VX_config.toml
Original file line number Diff line number Diff line change
Expand Up @@ -256,9 +256,13 @@ VX_CFG_TCU_WGMMA_ENABLE = false
VX_CFG_TCU_FEDP2K = false

[dxa]
# Cluster-level DXA engine count, decoupled from cores/socket size.
VX_CFG_NUM_DXA_CORES = "expr: max(1, up($VX_CFG_NUM_CORES / 8))"
VX_CFG_DXA_MEM_PORTS = "expr: min($VX_CFG_NUM_DXA_CORES, up($VX_CFG_NUM_CORES / $VX_CFG_SOCKET_SIZE) * $VX_CFG_L1_MEM_PORTS)"
# Socket-local DXA engine worker count. One DXA engine is instantiated per
# socket; this knob controls the number of workers inside each engine and is
# intentionally independent of the number of SM cores and sockets.
VX_CFG_NUM_DXA_CORES = 1
# Number of L2-facing request ports per socket-local engine. The cluster
# aggregate is derived in VX_gpu_pkg as NUM_SOCKETS times this value.
VX_CFG_DXA_MEM_PORTS = "expr: min($VX_CFG_NUM_DXA_CORES, $VX_CFG_L1_MEM_PORTS)"
VX_CFG_DXA_QUEUE_SIZE = 16
VX_CFG_DXA_MAX_INFLIGHT = 8

Expand Down
26 changes: 26 additions & 0 deletions ci/testcases/amo.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -54,3 +54,29 @@ tests:
shape: {cores: 4, l2cache: true, l3cache: true}
args: "-n8"
tier: full
# Shared-memory (LMEM) AMOADD.W: bank-serialized read-modify-write at the
# LMEM banks (tests/regression/smem_amo_mlp; closed-form old-value oracles).
- id: smem-contention
app: smem_amo_mlp
drivers: [simx, rtlsim]
args: "-m0 -n64"
touches: [hw/rtl/mem, sim/simx/mem]
- id: smem-directed
app: smem_amo_mlp
drivers: [simx, rtlsim]
args: "-m3"
touches: [hw/rtl/mem, sim/simx/mem]
- id: smem-banks
app: smem_amo_mlp
drivers: [simx]
args: "-m1 -n64"
touches: [hw/rtl/mem, sim/simx/mem]
tier: full
# 8 warps: deeper same-bank interleaving across warps.
- id: smem-warps8
app: smem_amo_mlp
drivers: [simx]
shape: {warps: 8}
args: "-m0 -n32"
touches: [hw/rtl/mem, sim/simx/mem]
tier: full
39 changes: 33 additions & 6 deletions ci/testcases/dxa.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -225,39 +225,39 @@ tests:
drivers:
- simx
app: sgemm2_dxa
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_UNITS=4
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_CORES=2
shape:
cores: 2
- id: sgemm2_dxa-6
via: blackbox
drivers:
- rtlsim
app: sgemm2_dxa
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_UNITS=4
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_CORES=2
shape:
cores: 2
- id: sgemm2_dxa-7
via: blackbox
drivers:
- simx
app: sgemm2_dxa
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_UNITS=4 -DVX_CFG_SOCKET_SIZE=2
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_CORES=2 -DVX_CFG_SOCKET_SIZE=2
shape:
cores: 4
- id: sgemm2_dxa-8
via: blackbox
drivers:
- rtlsim
app: sgemm2_dxa
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_UNITS=4 -DVX_CFG_SOCKET_SIZE=2
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_CORES=2 -DVX_CFG_SOCKET_SIZE=2
shape:
cores: 4
- id: sgemm2_dxa-9
via: blackbox
drivers:
- simx
app: sgemm2_dxa
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_UNITS=4
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_CORES=2
shape:
cores: 2
clusters: 2
Expand All @@ -266,7 +266,7 @@ tests:
drivers:
- rtlsim
app: sgemm2_dxa
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_UNITS=4
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_CORES=2
shape:
cores: 2
clusters: 2
Expand Down Expand Up @@ -437,3 +437,30 @@ tests:
app: sgemm_tcu_wg_dxa_mcast
args: -m 128 -n 128 -k 64
configs: -DVX_CFG_EXT_TCU_ENABLE -DVX_CFG_EXT_DXA_ENABLE
# ── Added with the socket-level DXA engine move ────────────────────────────
# Persistent-kernel lifetime: a resident CTA reissues DXA transfers across
# grid iterations. -b16 fits the default 4-warp x 4-thread single-core shape.
- id: dxa_copy_persist-1
via: blackbox
drivers:
- simx
app: dxa_copy_persist
configs: -DVX_CFG_EXT_DXA_ENABLE
args: -b16
# Double-buffered WGMMA+DXA pipeline.
- id: sgemm_tcu_wg_dxa_db-1
via: blackbox
drivers:
- simx
app: sgemm_tcu_wg_dxa_db
configs: -DVX_CFG_EXT_TCU_ENABLE -DVX_CFG_EXT_DXA_ENABLE
# Socket-level engine sizing: two workers per engine, two sockets.
- id: dxa_copy-socket2-workers2
via: blackbox
drivers:
- simx
app: dxa_copy
args: -d2
configs: -DVX_CFG_EXT_DXA_ENABLE -DVX_CFG_NUM_DXA_CORES=2 -DVX_CFG_SOCKET_SIZE=2
shape:
cores: 4
166 changes: 27 additions & 139 deletions hw/rtl/VX_cluster.sv
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@ module VX_cluster import VX_gpu_pkg::*;
cache_perf_t l2_perf;
sysmem_perf_t sysmem_perf_tmp;
`ifdef VX_CFG_EXT_DXA_ENABLE
dxa_perf_t dxa_core_perf;
dxa_perf_t per_socket_dxa_perf[NUM_SOCKETS];
`endif
`ifdef VX_CFG_EXT_TEX_ENABLE
tex_perf_t gfx_tex_perf;
Expand All @@ -75,9 +75,6 @@ module VX_cluster import VX_gpu_pkg::*;
always @(*) begin
sysmem_perf_tmp = sysmem_perf;
sysmem_perf_tmp.l2cache = l2_perf;
`ifdef VX_CFG_EXT_DXA_ENABLE
sysmem_perf_tmp.dxa = dxa_core_perf;
`endif
`ifdef VX_CFG_EXT_TEX_ENABLE
sysmem_perf_tmp.tex = gfx_tex_perf;
sysmem_perf_tmp.tcache = gfx_tcache_perf;
Expand Down Expand Up @@ -134,10 +131,10 @@ module VX_cluster import VX_gpu_pkg::*;
.TAG_WIDTH (L2_TAG_WIDTH)
) per_socket_mem_bus_if[L2_NUM_REQS]();

// Socket L1 output buses (pre-arb, original tag width)
// Socket outputs already include the socket-local L1-versus-DXA arb tag.
VX_mem_bus_if #(
.DATA_SIZE (`VX_CFG_L1_LINE_SIZE),
.TAG_WIDTH (L1_MEM_ARB_TAG_WIDTH)
.TAG_WIDTH (L2_TAG_WIDTH)
) socket_mem_bus_if[L2_SOCKET_REQS]();

`ifdef VX_CFG_EXT_TEX_ENABLE
Expand Down Expand Up @@ -182,27 +179,6 @@ module VX_cluster import VX_gpu_pkg::*;
) rtcache_l2_bus_if();
`endif

`ifdef VX_CFG_EXT_DXA_ENABLE
import VX_dxa_pkg::*;
VX_dxa_req_bus_if per_socket_dxa_req_bus_if[NUM_SOCKETS]();
VX_mem_bus_if #(
.DATA_SIZE (`VX_CFG_L1_LINE_SIZE),
.TAG_WIDTH (L1_MEM_ARB_TAG_WIDTH)
) dxa_gmem_bus_if[DXA_L2_GMEM_PORTS]();
VX_mem_bus_if #(
.DATA_SIZE (DXA_LMEM_WORD_SIZE),
.TAG_WIDTH (DXA_LMEM_OUT_TAG_W),
.ATTR_WIDTH (DXA_LMEM_ATTR_W),
.ADDR_WIDTH (DXA_LMEM_ADDR_W)
) dxa_lmem_bus_if[1]();
VX_mem_bus_if #(
.DATA_SIZE (DXA_LMEM_WORD_SIZE),
.TAG_WIDTH (DXA_LMEM_OUT_TAG_W),
.ATTR_WIDTH (DXA_LMEM_ATTR_W),
.ADDR_WIDTH (DXA_LMEM_ADDR_W)
) per_socket_dxa_lmem_bus_if[NUM_SOCKETS]();
`endif

VX_mem_bus_if #(
.DATA_SIZE (L2_SECTOR_SIZE),
.TAG_WIDTH (L2_MEM_TAG_WIDTH)
Expand Down Expand Up @@ -244,14 +220,15 @@ module VX_cluster import VX_gpu_pkg::*;
.mem_bus_if (l2_mem_bus_if)
);

// Cluster DCR distribution — declared before its consumers (DXA/sockets/gfx).
// Cluster DCR distribution. Each socket performs its own local fan-out to
// cores and the socket-owned DXA endpoint.
`ifdef EXT_GFX_ANY_ENABLE
localparam NUM_DCR_GFX = 1;
localparam DCR_GFX_IDX = NUM_SOCKETS + `VX_CFG_EXT_DXA_ENABLED;
localparam DCR_GFX_IDX = NUM_SOCKETS;
`else
localparam NUM_DCR_GFX = 0;
`endif
localparam NUM_DCR_REQS = NUM_SOCKETS + `VX_CFG_EXT_DXA_ENABLED + NUM_DCR_GFX;
localparam NUM_DCR_REQS = NUM_SOCKETS + NUM_DCR_GFX;
VX_dcr_bus_if per_socket_dcr_bus_if[NUM_DCR_REQS]();
VX_dcr_arb #(
.NUM_REQS (NUM_DCR_REQS),
Expand All @@ -263,111 +240,11 @@ module VX_cluster import VX_gpu_pkg::*;
.bus_out_if (per_socket_dcr_bus_if)
);

`ifdef VX_CFG_EXT_DXA_ENABLE
// Alias the DXA's DCR array element onto a scalar interface via signal
// assigns. A constant array index in a modport binding is rejected by
// sv2v; aliasing moves the index out of that context. Pure net joins.
VX_dcr_bus_if dxa_dcr_bus_if();
assign dxa_dcr_bus_if.req_valid = per_socket_dcr_bus_if[NUM_SOCKETS].req_valid;
assign dxa_dcr_bus_if.req_data = per_socket_dcr_bus_if[NUM_SOCKETS].req_data;
assign per_socket_dcr_bus_if[NUM_SOCKETS].rsp_valid = dxa_dcr_bus_if.rsp_valid;
assign per_socket_dcr_bus_if[NUM_SOCKETS].rsp_data = dxa_dcr_bus_if.rsp_data;

VX_dxa_core #(
.INSTANCE_ID (`SFORMATF(("%s-dxa-core", INSTANCE_ID))),
.NUM_REQS (NUM_SOCKETS),
.GMEM_OUT_PORTS (DXA_L2_GMEM_PORTS)
) dxa_core (
.clk (clk),
.reset (reset),
`ifdef PERF_ENABLE
.dxa_perf (dxa_core_perf),
`endif
.dcr_bus_if (dxa_dcr_bus_if),
.req_bus_if (per_socket_dxa_req_bus_if),
.smem_bus_if (dxa_lmem_bus_if),
.gmem_bus_if (dxa_gmem_bus_if),
`UNUSED_PIN (busy)
);

// Route DXA lmem requests to per-socket buses using core_id from tag.
// Tag value layout: {core_id[NC_BITS-1:0], engine_value[0]}
// socket_id = core_id[CORE_LOCAL_BITS +: SOCKET_SEL_BITS]
localparam DXA_LMEM_CORE_LOCAL_BITS = `CLOG2(`VX_CFG_SOCKET_SIZE);
localparam DXA_LMEM_SOCKET_SEL_BITS = `CLOG2(NUM_SOCKETS);
wire [`UP(DXA_LMEM_SOCKET_SEL_BITS)-1:0] dxa_lmem_socket_sel;
if (NUM_SOCKETS > 1) begin : g_dxa_lmem_sel
assign dxa_lmem_socket_sel = dxa_lmem_bus_if[0].req_data.tag.value[1 + DXA_LMEM_CORE_LOCAL_BITS +: DXA_LMEM_SOCKET_SEL_BITS];
end else begin : g_dxa_lmem_sel
assign dxa_lmem_socket_sel = '0;
end

VX_mem_bus_switch #(
.NUM_INPUTS (1),
.NUM_OUTPUTS (NUM_SOCKETS),
.DATA_SIZE (DXA_LMEM_WORD_SIZE),
.TAG_WIDTH (DXA_LMEM_OUT_TAG_W),
.ATTR_WIDTH (DXA_LMEM_ATTR_W),
.ADDR_WIDTH (DXA_LMEM_ADDR_W)
) dxa_lmem_socket_switch (
.clk (clk),
.reset (reset),
.bus_sel (dxa_lmem_socket_sel),
.bus_in_if (dxa_lmem_bus_if),
.bus_out_if (per_socket_dxa_lmem_bus_if)
);

// LSU+DXA arb: LSU gets priority ("P") to prevent DXA bulk traffic from
// starving core icache/dcache at L2 and the shared memory bus.
// Lower index = higher priority, so LSU is bound first.
VX_mem_bus_if #(
.DATA_SIZE (`VX_CFG_L1_LINE_SIZE),
.TAG_WIDTH (L1_MEM_ARB_TAG_WIDTH)
) l2_arb_in_if[2 * L2_SOCKET_REQS]();

// Bind LSU ports first (high priority, indices 0..L2_SOCKET_REQS-1)
for (genvar i = 0; i < L2_SOCKET_REQS; ++i) begin : g_lsu_l2_bind
`ASSIGN_VX_MEM_BUS_IF (l2_arb_in_if[i], socket_mem_bus_if[i]);
end

// Bind DXA gmem ports second (low priority, indices L2_SOCKET_REQS+..)
for (genvar i = 0; i < DXA_L2_GMEM_PORTS; ++i) begin : g_dxa_l2_bind
`ASSIGN_VX_MEM_BUS_IF (l2_arb_in_if[L2_SOCKET_REQS + i], dxa_gmem_bus_if[i]);
end

// Tie off unused DXA slots
for (genvar i = DXA_L2_GMEM_PORTS; i < L2_SOCKET_REQS; ++i) begin : g_dxa_l2_tieoff
assign l2_arb_in_if[L2_SOCKET_REQS + i].req_valid = 1'b0;
assign l2_arb_in_if[L2_SOCKET_REQS + i].req_data = '0;
assign l2_arb_in_if[L2_SOCKET_REQS + i].rsp_ready = 1'b1;
end

VX_mem_bus_arb #(
.NUM_INPUTS (2 * L2_SOCKET_REQS),
.NUM_OUTPUTS (L2_SOCKET_REQS),
.DATA_SIZE (`VX_CFG_L1_LINE_SIZE),
.TAG_WIDTH (L1_MEM_ARB_TAG_WIDTH),
.TAG_SEL_IDX (0),
.ARBITER ("P"),
// RSP_OUT_BUF=1: ensures the DXA port's buffer is empty (ready=1)
// after DXA completes, preventing stale sel_in from backpressuring
// the L2 bank when no response is pending.
.RSP_OUT_BUF (1)
) dxa_l2_priority_arb (
.clk (clk),
.reset (reset),
.bus_in_if (l2_arb_in_if),
// Drive only the socket ports; the upper indices of per_socket_mem_bus_if
// are the graphics cache ports, bound separately below.
.bus_out_if (per_socket_mem_bus_if[0 +: L2_SOCKET_REQS])
);

`else
// No DXA: direct socket → L2
for (genvar i = 0; i < L2_SOCKET_REQS; ++i) begin : g_no_dxa_l2
// Socket traffic, including socket-local DXA GMEM traffic, reaches L2 as
// one already-arbitrated stream per socket memory port.
for (genvar i = 0; i < L2_SOCKET_REQS; ++i) begin : g_socket_l2
`ASSIGN_VX_MEM_BUS_IF (per_socket_mem_bus_if[i], socket_mem_bus_if[i]);
end
`endif

for (genvar i = 0; i < L2_MEM_PORTS; ++i) begin : g_l2_mem_out
`ASSIGN_VX_MEM_BUS_IF (mem_bus_if[i], l2_mem_bus_if[i]);
Expand All @@ -393,6 +270,19 @@ module VX_cluster import VX_gpu_pkg::*;

for (genvar socket_id = 0; socket_id < NUM_SOCKETS; ++socket_id) begin : g_sockets

`ifdef PERF_ENABLE
// DXA counters are socket-owned. Feed each core only its socket's
// counters so runtime aggregation over one representative core per
// socket does not count the cluster total once for every socket.
sysmem_perf_t socket_sysmem_perf;
always @(*) begin
socket_sysmem_perf = sysmem_perf_tmp;
`ifdef VX_CFG_EXT_DXA_ENABLE
socket_sysmem_perf.dxa = per_socket_dxa_perf[socket_id];
`endif
end
`endif

VX_socket #(
.SOCKET_ID ((CLUSTER_ID * NUM_SOCKETS) + socket_id),
.INSTANCE_ID (`SFORMATF(("%s-socket%0d", INSTANCE_ID, socket_id)))
Expand All @@ -403,18 +293,16 @@ module VX_cluster import VX_gpu_pkg::*;
.reset (reset),

`ifdef PERF_ENABLE
.sysmem_perf (sysmem_perf_tmp),
.sysmem_perf (socket_sysmem_perf),
`ifdef VX_CFG_EXT_DXA_ENABLE
.dxa_perf (per_socket_dxa_perf[socket_id]),
`endif
`endif

.dcr_bus_if (per_socket_dcr_bus_if[socket_id]),

.mem_bus_if (socket_mem_bus_if[socket_id * L1_MEM_PORTS +: L1_MEM_PORTS]),

`ifdef VX_CFG_EXT_DXA_ENABLE
.dxa_req_bus_if (per_socket_dxa_req_bus_if[socket_id]),
.dxa_lmem_bus_if(per_socket_dxa_lmem_bus_if[socket_id +: 1]),
`endif

`ifdef VX_CFG_EXT_TEX_ENABLE
.per_socket_tex_bus_if (per_socket_tex_bus_if[socket_id]),
`endif
Expand Down
Loading
Loading