leadr64/database
0
1import gc2import hashlib3import os4from glob import glob5from pathlib import Path6 7import librosa8import torch9from diskcache import Cache10from qdrant_client import QdrantClient11from qdrant_client.http import models12from tqdm import tqdm13from transformers import ClapModel, ClapProcessor14 15from s3_utils import s3_auth, upload_file_to_bucket16from dotenv import load_dotenv17load_dotenv()18 19# PARAMETERS #######################################################################################20CACHE_FOLDER = '/home/arthur/data/music/demo_audio_search/audio_embeddings_cache_individual/'21KAGGLE_DB_PATH = '/home/arthur/data/kaggle/park-spring-2023-music-genre-recognition/train/train'22AWS_ACCESS_KEY_ID = os.environ['AWS_ACCESS_KEY_ID']23AWS_SECRET_ACCESS_KEY = os.environ['AWS_SECRET_ACCESS_KEY']24S3_BUCKET = "synthia-research"25S3_FOLDER = "huggingface_spaces_demo"26AWS_REGION = "eu-west-3"27 28s3 = s3_auth(AWS_ACCESS_KEY_ID, AWS_SECRET_ACCESS_KEY, AWS_REGION)29 30 31# Functions utils ##################################################################################32def get_md5(fpath):33 with open(fpath, "rb") as f:34 file_hash = hashlib.md5()35 while chunk := f.read(8192):36 file_hash.update(chunk)37 return file_hash.hexdigest()38 39 40def get_audio_embedding(model, audio_file, cache):41 # Compute a unique hash for the audio file42 file_key = f"{model.config._name_or_path}" + get_md5(audio_file)43 if file_key in cache:44 # If the embedding for this file is cached, retrieve it45 embedding = cache[file_key]46 else:47 # Otherwise, compute the embedding and cache it48 y, sr = librosa.load(audio_file, sr=48000)49 inputs = processor(audios=y, sampling_rate=sr, return_tensors="pt")50 embedding = model.get_audio_features(**inputs)[0]51 gc.collect()52 torch.cuda.empty_cache()53 cache[file_key] = embedding54 return embedding55 56 57 58# ################## Loading the CLAP model ###################59# loading the model60print("[INFO] Loading the model...")61model_name = "laion/larger_clap_general"62model = ClapModel.from_pretrained(model_name)63processor = ClapProcessor.from_pretrained(model_name)64 65# Initialize the cache66os.makedirs(CACHE_FOLDER, exist_ok=True)67cache = Cache(CACHE_FOLDER)68 69# Creating a qdrant collection #####################################################################70client = QdrantClient(os.environ['QDRANT_URL'], api_key=os.environ['QDRANT_KEY'])71print("[INFO] Client created...")72 73print("[INFO] Creating qdrant data collection...")74if not client.collection_exists("demo_spaces_db"):75 client.create_collection(76 collection_name="demo_spaces_db",77 vectors_config=models.VectorParams(78 size=model.config.projection_dim,79 distance=models.Distance.COSINE80 ),81 )82 83# Embed the audio files !84audio_files = [p for p in glob(os.path.join(KAGGLE_DB_PATH, '*/*.wav'))]85chunk_size, idx = 1, 086total_chunks = int(len(audio_files) / chunk_size)87 88# Use tqdm for a progress bar89print("Uploading on DB + S3")90for i in tqdm(range(0, len(audio_files), chunk_size),91 desc="[INFO] Uploading data records to data collection..."):92 chunk = audio_files[i:i + chunk_size] # Get a chunk of audio files93 records = []94 for audio_file in chunk:95 embedding = get_audio_embedding(model, audio_file, cache)96 file_obj = open(audio_file, 'rb')97 s3key = f'{S3_FOLDER}/{Path(audio_file).name}'98 upload_file_to_bucket(s3, file_obj, S3_BUCKET, s3key)99 records.append(100 models.PointStruct(101 id=idx, vector=embedding,102 payload={103 "audio_path": audio_file,104 "audio_s3url": f"https://{S3_BUCKET}.s3.amazonaws.com/{s3key}",105 "style": audio_file.split('/')[-1]}106 )107 )108 f"Uploaded s3 file : {idx}"109 idx += 1110 client.upload_points(111 collection_name="demo_spaces_db",112 points=records113 )114print("[INFO] Successfully uploaded data records to data collection!")115 116 117# It's a good practice to close the cache when done118cache.close()