| 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.