{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"markdown","source":"This kernel is written for visualizing hits data  points through KdTree to plot points with nearest neighbor with radius provided  .\nI am taking "},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"515ffc7e567656501246069162026890c14df08b"},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load in \n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the \"../input/\" directory.\n# For example, running this (by clicking run or pressing Shift+Enter) will list the files in the input directory\n\nimport os\nfrom trackml.dataset import load_event, load_dataset\nfrom trackml.score import score_event\nfrom mpl_toolkits.mplot3d import Axes3D\nfrom sklearn.neighbors import KDTree\nimport matplotlib.pyplot as plt\n\n\npath_to_test=\"../input/train_1/\"\nevent_prefix='event000001000'\n\nhits, cells, particles, truth = load_event(os.path.join(path_to_test, event_prefix))\n","execution_count":5,"outputs":[]},{"metadata":{"_uuid":"91da5a972aabe82bfa209809f8c82445ff134cda"},"cell_type":"markdown","source":"# helper functions"},{"metadata":{"trusted":true,"collapsed":true,"_uuid":"35e5049d6c54287a90aa3a435b65c79b7d320509"},"cell_type":"code","source":"def get_nn_tuples(all_nns):\n    key_value_map = dict()\n    for nns in all_nns:\n        sz = len(nns)\n        if sz > 0:\n            if sz in key_value_map.keys():\n                key_value_map[sz] = key_value_map[sz] + 1\n            else:\n                key_value_map[sz] = 1\n\n    print(key_value_map)\n    nn_tuples = []\n    for nns in all_nns:\n        nn_tuples.extend(list(totuple(nns)))\n\n    print(\"length of nn tuples \" + str(len(nn_tuples)))\n    return nn_tuples\n\n\ndef filter_z_component(hits, axis, threshold):\n    hits_z_component = hits.z\n    key_value_map = dict()\n\n    filter_set_max_hits = set()\n\n    # count the no of hits in one z point\n    for n in hits_z_component:\n        if n in key_value_map.keys():\n            key_value_map[n] = key_value_map[n] + 1\n        else:\n            key_value_map[n] = 1\n\n    sorted_key_value = [(k, key_value_map[k]) for k in sorted(key_value_map, key=key_value_map.get, reverse=True)]\n\n    # looking for disks having maximum hists\n    for k, v in sorted_key_value:\n        if k > axis and v > threshold:\n            #print((k, v))\n            filter_set_max_hits.add(k)\n\n    print(\"size after filter  \" + str(len(filter_set_max_hits)))\n\n    return filter_set_max_hits\n\n\ndef get_nn_tuples_list(tree, start, radius):\n    all_nn_indices = tree.query_radius(start, r=radius)\n\n    all_nns = [\n        [start[idx] for idx in nn_indices if idx != i]\n        for i, nn_indices in enumerate(all_nn_indices)\n    ]\n    return all_nns\n\n\ndef get_nn_tuples_representative(tree, start, radius):\n    all_nn_indices = tree.query_radius(start, r=radius)\n\n    all_nns = [\n        [start[idx] for idx in nn_indices if idx != i]\n        for i, nn_indices in enumerate(all_nn_indices)\n    ]\n\n    alls = []\n    for nns in all_nns:\n        if len(nns) > 0:\n            alls.append(nns.pop())\n\n    return alls\n\ndef totuple(a):\n    try:\n        return tuple(totuple(i) for i in a)\n    except TypeError:\n        return a\n\ndef plot_nn_tuples(nn_tuples):\n    dt = np.dtype('float,float,float')\n    xarr = np.array(nn_tuples, dtype=dt)\n    L = xarr['f0']\n    M = xarr['f1']\n    N = xarr['f2']\n\n    fig = plt.figure(1)\n    plt.suptitle('Plot of major points and their nns  ', fontsize=16)\n    ax = fig.add_subplot(111, projection='3d')\n\n    ax.scatter(L, M, N, zdir='z')\n\n    \n\n\ndef plot_all_data(hits):\n    fig = plt.figure(2)\n    plt.suptitle('Plot all data points ', fontsize=16)\n    ax = fig.add_subplot(111, projection='3d')\n\n    ax.plot(hits.x, hits.y, hits.z, zdir='z')\n\n    \n","execution_count":1,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"6b8339e9537992f0d1d16e35cef9fd116fee0708","collapsed":true},"cell_type":"code","source":"plot_all_data(hits)\n\n# look at axis means positive part or negative part in z axis and threshold is number of instances in one z point\nfilter_set_max_hits = filter_z_component(hits, 0, 900)\n\nhits_stacked_data = np.dstack((hits.x, hits.y, hits.z))\n\nhits_tuples = []\nfor row in hits_stacked_data:\n    hits_tuples.extend(totuple(row))\n\nprint(len(hits_tuples))\n\n# filter tuples which are present in points we select in z axis\nfiltered_hits_tuples = [t for t in hits_tuples if t[2] in filter_set_max_hits]\n\nprint(len(filtered_hits_tuples))\n\n# feed data to KDTree\nkdtree = KDTree(filtered_hits_tuples, leaf_size=5)\n\n# get nearest neighbor points from filtered_hits_tuples\n# all_nns are list of list of nn co ordinates\nall_nns = get_nn_tuples_list(kdtree, filtered_hits_tuples, 1.5)\n\n# print distance with no of instance found and return tuple\nnn_tuples = get_nn_tuples(all_nns)\nplot_nn_tuples(nn_tuples)\nplt.show()","execution_count":8,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.5","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}