{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":31254,"databundleVersionId":3103714,"sourceType":"competition"}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !pip install PyMySQL==1.1.0\n# !pip install openai==1.27.0\n# !pip install SQLAlchemy==2.0.30\n# !pip install tidb-vector==0.0.9\n# !pip install git+https://github.com/openai/CLIP.git\n# !pip install -U sentence-transformers","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-17T08:36:08.57331Z","iopub.execute_input":"2024-08-17T08:36:08.573622Z","iopub.status.idle":"2024-08-17T08:36:08.578244Z","shell.execute_reply.started":"2024-08-17T08:36:08.573595Z","shell.execute_reply":"2024-08-17T08:36:08.577385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sqlalchemy import Column, Integer, String, Text, create_engine, URL\nfrom sqlalchemy.orm import Session, declarative_base\nfrom tidb_vector.sqlalchemy import VectorType","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:36:08.583249Z","iopub.execute_input":"2024-08-17T08:36:08.583599Z","iopub.status.idle":"2024-08-17T08:36:08.832933Z","shell.execute_reply.started":"2024-08-17T08:36:08.583567Z","shell.execute_reply":"2024-08-17T08:36:08.83149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TIBD_HOST = \"gateway01.ap-southeast-1.prod.aws.tidbcloud.com\"\nPORT = 4000\nUSERNAME = \"44D4gZUd6CvBDnH.root\"\nPASSWORD = \"FBYgZ15057l9q81H\"\nDATABASE = \"test\"","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:36:08.834753Z","iopub.execute_input":"2024-08-17T08:36:08.835589Z","iopub.status.idle":"2024-08-17T08:36:08.840713Z","shell.execute_reply.started":"2024-08-17T08:36:08.835552Z","shell.execute_reply":"2024-08-17T08:36:08.839662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_db_url():\n    return URL(\n        drivername=\"mysql+pymysql\",\n        username=USERNAME,\n        password=PASSWORD,\n        host=TIBD_HOST,\n        port=PORT,\n        database=DATABASE,\n        query={\"ssl_verify_cert\": True, \"ssl_verify_identity\": True},\n    )\n\nengine = create_engine(get_db_url(), pool_recycle=300)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:36:08.845021Z","iopub.execute_input":"2024-08-17T08:36:08.845932Z","iopub.status.idle":"2024-08-17T08:36:08.953022Z","shell.execute_reply.started":"2024-08-17T08:36:08.845891Z","shell.execute_reply":"2024-08-17T08:36:08.952088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom tqdm.auto import tqdm \n\ncsv_path = \"/kaggle/input/h-and-m-personalized-fashion-recommendations/articles.csv\"\nimage_dir =  \"/kaggle/input/h-and-m-personalized-fashion-recommendations/images\"\n\ndf = pd.read_csv(csv_path)\n\n# Now make some filtering in the df\n\nrequired_columns = [\n    \"article_id\",\n    \n    # name of the product section \n    \"prod_name\",\n    \"product_type_name\",\n    \"product_group_name\",\n    \n    # Department\n    \"department_name\",\n    \"index_name\",\n    \"section_name\",\n    \n    # Detail Desc\n    \"detail_desc\",\n    \n    # Color section\n    \"graphical_appearance_name\",\n    \"colour_group_name\", \n    \"perceived_colour_value_name\",\n    \n]\n\ndf = df[required_columns]\n\n\n# Now add the images for the corresponding article id\n\nimage_path_dict = {}\nfor root, dirs, files in tqdm(os.walk(image_dir), total=len(os.listdir(image_dir))):\n    for file in files:\n        if file.endswith('.jpg'):\n            full_path = os.path.join(root, file)\n            article_id = str(int(file.split('.')[0]))\n            image_path_dict[article_id] = full_path\n\ndf['image_path'] = df['article_id'].astype(str).map(image_path_dict)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:36:08.954381Z","iopub.execute_input":"2024-08-17T08:36:08.954946Z","iopub.status.idle":"2024-08-17T08:37:32.243591Z","shell.execute_reply.started":"2024-08-17T08:36:08.954921Z","shell.execute_reply":"2024-08-17T08:37:32.242748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"description_template = \"\"\"\nProduct: {prod_name} | {product_type_name} | {product_group_name}\nDepartment: {department_name} | {index_name} | {section_name}\nProperties: {graphical_appearance_name} | {colour_group_name} | {perceived_colour_value_name}\nDescription: {detail_desc}\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:37:32.244908Z","iopub.execute_input":"2024-08-17T08:37:32.245287Z","iopub.status.idle":"2024-08-17T08:37:32.249925Z","shell.execute_reply.started":"2024-08-17T08:37:32.245252Z","shell.execute_reply":"2024-08-17T08:37:32.249059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def generate_description(row):\n    return description_template.format(\n        prod_name=row['prod_name'],\n        product_type_name=row['product_type_name'],\n        product_group_name=row['product_group_name'],\n        department_name=row['department_name'],\n        index_name=row['index_name'],\n        section_name=row['section_name'],\n        graphical_appearance_name=row['graphical_appearance_name'],\n        colour_group_name=row['colour_group_name'],\n        perceived_colour_value_name=row['perceived_colour_value_name'],\n        detail_desc=row['detail_desc']\n    ).strip()\n\ndf['desc'] = df.apply(generate_description, axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:37:32.251304Z","iopub.execute_input":"2024-08-17T08:37:32.251659Z","iopub.status.idle":"2024-08-17T08:37:38.170231Z","shell.execute_reply.started":"2024-08-17T08:37:32.251636Z","shell.execute_reply":"2024-08-17T08:37:38.169454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(df[\"desc\"][0])","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:37:38.171524Z","iopub.execute_input":"2024-08-17T08:37:38.171792Z","iopub.status.idle":"2024-08-17T08:37:38.17718Z","shell.execute_reply.started":"2024-08-17T08:37:38.17177Z","shell.execute_reply":"2024-08-17T08:37:38.176403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Now need to make the image and text embeddings\nimport torch\nfrom PIL import Image\nfrom torchvision.transforms import Compose, Resize, CenterCrop, ToTensor, Normalize\nfrom clip import clip\n\nmodel_name=\"ViT-B/32\"\ndevice=\"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel, preprocess = clip.load(model_name, device=device)\n\ndef get_image_embedding(image_path):\n    # Load the CLIP model\n    # Load and preprocess the image\n    image = Image.open(image_path).convert(\"RGB\")\n    image_input = preprocess(image).unsqueeze(0).to(device)\n    \n    # Generate the embedding\n    with torch.no_grad():\n        image_features = model.encode_image(image_input)\n    \n    # Normalize the embedding\n    image_embedding = image_features / image_features.norm(dim=-1, keepdim=True)\n    \n    return image_embedding.cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:37:38.178532Z","iopub.execute_input":"2024-08-17T08:37:38.178855Z","iopub.status.idle":"2024-08-17T08:37:45.931905Z","shell.execute_reply.started":"2024-08-17T08:37:38.17883Z","shell.execute_reply":"2024-08-17T08:37:45.931014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sentence_transformers import SentenceTransformer\n\n# Note:\n# There are two types of embedding: passage embedding and query embedding \n\ntext_embedding_model = SentenceTransformer(\n    \"dunzhang/stella_en_400M_v5\", trust_remote_code=True\n).cuda()\n\n\ndef get_text_embedding(text: str):\n    docs = [text]\n    doc_embeddings = text_embedding_model.encode(docs) \n    return doc_embeddings","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:38:11.620215Z","iopub.execute_input":"2024-08-17T08:38:11.620638Z","iopub.status.idle":"2024-08-17T08:38:18.902206Z","shell.execute_reply.started":"2024-08-17T08:38:11.620608Z","shell.execute_reply":"2024-08-17T08:38:18.901385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Base = declarative_base()\n\ndim_image_embedding_model = 512\ndim_text_embedding_model = 1024\n\nclass HandMProductEntity(Base):\n    __tablename__ = \"product\"\n    \n    article_id = Column(Integer, primary_key=True)\n    \n    prod_name = Column(Text)\n    product_type_name = Column(Text)\n    product_group_name = Column(Text)\n    \n    department_name = Column(Text)\n    index_name = Column(Text)\n    section_name = Column(Text)\n    \n    detail_desc = Column(Text)\n    \n    graphical_appearance_name = Column(Text)\n    colour_group_name = Column(Text)\n    perceived_colour_value_name = Column(Text)\n    \n    text_embedding = Column(\n        VectorType(dim=dim_text_embedding_model),\n        comment=\"hnsw(distance=l2)\"\n    )\n    image_embedding = Column(\n        VectorType(dim=dim_image_embedding_model),\n        comment=\"hnsw(distance=l2)\"\n    )","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:38:22.156907Z","iopub.execute_input":"2024-08-17T08:38:22.158448Z","iopub.status.idle":"2024-08-17T08:38:22.175164Z","shell.execute_reply.started":"2024-08-17T08:38:22.158381Z","shell.execute_reply":"2024-08-17T08:38:22.174006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Base.metadata.create_all(engine)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:37:46.427056Z","iopub.status.idle":"2024-08-17T08:37:46.42738Z","shell.execute_reply.started":"2024-08-17T08:37:46.427218Z","shell.execute_reply":"2024-08-17T08:37:46.427231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tqdm.pandas()\n\ndef fix_image_embeddings(data):\n    try:\n        return get_image_embedding(data).squeeze().tolist()\n    except Exception:\n        return [0.0] * 512\n\ndef fix_text_embeddings(data):\n    return get_text_embedding(data).squeeze().tolist()\n\nfirst_batch = df.iloc[:1000]\nfirst_batch[\"text_embedding\"] = first_batch[\"desc\"].progress_apply(fix_text_embeddings)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:50:20.372427Z","iopub.execute_input":"2024-08-17T08:50:20.37364Z","iopub.status.idle":"2024-08-17T08:51:11.758366Z","shell.execute_reply.started":"2024-08-17T08:50:20.373603Z","shell.execute_reply":"2024-08-17T08:51:11.757443Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"first_batch[\"image_embedding\"] = first_batch[\"image_path\"].progress_apply(fix_image_embeddings)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:51:18.160445Z","iopub.execute_input":"2024-08-17T08:51:18.160831Z","iopub.status.idle":"2024-08-17T08:52:06.129125Z","shell.execute_reply.started":"2024-08-17T08:51:18.1608Z","shell.execute_reply":"2024-08-17T08:52:06.128143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = 0\n\nfor _, row in first_batch.iterrows():\n    if row[\"image_embedding\"] == [0.0]*512: i += 1 \n\nprint(i)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:54:09.472787Z","iopub.execute_input":"2024-08-17T08:54:09.473767Z","iopub.status.idle":"2024-08-17T08:54:09.555636Z","shell.execute_reply.started":"2024-08-17T08:54:09.473723Z","shell.execute_reply":"2024-08-17T08:54:09.554634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Uploading everything to TIDB cloud\n\n\nfor _, row in tqdm(first_batch.iterrows(), total=len(first_batch)):\n    try:\n         with Session(engine) as session:\n            session.add(HandMProductEntity(\n                article_id=row[\"article_id\"],\n                \n                prod_name=row[\"prod_name\"],\n                product_type_name=row[\"product_type_name\"],\n                product_group_name=row[\"product_group_name\"],\n                \n                department_name=row[\"department_name\"],\n                index_name=row[\"index_name\"],\n                section_name=row[\"section_name\"],\n                \n                detail_desc=row[\"detail_desc\"],\n                \n                graphical_appearance_name=row[\"graphical_appearance_name\"],\n                colour_group_name=row[\"colour_group_name\"],\n                perceived_colour_value_name=row[\"perceived_colour_value_name\"],\n                \n                text_embedding=row[\"text_embedding\"],\n                image_embedding=row[\"image_embedding\"]\n            ))\n            session.commit()\n    except Exception as e:\n        print(f\"Hit exception, skipping ... Exception: {e}\")\n        continue ","metadata":{"execution":{"iopub.status.busy":"2024-08-17T08:57:04.063492Z","iopub.execute_input":"2024-08-17T08:57:04.063934Z","iopub.status.idle":"2024-08-17T09:08:49.154085Z","shell.execute_reply.started":"2024-08-17T08:57:04.063901Z","shell.execute_reply":"2024-08-17T09:08:49.152951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport shutil\n\nkaggle_image_dir = \"/kaggle/input/h-and-m-personalized-fashion-recommendations/images\"\nlocal_image_dir = \"/kaggle/working/downloaded_images\"\nos.makedirs(local_image_dir, exist_ok=True)\nfirst_1000_image_paths = first_batch['image_path']","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:09:14.026555Z","iopub.execute_input":"2024-08-17T09:09:14.027383Z","iopub.status.idle":"2024-08-17T09:09:14.032789Z","shell.execute_reply.started":"2024-08-17T09:09:14.027348Z","shell.execute_reply":"2024-08-17T09:09:14.031752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for image_path in first_1000_image_paths:\n    print(os.path.basename(image_path))\n    destination_path = os.path.join(local_image_dir, filename)\n    print(destination_path)\n    break","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:10:27.447769Z","iopub.execute_input":"2024-08-17T09:10:27.448975Z","iopub.status.idle":"2024-08-17T09:10:27.455568Z","shell.execute_reply.started":"2024-08-17T09:10:27.448934Z","shell.execute_reply":"2024-08-17T09:10:27.454596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for image_path in first_1000_image_paths:\n    try:\n        filename = str(os.path.basename(image_path))\n        destination_path = os.path.join(local_image_dir, filename)\n        shutil.copy(image_path, destination_path)\n    except Exception:\n        continue\n        \nprint(\"First 1000 images downloaded to:\", local_image_dir)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:11:27.72286Z","iopub.execute_input":"2024-08-17T09:11:27.723511Z","iopub.status.idle":"2024-08-17T09:11:29.955315Z","shell.execute_reply.started":"2024-08-17T09:11:27.723475Z","shell.execute_reply":"2024-08-17T09:11:29.954334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shutil.make_archive('/kaggle/working/downloaded_images', 'zip', local_image_dir)","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:12:45.411278Z","iopub.execute_input":"2024-08-17T09:12:45.412249Z","iopub.status.idle":"2024-08-17T09:12:57.979958Z","shell.execute_reply.started":"2024-08-17T09:12:45.412211Z","shell.execute_reply":"2024-08-17T09:12:57.978958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_query_embedding(query: str):\n    query_prompt_name = \"s2p_query\"\n    query_embeddings = text_embedding_model.encode(\n        [query], \n        prompt_name=query_prompt_name\n    )\n    return query_embeddings.squeeze().tolist()","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:44:03.38647Z","iopub.execute_input":"2024-08-17T09:44:03.387413Z","iopub.status.idle":"2024-08-17T09:44:03.393418Z","shell.execute_reply.started":"2024-08-17T09:44:03.387367Z","shell.execute_reply":"2024-08-17T09:44:03.392351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(get_query_embedding(\n    \"show me some cool red color tshirts\"\n))","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:44:19.320004Z","iopub.execute_input":"2024-08-17T09:44:19.3204Z","iopub.status.idle":"2024-08-17T09:44:19.381398Z","shell.execute_reply.started":"2024-08-17T09:44:19.320367Z","shell.execute_reply":"2024-08-17T09:44:19.380482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Time to search things\n\n\nquery = \"show me some cool red color tshirts\"\n\nquery_embedding = get_query_embedding(query)\n\nwith Session(engine) as session:\n    \n    entity = session.query(HandMProductEntity).order_by(\n        HandMProductEntity.text_embedding.cosine_distance(\n        query_embedding\n        )\n    ).limit(10).all()","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:52:06.457998Z","iopub.execute_input":"2024-08-17T09:52:06.458393Z","iopub.status.idle":"2024-08-17T09:52:07.514869Z","shell.execute_reply.started":"2024-08-17T09:52:06.45836Z","shell.execute_reply":"2024-08-17T09:52:07.513792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for article in entity:\n    print(article.detail_desc)\n    print()","metadata":{"execution":{"iopub.status.busy":"2024-08-17T09:52:47.434797Z","iopub.execute_input":"2024-08-17T09:52:47.435351Z","iopub.status.idle":"2024-08-17T09:52:47.443974Z","shell.execute_reply.started":"2024-08-17T09:52:47.435307Z","shell.execute_reply":"2024-08-17T09:52:47.442983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}