teticio/audio-diffusion
54
1import argparse2import os3import pickle4 5from datasets import load_dataset, load_from_disk6from tqdm.auto import tqdm7 8from audiodiffusion.audio_encoder import AudioEncoder9 10 11def main(args):12 audio_encoder = AudioEncoder.from_pretrained("teticio/audio-encoder")13 14 if args.dataset_name is not None:15 if os.path.exists(args.dataset_name):16 dataset = load_from_disk(args.dataset_name)["train"]17 else:18 dataset = load_dataset(19 args.dataset_name,20 args.dataset_config_name,21 cache_dir=args.cache_dir,22 use_auth_token=True if args.use_auth_token else None,23 split="train",24 )25 26 encodings = {}27 for audio_file in tqdm(dataset.to_pandas()["audio_file"].unique()):28 encodings[audio_file] = audio_encoder.encode([audio_file])29 pickle.dump(encodings, open(args.output_file, "wb"))30 31 32if __name__ == "__main__":33 parser = argparse.ArgumentParser(description="Create pickled audio encodings for dataset of audio files.")34 parser.add_argument("--dataset_name", type=str, default=None)35 parser.add_argument("--output_file", type=str, default="data/encodings.p")36 parser.add_argument("--use_auth_token", type=bool, default=False)37 args = parser.parse_args()38 main(args)39 