{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "bbfa9fa9-7275-4a00-ad8c-a156298b18e9",
   "metadata": {
    "tags": []
   },
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/home/ec2-user/anaconda3/envs/python3/lib/python3.10/site-packages/snowflake/connector/options.py:103: UserWarning: You have an incompatible version of 'pyarrow' installed (12.0.1), please install a version that adheres to: 'pyarrow<10.1.0,>=10.0.1; extra == \"pandas\"'\n",
      "  warn_incompatible_dep(\n",
      "DEBUG:absl:READY!!!\n"
     ]
    }
   ],
   "source": [
    "#snowflake connector and pytest\n",
    "!pip -q install snowflake-connector-python pytest pytest-sugar \n",
    "!pip -q install pyecharts absl-py\n",
    "!pip -q install statsmodels\n",
    "!pip -q install pmdarima\n",
    "\n",
    "# main imports\n",
    "import snowflake.connector\n",
    "import pandas as pd\n",
    "import pyecharts as echarts\n",
    "import os\n",
    "import boto3 \n",
    "\n",
    "import math\n",
    "import random\n",
    "import scipy\n",
    "import numpy as np\n",
    "import statsmodels.formula.api as smf\n",
    "import statsmodels.api as sm\n",
    "import pmdarima as pm \n",
    "import time\n",
    "\n",
    "from datetime import datetime, date, timedelta\n",
    "from statsmodels.tsa.statespace.sarimax import SARIMAX\n",
    "from statsmodels.tsa.arima.model import ARIMA\n",
    "from statsmodels.graphics.tsaplots import plot_acf, plot_pacf\n",
    "from statsmodels.tsa.seasonal import seasonal_decompose\n",
    "from statsmodels.tools.eval_measures import mse,rmse, meanabs\n",
    "from statsmodels.tsa.stattools import adfuller\n",
    "from statsmodels.tsa.statespace.tools import diff\n",
    "from scipy import fftpack\n",
    "from multiprocess import Pool\n",
    "\n",
    "# utils\n",
    "from getpass import getpass\n",
    "\n",
    "# logging and re\n",
    "from absl import logging\n",
    "import re\n",
    "\n",
    "\n",
    "log_level = \"DEBUG\"\n",
    "ticket_code = \"EXP_1\"\n",
    "\n",
    "logging.set_verbosity(log_level)\n",
    "logging.debug(\"READY!!!\")\n",
    "\n",
    "sec_id = 'dev/sagemaker-notebook-instance/SNOWFLAKE_PASSWORD'\n",
    "\n",
    "\n",
    "def get_secret_value(name, version=None):\n",
    "    \"\"\"Gets the value of a secret.\n",
    "\n",
    "    Version (if defined) is used to retrieve a particular version of\n",
    "    the secret.\n",
    "\n",
    "    \"\"\"\n",
    "    secrets_client = boto3.client(\"secretsmanager\")\n",
    "    kwargs = {'SecretId': name}\n",
    "    if version is not None:\n",
    "        kwargs['VersionStage'] = version\n",
    "    response = secrets_client.get_secret_value(**kwargs)\n",
    "    return response\n",
    "\n",
    "\n",
    "def get_snowflake_creds(username=\"SAGEMAKER\", account=\"orchard\",\n",
    "                        warehouse=\"DEV_OWS_ENGINEERING\"):\n",
    "    \"\"\"\n",
    "    Fetches and returns snowflake creds for connecting to snowflake\n",
    "\n",
    "    Please use this within the scope of a function if using this on a shared instance\n",
    "    This is so that the password is in memory only when its needed and gets dropped \n",
    "    once its no longer required.\n",
    "\n",
    "    returns:\n",
    "    - creds (dict) - a dictionary containing user creds\n",
    "\n",
    "    \"\"\"\n",
    "    creds = {\n",
    "      \"user\":  username,\n",
    "      \"password\": get_secret_value(sec_id)['SecretString'],\n",
    "      \"account\": \"orchard\",\n",
    "      \"warehouse\": warehouse,\n",
    "      \"protocol\": 'https'\n",
    "    }\n",
    "    return creds\n",
    "\n",
    "\n",
    "def snowflake_connector_factory(creds=None):\n",
    "    \"\"\"\n",
    "    A Factory for creating snowflake connectors.\n",
    "\n",
    "    This returns the cursor after opening a session with snowflake.\n",
    "\n",
    "    params:\n",
    "    - creds - snowflake credentials \n",
    "\n",
    "    returns:\n",
    "    - cursor - snowflake session cursor\n",
    "    \"\"\"\n",
    "    try:\n",
    "        if creds:\n",
    "            _creds = creds\n",
    "        else:\n",
    "            _creds = get_snowflake_creds()\n",
    "        return snowflake.connector.connect(**_creds).cursor()\n",
    "    except Exception as e:\n",
    "        logging.error(f\"Something went wrong - {str(e)}\")\n",
    "\n",
    "\n",
    "def _is_version_number(s):\n",
    "    \"Check and returns true if its a version number\"\n",
    "    return re.search(\"^[0-9][.0-9]*[0-9]$\", s) is not None\n",
    "\n",
    "\n",
    "def test_connection():\n",
    "    \"\"\" tests connection to snowflake \"\"\"\n",
    "    with snowflake_connector_factory() as cs:\n",
    "        try:\n",
    "            cs.execute(\"SELECT current_version()\")\n",
    "            one_row = cs.fetchone()\n",
    "            assert len(one_row) == 1\n",
    "            assert _is_version_number(one_row[0])\n",
    "            logging.info(f\"Your snowflake version - {one_row[0]} PASSED!\")\n",
    "        except Exception as e:\n",
    "          logging.error(f\"Something went wrong - {str(e)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "187557e4",
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "INFO:absl:Your snowflake version - 7.36.0 PASSED!\n"
     ]
    }
   ],
   "source": [
    "test_connection()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "17f21c5e",
   "metadata": {},
   "outputs": [],
   "source": [
    "with snowflake_connector_factory() as cs:\n",
    "    try:\n",
    "        cs.execute(\"USE WAREHOUSE DEV_PERFORMANCE_WAREHOUSE;\")\n",
    "        cs.execute(\"\"\"\n",
    "       select * from DEV_ENGINEERING.EIMPARA.TIKTOK_MOMENTS_V1\n",
    "        \"\"\")\n",
    "        rows = cs.fetchall()\n",
    "    except Exception as e:\n",
    "      logging.error(f\"Something went wrong - {str(e)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "20cb2c65",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(424200, 8)\n",
      "(424200, 8)\n"
     ]
    }
   ],
   "source": [
    "data_df = pd.DataFrame(rows, columns=map(lambda meta: meta[0], cs.description))\n",
    "df = data_df.drop_duplicates().copy()\n",
    "print(data_df.shape)\n",
    "print(df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "id": "f91073ff-ef2a-4826-a641-62c60faf803b",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "<class 'pandas.core.frame.DataFrame'>\n",
      "RangeIndex: 424200 entries, 0 to 424199\n",
      "Data columns (total 8 columns):\n",
      " #   Column         Non-Null Count   Dtype \n",
      "---  ------         --------------   ----- \n",
      " 0   ISRC           424200 non-null  object\n",
      " 1   ACTIVITY_DATE  424200 non-null  object\n",
      " 2   TRACK_NAME     424200 non-null  object\n",
      " 3   RELEASE_DATE   424200 non-null  object\n",
      " 4   ARTIST_NAME    424200 non-null  object\n",
      " 5   ARTIST_ID      424200 non-null  int64 \n",
      " 6   CREATIONS      424200 non-null  int64 \n",
      " 7   VIDEO_VIEWS    424200 non-null  int64 \n",
      "dtypes: int64(3), object(5)\n",
      "memory usage: 25.9+ MB\n"
     ]
    }
   ],
   "source": [
    "df.info()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "id": "69ea9b8d-eef2-4b2e-99b1-6472c8822767",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "3030"
      ]
     },
     "execution_count": 6,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "df['ISRC'].nunique()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "id": "7113a190",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Min activity date:  2023-05-16\n",
      "Max activity date:  2023-10-03\n"
     ]
    }
   ],
   "source": [
    "df['ACTIVITY_DATE'] = pd.to_datetime(df['ACTIVITY_DATE'])\n",
    "df['ACTIVITY_DATE'] = df['ACTIVITY_DATE'].dt.date\n",
    "\n",
    "print('Min activity date: ', df['ACTIVITY_DATE'].min())\n",
    "print('Max activity date: ', df['ACTIVITY_DATE'].max())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "id": "c4a8984d",
   "metadata": {},
   "outputs": [],
   "source": [
    "variable = 'CREATIONS'\n",
    "\n",
    "def check_one_track_df(one_track_df):\n",
    "    return len(one_track_df)==140\n",
    "\n",
    "def unique_isrc_df(isrc, df):\n",
    "    one_track_df = df[df['ISRC']==isrc].copy()\n",
    "    one_track_df['ACTIVITY_DATE'] = pd.to_datetime(one_track_df['ACTIVITY_DATE'])\n",
    "    one_track_df['ACTIVITY_DATE'] = one_track_df['ACTIVITY_DATE'].dt.date\n",
    "    one_track_df = one_track_df[['ISRC', 'ACTIVITY_DATE', variable]]\n",
    "    return one_track_df.sort_values(by = 'ACTIVITY_DATE')\n",
    "\n",
    "def fourier_table(one_track_df, loop=60):\n",
    "    y = np.array(one_track_df[variable])\n",
    "    y_fft_filtered = fftpack.fft(y).copy()\n",
    "    freq = fftpack.fftfreq(len(y), d = 1/len(y))\n",
    "    cut_off = len(y)/loop\n",
    "    y_fft_filtered[np.abs(freq) > cut_off]=0\n",
    "    Fourier = fftpack.ifft(y_fft_filtered)\n",
    "    y_diff = pd.Series(Fourier).diff()\n",
    "    y_diff_diff = y_diff.diff()\n",
    "    locs2 = abs(np.diff(np.sign(y_diff_diff)))\n",
    "    zero_diff_diff = np.where(np.logical_and(locs2 !=0, np.isnan(locs2)==False))\n",
    "    inflection_point = one_track_df.iloc[zero_diff_diff[0]]\n",
    "    Fourier = pd.DataFrame(Fourier, columns = ['Fourier'])\n",
    "    Fourier['ACTIVITY_DATE'] = np.array(one_track_df['ACTIVITY_DATE'])\n",
    "    streams = pd.DataFrame(one_track_df[[variable,'ACTIVITY_DATE']])\n",
    "    df_fourier = streams.merge(Fourier, on = ['ACTIVITY_DATE'])\n",
    "    df_fourier['Inflection_Point'] = np.where(df_fourier['ACTIVITY_DATE']\n",
    "                                              .isin(inflection_point['ACTIVITY_DATE']), 1,0)\n",
    "    df_fourier['Fourier_real_part'] =  np.array(df_fourier['Fourier']).real\n",
    "    df_fourier['ISRC'] = np.array(one_track_df['ISRC'])\n",
    "    return df_fourier\n",
    "\n",
    "def timing_val(func):\n",
    "    def wrapper(*arg, **kw):\n",
    "        '''source: http://www.daniweb.com/code/snippet368.html'''\n",
    "        t1 = time.time()\n",
    "        res = func(*arg, **kw)\n",
    "        t2 = time.time()\n",
    "        return (t2 - t1), res, func.__name__\n",
    "    return wrapper\n",
    "\n",
    "@timing_val\n",
    "def compile_fourier_table(df):\n",
    "    res = pd.DataFrame(columns=['ACTIVITY_DATE', variable, 'Fourier', 'Inflection_Point', 'ISRC'])\n",
    "    unique_isrcs = df['ISRC'].unique()\n",
    "    for isrc in unique_isrcs:\n",
    "        one_track_df = unique_isrc_df(isrc, df)\n",
    "        if check_one_track_df(one_track_df):\n",
    "            one_track_fourier = fourier_table(one_track_df, loop=60)\n",
    "            #res = res.append(one_track_fourier)\n",
    "            res = pd.concat([res,one_track_fourier])\n",
    "        else:\n",
    "            print('Error:',isrc,'data has length', len(one_track_df))\n",
    "    return res\n",
    "\n",
    "def save_dataframe_s3(df):\n",
    "    s3 = boto3.client('s3')\n",
    "    bucket_name = 'dev-cucumbers'\n",
    "    filepath = \"eimpara/Fourier/Fourier_table_TikTok.csv\"\n",
    "    csv_buffer = df.to_csv(index=False).encode('utf-8')\n",
    "    # Save the CSV file to S3\n",
    "    s3.put_object(Body=csv_buffer, Bucket=bucket_name, Key=filepath)\n",
    "    print(f\"Table saved to S3 bucket: {bucket_name}, with file name: {filepath}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 9,
   "id": "ce43cc04",
   "metadata": {},
   "outputs": [],
   "source": [
    "# IDs = df['ISRC'].unique().tolist()\n",
    "# sample_size = 100\n",
    "# random.seed(2023)\n",
    "# random_sample = random.sample(IDs, k=sample_size)\n",
    "# subset =df[df['ISRC'].isin(random_sample)]\n",
    "\n",
    "fourier_table_NEW_DATA = compile_fourier_table(df)\n",
    "FourierTable__NEW_DATA = pd.DataFrame(fourier_table_NEW_DATA[1])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 10,
   "id": "f33bbd8d-b76f-4247-a4d5-4f8c73fa419e",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "(424200, 6)"
      ]
     },
     "execution_count": 10,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "FourierTable__NEW_DATA.shape"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 11,
   "id": "ed24e6eb",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Table saved to S3 bucket: dev-cucumbers, with file name: eimpara/Fourier/Fourier_table_TikTok.csv\n"
     ]
    }
   ],
   "source": [
    "FourierTable__NEW_DATA ['ACTIVITY_DATE'] = pd.to_datetime(FourierTable__NEW_DATA ['ACTIVITY_DATE']).dt.date\n",
    "\n",
    "save_dataframe_s3(FourierTable__NEW_DATA)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 12,
   "id": "ab4179ca",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/html": [
       "<div>\n",
       "<style scoped>\n",
       "    .dataframe tbody tr th:only-of-type {\n",
       "        vertical-align: middle;\n",
       "    }\n",
       "\n",
       "    .dataframe tbody tr th {\n",
       "        vertical-align: top;\n",
       "    }\n",
       "\n",
       "    .dataframe thead th {\n",
       "        text-align: right;\n",
       "    }\n",
       "</style>\n",
       "<table border=\"1\" class=\"dataframe\">\n",
       "  <thead>\n",
       "    <tr style=\"text-align: right;\">\n",
       "      <th></th>\n",
       "      <th>ACTIVITY_DATE</th>\n",
       "      <th>CREATIONS</th>\n",
       "      <th>Fourier</th>\n",
       "      <th>Inflection_Point</th>\n",
       "      <th>ISRC</th>\n",
       "      <th>Fourier_real_part</th>\n",
       "    </tr>\n",
       "  </thead>\n",
       "  <tbody>\n",
       "    <tr>\n",
       "      <th>0</th>\n",
       "      <td>2023-05-16</td>\n",
       "      <td>0</td>\n",
       "      <td>1.222058+0.000000j</td>\n",
       "      <td>0</td>\n",
       "      <td>GBUQQ0732108</td>\n",
       "      <td>1.222058</td>\n",
       "    </tr>\n",
       "    <tr>\n",
       "      <th>1</th>\n",
       "      <td>2023-05-17</td>\n",
       "      <td>1</td>\n",
       "      <td>1.158867-0.000000j</td>\n",
       "      <td>0</td>\n",
       "      <td>GBUQQ0732108</td>\n",
       "      <td>1.158867</td>\n",
       "    </tr>\n",
       "    <tr>\n",
       "      <th>2</th>\n",
       "      <td>2023-05-18</td>\n",
       "      <td>0</td>\n",
       "      <td>1.085583+0.000000j</td>\n",
       "      <td>0</td>\n",
       "      <td>GBUQQ0732108</td>\n",
       "      <td>1.085583</td>\n",
       "    </tr>\n",
       "    <tr>\n",
       "      <th>3</th>\n",
       "      <td>2023-05-19</td>\n",
       "      <td>1</td>\n",
       "      <td>1.002930-0.000000j</td>\n",
       "      <td>0</td>\n",
       "      <td>GBUQQ0732108</td>\n",
       "      <td>1.002930</td>\n",
       "    </tr>\n",
       "    <tr>\n",
       "      <th>4</th>\n",
       "      <td>2023-05-20</td>\n",
       "      <td>0</td>\n",
       "      <td>0.911732-0.000000j</td>\n",
       "      <td>0</td>\n",
       "      <td>GBUQQ0732108</td>\n",
       "      <td>0.911732</td>\n",
       "    </tr>\n",
       "  </tbody>\n",
       "</table>\n",
       "</div>"
      ],
      "text/plain": [
       "  ACTIVITY_DATE CREATIONS             Fourier Inflection_Point          ISRC  \\\n",
       "0    2023-05-16         0  1.222058+0.000000j                0  GBUQQ0732108   \n",
       "1    2023-05-17         1  1.158867-0.000000j                0  GBUQQ0732108   \n",
       "2    2023-05-18         0  1.085583+0.000000j                0  GBUQQ0732108   \n",
       "3    2023-05-19         1  1.002930-0.000000j                0  GBUQQ0732108   \n",
       "4    2023-05-20         0  0.911732-0.000000j                0  GBUQQ0732108   \n",
       "\n",
       "   Fourier_real_part  \n",
       "0           1.222058  \n",
       "1           1.158867  \n",
       "2           1.085583  \n",
       "3           1.002930  \n",
       "4           0.911732  "
      ]
     },
     "execution_count": 12,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "FourierTable__NEW_DATA.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 13,
   "id": "0388daa6-ff20-4adc-91f6-fe6ab0ebcd68",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "Inflection_Point\n",
       "0    412714\n",
       "1     11486\n",
       "Name: count, dtype: int64"
      ]
     },
     "execution_count": 13,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "FourierTable__NEW_DATA['Inflection_Point'].value_counts()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 20,
   "id": "216e8605",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "77980 rows loaded in 0.17032909393310547 seconds\n",
      "Splitting work into 10 chunks\n"
     ]
    },
    {
     "ename": "TypeError",
     "evalue": "cannot concatenate object of type '<class 'list'>'; only Series and DataFrame objs are valid",
     "output_type": "error",
     "traceback": [
      "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
      "\u001b[0;31mTypeError\u001b[0m                                 Traceback (most recent call last)",
      "Cell \u001b[0;32mIn[20], line 171\u001b[0m\n\u001b[1;32m    169\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m rows loaded in \u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m seconds\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;241m.\u001b[39mformat(full_df\u001b[38;5;241m.\u001b[39mshape[\u001b[38;5;241m0\u001b[39m], timing))\n\u001b[1;32m    170\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mSplitting work into \u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m chunks\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;241m.\u001b[39mformat(parallelism))\n\u001b[0;32m--> 171\u001b[0m elapsed_time \u001b[38;5;241m=\u001b[39m \u001b[43msplit_work\u001b[49m\u001b[43m(\u001b[49m\u001b[43mfull_df\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mparallelism\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mwork_load\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    172\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mwork done in \u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m seconds\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;241m.\u001b[39mformat(elapsed_time))\n",
      "Cell \u001b[0;32mIn[20], line 6\u001b[0m, in \u001b[0;36mtiming_val.<locals>.wrapper\u001b[0;34m(*arg, **kw)\u001b[0m\n\u001b[1;32m      4\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21mwrapper\u001b[39m(\u001b[38;5;241m*\u001b[39marg, \u001b[38;5;241m*\u001b[39m\u001b[38;5;241m*\u001b[39mkw):\n\u001b[1;32m      5\u001b[0m     t1 \u001b[38;5;241m=\u001b[39m time\u001b[38;5;241m.\u001b[39mtime()\n\u001b[0;32m----> 6\u001b[0m     res \u001b[38;5;241m=\u001b[39m \u001b[43mfunc\u001b[49m\u001b[43m(\u001b[49m\u001b[38;5;241;43m*\u001b[39;49m\u001b[43marg\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;241;43m*\u001b[39;49m\u001b[38;5;241;43m*\u001b[39;49m\u001b[43mkw\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m      7\u001b[0m     t2 \u001b[38;5;241m=\u001b[39m time\u001b[38;5;241m.\u001b[39mtime()\n\u001b[1;32m      8\u001b[0m     \u001b[38;5;28;01mreturn\u001b[39;00m (t2 \u001b[38;5;241m-\u001b[39m t1), res, func\u001b[38;5;241m.\u001b[39m\u001b[38;5;18m__name__\u001b[39m\n",
      "Cell \u001b[0;32mIn[20], line 36\u001b[0m, in \u001b[0;36msplit_work\u001b[0;34m(full_df, n, f)\u001b[0m\n\u001b[1;32m     34\u001b[0m     df_chunk \u001b[38;5;241m=\u001b[39m full_df[full_df[\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mISRC\u001b[39m\u001b[38;5;124m'\u001b[39m]\u001b[38;5;241m.\u001b[39misin(chunk)]\n\u001b[1;32m     35\u001b[0m     \u001b[38;5;66;03m#splitted_df.append(df_chunk)\u001b[39;00m\n\u001b[0;32m---> 36\u001b[0m     \u001b[43mpd\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mconcat\u001b[49m\u001b[43m(\u001b[49m\u001b[43m[\u001b[49m\u001b[43msplitted_df\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mdf_chunk\u001b[49m\u001b[43m]\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mignore_index\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[38;5;28;43;01mTrue\u001b[39;49;00m\u001b[43m)\u001b[49m\n\u001b[1;32m     37\u001b[0m run_id \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mstr\u001b[39m(time\u001b[38;5;241m.\u001b[39mtime() \u001b[38;5;241m*\u001b[39m \u001b[38;5;241m1000000\u001b[39m)    \n\u001b[1;32m     38\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mISRCs have been splitted in \u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m dataframes for run id=\u001b[39m\u001b[38;5;132;01m{}\u001b[39;00m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;241m.\u001b[39mformat(n, run_id))\n",
      "File \u001b[0;32m~/anaconda3/envs/python3/lib/python3.10/site-packages/pandas/core/reshape/concat.py:372\u001b[0m, in \u001b[0;36mconcat\u001b[0;34m(objs, axis, join, ignore_index, keys, levels, names, verify_integrity, sort, copy)\u001b[0m\n\u001b[1;32m    369\u001b[0m \u001b[38;5;28;01melif\u001b[39;00m copy \u001b[38;5;129;01mand\u001b[39;00m using_copy_on_write():\n\u001b[1;32m    370\u001b[0m     copy \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;01mFalse\u001b[39;00m\n\u001b[0;32m--> 372\u001b[0m op \u001b[38;5;241m=\u001b[39m \u001b[43m_Concatenator\u001b[49m\u001b[43m(\u001b[49m\n\u001b[1;32m    373\u001b[0m \u001b[43m    \u001b[49m\u001b[43mobjs\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    374\u001b[0m \u001b[43m    \u001b[49m\u001b[43maxis\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43maxis\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    375\u001b[0m \u001b[43m    \u001b[49m\u001b[43mignore_index\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mignore_index\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    376\u001b[0m \u001b[43m    \u001b[49m\u001b[43mjoin\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mjoin\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    377\u001b[0m \u001b[43m    \u001b[49m\u001b[43mkeys\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mkeys\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    378\u001b[0m \u001b[43m    \u001b[49m\u001b[43mlevels\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mlevels\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    379\u001b[0m \u001b[43m    \u001b[49m\u001b[43mnames\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mnames\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    380\u001b[0m \u001b[43m    \u001b[49m\u001b[43mverify_integrity\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mverify_integrity\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    381\u001b[0m \u001b[43m    \u001b[49m\u001b[43mcopy\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43mcopy\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    382\u001b[0m \u001b[43m    \u001b[49m\u001b[43msort\u001b[49m\u001b[38;5;241;43m=\u001b[39;49m\u001b[43msort\u001b[49m\u001b[43m,\u001b[49m\n\u001b[1;32m    383\u001b[0m \u001b[43m\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    385\u001b[0m \u001b[38;5;28;01mreturn\u001b[39;00m op\u001b[38;5;241m.\u001b[39mget_result()\n",
      "File \u001b[0;32m~/anaconda3/envs/python3/lib/python3.10/site-packages/pandas/core/reshape/concat.py:462\u001b[0m, in \u001b[0;36m_Concatenator.__init__\u001b[0;34m(self, objs, axis, join, keys, levels, names, ignore_index, verify_integrity, copy, sort)\u001b[0m\n\u001b[1;32m    457\u001b[0m     \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;129;01mnot\u001b[39;00m \u001b[38;5;28misinstance\u001b[39m(obj, (ABCSeries, ABCDataFrame)):\n\u001b[1;32m    458\u001b[0m         msg \u001b[38;5;241m=\u001b[39m (\n\u001b[1;32m    459\u001b[0m             \u001b[38;5;124mf\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mcannot concatenate object of type \u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;132;01m{\u001b[39;00m\u001b[38;5;28mtype\u001b[39m(obj)\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124m; \u001b[39m\u001b[38;5;124m\"\u001b[39m\n\u001b[1;32m    460\u001b[0m             \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124monly Series and DataFrame objs are valid\u001b[39m\u001b[38;5;124m\"\u001b[39m\n\u001b[1;32m    461\u001b[0m         )\n\u001b[0;32m--> 462\u001b[0m         \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mTypeError\u001b[39;00m(msg)\n\u001b[1;32m    464\u001b[0m     ndims\u001b[38;5;241m.\u001b[39madd(obj\u001b[38;5;241m.\u001b[39mndim)\n\u001b[1;32m    466\u001b[0m \u001b[38;5;66;03m# get the sample\u001b[39;00m\n\u001b[1;32m    467\u001b[0m \u001b[38;5;66;03m# want the highest ndim that we have, and must be non-empty\u001b[39;00m\n\u001b[1;32m    468\u001b[0m \u001b[38;5;66;03m# unless all objs are empty\u001b[39;00m\n",
      "\u001b[0;31mTypeError\u001b[0m: cannot concatenate object of type '<class 'list'>'; only Series and DataFrame objs are valid"
     ]
    }
   ],
   "source": [
    "# from multiprocessing import Pool\n",
    "\n",
    "# def timing_val(func):\n",
    "#     def wrapper(*arg, **kw):\n",
    "#         t1 = time.time()\n",
    "#         res = func(*arg, **kw)\n",
    "#         t2 = time.time()\n",
    "#         return (t2 - t1), res, func.__name__\n",
    "#     return wrapper\n",
    "\n",
    "# # returns full df\n",
    "# @timing_val\n",
    "# def load_full_data():\n",
    "#     first_day_pred = df['ACTIVITY_DATE'].max() - timedelta(days=7)\n",
    "#     list_for_pred = FourierTable__NEW_DATA[(FourierTable__NEW_DATA['Inflection_Point']==1) & \n",
    "#                                              (FourierTable__NEW_DATA['ACTIVITY_DATE'] > first_day_pred)]['ISRC'].unique().tolist()\n",
    "#     subset_for_pred = df[df['ISRC'].isin(list_for_pred)].copy()\n",
    "#     return subset_for_pred\n",
    "\n",
    "# def split_universe(df, n):\n",
    "#     def chunks(l, n):    \n",
    "#         for i in range(0, len(l), n):\n",
    "#             yield l[i:i + n]\n",
    "#     isrcs = df['ISRC'].unique().tolist()\n",
    "#     isrcs.sort()\n",
    "#     isrc_chunks = chunks(isrcs, (len(isrcs) // n) + 1)\n",
    "#     return list(isrc_chunks)\n",
    "\n",
    "# @timing_val\n",
    "# def split_work(full_df, n, f):\n",
    "#     isrc_chunks = split_universe(full_df, n)\n",
    "#     splitted_df = []\n",
    "#     for chunk in isrc_chunks:\n",
    "#         df_chunk = full_df[full_df['ISRC'].isin(chunk)]\n",
    "#         #splitted_df.append(df_chunk)\n",
    "#         pd.concat([splitted_df, df_chunk], ignore_index=True)\n",
    "#     run_id = str(time.time() * 1000000)    \n",
    "#     print(\"ISRCs have been splitted in {} dataframes for run id={}\".format(n, run_id))\n",
    "#     with Pool(n) as p:\n",
    "#         splitted_df_with_index = list(enumerate(splitted_df))\n",
    "#         splitted_df_with_args = map(lambda x: (run_id, x[0], x[1]), splitted_df_with_index)\n",
    "#         # parallel computation happens here\n",
    "#         p.map(f, splitted_df_with_args)\n",
    "#     print(\"Moments calculation finished.\")\n",
    "\n",
    "# # Save one chunk of work    \n",
    "# # def save_chunk_s3(df, run_id, chunk_id):\n",
    "#     s3 = boto3.client('s3')\n",
    "#     bucket_name = 'dev-cucumbers'\n",
    "#     filepath = \"eimpara/ARIMA/TikTok/{}_{}.csv\".format(run_id, chunk_id)\n",
    "#     csv_buffer = df.to_csv(index=False).encode('utf-8')\n",
    "#     # Save the CSV file to S3\n",
    "#     s3.put_object(Body=csv_buffer, Bucket=bucket_name, Key=filepath)\n",
    "#     print(f\"Table saved to S3 bucket: {bucket_name}, with file name: {filepath}\")\n",
    "    \n",
    "\n",
    "# def auto_arima(df_isrc):\n",
    "#     cutoff_train_test_sets = 133\n",
    "#     (train, test)           = (df_isrc.iloc[:cutoff_train_test_sets], df_isrc.iloc[cutoff_train_test_sets:])\n",
    "#     train.index = pd.to_datetime(train['ACTIVITY_DATE'])\n",
    "#     train = train.sort_index(axis = 0)\n",
    "#     auto_df = train[[variable]].copy()\n",
    "#     auto_model = pm.auto_arima(auto_df, start_p=1, start_q=1,\n",
    "#                                 test='adf',\n",
    "#                                 max_p=3, max_q=3, m=7,\n",
    "#                                 start_P=0, seasonal=True,\n",
    "#                                 d=None, D=1, trace=False,\n",
    "#                                 error_action='ignore',  \n",
    "#                                 suppress_warnings=True, #stationary=False,\n",
    "#                                 stepwise=True)\n",
    "#     #return (auto_model.seasonal_order, auto_model.order)\n",
    "#     return auto_model\n",
    "\n",
    "# def iterate_group_by_key(df, col_for_key, sort_data=True):\n",
    "#     if sort_data:\n",
    "#         df = df.sort_values(by=col_for_key)\n",
    "#     def key_at(i): return np.array(df[col_for_key].iloc[i])\n",
    "#     index, size = (0, df.shape[0])\n",
    "#     while index < size:\n",
    "#         current_key = key_at(index)\n",
    "#         res = []\n",
    "#         while index < size and list(current_key) == list(key_at(index)):\n",
    "#              #res.append(df.iloc[index])\n",
    "#             pd.concat([res, df.iloc[index]], ignore_index=True)\n",
    "#             index =  index + 1\n",
    "#         resdf = pd.DataFrame(res, columns=df.columns)\n",
    "#         yield resdf\n",
    "\n",
    "# @timing_val\n",
    "# def arima_iscrs(df):\n",
    "#     # new_df = pd.DataFrame(columns=['ISRC','pred_type',\n",
    "#     #                                'mse', 'rmse','nrmse', 'avg_rmse',\n",
    "#     #                                'mae', 'avg_mae', \n",
    "#     #                                'sum_forecast_errors', 'avg_fe','len_df'])\n",
    "#     new_df = pd.DataFrame(columns=['ISRC','pred_type',\n",
    "#                                    'mse', 'avg_streams_train', 'avg_streams_test', 'median_streams_train',\n",
    "#                                    'median_streams_test', 'linear_gradient_train', 'linear_gradient_test',\n",
    "#                                    'sum_forecast_errors', 'len_df'])\n",
    "#     count = 1\n",
    "#     for subdf in iterate_group_by_key(df, ['ISRC']):\n",
    "#         df_isrc = subdf\n",
    "#         isrc    = df_isrc['ISRC'].iloc[0]\n",
    "#         df_isrc = df_isrc.sort_values(by=['ACTIVITY_DATE'])        \n",
    "\n",
    "#         try:\n",
    "#             cutoff_train_test_sets = 133\n",
    "#             start_pred = 133 #132?\n",
    "#             end_pred = 139\n",
    "#             (train, test)           = (df_isrc.iloc[:cutoff_train_test_sets], df_isrc.iloc[cutoff_train_test_sets:])\n",
    "#             (end_test, end_train)   = (len(test), len(train))\n",
    "#             #(seasonal_order, order) = auto_arima(train)\n",
    "#             model = auto_arima(train)\n",
    "\n",
    "#             #arima_model = ARIMA(train[variable], order=order, seasonal_order=seasonal_order,enforce_stationarity=False).fit()\n",
    "#             #pred        = arima_model.get_prediction(start= 1, end = (end_train + end_test-1),  dynamic=False)\n",
    "#             pred = model.predict(n_periods=7, return_conf_int=False)\n",
    "#             #forecast_errors =  np.subtract(np.array(test[variable]), np.array(pred.predicted_mean[start_pred:end_pred]))\n",
    "#             forecast_errors =  np.subtract(np.array(test[variable]), pred)\n",
    "#             mse     = np.square(forecast_errors).mean()\n",
    "#             avg_streams_train = train[variable].mean()\n",
    "#             avg_streams_test = test[variable].mean()\n",
    "#             median_streams_train = train[variable].median()\n",
    "#             median_streams_test = test[variable].median()\n",
    "#             linear_gradient_train = (train.iloc[-1][variable] - train.iloc[0, test.columns.get_loc(variable)])/len(train)\n",
    "#             linear_gradient_test = (test.iloc[-1][variable] - test.iloc[0, test.columns.get_loc(variable)])/len(test)\n",
    "#             sum_forecast_errors =  round(forecast_errors.sum(), 3)\n",
    "#             all_positives       = all(map(lambda x: x > 0, forecast_errors))\n",
    "#             all_negatives       = all(map(lambda x: x < 0, forecast_errors))        \n",
    "#             pred_categorical    = None\n",
    "#             if all_positives:\n",
    "#                 pred_categorical = 'actuals_above_predicted'\n",
    "#             elif all_negatives:\n",
    "#                 pred_categorical = 'actuals_below_predicted'\n",
    "#             else:\n",
    "#                 pred_categorical = 'actuals_crossing_predicted'\n",
    "#             new_df = new_df.append(pd.DataFrame([{\n",
    "#                 'ISRC': isrc, \n",
    "#                 'pred_type': pred_categorical,\n",
    "#                 'mse': mse,\n",
    "#                 'avg_streams_train': avg_streams_train,\n",
    "#                 'avg_streams_test': avg_streams_test,\n",
    "#                 'median_streams_train': median_streams_train,\n",
    "#                 'median_streams_test': median_streams_test,\n",
    "#                 'linear_gradient_train': linear_gradient_train,\n",
    "#                 'linear_gradient_test': linear_gradient_test,\n",
    "#                 'sum_forecast_errors': sum_forecast_errors,\n",
    "#                 'len_df': df_isrc.shape[0]\n",
    "#                 }], columns=new_df.columns))\n",
    "            \n",
    "#             if count % 100 ==0:\n",
    "#                 print(\"Worker processed {} ISRCs.\".format(count))\n",
    "#             count = count+1\n",
    "#         except Exception as e:\n",
    "#             print(\"Error while processing ISRC {}: '{}'\".format(isrc, e))\n",
    "#     return new_df\n",
    "    \n",
    "# def work_load(arg):\n",
    "#     (run_id, chunk_id, df) = arg\n",
    "#     print(\"Arguments {},{},{}\".format(run_id, chunk_id, type(df)))\n",
    "#     timing, res, _ = arima_iscrs(df)\n",
    "#     save_chunk_s3(res, run_id, chunk_id)\n",
    "#     return res\n",
    "\n",
    "# if __name__ == '__main__':\n",
    "\n",
    "#     #### Running the code\n",
    "#     parallelism = 10\n",
    "#     timing, full_df, _ = load_full_data()\n",
    "#     print(\"{} rows loaded in {} seconds\".format(full_df.shape[0], timing))\n",
    "#     print(\"Splitting work into {} chunks\".format(parallelism))\n",
    "#     elapsed_time = split_work(full_df, parallelism, work_load)\n",
    "#     print(\"work done in {} seconds\".format(elapsed_time))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "34371ce3",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9420629e",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "conda_python3",
   "language": "python",
   "name": "conda_python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.10.12"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
