clef / code /models /common /tests /modules /test_tensor_utils.py
tt-hous's picture
Add files using upload-large-folder tool
2415c4c verified
Raw History Blame Contribute Delete
11.7 kB
# SPDX-FileCopyrightText: Β© 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
import json
import pytest
import torch
import ttnn
from models.common.tensor_utils import (
get_rot_transformation_mat,
pad_dim_to_size,
pad_to_shape,
parse_shard_dims_from_mesh_mapper_config,
program_config_to_dict,
program_config_to_str,
zeros_like_kv_cache,
zeros_like_paged_cache,
)
def test_pad_dim_to_size(expect_error):
"""Test the pad_dim_to_size utility function."""
# Test padding on last dimension
x = torch.randn(1, 1, 32, 100)
padded = pad_dim_to_size(x, dim=-1, size=128)
assert padded.shape == (1, 1, 32, 128)
# Original data should be preserved
assert torch.allclose(padded[:, :, :, :100], x)
# Padding should be zeros
assert torch.allclose(padded[:, :, :, 100:], torch.zeros(1, 1, 32, 28))
# Test no padding needed
x2 = torch.randn(1, 1, 32, 128)
padded2 = pad_dim_to_size(x2, dim=-1, size=128)
assert torch.equal(padded2, x2)
# Test padding on different dimension
x3 = torch.randn(1, 1, 24, 128)
padded3 = pad_dim_to_size(x3, dim=-2, size=32)
assert padded3.shape == (1, 1, 32, 128)
# Test error when target size is smaller
with expect_error(ValueError, "smaller than current size"):
pad_dim_to_size(x, dim=-1, size=50)
def test_pad_dim_to_size_positive_dim():
"""Test pad_dim_to_size with positive dimension index."""
x = torch.randn(2, 3, 4, 5)
# Pad dim=0
padded = pad_dim_to_size(x, dim=0, size=4)
assert padded.shape == (4, 3, 4, 5)
assert torch.equal(padded[:2], x)
# Pad dim=1
padded = pad_dim_to_size(x, dim=1, size=8)
assert padded.shape == (2, 8, 4, 5)
assert torch.equal(padded[:, :3], x)
def test_pad_to_shape():
"""Test the pad_to_shape utility function."""
# Pad multiple dimensions at once
x = torch.randn(1, 2, 24, 100)
padded = pad_to_shape(x, (1, 4, 32, 128))
assert padded.shape == (1, 4, 32, 128)
# Original data preserved
assert torch.allclose(padded[:, :2, :24, :100], x)
# Padding is zeros
assert torch.allclose(padded[:, 2:, :, :], torch.zeros(1, 2, 32, 128))
assert torch.allclose(padded[:, :, 24:, :], torch.zeros(1, 4, 8, 128))
assert torch.allclose(padded[:, :, :, 100:], torch.zeros(1, 4, 32, 28))
def test_pad_to_shape_no_op():
"""Test pad_to_shape returns same tensor when no padding needed."""
x = torch.randn(1, 2, 32, 128)
padded = pad_to_shape(x, (1, 2, 32, 128))
assert padded is x # Should be the exact same object
def test_pad_to_shape_single_dim():
"""Test pad_to_shape with only one dimension needing padding."""
x = torch.randn(2, 3, 4, 5)
padded = pad_to_shape(x, (2, 3, 4, 8))
assert padded.shape == (2, 3, 4, 8)
assert torch.equal(padded[:, :, :, :5], x)
def test_pad_to_shape_error_on_smaller_target(expect_error):
"""Test pad_to_shape raises error when target is smaller than source."""
x = torch.randn(2, 3, 4, 5)
with expect_error(ValueError, "smaller than current size"):
pad_to_shape(x, (2, 3, 4, 3))
def test_parse_shard_dims_from_mesh_mapper_config():
"""Test parsing shard dims from MeshMapperConfig repr.
This test will fail if TTNN changes the repr format, alerting us to update the parser.
"""
# Single shard dimension
config1 = ttnn.MeshMapperConfig(
placements=[ttnn.PlacementShard(-1)],
mesh_shape_override=ttnn.MeshShape([8]),
)
assert parse_shard_dims_from_mesh_mapper_config(config1) == [-1]
# Different shard dimension
config2 = ttnn.MeshMapperConfig(
placements=[ttnn.PlacementShard(-2)],
mesh_shape_override=ttnn.MeshShape([4]),
)
assert parse_shard_dims_from_mesh_mapper_config(config2) == [-2]
# Positive dimension
config3 = ttnn.MeshMapperConfig(
placements=[ttnn.PlacementShard(0)],
mesh_shape_override=ttnn.MeshShape([2]),
)
assert parse_shard_dims_from_mesh_mapper_config(config3) == [0]
# Two dimensions sharded (2D mesh)
config4 = ttnn.MeshMapperConfig(
placements=[ttnn.PlacementShard(-2), ttnn.PlacementShard(-1)],
mesh_shape_override=ttnn.MeshShape([2, 4]),
)
assert parse_shard_dims_from_mesh_mapper_config(config4) == [-2, -1]
# Mixed: one sharded, one replicated (only shard dims should be returned)
config5 = ttnn.MeshMapperConfig(
placements=[ttnn.PlacementReplicate(), ttnn.PlacementShard(-1)],
mesh_shape_override=ttnn.MeshShape([2, 4]),
)
assert parse_shard_dims_from_mesh_mapper_config(config5) == [-1]
# All replicated (no shard dims)
config6 = ttnn.MeshMapperConfig(
placements=[ttnn.PlacementReplicate()],
mesh_shape_override=ttnn.MeshShape([8]),
)
assert parse_shard_dims_from_mesh_mapper_config(config6) == []
def test_get_rot_transformation_mat_tile_size():
"""Verify decode transformation matrix is TILE_SIZE x TILE_SIZE with correct pattern."""
mat = get_rot_transformation_mat(dhead=32)
assert mat.shape == (1, 1, 32, 32)
# Permutation pattern: even→odd = +1, odd→even = -1
assert mat[0, 0, 0, 1].item() == 1.0
assert mat[0, 0, 1, 0].item() == -1.0
assert mat[0, 0, 2, 3].item() == 1.0
assert mat[0, 0, 3, 2].item() == -1.0
# Diagonal is zero
assert mat[0, 0, 0, 0].item() == 0.0
assert mat[0, 0, 1, 1].item() == 0.0
def test_get_rot_transformation_mat_large():
"""Verify the matrix works for arbitrary dhead (e.g., head_dim=128 for prefill)."""
mat = get_rot_transformation_mat(dhead=128)
assert mat.shape == (1, 1, 128, 128)
# Pattern extends to last pair
assert mat[0, 0, 126, 127].item() == 1.0
assert mat[0, 0, 127, 126].item() == -1.0
# Off-pattern entries are zero
assert mat[0, 0, 0, 2].item() == 0.0
assert mat[0, 0, 0, 3].item() == 0.0
def test_get_rot_transformation_mat():
"""
Test that get_rot_transformation_mat produces the correct rotation matrix for RoPE.
The rotation transformation matrix is used by ttnn.experimental.rotary_embedding_llama.
It has the pattern:
- rot_emb_matrix[i, i+1] = 1 for even i
- rot_emb_matrix[i+1, i] = -1 for even i
"""
result = get_rot_transformation_mat()
# Validate shape
assert result.shape == (1, 1, 32, 32), f"Expected shape (1, 1, 32, 32), got {result.shape}"
# Validate specific known values
# Position (0, 1) should be 1
assert result[0, 0, 0, 1].item() == pytest.approx(1.0)
# Position (1, 0) should be -1
assert result[0, 0, 1, 0].item() == pytest.approx(-1.0)
# Position (0, 0) should be 0
assert result[0, 0, 0, 0].item() == pytest.approx(0.0)
# Position (2, 3) should be 1
assert result[0, 0, 2, 3].item() == pytest.approx(1.0)
# Position (3, 2) should be -1
assert result[0, 0, 3, 2].item() == pytest.approx(-1.0)
# Position (30, 31) should be 1
assert result[0, 0, 30, 31].item() == pytest.approx(1.0)
# Position (31, 30) should be -1
assert result[0, 0, 31, 30].item() == pytest.approx(-1.0)
# Validate that non-pattern positions are 0
assert result[0, 0, 0, 2].item() == pytest.approx(0.0)
assert result[0, 0, 1, 1].item() == pytest.approx(0.0)
def test_zeros_like_kv_cache():
"""Test zeros_like_kv_cache creates correct shape tensor."""
batch_size, n_kv_heads, max_seq_len, head_dim = 32, 8, 2048, 128
result = zeros_like_kv_cache(batch_size, n_kv_heads, max_seq_len, head_dim)
assert result.shape == (batch_size, n_kv_heads, max_seq_len, head_dim)
assert result.dtype == torch.float32
assert torch.all(result == 0)
def test_zeros_like_paged_cache():
"""Test zeros_like_paged_cache creates correct shape tensor."""
from dataclasses import dataclass
@dataclass
class MockPagedConfig:
max_num_blocks: int = 64
block_size: int = 64
paged_config = MockPagedConfig()
n_kv_heads = 8
head_dim = 128
result = zeros_like_paged_cache(paged_config, n_kv_heads, head_dim)
assert result.shape == (paged_config.max_num_blocks, n_kv_heads, paged_config.block_size, head_dim)
assert result.dtype == torch.float32
assert torch.all(result == 0)
@pytest.mark.parametrize("num_workers_per_dram_bank", [1, 2])
def test_program_config_to_dict_with_to_json(num_workers_per_dram_bank):
"""Test program_config_to_dict for a config that has to_json (matmul configs)."""
cfg = ttnn.MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig(
in0_block_w=4,
per_core_M=1,
per_core_N=2,
num_workers_per_dram_bank=num_workers_per_dram_bank,
)
d = program_config_to_dict(cfg)
assert isinstance(d, dict)
assert d["type"] == "MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig"
assert d["in0_block_w"] == 4
assert d["per_core_M"] == 1
assert d["per_core_N"] == 2
assert d["num_workers_per_dram_bank"] == num_workers_per_dram_bank
assert "fused_activation" in d
assert f"num_workers_per_dram_bank={num_workers_per_dram_bank}" in repr(cfg)
json_str = json.dumps(d, sort_keys=True)
roundtrip = json.loads(json_str)
assert roundtrip == d
def test_program_config_to_dict_without_to_json():
"""Test program_config_to_dict for a config that lacks to_json (SDPAProgramConfig)."""
cfg = ttnn.SDPAProgramConfig(
compute_with_storage_grid_size=ttnn.CoreCoord(8, 8),
q_chunk_size=256,
k_chunk_size=256,
)
d = program_config_to_dict(cfg)
assert isinstance(d, dict)
assert d["type"] == "SDPAProgramConfig"
assert "repr" in d
assert "SDPAProgramConfig" in d["repr"]
assert "q_chunk_size=256" in d["repr"]
assert "k_chunk_size=256" in d["repr"]
def test_program_config_to_str():
"""Test program_config_to_str returns valid sorted JSON."""
cfg = ttnn.MatmulMultiCoreReuseMultiCastDRAMShardedProgramConfig(in0_block_w=2, per_core_M=3, per_core_N=4)
result = program_config_to_str(cfg)
parsed = json.loads(result)
assert parsed["in0_block_w"] == 2
assert parsed["per_core_M"] == 3
assert parsed["per_core_N"] == 4
assert result == json.dumps(parsed, sort_keys=True)
if __name__ == "__main__":
test_pad_dim_to_size()
print(" βœ“ test_pad_dim_to_size")
test_pad_dim_to_size_positive_dim()
print(" βœ“ test_pad_dim_to_size_positive_dim")
test_pad_to_shape()
print(" βœ“ test_pad_to_shape")
test_pad_to_shape_no_op()
print(" βœ“ test_pad_to_shape_no_op")
test_pad_to_shape_single_dim()
print(" βœ“ test_pad_to_shape_single_dim")
test_pad_to_shape_error_on_smaller_target()
print(" βœ“ test_pad_to_shape_error_on_smaller_target")
test_parse_shard_dims_from_mesh_mapper_config()
print(" βœ“ test_parse_shard_dims_from_mesh_mapper_config")
test_get_rot_transformation_mat_tile_size()
print(" βœ“ test_get_rot_transformation_mat_tile_size")
test_get_rot_transformation_mat_large()
print(" βœ“ test_get_rot_transformation_mat_large")
test_get_rot_transformation_mat()
print(" βœ“ test_get_rot_transformation_mat")
test_zeros_like_kv_cache()
print(" βœ“ test_zeros_like_kv_cache")
test_zeros_like_paged_cache()
print(" βœ“ test_zeros_like_paged_cache")
test_program_config_to_dict_with_to_json()
print(" βœ“ test_program_config_to_dict_with_to_json")
test_program_config_to_dict_without_to_json()
print(" βœ“ test_program_config_to_dict_without_to_json")
test_program_config_to_str()
print(" βœ“ test_program_config_to_str")
print("\nAll tensor_utils tests passed! βœ“")