CoolFace
Apppublic

zachlopez/sample_3

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
parse_baselines.ipynb590 linesDownload Raw Back to human_annotation
1{2 "cells": [3  {4   "cell_type": "markdown",5   "metadata": {},6   "source": [7    "Baseline human labels for ours vs. other methods, with 3-per-row voting."8   ]9  },10  {11   "cell_type": "code",12   "execution_count": 1,13   "metadata": {},14   "outputs": [],15   "source": [16    "import csv\n",17    "import numpy as np\n",18    "import matplotlib.pyplot as plt\n",19    "from scipy import stats\n",20    "from collections import defaultdict\n",21    "\n",22    "MAX_FILES=2"23   ]24  },25  {26   "cell_type": "code",27   "execution_count": 2,28   "metadata": {},29   "outputs": [],30   "source": [31    "def get_data(filename):\n",32    "    csvfile = open(filename)\n",33    "    reader = csv.reader(csvfile)\n",34    "\n",35    "    data = []\n",36    "    for i, row in enumerate(reader):\n",37    "        if i == 0:\n",38    "            headers = row\n",39    "        else:\n",40    "            data.append(row)\n",41    "    csvfile.close()\n",42    "    return headers, data"43   ]44  },45  {46   "cell_type": "markdown",47   "metadata": {},48   "source": [49    "# Get stats\n",50    "\n",51    "Run these cells in order to:\n",52    "* get stats for ontopicness and fluency to copy/paste\n",53    "* save percents for each topic for plotting"54   ]55  },56  {57   "cell_type": "markdown",58   "metadata": {},59   "source": [60    "## topics"61   ]62  },63  {64   "cell_type": "code",65   "execution_count": 3,66   "metadata": {67    "scrolled": true68   },69   "outputs": [],70   "source": [71    "# for topics\n",72    "def decode(st):\n",73    "    ints = [int(s) for s in st.split('_')]\n",74    "    # Version 2\n",75    "    ii, j1, j2 = ints[0], np.mod(ints[1], MAX_FILES), np.mod(ints[2], MAX_FILES)\n",76    "    return ii, j1, j2\n",77    "\n",78    "# p-value of two binomial distributions\n",79    "# one sided tail\n",80    "def two_samp(x1, x2, n1, n2):\n",81    "    p1 = x1/n1\n",82    "    p2 = x2/n2\n",83    "    phat = (x1 + x2) / (n1 + n2)\n",84    "    z = (p1 - p2) / np.sqrt(phat * (1-phat) * (1/n1 + 1/n2))\n",85    "    return stats.norm.sf(np.abs(z))\n",86    "\n",87    "def print_info_t(scores, counts, single_pvalue=True):\n",88    "    pvalues = np.zeros((MAX_FILES, MAX_FILES))\n",89    "    for i in range(MAX_FILES):\n",90    "        for j in range(i, MAX_FILES):\n",91    "            dist_i = [1] * scores[i] + [0] * (counts[i] - scores[i])\n",92    "            dist_j = [1] * scores[j] + [0] * (counts[j] - scores[j])\n",93    "            pvalue = two_samp(scores[i], scores[j], counts[i], counts[j])\n",94    "            pvalues[i, j] = pvalue\n",95    "            pvalues[j, i] = pvalue\n",96    "    percs = scores / counts\n",97    "\n",98    "    print('total counts, on topic counts, percentages:')\n",99    "    for i in range(MAX_FILES):\n",100    "        if i == 0 and single_pvalue and MAX_FILES == 2:\n",101    "            print('{},{},{},{}'.format(counts[i], scores[i], percs[i], pvalues[0][1]))\n",102    "        else:\n",103    "            print('{},{},{}'.format(counts[i], scores[i], percs[i]))\n",104    "\n",105    "    if not (single_pvalue and MAX_FILES == 2):\n",106    "        for row in pvalues:\n",107    "            print('{},{}'.format(row[0],row[1]))\n",108    "\n",109    "def get_counts_indices(data, order_index, label_indices):\n",110    "    scores = np.zeros(MAX_FILES, dtype=int)\n",111    "    counts = np.zeros(MAX_FILES, dtype=int)\n",112    "    skipped = 0\n",113    "    for rownum, row in enumerate(data):\n",114    "        order = row[order_index]\n",115    "        for label_index in label_indices:\n",116    "            label = row[label_index].lower()\n",117    "            if len(order) > 0 and len(label) > 0:\n",118    "                a_cat, b_cat = decode(order)[1:]\n",119    "                # print(label, order, a_cat, b_cat)\n",120    "                if label == 'a' or label == 'both':\n",121    "                    scores[a_cat] += 1\n",122    "                if label == 'b' or label == 'both':\n",123    "                    scores[b_cat] += 1\n",124    "                counts[a_cat] += 1\n",125    "                counts[b_cat] += 1\n",126    "                if label not in ['a', 'b', 'both', 'neither']:\n",127    "                    print('******invalid label: {}'.format(label))\n",128    "            else:\n",129    "                #print('empty label; skipping', rownum)\n",130    "                skipped += 1\n",131    "    print('skipped {}'.format(skipped))\n",132    "    print_info_t(scores, counts)\n",133    "    return scores, counts\n",134    "\n",135    "# vote by row. each row contributes to one count (and 0 or 1 score based on majority vote)\n",136    "def get_counts_vote_row(data, order_index, label_indices):\n",137    "    scores = np.zeros(MAX_FILES, dtype=int)\n",138    "    counts = np.zeros(MAX_FILES, dtype=int)\n",139    "    skipped = 0\n",140    "    for rownum, row in enumerate(data):\n",141    "        order = row[order_index]\n",142    "        if len(order) == 0:\n",143    "            skipped += 1\n",144    "        else:\n",145    "            a_cat, b_cat = decode(order)[1:]\n",146    "            row_score_a, row_score_b, row_counts = 0, 0, 0\n",147    "            for label_index in label_indices:\n",148    "                label = row[label_index].lower()\n",149    "                if len(label) > 0:\n",150    "                    if label == 'a' or label == 'both':\n",151    "                        row_score_a += 1\n",152    "                    if label == 'b' or label == 'both':\n",153    "                        row_score_b += 1\n",154    "                    row_counts += 1\n",155    "                    if label not in ['a', 'b', 'both', 'neither']:\n",156    "                        print('******invalid label: {}'.format(label))\n",157    "                else:\n",158    "                    print('empty label for nonempty prompt', rownum)\n",159    "            # update big points\n",160    "            if row_counts == 3:\n",161    "                scores[a_cat] += row_score_a // 2\n",162    "                scores[b_cat] += row_score_b // 2\n",163    "                counts[a_cat] += 1\n",164    "                counts[b_cat] += 1\n",165    "            else:\n",166    "                print('incomplete row...')\n",167    "    print('skipped {}'.format(skipped))\n",168    "    print_info_t(scores, counts)\n",169    "    return scores, counts"170   ]171  },172  {173   "cell_type": "markdown",174   "metadata": {},175   "source": [176    "## fluency"177   ]178  },179  {180   "cell_type": "code",181   "execution_count": 4,182   "metadata": {183    "scrolled": true184   },185   "outputs": [],186   "source": [187    "def print_info_f_lists(scorelist, single_pvalue=True):\n",188    "    for i in range(MAX_FILES):\n",189    "        if len(scorelist[i]) == 0:\n",190    "            print('skipping; no data')\n",191    "            return\n",192    "\n",193    "    pvalues = np.zeros((MAX_FILES, MAX_FILES))\n",194    "    for i in range(MAX_FILES):\n",195    "        for j in range(i, MAX_FILES):\n",196    "            pvalue = stats.ttest_ind(scorelist[i], scorelist[j]).pvalue\n",197    "            pvalues[i, j] = pvalue\n",198    "            pvalues[j, i] = pvalue\n",199    "\n",200    "    print('mean, stdev, min, max, counts:')\n",201    "    for i in range(MAX_FILES):\n",202    "        if i == 0 and single_pvalue and len(scorelist) == 2:\n",203    "            print('{},{},{},{},{},{}'.format(np.mean(scorelist[i]), np.std(scorelist[i]),\n",204    "                np.min(scorelist[i]), np.max(scorelist[i]), len(scorelist[i]), pvalues[0][1]))\n",205    "        else:\n",206    "            print('{},{},{},{},{}'.format(np.mean(scorelist[i]), np.std(scorelist[i]),\n",207    "                np.min(scorelist[i]), np.max(scorelist[i]), len(scorelist[i])))\n",208    "    if not (single_pvalue and len(scorelist) == 2):\n",209    "        print('p-values')\n",210    "        for row in pvalues:\n",211    "            print('{},{}'.format(row[0],row[1]))\n",212    "\n",213    "def get_fluencies_indices(data, order_index, label_indices):\n",214    "    scorelist = [[], []]\n",215    "    skipped = 0\n",216    "    for r, row in enumerate(data):\n",217    "        order = row[order_index]\n",218    "        if len(order) == 0:\n",219    "            continue\n",220    "        for label_ind_pair in label_indices:\n",221    "            #a_cat, b_cat = decode(order)[1:]\n",222    "            cats = decode(order)[1:]\n",223    "            for i, ind in enumerate(label_ind_pair):\n",224    "                label = row[ind]\n",225    "                if len(label) > 0:\n",226    "                    scorelist[cats[i]].append(int(label))\n",227    "                else:\n",228    "                    skipped += 1\n",229    "    print('skipped {}'.format(skipped))\n",230    "    print_info_f_lists(scorelist)\n",231    "    return scorelist"232   ]233  },234  {235   "cell_type": "markdown",236   "metadata": {},237   "source": [238    "## Run on all files"239   ]240  },241  {242   "cell_type": "code",243   "execution_count": 7,244   "metadata": {},245   "outputs": [],246   "source": [247    "# aggregated human labeled everything\n",248    "dirname = 'ctrl_wd_openai_csvs/'\n",249    "# comment out any of the below if you don't want to include them in \"all\"\n",250    "file_info = [\n",251    "    'ctrl_legal.csv',\n",252    "    'ctrl_politics.csv',\n",253    "    'ctrl_religion.csv',\n",254    "    'ctrl_science.csv',\n",255    "    'ctrl_technologies.csv',\n",256    "    'ctrl_positive.csv',\n",257    "    'ctrl_negative.csv',\n",258    "    'openai_positive.csv',\n",259    "    'greedy_legal.csv',\n",260    "    'greedy_military.csv',\n",261    "    'greedy_politics.csv',\n",262    "    'greedy_religion.csv',\n",263    "    'greedy_science.csv',\n",264    "    'greedy_space.csv',\n",265    "    'greedy_technologies.csv',\n",266    "    'greedy_positive.csv',\n",267    "    'greedy_negative.csv',\n",268    "]"269   ]270  },271  {272   "cell_type": "code",273   "execution_count": 9,274   "metadata": {275    "scrolled": false276   },277   "outputs": [278    {279     "name": "stdout",280     "output_type": "stream",281     "text": [282      "ctrl_legal.csv\n",283      "skipped 0\n",284      "total counts, on topic counts, percentages:\n",285      "20,7,0.35,0.24507648020791256\n",286      "20,5,0.25\n",287      "\n",288      "ctrl_politics.csv\n",289      "skipped 0\n",290      "total counts, on topic counts, percentages:\n",291      "20,7,0.35,0.16864350736717681\n",292      "20,10,0.5\n",293      "\n",294      "ctrl_religion.csv\n",295      "skipped 0\n",296      "total counts, on topic counts, percentages:\n",297      "20,12,0.6,0.000782701129001274\n",298      "20,20,1.0\n",299      "\n",300      "ctrl_science.csv\n",301      "skipped 0\n",302      "total counts, on topic counts, percentages:\n",303      "20,15,0.75,0.012580379600204389\n",304      "20,8,0.4\n",305      "\n",306      "ctrl_technologies.csv\n",307      "skipped 0\n",308      "total counts, on topic counts, percentages:\n",309      "20,15,0.75,0.005502076588434386\n",310      "20,7,0.35\n",311      "\n",312      "ctrl_positive.csv\n",313      "skipped 0\n",314      "total counts, on topic counts, percentages:\n",315      "15,13,0.8666666666666667,0.312103057383203\n",316      "15,12,0.8\n",317      "\n",318      "ctrl_negative.csv\n",319      "skipped 0\n",320      "total counts, on topic counts, percentages:\n",321      "15,8,0.5333333333333333,0.12785217497142026\n",322      "15,11,0.7333333333333333\n",323      "\n",324      "openai_positive.csv\n",325      "skipped 0\n",326      "total counts, on topic counts, percentages:\n",327      "45,38,0.8444444444444444,7.502148606340828e-12\n",328      "45,6,0.13333333333333333\n",329      "\n",330      "greedy_legal.csv\n",331      "skipped 0\n",332      "total counts, on topic counts, percentages:\n",333      "60,26,0.43333333333333335,0.014054020073575932\n",334      "60,38,0.6333333333333333\n",335      "\n",336      "greedy_military.csv\n",337      "skipped 0\n",338      "total counts, on topic counts, percentages:\n",339      "60,21,0.35,0.423683196354148\n",340      "60,20,0.3333333333333333\n",341      "\n",342      "greedy_politics.csv\n",343      "skipped 0\n",344      "total counts, on topic counts, percentages:\n",345      "60,20,0.3333333333333333,0.423683196354148\n",346      "60,21,0.35\n",347      "\n",348      "greedy_religion.csv\n",349      "skipped 0\n",350      "total counts, on topic counts, percentages:\n",351      "60,31,0.5166666666666667,0.004543733726219588\n",352      "60,17,0.2833333333333333\n",353      "\n",354      "greedy_science.csv\n",355      "skipped 0\n",356      "total counts, on topic counts, percentages:\n",357      "60,33,0.55,0.04996165925796605\n",358      "60,24,0.4\n",359      "\n",360      "greedy_space.csv\n",361      "skipped 0\n",362      "total counts, on topic counts, percentages:\n",363      "60,34,0.5666666666666667,2.9438821372586324e-08\n",364      "60,6,0.1\n",365      "\n",366      "greedy_technologies.csv\n",367      "skipped 0\n",368      "total counts, on topic counts, percentages:\n",369      "60,36,0.6,0.014229868458155282\n",370      "60,24,0.4\n",371      "\n",372      "greedy_positive.csv\n",373      "skipped 0\n",374      "total counts, on topic counts, percentages:\n",375      "45,37,0.8222222222222222,6.07065790526639e-09\n",376      "45,10,0.2222222222222222\n",377      "\n",378      "greedy_negative.csv\n",379      "skipped 0\n",380      "total counts, on topic counts, percentages:\n",381      "45,18,0.4,0.0048164878862943334\n",382      "45,7,0.15555555555555556\n",383      "\n",384      "all:\n",385      "total counts, on topic counts, percentages:\n",386      "685,371,0.5416058394160584,5.6920836882984375e-12\n",387      "685,246,0.35912408759124087\n",388      "\n",389      "------------\n",390      "\n",391      "ctrl_legal.csv\n",392      "skipped 0\n",393      "mean, stdev, min, max, counts:\n",394      "3.35,0.6538348415311009,2,5,60,0.21268659490448816\n",395      "3.183333333333333,0.7851043808875918,2,5,60\n",396      "\n",397      "ctrl_politics.csv\n",398      "skipped 0\n",399      "mean, stdev, min, max, counts:\n",400      "3.6333333333333333,0.682316316348624,2,5,60,0.5620319695586566\n",401      "3.7,0.5567764362830021,2,5,60\n",402      "\n",403      "ctrl_religion.csv\n",404      "skipped 0\n",405      "mean, stdev, min, max, counts:\n",406      "3.5833333333333335,0.7369230323144715,2,5,60,0.025496401986981814\n",407      "3.8666666666666667,0.6182412330330469,2,5,60\n",408      "\n",409      "ctrl_science.csv\n",410      "skipped 0\n",411      "mean, stdev, min, max, counts:\n",412      "3.9166666666666665,0.7139483330201298,2,5,60,0.11926537531844811\n",413      "3.7333333333333334,0.5436502143433364,3,5,60\n",414      "\n",415      "ctrl_technologies.csv\n",416      "skipped 0\n",417      "mean, stdev, min, max, counts:\n",418      "3.566666666666667,0.8239471396205517,2,5,60,0.41405751072305697\n",419      "3.683333333333333,0.7186020379103366,1,5,60\n",420      "\n",421      "ctrl_positive.csv\n",422      "skipped 0\n",423      "mean, stdev, min, max, counts:\n",424      "3.7777777777777777,0.5921294486432991,2,5,45,0.2770324945551848\n",425      "3.911111111111111,0.5506449641495051,3,5,45\n",426      "\n",427      "ctrl_negative.csv\n",428      "skipped 0\n",429      "mean, stdev, min, max, counts:\n",430      "2.933333333333333,0.7999999999999999,1,4,45,0.15456038547144507\n",431      "3.1777777777777776,0.7969076034240491,1,4,45\n",432      "\n",433      "openai_positive.csv\n",434      "skipped 0\n",435      "mean, stdev, min, max, counts:\n",436      "3.6814814814814816,0.83134742794656,2,5,135,0.00044715078341087973\n",437      "3.3185185185185184,0.8402103074636584,1,5,135\n",438      "\n",439      "greedy_legal.csv\n",440      "skipped 0\n",441      "mean, stdev, min, max, counts:\n",442      "3.861111111111111,0.6033599339337534,2,5,180,2.0680624898872873e-09\n",443      "3.3722222222222222,0.8757888859990781,1,5,180\n",444      "\n",445      "greedy_military.csv\n",446      "skipped 0\n",447      "mean, stdev, min, max, counts:\n",448      "3.988888888888889,0.7148340047318592,1,5,180,2.9929380752302575e-05\n",449      "3.6222222222222222,0.9138171197756484,1,5,180\n",450      "\n",451      "greedy_politics.csv\n",452      "skipped 0\n",453      "mean, stdev, min, max, counts:\n",454      "3.8222222222222224,0.684393937530422,2,5,180,0.0007209758186600587\n",455      "3.522222222222222,0.957169180190301,1,5,180\n",456      "\n",457      "greedy_religion.csv\n",458      "skipped 0\n",459      "mean, stdev, min, max, counts:\n",460      "3.8333333333333335,0.8975274678557507,1,5,180,4.066885786996924e-09\n",461      "3.2111111111111112,1.0487499908032782,1,5,180\n",462      "\n",463      "greedy_science.csv\n",464      "skipped 0\n",465      "mean, stdev, min, max, counts:\n",466      "3.8777777777777778,0.5836242660741733,2,5,180,0.0006552437647639663\n",467      "3.6166666666666667,0.8318319808978519,1,5,180\n",468      "\n",469      "greedy_space.csv\n",470      "skipped 0\n",471      "mean, stdev, min, max, counts:\n",472      "3.716666666666667,0.8251262529657709,1,5,180,0.0954991700009854\n",473      "3.577777777777778,0.7450246495217872,1,5,180\n",474      "\n",475      "greedy_technologies.csv\n",476      "skipped 0\n",477      "mean, stdev, min, max, counts:\n",478      "4.011111111111111,0.5476098457934048,2,5,180,6.183501355237109e-11\n",479      "3.4555555555555557,0.9563949801335698,1,5,180\n",480      "\n",481      "greedy_positive.csv\n",482      "skipped 0\n",483      "mean, stdev, min, max, counts:\n",484      "3.740740740740741,0.7299209796192601,1,5,135,0.7085851710838819\n",485      "3.7777777777777777,0.8833158628600795,1,5,135\n",486      "\n",487      "greedy_negative.csv\n",488      "skipped 0\n",489      "mean, stdev, min, max, counts:\n",490      "3.762962962962963,0.5863560719159496,2,5,135,0.02527504830979518\n",491      "3.5555555555555554,0.8916623398995057,1,5,135\n",492      "\n",493      "all:\n",494      "mean, stdev, min, max, counts:\n",495      "3.78345498783455,0.735818880968324,1,5,2055,6.013335996010963e-25\n",496      "3.5206812652068127,0.8798559510028829,1,5,2055\n",497      "total counts\n",498      "2055\n",499      "2055\n"500     ]501    },502    {503     "name": "stderr",504     "output_type": "stream",505     "text": [506      "/Users/rosanne/anaconda3/envs/py36/lib/python3.6/site-packages/ipykernel_launcher.py:14: RuntimeWarning: invalid value encountered in double_scalars\n",507      "  \n"508     ]509    }510   ],511   "source": [512    "# hardcoded indices\n",513    "category_index = -1 # index of encoded seed and methods\n",514    "topic_indices = [2, 6, 10]\n",515    "fluency_indices = [(3,4), (7,8), (11,12)]\n",516    "\n",517    "all_scores = np.zeros(MAX_FILES, dtype=int)\n",518    "all_counts = np.zeros(MAX_FILES, dtype=int)\n",519    "percs_ordered = np.zeros((len(file_info), MAX_FILES)) # percents saved in same order as file names\n",520    "for i, fname in enumerate(file_info):\n",521    "    filename = dirname + fname\n",522    "    headers, data = get_data(filename)\n",523    "    print(fname)\n",524    "    scores, counts = get_counts_vote_row(data, category_index, topic_indices)\n",525    "    all_scores += scores\n",526    "    all_counts += counts\n",527    "    percs_ordered[i] = 100 * scores / counts\n",528    "    print()\n",529    "print('all:')\n",530    "print_info_t(all_scores, all_counts)\n",531    "print('\\n------------\\n')\n",532    "\n",533    "# uber labeled fluencies\n",534    "all_fluencies = [[], []]\n",535    "for fname in file_info:\n",536    "    filename = dirname + fname\n",537    "    headers, data = get_data(filename)\n",538    "    print(fname)\n",539    "    new_scores = get_fluencies_indices(data, category_index, fluency_indices)\n",540    "    for i in range(len(all_fluencies)):\n",541    "        all_fluencies[i].extend(new_scores[i])\n",542    "    print()\n",543    "print('all:')\n",544    "print_info_f_lists(all_fluencies)\n",545    "print('total counts')\n",546    "\n",547    "for x in all_fluencies:\n",548    "    print(len(x))\n",549    "    \n",550    "all_scores_hist = all_fluencies"551   ]552  },553  {554   "cell_type": "code",555   "execution_count": null,556   "metadata": {},557   "outputs": [],558   "source": []559  },560  {561   "cell_type": "code",562   "execution_count": null,563   "metadata": {},564   "outputs": [],565   "source": []566  }567 ],568 "metadata": {569  "kernelspec": {570   "display_name": "Python 3",571   "language": "python",572   "name": "python3"573  },574  "language_info": {575   "codemirror_mode": {576    "name": "ipython",577    "version": 3578   },579   "file_extension": ".py",580   "mimetype": "text/x-python",581   "name": "python",582   "nbconvert_exporter": "python",583   "pygments_lexer": "ipython3",584   "version": "3.6.7"585  }586 },587 "nbformat": 4,588 "nbformat_minor": 2589}590