etomoscow's picture
download
raw
552 Bytes
from __future__ import annotations
from typing import Any
import torch
from torch import nn
def resolve_linear_layer(model: nn.Module, layer_name: str) -> nn.Linear:
layer = dict(model.named_modules()).get(layer_name)
if not isinstance(layer, nn.Linear):
raise TypeError(f"{layer_name!r} is not nn.Linear")
return layer
def move_batch_to_device(batch: Any, device: torch.device) -> Any:
if isinstance(batch, dict):
return {k: (v.to(device) if torch.is_tensor(v) else v) for k, v in batch.items()}
return batch

Xet Storage Details

Size:
552 Bytes
·
Xet hash:
11209671123fff40774d8359a0a52242f8785816ef409a60a4c8c1bb8b7f2fcc

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.