{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 12,
   "metadata": {},
   "outputs": [],
   "source": [
    "import mesa\n",
    "import math\n",
    "from enum import Enum\n",
    "import networkx as nx\n",
    "import matplotlib.pyplot as plt\n",
    "import numpy as np\n",
    "import pandas as pd\n",
    "import os\n",
    "import pickle\n",
    "from tqdm import tqdm\n",
    "ITERATIONS = 100000\n",
    "%matplotlib inline"
   ]
  },
  {
   "attachments": {},
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "echo_party is the majority party. To get liberal majority results. Run with echo_party = 'liberal'. For conservative results, run with echo_party = 'conservative'"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 13,
   "metadata": {},
   "outputs": [],
   "source": [
    "echo_party = 'conservative'\n",
    "opposition_n = 1"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:318: UserWarning: Trying to unpickle estimator LogisticRegression from version 1.1.2 when using version 1.2.2. This might lead to breaking code or invalid results. Use at your own risk. For more info please refer to:\n",
      "https://scikit-learn.org/stable/model_persistence.html#security-maintainability-limitations\n",
      "  warnings.warn(\n"
     ]
    }
   ],
   "source": [
    "# load in the logistic regression model to map morals to emotions\n",
    "lrmodel = pickle.load(open('./moral2emotelr.pkl', 'rb'))\n",
    "moral_categories = ['purity', 'authority', 'fairness', 'degradation', 'care', \n",
    "                    'loyalty', 'subversion', 'cheating', 'harm', 'betrayal']\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 15,
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "array([[ 2.14834134,  0.37811929,  1.58736116, -1.36935101,  1.95219789,\n",
       "         1.8168247 , -1.1393514 , -2.13577438, -2.34801116, -0.9017263 ]])"
      ]
     },
     "execution_count": 15,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "lrmodel.coef_"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 16,
   "metadata": {},
   "outputs": [],
   "source": [
    "## generate graphs and centrality statistics at the beginning of each sim.\n",
    "class GraphHolder:\n",
    "    def __init__(\n",
    "        self,\n",
    "        num_nodes=10, \n",
    "        avg_node_degree=3, \n",
    "        iter = ITERATIONS,\n",
    "    ):\n",
    "        self.G_ = []\n",
    "        self.degree_ = []\n",
    "        self.eigen_ = []\n",
    "        self.betweenness_ = []\n",
    "        \n",
    "        for _ in range(iter):\n",
    "            G = nx.erdos_renyi_graph(n=num_nodes, p=1/avg_node_degree)\n",
    "            #G = nx.complete_graph(num_nodes)\n",
    "            self.G_.append(G)\n",
    "            deg = nx.degree_centrality(G)\n",
    "            # networkx degree normalizes degree by dividing by |V|-1. Undoing that normalization:\n",
    "            degree_renorm = {key:(val*(num_nodes - 1)) for key, val in deg.items()}\n",
    "            self.degree_.append(degree_renorm)\n",
    "            #self.eigen_.append(nx.eigenvector_centrality(G))\n",
    "            #self.betweenness_.append(nx.betweenness_centrality(G))\n",
    "            \n",
    "        self.G = {idx: i for idx, i in enumerate(self.G_)}\n",
    "        self.degree = {idx: i for idx, i in enumerate(self.degree_)}\n",
    "        #self.eigen = {idx: i for idx, i in enumerate(self.eigen_)}\n",
    "        #self.betweenness = {idx: i for idx, i in enumerate(self.betweenness_)}\n",
    "    #self.degre\n",
    "\n",
    "class State(Enum):\n",
    "    READ = 1\n",
    "    UNREAD = 0"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 17,
   "metadata": {},
   "outputs": [],
   "source": [
    "GH = GraphHolder()"
   ]
  },
  {
   "attachments": {},
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Construct the Agent Class"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 18,
   "metadata": {},
   "outputs": [],
   "source": [
    "class EmoAgent(mesa.Agent):\n",
    "    def __init__(\n",
    "        self,\n",
    "        unique_id,\n",
    "        model,\n",
    "        read_state,\n",
    "        initial_state,\n",
    "        degree_centrality,\n",
    "        initial_message,\n",
    "        party,\n",
    "    ):\n",
    "        super().__init__(unique_id, model)\n",
    "\n",
    "        self.initial_state = initial_state\n",
    "        self.current_state = initial_state\n",
    "        self.state_cache = list(initial_state)\n",
    "        self.degree = degree_centrality\n",
    "        self.susceptibility = 1/(degree_centrality+1)\n",
    "        self.read_state = read_state\n",
    "        self.initial_message = initial_message\n",
    "        self.party = party\n",
    "    \n",
    "    def look_at_neighbors(self):\n",
    "        neighbors_nodes = self.model.grid.get_neighbors(self.pos, include_center=False)\n",
    "        read_neighbors = [\n",
    "            agent\n",
    "            for agent in self.model.grid.get_cell_list_contents(neighbors_nodes)\n",
    "            if agent.read_state is State.READ\n",
    "        ]\n",
    "        \n",
    "        comments = [a.current_state for a in read_neighbors]\n",
    "        if len(comments) == 0:\n",
    "            mean_comments = np.NAN\n",
    "        elif len(comments) == 1:\n",
    "            mean_comments = comments[0]\n",
    "        else:\n",
    "            mean_comments = np.mean(comments, axis = 0)\n",
    "        return mean_comments                    \n",
    "            \n",
    "    def read_messages_and_post(self):\n",
    "        mean_comments = self.look_at_neighbors()\n",
    "        if mean_comments is np.NAN:\n",
    "            self.current_state = np.tanh(self.initial_state + (self.susceptibility*self.initial_message)).reshape(1,10)\n",
    "            self.read_state = State.READ\n",
    "        else:\n",
    "            self.current_state = np.tanh(self.initial_state + (self.susceptibility*self.initial_message + mean_comments)).reshape(1,10)\n",
    "            self.read_state = State.READ\n",
    "        self.state_cache.append(self.current_state)\n",
    "    \n",
    "    def step(self):\n",
    "        self.read_messages_and_post()"
   ]
  },
  {
   "attachments": {},
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Helper functions and the model"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 19,
   "metadata": {},
   "outputs": [],
   "source": [
    "def opposition_party(echo_party):\n",
    "    if echo_party == 'conservative':\n",
    "        return 'liberal'\n",
    "    elif echo_party == 'liberal':\n",
    "        return 'conservative'\n",
    "\n",
    "def gen_political_agent(echo_party):\n",
    "    if echo_party == 'conservative':\n",
    "        # Rep - higher authority & Purity\n",
    "        poli_init = np.concatenate([np.random.uniform(low=0.5, high=1, size=(1)), # purity\n",
    "                    np.random.uniform(low=0.5, high =1, size= (1)), #authority\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), #fairness\n",
    "                    np.random.uniform(low=0.5, high=1, size=(1)), # degradation\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), # care\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), # loyalty\n",
    "                    np.random.uniform(low=0.5, high=1, size=(1)), # subversion\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), # cheating\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), # harm\n",
    "                    np.random.uniform(low=0, high=1, size=(1))]).reshape(1,10) #betrayal\n",
    "    elif echo_party == 'liberal':\n",
    "        # Dem: higher fairness/harm & fairness/cheating\n",
    "        poli_init = np.concatenate([np.random.uniform(low=0, high=1, size=(1)), # purity\n",
    "                    np.random.uniform(low=0, high =1, size= (1)), #authority\n",
    "                    np.random.uniform(low=0.5, high=1, size=(1)), #fairness\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), # degradation\n",
    "                    np.random.uniform(low=0.5, high=1, size=(1)), # care\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), # loyalty\n",
    "                    np.random.uniform(low=0, high=1, size=(1)), # subversion\n",
    "                    np.random.uniform(low=0.5, high=1, size=(1)), # cheating\n",
    "                    np.random.uniform(low=0.5, high=1, size=(1)), # harm\n",
    "                    np.random.uniform(low=0, high=1, size=(1))]).reshape(1,10) #betrayal\n",
    "    else:\n",
    "        TypeError('Must be liberal or conservative')\n",
    "    return poli_init\n",
    "\n",
    "def number_state(model, state):\n",
    "    return sum(1 for a in model.grid.get_all_cell_contents() if a.read_state is state)\n",
    "\n",
    "def number_read(model):\n",
    "    return number_state(model, State.READ)\n",
    "\n",
    "def mean_morality(model):\n",
    "    morality = [a.current_state for a in model.grid.get_all_cell_contents()]\n",
    "    mean_morality = np.mean(morality, axis = 0)\n",
    "    return mean_morality\n",
    "\n",
    "class MoralPosts(mesa.Model):\n",
    "    \"\"\"A virus model with some number of agents\"\"\"\n",
    "\n",
    "    def __init__(\n",
    "        self,\n",
    "        id,\n",
    "        initial_message,\n",
    "        G,\n",
    "        degree,\n",
    "        echo_party,\n",
    "    ):\n",
    "        self.degree = degree\n",
    "        self.id = id\n",
    "        self.G = G\n",
    "        self.initial_message = initial_message\n",
    "        self.n_iter = 10\n",
    "        self.iter_counter = 0\n",
    "        self.grid = mesa.space.NetworkGrid(self.G)\n",
    "        self.schedule = mesa.time.RandomActivation(self)\n",
    "        self.echo_party = echo_party\n",
    "        self.opposition = opposition_party(self.echo_party)\n",
    "        \n",
    "        self.datacollector = mesa.DataCollector(\n",
    "            model_reporters={\"Read\": number_read,\n",
    "                             'Mean_morality':mean_morality},\n",
    "            agent_reporters={\"Morals\": lambda _: _.current_state,\n",
    "                             \"Degree\": lambda _: _.degree,\n",
    "                             \"Party\": lambda _: _.party},\n",
    "        )\n",
    "\n",
    "        opposition_member = list(np.random.choice(list(range(10)), size = opposition_n, replace = False))\n",
    "        #opposition_member = np.random.randint(low=0, high=10, size = None)\n",
    "        for i, node in enumerate(self.G.nodes()):\n",
    "            if i in opposition_member:\n",
    "                a = EmoAgent(\n",
    "                    i,\n",
    "                    self,\n",
    "                    State.UNREAD,\n",
    "                    gen_political_agent(self.opposition),\n",
    "                    self.degree[i],\n",
    "                    self.initial_message,\n",
    "                    self.opposition\n",
    "                )\n",
    "                self.schedule.add(a)\n",
    "                # Add the agent to the node\n",
    "                self.grid.place_agent(a, node)            \n",
    "            else:\n",
    "                a = EmoAgent(\n",
    "                    i,\n",
    "                    self,\n",
    "                    State.UNREAD,\n",
    "                    gen_political_agent(self.echo_party),\n",
    "                    self.degree[i],\n",
    "                    self.initial_message,\n",
    "                    self.echo_party\n",
    "                )\n",
    "                self.schedule.add(a)\n",
    "                # Add the agent to the node\n",
    "                self.grid.place_agent(a, node)\n",
    "            \n",
    "        self.running = True\n",
    "        self.datacollector.collect(self)\n",
    "\n",
    "        #nodes = list(self.G)\n",
    "        #angry_nodes = self.random.sample(nodes, 1)\n",
    "        #for a in self.grid.get_cell_list_contents(angry_nodes):\n",
    "        #    a.state = State.ANGRY\n",
    "    \n",
    "    def new_post(self):\n",
    "        for a in self.grid.get_cell_list_contents(self.G.nodes()):\n",
    "            a.read_state = State.UNREAD\n",
    "\n",
    "    def step(self):\n",
    "        self.schedule.step()\n",
    "        self.datacollector.collect(self)\n",
    "        #self.new_post()\n",
    "\n",
    "    def run_model(self):\n",
    "        self.step()"
   ]
  },
  {
   "attachments": {},
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Run the simulation"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 20,
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "  0%|          | 0/10 [00:00<?, ?it/s]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 10%|█         | 1/10 [03:00<27:03, 180.33s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 20%|██        | 2/10 [05:54<23:34, 176.87s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 30%|███       | 3/10 [08:49<20:29, 175.68s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 40%|████      | 4/10 [11:47<17:40, 176.74s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 50%|█████     | 5/10 [14:44<14:44, 176.96s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 60%|██████    | 6/10 [17:43<11:50, 177.57s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 70%|███████   | 7/10 [20:45<08:57, 179.16s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 80%|████████  | 8/10 [23:47<05:59, 179.95s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      " 90%|█████████ | 9/10 [26:45<02:59, 179.44s/it]/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "/Users/evanwilliams/.pyenv/versions/3.9.1/lib/python3.9/site-packages/sklearn/base.py:439: UserWarning: X does not have valid feature names, but LogisticRegression was fitted with feature names\n",
      "  warnings.warn(\n",
      "100%|██████████| 10/10 [29:46<00:00, 178.62s/it]\n"
     ]
    }
   ],
   "source": [
    "final_stats = []\n",
    "graph_list = GH.G\n",
    "degree_list = GH.degree\n",
    "\n",
    "for initial_message in tqdm((np.eye(10))):\n",
    "    initial_message = initial_message.reshape(1, 10)\n",
    "    initial_morality = []\n",
    "    final_morality = []\n",
    "    initial_mean_morality = []\n",
    "    final_mean_morality = []\n",
    "    degree_cache = []\n",
    "    party_cache = []\n",
    "    for idx in range(ITERATIONS):\n",
    "        MP = MoralPosts(0, initial_message, G=graph_list[idx], degree = degree_list[idx], echo_party=echo_party)\n",
    "        MP.run_model()\n",
    "        # collect results\n",
    "        mean_moral_states = MP.datacollector.get_model_vars_dataframe()\n",
    "        initial_mean_morality.append(mean_moral_states.iloc[0].Mean_morality)\n",
    "        final_mean_morality.append(mean_moral_states.iloc[1].Mean_morality)\n",
    "        agent_states = pd.DataFrame(MP.datacollector.get_agent_vars_dataframe().to_records())\n",
    "        initial_morality.append(np.concatenate(agent_states[agent_states.Step == 0]['Morals'].tolist()))\n",
    "        final_morality.append(np.concatenate(agent_states[agent_states.Step == 1]['Morals'].tolist()))\n",
    "        degree_cache.append(agent_states[agent_states.Step == 0]['Degree'].tolist())\n",
    "        party_cache.append(agent_states[agent_states.Step == 0]['Party'].tolist())\n",
    "    \n",
    "    initial_morality = np.concatenate(initial_morality)\n",
    "    final_morality = np.concatenate(final_morality)\n",
    "    im_pred = lrmodel.predict(initial_morality)\n",
    "    fn_pred = lrmodel.predict(final_morality)\n",
    "    degree_cache = [item for sublist in degree_cache for item in sublist]\n",
    "    party_cache = [item for sublist in party_cache for item in sublist]\n",
    "\n",
    "    uh = pd.DataFrame({'initial':im_pred, 'final':fn_pred})\n",
    "    uh['prod'] = uh.final * uh.initial\n",
    "    uh['changed'] = np.where(uh['prod'] == -1, 1, 0)\n",
    "    uh['positive_change'] = np.where(uh['changed'] ==uh['final'], 1, 0)\n",
    "    uh['neg_change'] = 1*(uh.changed != uh.positive_change)\n",
    "    uh['degree'] = degree_cache\n",
    "    uh['party'] = party_cache\n",
    "    uh['echo'] = echo_party\n",
    "    minority = uh[uh['party'] != uh['echo']].copy()\n",
    "    minority = minority[minority['changed'] == 1].copy()\n",
    "    \n",
    "    final_stats.append({'trials': uh.shape[0],\n",
    "                        'initial_pos':(uh['initial']==1).sum(),\n",
    "                        'initial_neg':(uh['initial']==-1).sum(),\n",
    "                        'final_pos':(uh['final']==1).sum(),\n",
    "                        'final_neg':(uh['final']==-1).sum(),\n",
    "                        'changes':uh['changed'].sum(),\n",
    "                        'pos_changes':uh['positive_change'].sum(),\n",
    "                        'neg_changes':uh['changed'].sum() - uh['positive_change'].sum(),\n",
    "                        'pos_correlation':uh['degree'].corr(uh['positive_change']),\n",
    "                        'neg_correlation':uh['degree'].corr(uh['neg_change']),\n",
    "                        'minority_pos_changes':minority.positive_change.sum(),\n",
    "                        'minority_neg_changes':minority.neg_change.sum(),})"
   ]
  },
  {
   "attachments": {},
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Export results"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 21,
   "metadata": {},
   "outputs": [],
   "source": [
    "out = pd.DataFrame(final_stats)\n",
    "out['morals'] = moral_categories\n",
    "out.to_csv('../results/conservative_maj_100k.csv')"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3.9.1 64-bit ('3.9.1')",
   "language": "python",
   "name": "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.9.1"
  },
  "orig_nbformat": 4,
  "vscode": {
   "interpreter": {
    "hash": "428fe311e7c18bbcbe168afadaa7ed4121f27507cb52005f46617a2bd5ad930b"
   }
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
