diff --git a/hyperswap/README.md b/hyperswap/README.md index 50d43db..ab4b173 100644 --- a/hyperswap/README.md +++ b/hyperswap/README.md @@ -31,6 +31,7 @@ file_pattern = .datasets/vggface2/**/*.jpg convert_template = vggfacehq_512_to_arcface_128 transform_size = 256 usage_mode = both +usage_ratio = 1 batch_mode = same batch_ratio = 0.2 ``` diff --git a/hyperswap/config.ini b/hyperswap/config.ini index 6e709b1..965860a 100644 --- a/hyperswap/config.ini +++ b/hyperswap/config.ini @@ -3,6 +3,7 @@ file_pattern = convert_template = transform_size = usage_mode = +usage_ratio = batch_mode = batch_ratio = diff --git a/hyperswap/src/training.py b/hyperswap/src/training.py index f9f83d2..41a952c 100644 --- a/hyperswap/src/training.py +++ b/hyperswap/src/training.py @@ -208,6 +208,7 @@ def prepare_datasets(config_parser : ConfigParser) -> List[Dataset[Tensor]]: for config_section in config_parser.sections(): if config_section.startswith('training.dataset'): + config_usage_ratio = config_parser.getint(config_section, 'usage_ratio') __config_parser__ = deepcopy(config_parser) __config_parser__.remove_section(config_section) __config_parser__.add_section('training.dataset.current') @@ -215,7 +216,8 @@ def prepare_datasets(config_parser : ConfigParser) -> List[Dataset[Tensor]]: for key, value in config_parser.items(config_section): __config_parser__.set('training.dataset.current', key, value) - datasets.append(DynamicDataset(__config_parser__)) + dynamic_dataset = DynamicDataset(__config_parser__) + datasets.extend([ dynamic_dataset ] * config_usage_ratio) return datasets