Create a AutoAugment transform :param config_str: String defining configuration of auto augmentation. Consists of multiple sections separated by dashes ('-'). The first section defines the AutoAugment policy (one of 'v0', 'v0r', 'original', 'originalr'). The remaining sections, not
(config_str, hparams)
| 530 | |
| 531 | |
| 532 | def auto_augment_transform(config_str, hparams): |
| 533 | """ |
| 534 | Create a AutoAugment transform |
| 535 | |
| 536 | :param config_str: String defining configuration of auto augmentation. Consists of multiple sections separated by |
| 537 | dashes ('-'). The first section defines the AutoAugment policy (one of 'v0', 'v0r', 'original', 'originalr'). |
| 538 | The remaining sections, not order sepecific determine |
| 539 | 'mstd' - float std deviation of magnitude noise applied |
| 540 | Ex 'original-mstd0.5' results in AutoAugment with original policy, magnitude_std 0.5 |
| 541 | |
| 542 | :param hparams: Other hparams (kwargs) for the AutoAugmentation scheme |
| 543 | |
| 544 | :return: A PyTorch compatible Transform |
| 545 | """ |
| 546 | config = config_str.split('-') |
| 547 | policy_name = config[0] |
| 548 | config = config[1:] |
| 549 | for c in config: |
| 550 | cs = re.split(r'(\d.*)', c) |
| 551 | if len(cs) < 2: |
| 552 | continue |
| 553 | key, val = cs[:2] |
| 554 | if key == 'mstd': |
| 555 | # noise param injected via hparams for now |
| 556 | hparams.setdefault('magnitude_std', float(val)) |
| 557 | else: |
| 558 | assert False, 'Unknown AutoAugment config section' |
| 559 | aa_policy = auto_augment_policy(policy_name, hparams=hparams) |
| 560 | return AutoAugment(aa_policy) |
| 561 | |
| 562 | |
| 563 | _RAND_TRANSFORMS = [ |
no test coverage detected