Team Ai
Apppublic

leadr64/database

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
database.py118 linesDownload Raw Back to root
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()