| 1 | class Swizzle: |
| 2 | def __init__(self, num_bits: int, num_base: int, num_shft: int): |
| 3 | self.num_bits = num_bits |
| 4 | self.num_base = num_base |
| 5 | self.num_shft = num_shft |
| 6 | |
| 7 | assert self.num_bits >= 0, "MBase must be positive." |
| 8 | assert self.num_bits >= 0, "BBits must be positive." |
| 9 | assert ( |
| 10 | abs(self.num_shft) >= self.num_bits |
| 11 | ), "abs(SShift) must be more than BBits." |
| 12 | |
| 13 | self.bit_msk = (1 << self.num_bits) - 1 |
| 14 | self.yyy_msk = self.bit_msk << (self.num_base + max(0, self.num_shft)) |
| 15 | self.zzz_msk = self.bit_msk << (self.num_base - min(0, self.num_shft)) |
| 16 | self.msk_sft = self.num_shft |
| 17 | |
| 18 | self.swizzle_code = self.yyy_msk | self.zzz_msk |
| 19 | |
| 20 | def apply(self, offset): |
| 21 | if self.msk_sft >= 0: |
| 22 | return offset ^ ((offset & self.yyy_msk) >> self.msk_sft) |
| 23 | else: |
| 24 | return offset ^ ((offset & self.yyy_msk) << -self.msk_sft) |
| 25 | |
| 26 | def __call__(self, offset): |
| 27 | return self.apply(offset) |
| 28 | |
| 29 | |
| 30 | def test_swizzle(): |