More tweaks

This commit is contained in:
henryruhs
2025-02-21 09:45:44 +01:00
parent 4078681031
commit c1bed34c27
3 changed files with 10 additions and 8 deletions
+1 -1
View File
@@ -27,7 +27,7 @@ This `config.ini` utilizes the MegaFace dataset to train the Embedding Converter
```
[training.dataset]
file_pattern = .datasets/images/{}/*.*g
file_pattern = .datasets/megaface/**/*.jpg
```
```
+4 -4
View File
@@ -11,17 +11,17 @@ from .types import Batch
class DynamicDataset(Dataset[Tensor]):
def __init__(self, dataset_file_pattern : str) -> None:
self.image_paths = glob.glob(dataset_file_pattern)
def __init__(self, file_pattern : str) -> None:
self.file_paths = glob.glob(file_pattern)
self.transforms = self.compose_transforms()
def __getitem__(self, index : int) -> Batch:
image_path = random.choice(self.image_paths)
image_path = random.choice(self.file_paths)
vision_frame = cv2.imread(image_path)
return self.transforms(vision_frame)
def __len__(self) -> int:
return len(self.image_paths)
return len(self.file_paths)
@staticmethod
def compose_transforms() -> transforms:
+5 -3
View File
@@ -24,13 +24,15 @@ class EmbeddingConverterTrainer(lightning.LightningModule):
super(EmbeddingConverterTrainer, self).__init__()
source_path = CONFIG.get('training.model', 'source_path')
target_path = CONFIG.get('training.model', 'target_path')
self.lr = CONFIG.getfloat('training.trainer', 'learning_rate')
learning_rate = CONFIG.getfloat('training.trainer', 'learning_rate')
self.embedding_converter = EmbeddingConverter()
self.source_embedder = torch.jit.load(source_path) # type:ignore[no-untyped-call]
self.target_embedder = torch.jit.load(target_path) # type:ignore[no-untyped-call]
self.source_embedder.eval()
self.target_embedder.eval()
self.mse_loss = nn.MSELoss()
self.lr = learning_rate
def forward(self, source_embedding : Embedding) -> Embedding:
return self.embedding_converter(source_embedding)
@@ -84,8 +86,8 @@ def create_loaders(dataset : Dataset[Tensor]) -> Tuple[DataLoader[Tensor], DataL
def split_dataset(dataset : Dataset[Tensor]) -> Tuple[Dataset[Tensor], Dataset[Tensor]]:
loader_split_ratio = CONFIG.getfloat('training.loader', 'split_ratio')
dataset_size = len(dataset) # type:ignore[arg-type]
training_size = dataset_size * loader_split_ratio
validation_size = dataset_size - training_size
training_size = int(dataset_size * loader_split_ratio)
validation_size = int(dataset_size - training_size)
training_dataset, validate_dataset = random_split(dataset, [ training_size, validation_size ])
return training_dataset, validate_dataset