File size: 1,195 Bytes
be3ecc8 | 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 | # 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_
|