Skip to content

initializers

initializers ¤

DEFAULT_INITIALIZER_COMPILATION_RULES = {ConstantTensorInitializer: compile_constant_tensor_initializer, UniformInitializer: compile_uniform_initializer, NormalInitializer: compile_normal_initializer, DirichletInitializer: compile_dirichlet_initializer} module-attribute ¤

compile_constant_tensor_initializer(compiler, init) ¤

Source code in cirkit/backend/torch/rules/initializers.py
46
47
48
49
50
51
def compile_constant_tensor_initializer(
    compiler: "TorchCompiler", init: ConstantTensorInitializer
) -> InitializerFunc:
    if isinstance(init.value, np.ndarray):
        return functools.partial(copy_from_ndarray_, array=init.value)
    return functools.partial(torch.fill_, value=init.value)

compile_dirichlet_initializer(compiler, init) ¤

Source code in cirkit/backend/torch/rules/initializers.py
76
77
78
79
80
def compile_dirichlet_initializer(
    compiler: "TorchCompiler", init: DirichletInitializer
) -> InitializerFunc:
    axis = init.axis if init.axis < 0 else init.axis + 1
    return functools.partial(dirichlet_, alpha=init.alpha, dim=axis)

compile_normal_initializer(compiler, init) ¤

Source code in cirkit/backend/torch/rules/initializers.py
65
66
67
68
69
70
71
72
73
def compile_normal_initializer(
    compiler: "TorchCompiler", init: NormalInitializer
) -> InitializerFunc:
    if init.convex:
        return normalize_initializer(
            functools.partial(nn.init.normal_, mean=init.mean, std=init.stddev)
        )
    else:
        return functools.partial(nn.init.normal_, mean=init.mean, std=init.stddev)

compile_uniform_initializer(compiler, init) ¤

Source code in cirkit/backend/torch/rules/initializers.py
54
55
56
57
58
59
60
61
62
def compile_uniform_initializer(
    compiler: "TorchCompiler", init: UniformInitializer
) -> InitializerFunc:
    if init.convex:
        return normalize_initializer(
            functools.partial(nn.init.uniform_, a=init.a, b=init.b)
        )
    else:
        return functools.partial(nn.init.uniform_, a=init.a, b=init.b)

normalize_initializer(init) ¤

Modify an initializer to normalize the parameter to a convex sum.

Parameters:

Name Type Description Default
init Callable[[Tensor], Tensor]

initializer function (can be partial).

required

Returns:

Type Description

Normalized initializer function

Source code in cirkit/backend/torch/rules/initializers.py
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
def normalize_initializer(init: Callable[[torch.Tensor], torch.Tensor]):
    """Modify an initializer to normalize the parameter to a convex sum.

    Args:
        init: initializer function (can be partial).

    Returns:
        Normalized initializer function
    """

    def norm_init(tensor: torch.Tensor):
        init(tensor)
        tensor.copy_(tensor.softmax(dim=-1))
        return tensor

    return norm_init