File size: 27,783 Bytes
9aa90e0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 | # SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import math
from typing import ClassVar
import torch
import ttnn
from .module import Module, Parameter
class RMSNorm(Module):
def __init__(
self,
embedding_dim,
norm_eps=1e-5,
norm_elementwise_affine=True,
bias=True,
mesh_device=None,
dtype=ttnn.bfloat16,
fused_activation=None,
):
super().__init__()
# https://github.com/tenstorrent/tt-metal/issues/31216
assert embedding_dim % 32 == 0, "embedding_dim must be divisible by tile size"
self.embedding_dim = embedding_dim
self.norm_eps = norm_eps
self.norm_elementwise_affine = norm_elementwise_affine
self.mesh_device = mesh_device
self.use_bias = norm_elementwise_affine and bias
self.fused_activation = fused_activation
if norm_elementwise_affine:
self.weight = Parameter(total_shape=[1, embedding_dim], device=mesh_device, dtype=dtype)
self.bias = Parameter(total_shape=[1, embedding_dim], device=mesh_device, dtype=dtype) if bias else None
else:
self.weight = None
self.bias = None
def forward(
self,
x: ttnn.Tensor,
*,
compute_kernel_config=None,
program_config: ttnn.LayerNormDefaultProgramConfig | ttnn.LayerNormShardedMultiCoreProgramConfig | None = None,
) -> ttnn.Tensor:
return ttnn.experimental.dit_rms_norm_unary_fused(
x,
weight=self.weight.data if self.weight is not None else None,
bias=self.bias.data if self.bias is not None else None,
epsilon=self.norm_eps,
program_config=program_config,
compute_kernel_config=compute_kernel_config,
activation=self.fused_activation,
)
def _prepare_torch_state(self, state: dict[str, torch.Tensor]) -> None:
if "weight" in state:
state["weight"] = state["weight"].unsqueeze(0)
if "bias" in state:
state["bias"] = state["bias"].unsqueeze(0)
class LayerNorm(Module):
def __init__(
self,
embedding_dim,
norm_eps=1e-5,
norm_elementwise_affine=True,
bias=True,
mesh_device=None,
use_row_major_workaround=False, # Issue #20789
):
super().__init__()
assert embedding_dim % 32 == 0, "embedding_dim must be divisible by tile size"
self.embedding_dim = embedding_dim
self.norm_eps = norm_eps
self.norm_elementwise_affine = norm_elementwise_affine
self.mesh_device = mesh_device
self.use_bias = norm_elementwise_affine and bias
self.use_row_major_workaround = use_row_major_workaround
self.compute_kernel_config = ttnn.init_device_compute_kernel_config(
self.mesh_device.arch(),
math_fidelity=ttnn.MathFidelity.HiFi4,
math_approx_mode=False,
fp32_dest_acc_en=True,
packer_l1_acc=False,
)
shape = [embedding_dim // 32, 32] if use_row_major_workaround else [1, embedding_dim]
layout = ttnn.ROW_MAJOR_LAYOUT if self.use_row_major_workaround else ttnn.TILE_LAYOUT
self.weight = (
Parameter(total_shape=shape, layout=layout, device=mesh_device)
if norm_elementwise_affine or self.use_row_major_workaround
else None
)
self.bias = Parameter(total_shape=shape, layout=layout, device=mesh_device) if self.use_bias else None
def _prepare_torch_state(self, state: dict[str, torch.Tensor]) -> None:
weight = state.pop("weight", None)
bias = state.pop("bias", None)
# When using the row-major workaround, ensure that dummy weight/bias are created
if self.use_row_major_workaround:
assert self.norm_elementwise_affine == (weight is not None)
assert self.use_bias == (bias is not None)
if weight is None:
weight = torch.ones(self.embedding_dim)
if self.use_bias and bias is None:
bias = torch.zeros(self.embedding_dim)
if weight is not None:
state["weight"] = weight.reshape(-1, 32) if self.use_row_major_workaround else weight.unsqueeze(0)
if bias is not None:
state["bias"] = bias.reshape(-1, 32) if self.use_row_major_workaround else bias.unsqueeze(0)
def forward(self, x: ttnn.Tensor) -> ttnn.Tensor:
return ttnn.layer_norm(
x,
weight=self.weight.data if self.weight is not None else None,
bias=self.bias.data if self.bias is not None else None,
epsilon=self.norm_eps,
compute_kernel_config=self.compute_kernel_config,
)
class DistributedRMSNorm(Module):
"""
Implements RMSNorm on an activation sharded on the reduction dimension.
"""
def __init__(
self,
embedding_dim,
norm_eps=1e-5,
norm_elementwise_affine=True,
bias=False,
mesh_axis=0,
mesh_device=None,
ccl_manager=None,
):
super().__init__()
assert not bias, "bias is not supported for DistributedRMSNorm"
self.embedding_dim = embedding_dim
self.norm_eps = norm_eps
self.norm_elementwise_affine = norm_elementwise_affine
self.mesh_axis = mesh_axis
self.mesh_device = mesh_device
self.ccl_manager = ccl_manager
self.mesh_width = tuple(mesh_device.shape)[mesh_axis]
self.TILE_SIZE = 32
self.compute_kernel_config = ttnn.init_device_compute_kernel_config(
self.mesh_device.arch(),
math_fidelity=ttnn.MathFidelity.HiFi4,
math_approx_mode=False,
fp32_dest_acc_en=True,
packer_l1_acc=False,
)
n = self.TILE_SIZE * self.mesh_width
# https://github.com/tenstorrent/tt-metal/issues/31216
assert embedding_dim % n == 0, "embedding_dim must be divisible by tile size times mesh width"
self.weight = (
Parameter(
total_shape=[1, embedding_dim],
layout=ttnn.TILE_LAYOUT,
device=mesh_device,
mesh_axes=[None, mesh_axis],
)
if norm_elementwise_affine
else None
)
def _prepare_torch_state(self, state: dict[str, torch.Tensor]) -> None:
if "weight" in state:
state["weight"] = state["weight"].reshape(1, self.embedding_dim)
def forward(
self,
x: ttnn.Tensor,
num_heads_per_device=1,
compute_kernel_config=None,
rope_cos=None,
rope_sin=None,
trans_mat=None,
dtype=None,
dynamic_weight=None,
dynamic_bias=None,
per_head_norm=False,
) -> ttnn.Tensor:
# per_head_norm selects the normalization semantics when the activation is
# head-split (num_heads_per_device > 1):
# True -> RMSNorm INDEPENDENTLY over each head's head_dim (per-head QK-norm, e.g.
# Ideogram4). No cross-device all-gather (each head is device-local).
# False -> one RMSNorm over the FULL per-device row, then reshape to heads
# (WAN2.2 / LTX "norm before splitting heads").
# The two are NOT equivalent. Default is whole-row (False); per-head models must opt
# in with per_head_norm=True at the call site (e.g. the Ideogram4 QK-norm).
expected_dim = self.embedding_dim // self.mesh_width
if x.shape[-1] != expected_dim:
msg = (
f"last dimension of input tensor with shape {tuple(x.shape)} should match "
f"embedding_dim / mesh_width = {expected_dim}"
)
raise ValueError(msg)
# Effective affine weight: the static per-channel weight, optionally modulated by a
# per-sample dynamic weight (e.g. an adaLN (1 + scale) factor). Folding it in here
# applies the modulation inside the fused op (fp32 internals) and removes the separate
# elementwise scale op the caller would otherwise need. RMSNorm has no bias term.
weight = self.weight.data if self.weight is not None else None
if dynamic_weight is not None:
weight = dynamic_weight if weight is None else ttnn.multiply(weight, dynamic_weight)
weight_key = tuple(weight.shape) if weight is not None else None
# dynamic_bias is the additive half of an adaLN modulation (the `shift`), folded into the same
# fused op as the scale so the caller needs no separate elementwise add. The device op accepts
# a per-token bias of shape [.., N, H] alongside a per-token weight; it requires a weight
# whenever a bias is given. Unlike the weight it does not reach create_stats_buffer, because
# only the weight can change the stats scratch geometry -- hence it is absent from the cache key.
if dynamic_bias is not None and weight is None:
msg = "dynamic_bias requires a weight: pass dynamic_weight or build the norm with affine=True"
raise ValueError(msg)
# Fused distributed RMSNorm device op (PRE sum-of-squares + fabric ring AG + POST
# normalize, with optional fused RoPE / per-head norm).
return ttnn.experimental.dit_fused_distributed_rmsnorm(
x,
self.mesh_axis,
self.mesh_device,
self.ccl_manager.get_ag_ping_pong_semaphore(self.mesh_axis),
topology=self.ccl_manager.topology,
persistent_output_buffer=self.ccl_manager.get_fused_norm_stats_buffer(
# Key includes everything that changes the stats-buffer geometry:
# shape, heads-per-device, RoPE presence, and weight presence (weight is
# forwarded to create_stats_buffer and affects its sizing). Guards against
# a shared-cache collision between two same-shape modules differing only
# in affine geometry.
# per_head_norm changes the stats geometry (per-head reduces locally -> no
# all-gather scratch), so it MUST be part of the key and forwarded to
# create_stats_buffer (which returns None for the per-head/local path).
("rms", tuple(x.shape), num_heads_per_device, per_head_norm, rope_cos is not None, weight_key),
lambda: ttnn.experimental.dit_fused_distributed_rmsnorm_create_stats_buffer(
x,
self.mesh_axis,
self.mesh_device,
num_heads_per_device=num_heads_per_device,
per_head_norm=per_head_norm,
num_links=self.ccl_manager.num_links,
weight=weight,
transformation_mat=trans_mat,
rope_cos=rope_cos,
rope_sin=rope_sin,
),
),
epsilon=self.norm_eps,
num_heads_per_device=num_heads_per_device,
per_head_norm=per_head_norm,
weight=weight,
bias=dynamic_bias,
compute_kernel_config=compute_kernel_config or self.compute_kernel_config,
num_preferred_links=self.ccl_manager.num_links, # must match create_stats_buffer above
transformation_mat=trans_mat,
rope_cos=rope_cos,
rope_sin=rope_sin,
dtype=dtype,
)
class DistributedLayerNorm(Module):
"""
Implements LayerNorm on an activation sharded on the reduction dimension.
"""
# The fused-op reciprocal LUT depends only on (device, width_per_device), so it is shared
# across DistributedLayerNorm instances of the same shape instead of each allocating its own.
_fused_ln_recip_cache: ClassVar[dict[tuple[int, int], "ttnn.Tensor"]] = {}
def __init__(
self,
embedding_dim,
norm_eps=1e-5,
norm_elementwise_affine=True,
bias=True,
mesh_axis=0,
mesh_device=None,
ccl_manager=None,
):
super().__init__()
self.embedding_dim = embedding_dim
self.norm_eps = norm_eps
self.norm_elementwise_affine = norm_elementwise_affine
self.use_bias = norm_elementwise_affine and bias
self.mesh_axis = mesh_axis
self.mesh_device = mesh_device
self.ccl_manager = ccl_manager
self.mesh_width = tuple(mesh_device.shape)[mesh_axis]
self.TILE_SIZE = 32
self.compute_kernel_config = ttnn.init_device_compute_kernel_config(
self.mesh_device.arch(),
math_fidelity=ttnn.MathFidelity.HiFi4,
math_approx_mode=False,
fp32_dest_acc_en=True,
packer_l1_acc=False,
)
n = self.TILE_SIZE * self.mesh_width
assert embedding_dim % n == 0, "embedding_dim must be divisible by tile size times mesh width"
# Static affine weight/bias are TILE [1, embedding_dim] sharded on the reduction axis —
# the broadcast layout the fused dit_fused_distributed_rmsnorm op consumes (per-device
# [1, H/mesh_width]). adaLN passes dynamic weight/bias at forward instead.
self.weight = (
Parameter(
total_shape=[1, embedding_dim], layout=ttnn.TILE_LAYOUT, mesh_axes=[None, mesh_axis], device=mesh_device
)
if self.norm_elementwise_affine
else None
)
self.bias = (
Parameter(
total_shape=[1, embedding_dim], layout=ttnn.TILE_LAYOUT, mesh_axes=[None, mesh_axis], device=mesh_device
)
if self.use_bias
else None
)
def _prepare_torch_state(self, state: dict[str, torch.Tensor]) -> None:
weight = state.pop("weight", None)
bias = state.pop("bias", None)
assert (weight is not None) == self.norm_elementwise_affine
assert (bias is not None) == self.use_bias
# TILE [1, embedding_dim] sharded on the reduction axis (matches DistributedRMSNorm).
if self.norm_elementwise_affine:
state["weight"] = weight.reshape(1, self.embedding_dim)
if self.use_bias:
state["bias"] = bias.reshape(1, self.embedding_dim)
def _ensure_fused_ln_recip(self, x: ttnn.Tensor) -> ttnn.Tensor:
"""Lazy-allocate the row-major fp32 reciprocal LUT the fused op consumes.
The fused op's reader NoC-reads a ROW_MAJOR [1,1,1,width_per_device] DRAM tensor
holding [1/1..1/width] (replicated per device). Cached per (device, width).
"""
width = self.embedding_dim // self.mesh_width
key = (self.mesh_device.id(), width)
cached = DistributedLayerNorm._fused_ln_recip_cache.get(key)
if cached is not None:
return cached
recip = torch.tensor([1.0 / (i + 1) for i in range(width)], dtype=torch.float32).reshape(1, 1, 1, width)
tensor = ttnn.from_torch(
recip,
dtype=ttnn.float32,
layout=ttnn.ROW_MAJOR_LAYOUT,
device=self.mesh_device,
mesh_mapper=ttnn.ReplicateTensorToMesh(self.mesh_device),
)
DistributedLayerNorm._fused_ln_recip_cache[key] = tensor
return tensor
def forward(
self, x: ttnn.Tensor, dynamic_weight=None, dynamic_bias=None, compute_kernel_config=None, dtype=None
) -> ttnn.Tensor:
assert (dynamic_weight is None) == (
dynamic_bias is None
), "dynamic_weight and dynamic_bias must be either both provided or both None"
if dynamic_weight is not None:
assert (
not self.norm_elementwise_affine
), "Module must not have weight and bias parameters when dynamic_weight and dynamic_bias are provided"
weight = dynamic_weight
bias = dynamic_bias
else:
weight = self.weight.data if self.weight is not None else None
bias = self.bias.data if self.bias is not None else None
# Fused Welford LayerNorm device op. weight/bias (static or adaLN, bf16 or fp32) are
# consumed natively in-op — fp32 affine keeps the modulation precision adaLN needs.
return ttnn.experimental.dit_fused_distributed_layernorm(
x,
self.mesh_axis,
self.mesh_device,
self.ccl_manager.get_ag_ping_pong_semaphore(self.mesh_axis),
topology=self.ccl_manager.topology,
persistent_output_buffer=self.ccl_manager.get_fused_norm_stats_buffer(
("ln", tuple(x.shape)),
lambda: ttnn.experimental.dit_fused_distributed_layernorm_create_stats_buffer(
x,
self.mesh_axis,
self.mesh_device,
num_links=self.ccl_manager.num_links,
),
),
epsilon=self.norm_eps,
weight=weight,
bias=bias,
compute_kernel_config=compute_kernel_config or self.compute_kernel_config,
num_preferred_links=self.ccl_manager.num_links, # must match create_stats_buffer above
dtype=dtype,
reciprocals=self._ensure_fused_ln_recip(x),
)
"""
Groupnorm that supports data parallel computation.
The number of channels and groups will be updated to match the distribution of the data across the mesh.
Set mesh_axis to None to disable data parallelism.
"""
class GroupNorm(Module):
default_num_out_blocks = {
# (Batch, Height, Width, Channels): num_out_blocks
} # overrides the num_out_blocks computed from the input shape.
def __init__(
self,
num_channels: int,
num_groups: int,
*,
eps: float = 1e-5,
mesh_device: ttnn.MeshDevice,
mesh_axis: int | None = None,
core_grid: ttnn.CoreGrid | None = None,
) -> None:
"""
Args:
num_channels: Number of channels in the input tensor.
num_groups: Number of groups.
eps: Epsilon value for numerical stability.
mesh_device: The device to use.
mesh_axis: The mesh axis to use for sharding.
core_grid: The core grid to use.
num_out_blocks: The number of output blocks to use.
"""
super().__init__()
self.eps = eps
self.mesh_device = mesh_device
self.mesh_axis = mesh_axis
self.num_devices = tuple(mesh_device.shape)[mesh_axis] if mesh_axis is not None else 1
self.core_grid = core_grid or ttnn.CoreGrid(x=8, y=8) # mesh_device.core_grid # Issue on 6U 8x9 grid
assert num_channels % num_groups == 0, "num_channels must be divisible by num_groups"
assert num_groups % self.num_devices == 0, "num_groups must be divisible by num_devices"
num_local_channels = num_channels // self.num_devices
num_padded_channels = math.ceil(num_local_channels / 32) * 32
assert num_padded_channels % num_local_channels == 0, "padded channels must be divisible by channels"
num_padded_groups = num_groups // self.num_devices * num_padded_channels // num_local_channels
self.num_virtual_cols = ttnn.operations.normalization.dram_group_norm_virtual_columns(
mesh_device.core_grid, num_padded_channels, num_padded_groups
)
weight_shape = [
self.num_devices,
1,
math.ceil(num_padded_channels // self.num_virtual_cols / 32) * self.num_virtual_cols,
32,
]
block_wt = ttnn.operations.normalization.find_max_tile_span(
num_padded_channels, num_padded_channels // num_padded_groups, 32
)
mask_shape = [1, num_padded_groups, 32, 32 * block_wt]
self.weight = Parameter(
total_shape=weight_shape,
layout=ttnn.ROW_MAJOR_LAYOUT,
mesh_axes=[mesh_axis, None, None, None],
device=self.mesh_device,
)
self.bias = Parameter(
total_shape=weight_shape,
layout=ttnn.ROW_MAJOR_LAYOUT,
mesh_axes=[mesh_axis, None, None, None],
device=self.mesh_device,
)
self.mask = Parameter(total_shape=mask_shape, device=self.mesh_device)
self.num_local_channels = num_local_channels
self.num_padded_channels = num_padded_channels
self.num_padded_groups = num_padded_groups
@classmethod
def from_torch(
cls,
torch_ref: torch.nn.GroupNorm,
*,
mesh_device: ttnn.MeshDevice,
mesh_axis: int | None = None,
core_grid: ttnn.CoreGrid | None = None,
) -> GroupNorm:
module = cls(
num_channels=torch_ref.num_channels,
num_groups=torch_ref.num_groups,
eps=torch_ref.eps,
mesh_device=mesh_device,
mesh_axis=mesh_axis,
core_grid=core_grid,
)
module.load_torch_state_dict(torch_ref.state_dict())
return module
def _prepare_torch_state(self, state: dict[str, torch.Tensor]) -> None:
if "weight" in state:
state["weight"] = self._prepare_param(state["weight"])
if "bias" in state:
state["bias"] = self._prepare_param(state["bias"])
input_mask = ttnn.create_group_norm_input_mask(
self.num_padded_channels, self.num_padded_groups, self.num_virtual_cols
)
state["mask"] = ttnn.to_torch(input_mask)
def _prepare_param(self, param: torch.Tensor) -> torch.Tensor:
expected_shape = (self.num_local_channels * self.num_devices,)
assert param.shape == expected_shape, f"expected shape {expected_shape}, got {param.shape}"
padding = self.num_padded_channels - self.num_local_channels
params = [torch.nn.functional.pad(t, (0, padding)) for t in param.chunk(self.num_devices)]
torch_sharded_lst = [
ttnn.create_group_norm_weight_bias_rm(t, self.num_padded_channels, self.num_virtual_cols) for t in params
]
return torch.cat(torch_sharded_lst, dim=0)
def forward(self, x: ttnn.Tensor, num_out_blocks=-1, compute_kernel_config=None) -> ttnn.Tensor:
batch_size, height, width, channels = x.shape
x = x.reshape([batch_size, 1, width * height, channels])
kwargs = dict(
weight=self.weight.data,
bias=self.bias.data,
input_mask=self.mask.data,
num_groups=self.num_padded_groups,
epsilon=self.eps,
core_grid=self.core_grid,
inplace=False,
num_out_blocks=num_out_blocks,
output_layout=ttnn.TILE_LAYOUT,
)
if compute_kernel_config is not None:
kwargs["compute_kernel_config"] = compute_kernel_config
x = ttnn.group_norm(x, **kwargs)
x = x.reshape([batch_size, height, width, channels])
return x
class GroupNorm3D(Module):
"""``torch.nn.GroupNorm(num_groups, num_channels)`` on a 5D BTHWC tensor, dims=3
semantics (statistics pool over ``channels-in-group x T x H x W`` per batch).
Routes through the DRAM-interleaved ``ttnn.group_norm``. The grid is pinned at
construction from ``input_nhw``/``num_batches`` via
``determine_expected_group_norm_dram_grid_size`` (uniform multicast groups; avoids
the mcast deadlock at small spatial sizes), so gamma/beta are static
``Parameter``s and round-trip through ``Module.save``/``load``.
"""
def __init__(
self,
num_channels: int,
num_groups: int,
*,
input_nhw: int,
num_batches: int = 1,
eps: float = 1e-5,
mesh_device: ttnn.MeshDevice,
dtype: ttnn.DataType = ttnn.bfloat16,
use_welford: bool = True,
) -> None:
super().__init__()
assert num_channels % 32 == 0, "num_channels must be divisible by tile size"
assert num_channels % num_groups == 0, "num_channels must be divisible by num_groups"
self.num_channels = num_channels
self.num_groups = num_groups
self.eps = eps
self.mesh_device = mesh_device
self.dtype = dtype
# Welford avoids the (E[x^2]-E[x]^2) precision loss when groups have nonzero mean.
self.use_welford = use_welford
self.input_nhw = input_nhw
self.num_batches = num_batches
self.core_grid = ttnn.determine_expected_group_norm_dram_grid_size(
device=mesh_device,
num_channels=num_channels,
num_groups=num_groups,
input_nhw=input_nhw,
num_batches=num_batches,
)
self.num_virtual_cols = ttnn.operations.normalization.dram_group_norm_virtual_columns(
self.core_grid, num_channels, num_groups
)
weight_shape = [1, 1, math.ceil(num_channels // self.num_virtual_cols / 32) * self.num_virtual_cols, 32]
self.weight = Parameter(total_shape=weight_shape, layout=ttnn.ROW_MAJOR_LAYOUT, dtype=dtype, device=mesh_device)
self.bias = Parameter(total_shape=weight_shape, layout=ttnn.ROW_MAJOR_LAYOUT, dtype=dtype, device=mesh_device)
@classmethod
def from_torch(
cls,
torch_ref: torch.nn.GroupNorm,
*,
input_nhw: int,
num_batches: int = 1,
mesh_device: ttnn.MeshDevice,
dtype: ttnn.DataType = ttnn.bfloat16,
) -> GroupNorm3D:
module = cls(
num_channels=torch_ref.num_channels,
num_groups=torch_ref.num_groups,
input_nhw=input_nhw,
num_batches=num_batches,
eps=torch_ref.eps,
mesh_device=mesh_device,
dtype=dtype,
)
module.load_torch_state_dict(torch_ref.state_dict())
return module
def _prepare_torch_state(self, state: dict[str, torch.Tensor]) -> None:
if "weight" in state:
state["weight"] = ttnn.create_group_norm_weight_bias_rm(
state["weight"], self.num_channels, self.num_virtual_cols
)
if "bias" in state:
state["bias"] = ttnn.create_group_norm_weight_bias_rm(
state["bias"], self.num_channels, self.num_virtual_cols
)
def forward(self, x_BTHWC: ttnn.Tensor) -> ttnn.Tensor:
B, T, H, W, C = x_BTHWC.shape
# dims=3: pool over (channels-in-group, T, H, W) — frames share one group statistic.
THW = T * H * W
assert B == self.num_batches and B * THW == self.input_nhw, (
f"GroupNorm3D built for input_nhw={self.input_nhw}, num_batches={self.num_batches}; "
f"got B={B}, T*H*W={THW}"
)
if x_BTHWC.layout != ttnn.ROW_MAJOR_LAYOUT:
x_BTHWC = ttnn.to_layout(x_BTHWC, ttnn.ROW_MAJOR_LAYOUT)
x = ttnn.tilize_with_zero_padding(ttnn.reshape(x_BTHWC, (B, 1, THW, C)), use_multicore=True)
out = ttnn.group_norm(
x,
num_groups=self.num_groups,
# -1 = built-in chunk heuristic. Default 1 (with pinned core_grid) overflows L1
# at large gathered spatial.
num_out_blocks=-1,
weight=self.weight.data,
bias=self.bias.data,
epsilon=self.eps,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
output_layout=ttnn.TILE_LAYOUT,
core_grid=self.core_grid,
inplace=False,
use_welford=self.use_welford,
)
out = ttnn.to_layout(out, ttnn.ROW_MAJOR_LAYOUT)
return ttnn.reshape(out, (B, T, H, W, C))
|