clef / code /models /common /helper_funcs.py
tt-hous's picture
Add files using upload-large-folder tool
be3ecc8 verified
Raw History Blame Contribute Delete
1.2 kB
# SPDX-FileCopyrightText: © 2023 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
import ttnn
def Linear(
in_features: int,
out_features: int,
weight: ttnn.Tensor,
bias: Optional[ttnn.Tensor] = None,
output_mem_config=ttnn.DRAM_MEMORY_CONFIG,
):
"""
Returns a function that performs a Linear operation with optional bias.
``weight`` must be tt_tensor.
"""
assert weight.padded_shape == [
1,
1,
out_features,
in_features,
], "weight does not have the expected shape"
if bias is not None:
assert bias.padded_shape[-1] == out_features, "bias does not have the expected shape"
weight = weight
bias = bias
weight_T = ttnn.transpose(weight, -2, -1)
def linear_(activation):
nonlocal bias
assert activation.padded_shape[-1] == in_features, "activation tensor do not have the expected shape"
if bias is not None and bias.get_layout() != ttnn.TILE_LAYOUT:
bias = ttnn.to_layout(bias, ttnn.TILE_LAYOUT)
return ttnn.linear(activation, weight_T, bias=bias, memory_config=output_mem_config)
return linear_