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_