Views
No views yet
«оттенок в графической библиотеке тензоров»
PyTorch, вы можете скачать файл model.pth напрямую из этого репозитория и инициализировать архитектуру следующим кодом.1import torch
2from torch import nn
3from huggingface_hub import hf_hub_download
4from safetensors.torch import load_file
5
6# архитектура
7class VideoNet128(nn.Module):
8 def __init__(self):
9 super().__init__()
10 self.enc = nn.Sequential(
11 nn.Conv3d(3, 32, 3, padding=1), nn.ReLU(True), nn.MaxPool3d((1, 2, 2)),
12 nn.Conv3d(32, 64, 3, padding=1), nn.ReLU(True), nn.MaxPool3d((1, 2, 2)),
13 nn.Conv3d(64, 128, 3, padding=1), nn.ReLU(True), nn.MaxPool3d((1, 2, 2))
14 )
15 self.dec = nn.Sequential(
16 nn.Upsample(scale_factor=(1, 2, 2)), nn.Conv3d(128, 64, 3, padding=1), nn.ReLU(True),
17 nn.Upsample(scale_factor=(1, 2, 2)), nn.Conv3d(64, 32, 3, padding=1), nn.ReLU(True),
18 nn.Upsample(scale_factor=(1, 2, 2)), nn.Conv3d(32, 16, 3, padding=1), nn.ReLU(True),
19 nn.Conv3d(16, 3, 3, padding=1), nn.Tanh()
20 )
21 def forward(self, x):
22 return self.dec(self.enc(x))
23
24# выбираем девайс
25device = "cuda" if torch.cuda.is_available() else "cpu"
26
27# скачиваем веса HueGLoT в новом формате .safetensors
28weights_path = hf_hub_download(repo_id="prostochel097/HueGLoT", filename="model.safetensors")
29
30# инициализируем модель и загружаем безопасные веса
31model = VideoNet128().to(device)
32
33try:
34 # загружаем тензоры напрямую на нужное устройство
35 weights = load_file(weights_path, device=device)
36 model.load_state_dict(weights)
37 model.eval()
38 print("модель HueGLoT успешно загружена из .safetensors и готова к работе!")
39except Exception as e:
40 print(f"ошибка загрузки .safetensors: {e}!")
41 print("проверьте, что файл 'model.safetensors' залит в репозиторий hugging face.")