from typing import List, Tuple import numpy as np import pandas as pd from sklearn.metrics.pairwise import cosine_similarity import json import sys from utils.snowflake.snowflake_client import get_snowflake_session # type: ignore from utils.util import get_environment, init_spark_context # type: ignore from awsglue.utils import getResolvedOptions # type: ignore args = getResolvedOptions( sys.argv, ["JOB_NAME", "partition_date", "partition_hour"], ) p_date = args["partition_date"] p_hour = args["partition_hour"] sc, glueContext, spark, job = init_spark_context() print(f"p_date: {p_date}, p_hour: {p_hour}") environment = get_environment() warehouse = "SUNO_PROD_GLUE_HOURLY_X_SMALL" if environment == "PROD" else "SUNO_STAGING_GLUE_HOURLY_X_SMALL" database = f"SUNO_{environment}" schema = environment def parse_array(value): """Parse Snowflake VARIANT/ARRAY to Python list""" if value is None: return [] if isinstance(value, str): try: return json.loads(value) except json.JSONDecodeError: # Try to handle string representation of lists return eval(value) if value.startswith('[') else [] elif isinstance(value, list): return value else: # Handle Snowflake array types or other iterables try: return list(value) except TypeError: return [] def load_genre_data_from_snowflake() -> Tuple[np.ndarray, np.ndarray, np.ndarray]: session = None try: session = get_snowflake_session(warehouse=warehouse, database=database, schema=schema) df = session.sql( """ SELECT LABEL, EMBEDDING, QUANTILE FROM GENRE_DATA """ ) rows = df.collect() if not rows: raise RuntimeError("No genre rows found in GENRE_DATA") genre_labels: List[str] = [] embeddings_list: List[np.ndarray] = [] quantiles_list: List[np.ndarray] = [] for r in rows: genre_labels.append(str(r["LABEL"])) # Convert VARIANT/ARRAY to numpy arrays emb = np.array(parse_array(r["EMBEDDING"]), dtype=float) q = np.array(parse_array(r["QUANTILE"]), dtype=float) # Validation if emb.size == 0: raise ValueError(f"Empty embedding for genre: {r['LABEL']}") if q.size == 0: raise ValueError(f"Empty quantile for genre: {r['LABEL']}") embeddings_list.append(emb) quantiles_list.append(q) # Stack embeddings to 2D (n_genres, embedding_dim) genre_embeddings = np.stack(embeddings_list, axis=0) # Quantiles: if per-genre vectors, stack to (n_genres, n_q) quantiles = np.stack(quantiles_list, axis=0) print(f"Loaded {len(genre_labels)} genres with embedding dim {genre_embeddings.shape[1]}") return np.array(genre_labels), genre_embeddings, quantiles finally: if session: session.close() def query_recent_clips_from_snowflake(): session = None try: session = get_snowflake_session(warehouse=warehouse, database=database, schema=schema) sql = f""" SELECT CLIP_ID, DITTO_GENRE_EMBEDDING, CREATED_AT AS CREATED_AT FROM RDS_CLIP_EMBEDDING WHERE P_DATE = '{p_date}' AND P_HOUR = {p_hour} AND DITTO_GENRE_EMBEDDING IS NOT NULL """ df = session.sql(sql) rows = df.collect() hits = [] for r in rows: emb = r["DITTO_GENRE_EMBEDDING"] emb_list = parse_array(emb) if not emb_list: # Skip if empty continue created_at_val = r["CREATED_AT"] hits.append({ "_id": str(r["CLIP_ID"]), "_source": { "ditto_genre_vector_v2": emb_list, "created_at": created_at_val, } }) print(f"Fetched {len(hits)} clips from Snowflake") return hits finally: if session: session.close() def get_predictions(CDF, temperature=1): """Convert CDF probabilities to genre predictions (descending order - best first)""" epsilon = 1e-6 P = np.exp(np.log(CDF + epsilon) / temperature) sorted_indices = P.argsort(axis=1)[:, ::-1] return sorted_indices def classify_clips(clip_hits, genre_labels, genre_embeddings, quantiles, max_n_matches=3, threshold=0.0): """Classify clips using the genre classification pipeline.""" print(f"Classifying {len(clip_hits)} clips...") clip_data = [] clip_embeddings = [] for hit in clip_hits: clip_id = hit["_id"] src = hit.get("_source", {}) embedding = src.get("ditto_genre_vector_v2") created_at = src.get("created_at") if embedding and len(embedding) > 0: clip_data.append({"clip_id": clip_id, "created_at": created_at}) clip_embeddings.append(embedding) if not clip_embeddings: print("No valid clip embeddings found!") return [] clip_embeddings_array = np.array(clip_embeddings, dtype=float) # Validate dimensions expected_dim = genre_embeddings.shape[1] if clip_embeddings_array.shape[1] != expected_dim: raise ValueError(f"Clip embedding dimension {clip_embeddings_array.shape[1]} does not match genre embedding dimension {expected_dim}") cosine_sim = cosine_similarity(clip_embeddings_array, genre_embeddings) # Compute CDF using quantiles: fraction of quantile thresholds exceeded by similarity cdf = np.mean(cosine_sim[:, :, None] > quantiles[None, :, :], axis=2) sorted_indices = get_predictions(cdf) results = [] for i, clip_info in enumerate(clip_data): sorted_idx = sorted_indices[i][:max_n_matches] genres = [] for genre_index in sorted_idx: if cosine_sim[i, genre_index] > threshold: genres.append(str(genre_labels[genre_index]).lower().strip()) if len(genres) == 0: genres = ["other"] results.append( { "clip_id": clip_info["clip_id"], "created_at": clip_info["created_at"], "genre_tags": genres, } ) return results def save_results_to_snowflake(results): session = None try: session = get_snowflake_session(warehouse=warehouse, database=database, schema=schema) # Normalize rows (same as before) rows = [] for r in results: tags = r.get("genre_tags", r.get("genre_tag", [])) if not isinstance(tags, list): tags = [str(tags)] dt_raw = r.get("created_at") dt = pd.to_datetime(dt_raw, errors="coerce", utc=True) # Derive safely from the Timestamp if isinstance(dt, pd.Timestamp) and not pd.isna(dt): created_at_iso = dt.isoformat() # string for TIMESTAMP_TZ cast p_date = dt.date().isoformat() # 'YYYY-MM-DD' p_hour = int(dt.hour) # 0..23 else: created_at_iso = None p_date = None p_hour = None rows.append({ "CLIP_ID": str(r["clip_id"]), "GENRE_TAG": ",".join(map(str, tags)), "CREATED_AT": created_at_iso, "P_DATE": p_date, "P_HOUR": p_hour, }) sf_df = pd.DataFrame(rows, columns=["CLIP_ID", "GENRE_TAG", "CREATED_AT", "P_DATE", "P_HOUR"]) tmp_table = f"{database}.{schema}.TMP_CLIP_GENRE_UPSERT" sp_df = session.create_dataframe(sf_df) sp_df.write.mode("overwrite").save_as_table(tmp_table, table_type="temporary") fq_target = 'clip_genre' # MERGE on CLIP_ID; derive P_DATE and P_HOUR from CREATED_AT merge_sql = f""" MERGE INTO {fq_target} AS T USING {tmp_table} AS S ON T.CLIP_ID = S.CLIP_ID and T.P_DATE = S.P_DATE and T.P_HOUR = S.P_HOUR WHEN MATCHED THEN UPDATE SET T.GENRE_TAG = S.GENRE_TAG, T.CREATED_AT = TO_TIMESTAMP_TZ(S.CREATED_AT), T.P_DATE = TO_DATE(TO_TIMESTAMP_TZ(S.CREATED_AT)), T.P_HOUR = EXTRACT(HOUR FROM TO_TIMESTAMP_TZ(S.CREATED_AT)) WHEN NOT MATCHED THEN INSERT (CLIP_ID, GENRE_TAG, CREATED_AT, P_DATE, P_HOUR) VALUES ( S.CLIP_ID, S.GENRE_TAG, TO_TIMESTAMP_TZ(S.CREATED_AT), TO_DATE(TO_TIMESTAMP_TZ(S.CREATED_AT)), EXTRACT(HOUR FROM TO_TIMESTAMP_TZ(S.CREATED_AT)) ); """ session.sql(merge_sql).collect() print(f"Upserted {len(sf_df)} rows into {fq_target}") finally: if session: session.close() # Main execution try: # Load genre data from Snowflake genre_labels, genre_embeddings, quantiles = load_genre_data_from_snowflake() # Fetch recent clip embeddings from Snowflake clip_hits = query_recent_clips_from_snowflake() if not clip_hits: print("No clips found in the past window!") raise Exception("No clips found in the past window!") # Classify results = classify_clips( clip_hits, genre_labels, genre_embeddings, quantiles, max_n_matches=3, threshold=0.0 ) # Save results to Snowflake if results: save_results_to_snowflake(results) print(f"Job completed successfully. Processed {len(results)} clips.") else: print("No results to save.") finally: job.commit()