Skip to content
Merged
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
25 changes: 24 additions & 1 deletion src/target/sunmmio_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
#include <tvm/ffi/string.h>
#include <tvm/ir/expr.h>
#include <tvm/runtime/data_type.h>
#include <tvm/runtime/logging.h>
#include <tvm/target/target.h>

namespace tvm {
Expand All @@ -29,7 +30,7 @@ struct SunmmioTileProcessorConfig {
int register_bits;
int block_height;
int block_width;
/// Minimum byte-alignment for RSRAM tile rows (DMA constraint).
/// Minimum byte-alignment for RSRAM vector memory accesses.
int rsram_align_bytes;
};

Expand All @@ -56,6 +57,28 @@ inline bool IsSunmmioSramScope(const ffi::String &scope) {
scope == kSunmmioScopeRSRAM;
}

/*!
* \brief Convert an RSRAM byte-alignment requirement into element count.
*/
inline int GetSunmmioRsramAlignmentElems(int rsram_align_bytes,
DataType dtype) {
if (rsram_align_bytes <= 0) {
return 1;
}
ICHECK_EQ(dtype.lanes(), 1)
<< "Sunmmio RSRAM alignment expects scalar element dtypes, but got "
<< dtype << ".";
int element_bits = dtype.bits();
int align_bits = rsram_align_bytes * 8;
if (align_bits <= element_bits) {
return 1;
}
ICHECK_EQ(align_bits % element_bits, 0)
<< "RSRAM alignment " << rsram_align_bytes
<< " bytes is not divisible by element bit-width " << element_bits << ".";
return align_bits / element_bits;
}

} // namespace tl
} // namespace tvm

Expand Down
12 changes: 2 additions & 10 deletions src/tileview/tileview.cc
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ TileViewNode::TileViewNode(Array<PrimExpr> buffer_shape,

// Compute tiled_buffer_shape
// For each original dimension:
// - If tiled: replace with num_tiles (buffer_dim / tile_dim)
// - If tiled: replace with num_tiles (ceildiv(buffer_dim, tile_dim))
// - If not tiled: keep as is
// Then append all tile dimensions at the end

Expand All @@ -55,15 +55,7 @@ TileViewNode::TileViewNode(Array<PrimExpr> buffer_shape,
PrimExpr buf_dim = buffer_shape_[d];
PrimExpr tile_dim = tile_shape_[tile_idx];

// Check divisibility
PrimExpr remainder = floormod(buf_dim, tile_dim);
ICHECK(analyzer.CanProve(remainder == 0))
<< "Buffer dimension " << d << " (size=" << buf_dim
<< ") must be divisible by tile dimension " << tile_idx
<< " (size=" << tile_dim << ")";

// num_tiles = buf_dim / tile_dim
PrimExpr num_tiles = analyzer.Simplify(floordiv(buf_dim, tile_dim));
PrimExpr num_tiles = analyzer.Simplify(ceildiv(buf_dim, tile_dim));
tiled_shape.push_back(num_tiles);
} else {
tiled_shape.push_back(buffer_shape_[d]);
Expand Down
3 changes: 2 additions & 1 deletion src/tileview/tileview.h
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,8 @@ class TileView;
* A TileView captures:
* - tile_shape: The shape of each tile (e.g., (16, 32) for a 2D tile)
* - index_map: Which dimensions of the buffer are tiled
* - tiled_buffer_shape: The shape of tiled buffer
* - tiled_buffer_shape: The shape of tiled buffer, using ceildiv for tiled
* dimensions so tail tiles are represented explicitly.
*
* For Sunmmio target, the Tile unit processes fixed-size 2D tiles with
* constraints: width must be 32, height can be 8/16/32.
Expand Down
Loading
Loading