GPUCodeForces/S1 codes/Ljy123_#112/torchcode.py

29 lines
867 B
Python

import torch
import torch.nn as nn
class Model(nn.Module):
def __init__(self, scale: torch.Tensor, bias: torch.Tensor, alpha: float, beta: float):
super(Model, self).__init__()
self.register_buffer("scale", scale)
self.register_buffer("bias", bias)
self.register_buffer("alpha", torch.tensor(float(alpha), dtype=torch.float32))
self.register_buffer("beta", torch.tensor(float(beta), dtype=torch.float32))
def forward(self, x: torch.Tensor) -> torch.Tensor:
z = x * self.scale + self.bias
m = z / (1.0 + z * z)
g = torch.sigmoid(self.alpha * m + self.beta)
return x * g
batch_size = 16
dim = 16384
def get_inputs():
x = torch.randn(batch_size, dim)
return [x]
def get_init_inputs():
scale = torch.randn(dim)
bias = torch.randn(dim)
return [scale, bias, 1.0, 0.0]