{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ================== Notebook 4: 最终提交（重构版） ==================\n# 改进点：\n#   - 所有字典在提交窗口内重建（不再依赖 notebook1 的过期字典）\n#   - 与 notebook 2 相同的 fallback 召回链\n#   - 扩展历史（16周）用于召回\n#   - 新增回购特征、trending、recent_30d 策略\n#   - 性能优化：预计算 W2V 数组、向量化回购特征、字典化品类查找\n#   - 完全复现训练时的特征工程（时序、TF-IDF、生命周期、用户偏好等）\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport pickle, datetime, gc\nfrom collections import defaultdict\nfrom tqdm import tqdm\nfrom gensim.models import Word2Vec\nimport lightgbm as lgb\n\nBASE_PATH = \"../input/competitions/h-and-m-personalized-fashion-recommendations/\"\nTRAIN_END  = datetime.datetime(2020, 9, 22)\nWEEKS      = 4\nEXT_WEEKS  = 16\nTOPK_RECALL = 300\nTOPK_PRED   = 12\nSEED        = 610\n\n# ==================== 加载基础数据 ====================\nwith open('/kaggle/input/notebooks/yeahyiiiiii/ccglm1/recall_dicts.pkl', 'rb') as f:\n    recall = pickle.load(f)\ncust_pd       = recall['cust_pd']\nart_pd        = recall['art_pd']\nuser_age_group = recall['user_age_group']\n\nmodel = lgb.Booster(model_file='/kaggle/input/notebooks/yeahyiiiiii/ccglm2/lgb_model.txt')\n\nwith open('/kaggle/input/notebooks/yeahyiiiiii/ccglm2/feature_metadata.pkl', 'rb') as f:\n    meta = pickle.load(f)\nnum_features = meta['num_features']\ncat_features = meta['cat_features']\nall_features = meta['all_features']\nstat_cols    = meta['stat_cols']\nSTRATEGIES = meta.get('STRATEGIES', ['pair', 'repurchase', 'trending', 'recent_30d', \n                                     'cat_pop', 'i2i', 'product', 'dept', 'price', \n                                     'vec', 'vec_recent', 'age_pop', 'global'])\ntfidf_vectorizer = meta.get('tfidf_vectorizer')\nif tfidf_vectorizer is None:\n    raise ValueError(\"tfidf_vectorizer not found in metadata, please retrain with saving it.\")\n\ncat_encoders = meta.get('cat_encoders')\nif cat_encoders is None:\n    raise ValueError(\"cat_encoders not found in metadata, please retrain ccglm2 with saving it.\")\n\n# ==================== 加载交易数据 ====================\ntx = pl.read_csv(\n    BASE_PATH + \"transactions_train.csv\",\n    columns=['t_dat', 'customer_id', 'article_id', 'price'],\n    try_parse_dates=True,\n    schema_overrides={'article_id': pl.Utf8, 'customer_id': pl.Utf8}\n)\narticles = pl.read_csv(\n    BASE_PATH + \"articles.csv\",\n    columns=['article_id', 'product_type_no', 'colour_group_code',\n             'department_no', 'garment_group_no', 'index_group_no', 'index_group_name',\n             'graphical_appearance_no'], \n    schema_overrides={'article_id': pl.Utf8}\n)\ncustomers = pl.read_csv(\n    BASE_PATH + \"customers.csv\",\n    columns=['customer_id', 'age', 'FN', 'Active'],\n    schema_overrides={'customer_id': pl.Utf8}\n)\ntx = tx.join(articles, on='article_id', how='left')\n\n# ==================== 双窗口（提交窗口）====================\nsub_ft_start = TRAIN_END - datetime.timedelta(weeks=WEEKS)\nsub_ft_data = tx.filter((pl.col('t_dat') >= sub_ft_start) & (pl.col('t_dat') < TRAIN_END))\n\nsub_ext_start = TRAIN_END - datetime.timedelta(weeks=EXT_WEEKS)\nsub_ext_data = tx.filter((pl.col('t_dat') >= sub_ext_start) & (pl.col('t_dat') < TRAIN_END))\n\nprint(f\"短窗口: {sub_ft_data.shape[0]:,} 行, 扩展窗口: {sub_ext_data.shape[0]:,} 行\")\n\n# ==================== 物品生命周期特征（全历史，使用 tx）====================\nprint(\">>> 计算提交窗口物品生命周期特征（全历史）...\")\nitem_lifecycle = tx.group_by('article_id').agg([\n    pl.col('t_dat').min().alias('first_sale'),\n    pl.col('t_dat').max().alias('last_sale'),\n    pl.col('price').mean().alias('i_price_hist_avg'),\n])\nitem_lifecycle = item_lifecycle.with_columns([\n    ((TRAIN_END - pl.col('first_sale')).dt.total_days().cast(pl.Int32)).alias('days_since_first_sale_item'),\n    (pl.col('first_sale') >= (TRAIN_END - datetime.timedelta(days=7))).cast(pl.Int8).alias('is_new_this_week'),\n    ((pl.col('last_sale') < (TRAIN_END - datetime.timedelta(days=14))) &\n     (pl.col('last_sale') > pl.col('first_sale'))).cast(pl.Int8).alias('is_discontinued'),\n])\nitem_lifecycle_pd_sub = item_lifecycle.select([\n    'article_id', 'days_since_first_sale_item', 'is_new_this_week', 'is_discontinued', 'i_price_hist_avg'\n]).to_pandas().set_index('article_id')\ndel item_lifecycle; gc.collect()\n\nprint(\">>> 计算去年同期销量...\")\nly_start = TRAIN_END - datetime.timedelta(days=365)\nly_end = ly_start + datetime.timedelta(days=7)\nlast_year_sales_sub = tx.filter(\n    (pl.col('t_dat') >= ly_start) & (pl.col('t_dat') < ly_end)\n).group_by('article_id').agg(pl.len().alias('last_year_same_week_sales')).to_pandas().set_index('article_id')\n\n# 释放 tx（不再需要全量）\ndel tx, articles\ngc.collect()\n\n# ==================== 重建所有召回字典（提交窗口）====================\n# 扩展窗口用户历史\nsub_user_rep_ext = sub_ext_data.group_by('customer_id').agg(\n    pl.col('article_id').unique()\n).to_dict(as_series=False)\nsub_user_rep_ext = {r[0]: r[1] for r in zip(sub_user_rep_ext['customer_id'], sub_user_rep_ext['article_id'])}\n\n# 用户品类/部门/颜色缓存（扩展窗口）\nuser_cats_ext, user_depts_ext, user_colors_ext = {}, {}, {}\nfor row in sub_ext_data.group_by('customer_id').agg(pl.col('product_type_no').unique().alias('cats')).iter_rows(named=True):\n    user_cats_ext[row['customer_id']] = row['cats']\nfor row in sub_ext_data.group_by('customer_id').agg(pl.col('department_no').unique().alias('depts')).iter_rows(named=True):\n    user_depts_ext[row['customer_id']] = row['depts']\nfor row in sub_ext_data.group_by('customer_id').agg(pl.col('colour_group_code').unique().alias('colors')).iter_rows(named=True):\n    user_colors_ext[row['customer_id']] = row['colors']\n\n# W2V（扩展窗口）\nprint(\"训练提交窗口 W2V...\")\nsub_seqs = (sub_ext_data.sort('t_dat').group_by('customer_id')\n            .agg(pl.col('article_id').tail(80).alias('seq')))['seq'].to_list()\nsub_w2v = Word2Vec(sub_seqs, vector_size=64, window=5, min_count=3,\n                   workers=4, negative=10, hs=0, epochs=8, seed=SEED)\nsub_w2v_vectors = {w: sub_w2v.wv[w] for w in sub_w2v.wv.index_to_key}\n\n# 预计算 W2V 数组\nprint(\"预计算 W2V 数组...\")\nALL_ITEMS = list(sub_w2v_vectors.keys())\nALL_VECS = np.array([sub_w2v_vectors[it] for it in ALL_ITEMS])\n\n# 预计算 article_id → product_type_no 映射\nart_type_dict = art_pd['product_type_no'].to_dict()\n\n# I2I\nsub_i2i = defaultdict(list)\nfor art in tqdm(sub_w2v.wv.index_to_key, desc=\"I2I\"):\n    for s, _ in sub_w2v.wv.most_similar(art, topn=50):\n        sub_i2i[art].append(s)\n\n# 流行度\nsub_pop = sub_ft_data.with_columns(\n    (1.0 / ((TRAIN_END - pl.col('t_dat')).dt.total_days() + 1)).alias('t_decay')\n)\n\n# 品类热门\nsub_cat_pop = sub_pop.group_by(['product_type_no', 'article_id']).agg(pl.sum('t_decay').alias('score'))\nsub_cat_top = sub_cat_pop.sort(['product_type_no', 'score'], descending=[False, True])\nsub_cat_dict = sub_cat_top.group_by('product_type_no').agg(pl.col('article_id').head(100)).to_dict(as_series=False)\nsub_cat_dict = {r[0]: r[1] for r in zip(sub_cat_dict['product_type_no'], sub_cat_dict['article_id'])}\n\n# 全局热门（取前300）\nsub_global_top = sub_pop.group_by('article_id').agg(pl.sum('t_decay').alias('pop')) \\\n                       .sort('pop', descending=True).head(400)['article_id'].to_list()\n\n# Trending（过去1周）\nsub_recent_week = sub_ft_data.filter(pl.col('t_dat') >= (TRAIN_END - datetime.timedelta(weeks=1)))\nsub_trending_top = sub_recent_week.group_by('article_id').agg(pl.len().alias('cnt')) \\\n                    .sort('cnt', descending=True).head(200)['article_id'].to_list()\n\n# 最近30天热门\nsub_recent_30d_start = TRAIN_END - datetime.timedelta(days=30)\nsub_recent_30d_data = sub_ft_data.filter(pl.col('t_dat') >= sub_recent_30d_start)\nsub_recent_30d_top = sub_recent_30d_data.group_by('article_id').agg(\n    pl.len().alias('cnt')\n).sort('cnt', descending=True).head(200)['article_id'].to_list()\n\n# 时间感知 Pairs\nprint(\"构建提交窗口时间感知 Pairs...\")\nsub_recent_articles = sub_recent_week['article_id'].unique().to_list()\nsub_recent_tx = sub_ft_data.filter(pl.col('article_id').is_in(sub_recent_articles))\nsub_pairs = sub_recent_tx.join(sub_recent_tx, on=['customer_id', 't_dat']).filter(\n    pl.col('article_id') != pl.col('article_id_right')\n)\nsub_pair_counts = sub_pairs.group_by(['article_id', 'article_id_right']).agg(\n    pl.len().alias('pair_strength')\n).filter(pl.col('pair_strength') >= 2)\nsub_pair_counts = sub_pair_counts.sort(['article_id', 'pair_strength'], descending=[False, True])\nsub_pair_groups = sub_pair_counts.group_by('article_id').agg([\n    pl.col('article_id_right').head(5).alias('candidates'),\n    pl.col('pair_strength').head(5).alias('strengths')\n])\nsub_pair_dict = {}\nfor row in sub_pair_groups.iter_rows(named=True):\n    sub_pair_dict[row['article_id']] = {'candidates': row['candidates'], 'strengths': row['strengths']}\ndel sub_recent_tx, sub_pairs, sub_pair_counts, sub_pair_groups; gc.collect()\n\n# 换购/部门/年龄段热门（在提交窗口重建）\nprint(\"重建换购/部门/年龄段热门字典...\")\nsub_pt_pop = sub_pop.group_by(['product_type_no', 'article_id']).agg(pl.sum('t_decay').alias('score')) \\\n    .sort(['product_type_no', 'score'], descending=[False, True])\nsub_product_type_dict = sub_pt_pop.group_by('product_type_no').agg(pl.col('article_id').head(30)).to_dict(as_series=False)\nsub_product_type_dict = {r[0]: r[1] for r in zip(sub_product_type_dict['product_type_no'], sub_product_type_dict['article_id'])}\n\nsub_dept_pop = sub_pop.group_by(['department_no', 'article_id']).agg(pl.sum('t_decay').alias('score')) \\\n    .sort(['department_no', 'score'], descending=[False, True])\nsub_department_dict = sub_dept_pop.group_by('department_no').agg(pl.col('article_id').head(30)).to_dict(as_series=False)\nsub_department_dict = {r[0]: r[1] for r in zip(sub_department_dict['department_no'], sub_department_dict['article_id'])}\n\n# 年龄组\nage_med = customers['age'].median()\ncustomers = customers.with_columns([\n    pl.col('age').fill_null(age_med).cast(pl.Float32),\n    pl.when(pl.col('age').is_between(0, 24)).then(0)\n     .when(pl.col('age').is_between(25, 35)).then(1)\n     .when(pl.col('age').is_between(36, 50)).then(2)\n     .otherwise(3).alias('age_group').cast(pl.Int8)\n])\nsub_user_age_group = dict(zip(customers['customer_id'].to_list(), customers['age_group'].to_list()))\n\ntx_with_age = sub_ft_data.join(customers.select(['customer_id', 'age_group']), on='customer_id', how='left')\nsub_age_pop = tx_with_age.with_columns(\n    (1.0 / ((TRAIN_END - pl.col('t_dat')).dt.total_days() + 1)).alias('t_decay')\n).group_by(['age_group', 'article_id']).agg(pl.sum('t_decay').alias('score')) \\\n .sort(['age_group', 'score'], descending=[False, True])\nsub_age_group_dict = sub_age_pop.group_by('age_group').agg(pl.col('article_id').head(100)).to_dict(as_series=False)  # 扩大到100\nsub_age_group_dict = {r[0]: r[1] for r in zip(sub_age_group_dict['age_group'], sub_age_group_dict['article_id'])}\n\n# ==================== 统计特征（与训练时完全一致）====================\nsub_user_stats = sub_ft_data.group_by('customer_id').agg([\n    pl.col('article_id').count().alias('u_cnt'),\n    ((TRAIN_END - pl.col('t_dat').max()).dt.total_days()).alias('u_last'),\n    pl.col('price').mean().alias('u_avg_price'),\n    pl.col('price').min().alias('u_price_min'),\n    pl.col('price').max().alias('u_price_max'),\n]).to_pandas().set_index('customer_id').astype({'u_cnt':'int32','u_last':'float32','u_avg_price':'float32',\n                                                  'u_price_min':'float32','u_price_max':'float32'})\n\nsub_item_stats = sub_ft_data.with_columns(\n    (1.0 / ((TRAIN_END - pl.col('t_dat')).dt.total_days() + 1)).alias('w')\n).group_by('article_id').agg([\n    pl.col('customer_id').count().alias('i_cnt'),\n    pl.sum('w').alias('i_pop_w'),\n    pl.col('price').mean().alias('i_price'),\n    pl.col('price').min().alias('i_price_min'),\n    pl.col('price').max().alias('i_price_max'),\n]).to_pandas().set_index('article_id').astype({'i_cnt':'int32','i_pop_w':'float32','i_price':'float32',\n                                                 'i_price_min':'float32','i_price_max':'float32'})\n\n# 性别推断\nsub_gender_tendency = (\n    sub_ft_data.with_columns(\n        pl.col('index_group_name').replace_strict({'Ladieswear': 0, 'Menswear': 1}, default=2).cast(pl.Int8).alias('gender_code')\n    ).group_by('customer_id')\n    .agg(pl.col('gender_code').mode().first().alias('pred_gender'))\n    .to_pandas().set_index('customer_id')\n)\n\n# 价格区间索引\nprice_item = sub_item_stats[['i_price', 'i_pop_w']].copy().sort_values('i_price')\nprices, popws = price_item.i_price.values, price_item.i_pop_w.values\nart_ids_price = price_item.index.values\n\n# 销量趋势\nw1 = TRAIN_END - datetime.timedelta(weeks=2)\nw2 = TRAIN_END - datetime.timedelta(weeks=4)\ns2w = sub_ft_data.filter(pl.col('t_dat') >= w1).group_by('article_id').agg(pl.len().alias('i_sales_2w'))\nsp2w = sub_ft_data.filter((pl.col('t_dat') >= w2) & (pl.col('t_dat') < w1)).group_by('article_id').agg(pl.len().alias('i_sales_prev_2w'))\nstrend = s2w.join(sp2w, on='article_id', how='full').fill_null(0.1)\nstrend = strend.with_columns(((pl.col('i_sales_2w')+0.1)/(pl.col('i_sales_prev_2w')+0.1)).alias('i_sales_trend'))\nsub_sales_pd = strend.to_pandas().set_index('article_id')[['i_sales_2w','i_sales_prev_2w','i_sales_trend']]\ndel s2w, sp2w, strend; gc.collect()\n\n# UI 交互\nsub_ui = sub_ft_data.group_by(['customer_id','article_id']).agg([\n    pl.len().alias('u_i_buy_cnt'),\n    ((TRAIN_END - pl.col('t_dat').max()).dt.total_days()).alias('u_i_last_buy_days')\n]).to_pandas().set_index(['customer_id','article_id'])\n\n# 回购集合（预计算为 set 用于快速查找）\nuser_bought_set_sub = {}\nfor row in sub_ft_data.group_by('customer_id').agg(pl.col('article_id').unique().alias('bought')).iter_rows(named=True):\n    user_bought_set_sub[row['customer_id']] = set(row['bought'])\nrepurchase_pairs = set()\nfor cid, arts in user_bought_set_sub.items():\n    for a in arts:\n        repurchase_pairs.add((cid, a))\n\n# ==================== 用户品类偏好计数（Group 1）====================\nprint(\"计算提交窗口用户品类偏好计数...\")\nuser_pref_dfs_sub = {}\nfor dim_col, feat_name in [\n    ('product_type_no', 'user_pref_cnt_ptype'),\n    ('department_no', 'user_pref_cnt_dept'),\n    ('garment_group_no', 'user_pref_cnt_garment'),\n    ('colour_group_code', 'user_pref_cnt_color'),\n    ('graphical_appearance_no', 'user_pref_cnt_graphical'),\n]:\n    pdf = (sub_ft_data.group_by(['customer_id', dim_col])\n           .agg(pl.len().alias(feat_name))\n           .to_pandas())\n    pdf[feat_name] = pdf[feat_name].astype('int16')\n    user_pref_dfs_sub[feat_name] = pdf\n\n# ==================== 时间衰减亲和度（Group 2）====================\nprint(\"计算提交窗口时间衰减亲和度...\")\nTAU = 21.0\nft_td_sub = sub_ft_data.with_columns(\n    (1.0 / (((TRAIN_END - pl.col('t_dat')).dt.total_days().cast(pl.Float32) / TAU)).exp()).alias('decay')\n)\ntd_dfs_sub = {}\nfor dim_col, feat_name in [('product_type_no', 'td_affinity_ptype'), ('department_no', 'td_affinity_dept')]:\n    pdf = (ft_td_sub.group_by(['customer_id', dim_col])\n           .agg(pl.sum('decay').alias(feat_name))\n           .to_pandas())\n    pdf[feat_name] = pdf[feat_name].astype('float32')\n    td_dfs_sub[feat_name] = pdf\ndel ft_td_sub; gc.collect()\n\n# ==================== TF-IDF 文本特征（预计算所有商品）====================\nprint(\"预计算商品 TF-IDF 特征...\")\narticles_desc = pl.read_csv(\n    BASE_PATH + \"articles.csv\",\n    columns=['article_id', 'detail_desc'],\n    schema_overrides={'article_id': pl.Utf8}\n).to_pandas().set_index('article_id')\narticles_desc['detail_desc'] = articles_desc['detail_desc'].fillna('')\ndesc_tfidf_matrix = tfidf_vectorizer.transform(articles_desc['detail_desc'])\ndesc_tfidf_df = pd.DataFrame(\n    desc_tfidf_matrix.toarray(),\n    index=articles_desc.index,\n    columns=[f'desc_tfidf_{i}' for i in range(desc_tfidf_matrix.shape[1])]\n).astype('float32')\ndel articles_desc, desc_tfidf_matrix; gc.collect()\n\n# 释放不再需要的 sub_ext_data? 但之后不再需要，可以释放\ndel sub_ext_data, sub_pop, sub_ft_data, sub_cat_pop, sub_cat_top, sub_pt_pop, sub_dept_pop, tx_with_age, sub_age_pop\ngc.collect()\n\n# ===== 预合并用户特征和商品特征（减少每批 merge 次数） =====\nprint(\"预合并用户和商品特征表...\")\nuser_feat_all = cust_pd.join(sub_user_stats, how='left').join(sub_gender_tendency, how='left')\nfor col in user_feat_all.columns:\n    user_feat_all[col] = user_feat_all[col].fillna(0)\n\nitem_feat_all = art_pd.join(sub_item_stats, how='left') \\\n                       .join(sub_sales_pd, how='left') \\\n                       .join(item_lifecycle_pd_sub, how='left') \\\n                       .join(last_year_sales_sub, how='left') \\\n                       .join(desc_tfidf_df, how='left')\nfor col in item_feat_all.columns:\n    item_feat_all[col] = item_feat_all[col].fillna(0)\n\n# ==================== 召回函数（与 notebook 2 完全一致）====================\ndef generate_candidates_sub(user, K=TOPK_RECALL):\n    cand = {}\n    # 无历史用户直接返回全局热门\n    if user not in sub_user_rep_ext:\n        final_items = sub_global_top[:K]\n        actual_len = len(final_items)\n        strat_lists = {s: [-1]*actual_len for s in STRATEGIES}\n        strat_lists['global'] = list(range(actual_len))\n        strat_counts = {s: 0 for s in STRATEGIES}\n        strat_counts['global'] = actual_len   # 用实际长度\n        str_lists = {'pair_str': [0]*actual_len, 'prod_str': [0]*actual_len, 'dept_str': [0]*actual_len}\n        return final_items, strat_lists, str_lists, strat_counts, None, None\n\n    def add(article, strategy, rank, strength=0):\n        if article not in cand:\n            cand[article] = {s: -1 for s in STRATEGIES}\n            cand[article].update({'pair_str': 0, 'prod_str': 0, 'dept_str': 0})\n        if cand[article][strategy] == -1:\n            cand[article][strategy] = rank\n        if strategy == 'pair':\n            cand[article]['pair_str'] = max(cand[article]['pair_str'], strength)\n        elif strategy == 'product':\n            cand[article]['prod_str'] = max(cand[article]['prod_str'], rank)\n        elif strategy == 'dept':\n            cand[article]['dept_str'] = max(cand[article]['dept_str'], rank)\n\n    has_history = user in sub_user_rep_ext\n    user_vec_final = None\n    recent_vec_final = None\n\n    # 1. 时间感知 Pairs\n    if has_history:\n        for a in sub_user_rep_ext[user][-20:]:\n            if a in sub_pair_dict:\n                for i, (b, s) in enumerate(zip(sub_pair_dict[a]['candidates'], sub_pair_dict[a]['strengths'])):\n                    add(b, 'pair', i, s)\n\n    # 2. 回购\n    if has_history:\n        for i, a in enumerate(sub_user_rep_ext[user][-80:]):\n            add(a, 'repurchase', i)\n\n    # 3. 最近30天热门\n    for i, a in enumerate(sub_recent_30d_top[:200]):\n        add(a, 'recent_30d', i)\n\n    # 4. Trending\n    for i, a in enumerate(sub_trending_top[:200]):\n        add(a, 'trending', i)\n\n    # 5. 品类流行度\n    if has_history:\n        for cat in user_cats_ext.get(user, []):\n            for i, a in enumerate(sub_cat_dict.get(cat, [])[:50]):\n                add(a, 'cat_pop', i)\n\n    # 6. I2I\n    if has_history:\n        for a in set(sub_user_rep_ext[user]):\n            if a in sub_i2i:\n                for i, b in enumerate(sub_i2i[a][:30]):\n                    add(b, 'i2i', i)\n\n    # 7. 换购\n    if has_history:\n        user_prod_types = set()\n        for a in sub_user_rep_ext[user]:\n            ptype = art_type_dict.get(a)\n            if ptype is not None:\n                user_prod_types.add(ptype)\n        for ptype in user_prod_types:\n            for i, b in enumerate(sub_product_type_dict.get(ptype, [])[:20]):\n                add(b, 'product', i, 1)\n\n    # 8. 部门热门\n    if has_history:\n        for dept in user_depts_ext.get(user, []):\n            for i, b in enumerate(sub_department_dict.get(dept, [])[:20]):\n                add(b, 'dept', i, 1)\n\n    # 9. 价格区间\n    if user in sub_user_stats.index:\n        avg_p = sub_user_stats.loc[user, 'u_avg_price']\n        left = np.searchsorted(prices, 0.7*avg_p, side='left')\n        right = np.searchsorted(prices, 1.5*avg_p, side='right')\n        if left < right:\n            c_arts, c_pops = art_ids_price[left:right], popws[left:right]\n            top_idx = np.argsort(c_pops)[::-1][:50]\n            for rank, idx in enumerate(top_idx):\n                add(c_arts[idx], 'price', rank)\n\n    # 10. 全部历史偏好向量\n    if has_history:\n        vecs = [sub_w2v_vectors[a] for a in sub_user_rep_ext[user] if a in sub_w2v_vectors]\n        if vecs:\n            user_vec_final = np.mean(vecs, axis=0)\n            sims = ALL_VECS @ user_vec_final\n            top_indices = np.argpartition(-sims, 60)[:60]\n            top_indices = top_indices[np.argsort(-sims[top_indices])]\n            for rank, idx in enumerate(top_indices):\n                add(ALL_ITEMS[idx], 'vec', rank)\n\n    # 11. 近期偏好向量\n    if has_history:\n        recent = sub_user_rep_ext[user][-10:]\n        rvecs = [sub_w2v_vectors[a] for a in recent if a in sub_w2v_vectors]\n        if rvecs:\n            recent_vec_final = np.mean(rvecs, axis=0)\n            sims = ALL_VECS @ recent_vec_final\n            top_indices = np.argpartition(-sims, 40)[:40]\n            top_indices = top_indices[np.argsort(-sims[top_indices])]\n            for rank, idx in enumerate(top_indices):\n                add(ALL_ITEMS[idx], 'vec_recent', rank)\n\n    # 12. 冷启动 — 年龄段\n    if len(cand) < K:\n        age = sub_user_age_group.get(user, 3)\n        for i, a in enumerate(sub_age_group_dict.get(age, [])[:200]):\n            add(a, 'age_pop', i)\n\n    # 13. 全局兜底\n    if len(cand) < K:\n        for i, a in enumerate(sub_global_top[:300]):\n            add(a, 'global', i)\n\n    final_items = list(cand.keys())[:K]\n    strat_lists = {s: [cand[a][s] for a in final_items] for s in STRATEGIES}\n    # 计算每个策略的计数（用于 normrank）\n    strat_counts = {s: sum(1 for a in final_items if cand[a][s] != -1) for s in STRATEGIES}\n    str_lists = {\n        'pair_str': [cand[a]['pair_str'] for a in final_items],\n        'prod_str': [cand[a]['prod_str'] for a in final_items],\n        'dept_str': [cand[a]['dept_str'] for a in final_items],\n    }\n    return final_items, strat_lists, str_lists, strat_counts, user_vec_final, recent_vec_final\n\n# ==================== 分批预测 ====================\nsub = pd.read_csv(BASE_PATH + \"sample_submission.csv\", dtype={'customer_id': str})\nbatch_size = 15000   # 减小 batch 防止内存溢出\npreds = []\n\n# 预计算用户偏好和亲和度合并所需的辅助 DataFrame（为了快速合并）\n# 我们将在每批中合并，但为了效率，可以预先转换为字典？这里直接按批合并\n\nfor start in tqdm(range(0, len(sub), batch_size), desc=\"分批预测\"):\n    batch_users = sub['customer_id'].iloc[start:start+batch_size].tolist()\n    batch_rows = []\n    for u in batch_users:\n        items, strats, strs, strat_counts, user_vec, recent_vec = generate_candidates_sub(u)\n        n = len(items)\n        for i in range(n):\n            row = [u, items[i]]\n            for s in STRATEGIES:\n                row.append(strats[s][i])\n            row.extend([strs['pair_str'][i], strs['prod_str'][i], strs['dept_str'][i]])\n            for s in STRATEGIES:\n                cnt = strat_counts[s]\n                if cnt > 0:\n                    row.append(strats[s][i] / cnt)\n                else:\n                    row.append(-1.0)\n            v_item = sub_w2v_vectors.get(items[i])\n            sim_top = 0.0\n            sim_recent = 0.0\n            if v_item is not None:\n                if user_vec is not None:\n                    n1 = np.linalg.norm(v_item) * np.linalg.norm(user_vec)\n                    sim_top = float(np.dot(v_item, user_vec) / n1) if n1 > 1e-9 else 0.0\n                if recent_vec is not None:\n                    n2 = np.linalg.norm(v_item) * np.linalg.norm(recent_vec)\n                    sim_recent = float(np.dot(v_item, recent_vec) / n2) if n2 > 1e-9 else 0.0\n            row.append(sim_top)\n            row.append(sim_recent)\n            batch_rows.append(row)\n\n    cols = (['customer_id', 'article_id'] +\n            [f'strategy_{s}' for s in STRATEGIES] +\n            ['pair_strength', 'product_strength', 'department_strength'] +\n            [f'strategy_{s}_normrank' for s in STRATEGIES] +\n            ['w2v_top_sim', 'w2v_last5_sim'])\n    bdf = pd.DataFrame(batch_rows, columns=cols)\n\n    # 合并用户特征和商品特征\n    bdf = bdf.merge(user_feat_all, on='customer_id', how='left')\n    bdf = bdf.merge(item_feat_all, on='article_id', how='left')\n\n    # ---- 用户-物品交互特征（UI） ----\n    bdf = bdf.merge(sub_ui, left_on=['customer_id', 'article_id'], right_index=True, how='left')\n    bdf['u_i_buy_cnt'] = bdf['u_i_buy_cnt'].fillna(0).astype('int32')\n    bdf['u_i_last_buy_days'] = bdf['u_i_last_buy_days'].fillna(999).astype('float32')\n\n    # 回购特征\n    bdf['has_repurchased'] = [\n        1 if (cid, aid) in repurchase_pairs else 0\n        for cid, aid in zip(bdf['customer_id'], bdf['article_id'])\n    ]\n    bdf['has_repurchased'] = bdf['has_repurchased'].astype('int8')\n\n    # 物品生命周期 + 去年同期销量（已在 item_feat_all 中，但可再填充一次）\n    bdf['days_since_first_sale_item'] = bdf['days_since_first_sale_item'].fillna(999).astype('int32')\n    bdf['is_new_this_week'] = bdf['is_new_this_week'].fillna(0).astype('int8')\n    bdf['is_discontinued'] = bdf['is_discontinued'].fillna(0).astype('int8')\n    bdf['i_price_hist_avg'] = bdf['i_price_hist_avg'].fillna(bdf['i_price']).astype('float32')\n    bdf['i_price_to_hist_avg'] = (bdf['i_price'] + 1e-5) / (bdf['i_price_hist_avg'] + 1e-5)\n    bdf['is_on_sale'] = (bdf['i_price'] < 0.9 * bdf['i_price_hist_avg']).astype('int8')\n    bdf['last_year_same_week_sales'] = bdf['last_year_same_week_sales'].fillna(0).astype('int32')\n\n    # 基本价格特征\n    bdf['i_price'] = bdf['i_price'].astype('float32')\n    bdf['u_avg_price'] = bdf['u_avg_price'].astype('float32')\n    bdf['price_dev'] = np.abs(bdf['i_price'] - bdf['u_avg_price'])\n    bdf['price_ratio'] = (bdf['i_price'] + 1e-5) / (bdf['u_avg_price'] + 1e-5)\n    bdf['price_is_cheaper'] = (bdf['i_price'] < bdf['u_avg_price']).astype('int8')\n    bdf['price_to_u_min'] = (bdf['i_price'] + 1e-5) / (bdf['u_price_min'] + 1e-5)\n    bdf['price_to_u_max'] = (bdf['i_price'] + 1e-5) / (bdf['u_price_max'] + 1e-5)\n\n    # 销量趋势（确保类型）\n    bdf['i_sales_2w'] = bdf['i_sales_2w'].fillna(0).astype('int32')\n    bdf['i_sales_prev_2w'] = bdf['i_sales_prev_2w'].fillna(0).astype('int32')\n    bdf['i_sales_trend'] = bdf['i_sales_trend'].fillna(1.0).astype('float32')\n\n    # 时序增强特征\n    for col in ['u_last', 'u_cnt', 'u_avg_price', 'u_price_min', 'u_price_max']:\n        if col in bdf.columns:\n            bdf[col] = bdf[col].fillna(0)\n    bdf['u_last_day_of_week'] = (bdf['u_last'].astype('int') % 7).astype('int8')\n    bdf['u_last_month'] = np.clip((bdf['u_last'].astype('int') // 30), 0, 11).astype('int8')\n    bdf['u_daily_freq'] = (bdf['u_cnt'] / (bdf['u_last'] + 1.0)).astype('float32')\n    bdf['days_since_first_sale_log'] = np.log1p(bdf['days_since_first_sale_item']).astype('float32')\n    bdf['i_sold_last_2w'] = (bdf['i_sales_2w'] > 0).astype('int8')\n    bdf['price_match_score'] = (1.0 / (1.0 + np.abs(bdf['price_ratio'] - 1.0))).astype('float32')\n    bdf['price_dev_to_range'] = (\n        bdf['price_dev'] / (bdf['u_price_max'] - bdf['u_price_min'] + 1e-5)\n    ).astype('float32')\n    bdf['price_dev_to_range'] = bdf['price_dev_to_range'].clip(0, 10)\n\n    # 用户品类偏好计数（Group 1）\n    for feat_name, dim_col in [\n        ('user_pref_cnt_ptype', 'product_type_no'),\n        ('user_pref_cnt_dept', 'department_no'),\n        ('user_pref_cnt_garment', 'garment_group_no'),\n        ('user_pref_cnt_color', 'colour_group_code'),\n        ('user_pref_cnt_graphical', 'graphical_appearance_no'),\n    ]:\n        pdf = user_pref_dfs_sub[feat_name].copy()\n        bdf = bdf.merge(pdf, left_on=['customer_id', dim_col], right_on=['customer_id', dim_col], how='left')\n        bdf[feat_name] = bdf[feat_name].fillna(0).astype('int16')\n        del pdf\n\n    # 时间衰减亲和度（Group 2）\n    for feat_name, dim_col in [('td_affinity_ptype', 'product_type_no'), ('td_affinity_dept', 'department_no')]:\n        pdf = td_dfs_sub[feat_name].copy()\n        bdf = bdf.merge(pdf, left_on=['customer_id', dim_col], right_on=['customer_id', dim_col], how='left')\n        bdf[feat_name] = bdf[feat_name].fillna(0.0).astype('float32')\n        del pdf\n\n    # 其他 stat_cols 填充（只填充存在的列）\n    for c in stat_cols:\n        if c in bdf.columns:\n            bdf[c] = bdf[c].fillna(0)\n\n    # 类别特征编码（使用训练时的编码映射）\n    for col in cat_features:\n        bdf[col] = bdf[col].map(cat_encoders[col]).fillna(-1).astype('int8')\n\n    # 确保所有特征都存在（缺失补0）\n    for f in all_features:\n        if f not in bdf.columns:\n            bdf[f] = 0.0\n\n    # 预测\n    X_batch = bdf[all_features]\n    bdf['score'] = model.predict(X_batch)\n\n    # 取每个用户 top12\n    bdf = bdf.sort_values(['customer_id', 'score'], ascending=[True, False])\n    top12 = bdf.groupby('customer_id')['article_id'].apply(lambda x: ' '.join(x.head(TOPK_PRED))).reset_index()\n    preds.append(top12)\n    del batch_rows, bdf, X_batch, top12\n    gc.collect()\n\n# ==================== 合并 + 后处理 ====================\nfinal_pred = pd.concat(preds).reset_index(drop=True)\nsub = sub[['customer_id']].merge(final_pred, on='customer_id', how='left')\nsub['prediction'] = sub['article_id'].fillna(' '.join(sub_global_top[:12]))\n\nbackup = ' '.join(sub_global_top[:12])\ndef fill_to_12(pred_str):\n    if pd.isna(pred_str):\n        return backup\n    items = pred_str.split()\n    if len(items) < 12:\n        for b in backup.split():\n            if b not in items:\n                items.append(b)\n            if len(items) == 12:\n                break\n        return ' '.join(items[:12])\n    return pred_str\n\nsub['prediction'] = sub['prediction'].apply(fill_to_12)\nsub = sub[['customer_id', 'prediction']]\nsub.to_csv(\"submission.csv\", index=False)\nprint(\"提交完成: submission.csv\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-06-20T16:15:35.518124Z","iopub.execute_input":"2026-06-20T16:15:35.518511Z"}},"outputs":[],"execution_count":null}]}