zachlopez/sample_3
0
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 