Fixed bugs in auto-training process, removed deprecated arguments [skip ci]

This commit is contained in:
David Scripka 2023-10-07 21:27:43 -04:00
parent 3594e594c7
commit ef50fcf1c1
2 changed files with 35 additions and 34 deletions

View file

@ -2,16 +2,12 @@
# The name of the model (will be used when creating directoires and when saving the final .onnx and .tflite files)
model_name: "my_model"
# The target word/phrase to be detected by the model. Adding multiple unique words/phrases will
# still only train a binary model detection model, but it will activate on any one of the provided words/phrases.
target_phrase:
- "hey jarvis"
# The total length (in samples @ 16khz) of the positive clips used for training, after augmentations.
# Should be large enough to contain the entire target word/phrase, with at least 0.75 seconds
# before the start of the word/phrase, and at least 0.2 seconds after the end of the word/phrase.
total_length: 32000
# Specific phrases that you do *not* want the model to activate on, outside of those generated automatically via phoneme overlap
# This can be a good way to reduce false positives if you notice that, in practice, certain words or phrases are problematic
custom_negative_phrases: []
@ -94,12 +90,12 @@ layer_size: 32
# Define training parameters. The values below are recommended defaults for most applications,
# but unique deployment environments will likely require testing to determine which values
# are the most appropriate. Note that all "target_" values are determined from the validation data,
# and since early-stopping is utilized, the final performance of the trained model
# may be slighly overfit to the validation data.
# are the most appropriate.
steps: 50000 # the maximum number of steps when training the model
max_negative_weight: 1500 # the maximum weight to give negative samples during training to reduce false positives
target_accuracy: 0.7 # the target validation set accuracy for wake word/phrase detection
target_recall: 0.5 # the target validation recall for wake word/phrase detection
target_false_positives_per_hour: 0.2 # the maximum validation false positive rate per hour
# The maximum number of steps to train the model
steps: 50000
# The maximum negative weight and target false positives per hour, used to control the auto training process
# The target false positive rate may not be achieved, and adjusting the maximum negative weight may be necessary
max_negative_weight: 1500
target_false_positives_per_hour: 0.2

View file

@ -8,6 +8,7 @@ import sys
import tempfile
import uuid
import numpy as np
import scipy
import collections
import argparse
import logging
@ -192,7 +193,7 @@ class Model(nn.Module):
return averaged_model
def auto_train(self, X_train, X_val, false_positive_val_data, steps=50000, max_negative_weight=1000,
target_val_accuracy=0.7, target_val_recall=0.5, target_val_fp_per_hour=0.2):
target_fp_per_hour=0.2):
"""A sequence of training steps that produce relatively strong models
automatically, based on validation data and performance targets provided.
After training merges the best checkpoints and returns a single model.
@ -213,8 +214,7 @@ class Model(nn.Module):
max_steps=steps,
negative_weight_schedule=weights,
val_steps=val_steps, warmup_steps=steps//5,
hold_steps=steps//3, lr=lr, max_val_fp_per_hr=target_val_fp_per_hour, val_set_hrs=val_set_hrs,
target_val_accuracy=target_val_accuracy, target_val_recall=target_val_recall)
hold_steps=steps//3, lr=lr, val_set_hrs=val_set_hrs)
# Sequence 2
logging.info("#"*50 + "\nStarting training sequence 2...\n" + "#"*50)
@ -222,7 +222,7 @@ class Model(nn.Module):
steps = steps/10
# Adjust weights as needed based on false positive per hour performance from first sequence
if self.best_val_fp > target_val_fp_per_hour:
if self.best_val_fp > target_fp_per_hour:
max_negative_weight = max_negative_weight*2
logging.info("Increasing weight on negative examples to reduce false positives...")
@ -235,15 +235,14 @@ class Model(nn.Module):
max_steps=steps,
negative_weight_schedule=weights,
val_steps=val_steps, warmup_steps=steps//5,
hold_steps=steps//3, lr=lr, max_val_fp_per_hr=target_val_fp_per_hour, val_set_hrs=val_set_hrs,
target_val_accuracy=target_val_accuracy, target_val_recall=target_val_recall)
hold_steps=steps//3, lr=lr, val_set_hrs=val_set_hrs)
# Sequence 3
logging.info("#"*50 + "\nStarting training sequence 3...\n" + "#"*50)
lr = lr/10
# Adjust weights as needed based on false positive per hour performance from second sequence
if self.best_val_fp > target_val_fp_per_hour:
if self.best_val_fp > target_fp_per_hour:
max_negative_weight = max_negative_weight*2
logging.info("Increasing weight on negative examples to reduce false positives...")
@ -256,8 +255,7 @@ class Model(nn.Module):
max_steps=steps,
negative_weight_schedule=weights,
val_steps=val_steps, warmup_steps=steps//5,
hold_steps=steps//3, lr=lr, max_val_fp_per_hr=target_val_fp_per_hour, val_set_hrs=val_set_hrs,
target_val_accuracy=target_val_accuracy, target_val_recall=target_val_recall)
hold_steps=steps//3, lr=lr, val_set_hrs=val_set_hrs)
# Merge best models
logging.info("Merging checkpoints above the 90th percentile into single model...")
@ -265,12 +263,6 @@ class Model(nn.Module):
recall_percentile = np.percentile(self.history["val_recall"], 90)
fp_percentile = np.percentile(self.history["val_fp_per_hr"], 10)
# Show warning if 90th percentile is above/below targets
if accuracy_percentile < target_val_accuracy or recall_percentile < target_val_recall or \
fp_percentile > target_val_fp_per_hour:
logging.warning("\nWARNING!\nNo checkpoint with metrics better than the target values was found!\n"
"Consider generating more positive and negative examples for training or reducing target metrics.")
# Get models above the 90th percentile
models = []
for model, score in zip(self.best_models, self.best_model_scores):
@ -324,7 +316,7 @@ class Model(nn.Module):
def train_model(self, X, max_steps, warmup_steps, hold_steps, X_val=None,
false_positive_val_data=None,
negative_weight_schedule=[1],
val_steps=[250], lr=0.0001, max_val_fp_per_hr=0.1, target_val_accuracy=0.7, target_val_recall=0.5, val_set_hrs=1):
val_steps=[250], lr=0.0001, val_set_hrs=1):
# Move models and main class to target device
self.to(self.device)
self.model.to(self.device)
@ -560,7 +552,7 @@ if __name__ == '__main__':
for target_phrase in config["target_phrase"]:
adversarial_texts.extend(generate_adversarial_texts(
input_text=target_phrase,
N=config["n_samples"]//len(config["target_phrase"]),
N=config["n_samples"],
include_partial_phrase=1.0,
include_input_words=0.2))
generate_samples(text=adversarial_texts, max_samples=config["n_samples"]-n_current_samples,
@ -583,7 +575,7 @@ if __name__ == '__main__':
for target_phrase in config["target_phrase"]:
adversarial_texts.extend(generate_adversarial_texts(
input_text=target_phrase,
N=config["n_samples_val"]//len(config["target_phrase"]),
N=config["n_samples_val"],
include_partial_phrase=1.0,
include_input_words=0.2))
generate_samples(text=adversarial_texts, max_samples=config["n_samples_val"]-n_current_samples,
@ -594,6 +586,21 @@ if __name__ == '__main__':
else:
logging.warning(f"Skipping generation of negative clips for testing, as ~{config['n_samples_val']} already exist")
# Set the total length of the training clips based on the ~median generated clip duration, rounding to the nearest 1000 samples
# and setting to 32000 when the median + 750 ms is close to that, as it's a good default value
n = 50 # sample size
positive_clips = [str(i) for i in Path(positive_test_output_dir).glob("*.wav")]
duration_in_samples = []
for i in range(n):
sr, dat = scipy.io.wavfile.read(positive_clips[np.random.randint(0, len(positive_clips))])
duration_in_samples.append(len(dat))
config["total_length"] = int(round(np.median(duration_in_samples)/1000)*1000) + 12000 # add 750 ms to clip duration as buffer
if config["total_length"] < 32000:
config["total_length"] = 32000 # set a minimum of 32000 samples (2 seconds)
elif abs(config["total_length"] - 32000) <= 4000:
config["total_length"] = 32000
# Do Data Augmentation
if args.augment_clips is True:
if not os.path.exists(os.path.join(feature_save_dir, "positive_features_train.npy")) or args.overwrite is True:
@ -735,9 +742,7 @@ if __name__ == '__main__':
false_positive_val_data=X_val_fp,
steps=config["steps"],
max_negative_weight=config["max_negative_weight"],
target_val_accuracy=config["target_accuracy"],
target_val_recall=config["target_recall"],
target_val_fp_per_hour=config["target_false_positives_per_hour"]
target_fp_per_hour=config["target_false_positives_per_hour"],
)
# Export the trained model to onnx