51 KiB
51 KiB
In [1]:
%matplotlib inline
import time
import numpy as np
import tensorflow as tf
from matplotlib import image
import matplotlib.pyplot as plt
from sklearn.cluster import KMeans
from IPython.display import display
np.random.seed(2021)/Users/mhjensen/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/_src/lib/__init__.py:32: UserWarning: JAX on Mac ARM machines is experimental and minimally tested. Please see https://github.com/google/jax/issues/5501 in the event of problems.
warnings.warn("JAX on Mac ARM machines is experimental and minimally tested. "
[0;31m---------------------------------------------------------------------------[0m [0;31mAttributeError[0m Traceback (most recent call last) Input [0;32mIn [1][0m, in [0;36m<cell line: 5>[0;34m()[0m [1;32m 3[0m [38;5;28;01mimport[39;00m [38;5;21;01mtime[39;00m [1;32m 4[0m [38;5;28;01mimport[39;00m [38;5;21;01mnumpy[39;00m [38;5;28;01mas[39;00m [38;5;21;01mnp[39;00m [0;32m----> 5[0m [38;5;28;01mimport[39;00m [38;5;21;01mtensorflow[39;00m [38;5;28;01mas[39;00m [38;5;21;01mtf[39;00m [1;32m 6[0m [38;5;28;01mfrom[39;00m [38;5;21;01mmatplotlib[39;00m [38;5;28;01mimport[39;00m image [1;32m 7[0m [38;5;28;01mimport[39;00m [38;5;21;01mmatplotlib[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpyplot[39;00m [38;5;28;01mas[39;00m [38;5;21;01mplt[39;00m File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/__init__.py:51[0m, in [0;36m<module>[0;34m[0m [1;32m 49[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m autograph [1;32m 50[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m bitwise [0;32m---> 51[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m compat [1;32m 52[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m config [1;32m 53[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m data File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/__init__.py:37[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Compatibility functions.[39;00m [1;32m 4[0m [1;32m 5[0m [38;5;124;03mThe `tf.compat` module contains two sets of compatibility functions.[39;00m [0;32m (...)[0m [1;32m 32[0m [1;32m 33[0m [38;5;124;03m"""[39;00m [1;32m 35[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m---> 37[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v1 [1;32m 38[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v2 [1;32m 39[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m [38;5;28;01mimport[39;00m forward_compatibility_horizon File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/__init__.py:30[0m, in [0;36m<module>[0;34m[0m [1;32m 28[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m autograph [1;32m 29[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m bitwise [0;32m---> 30[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m compat [1;32m 31[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m config [1;32m 32[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m data File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/compat/__init__.py:37[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Compatibility functions.[39;00m [1;32m 4[0m [1;32m 5[0m [38;5;124;03mThe `tf.compat` module contains two sets of compatibility functions.[39;00m [0;32m (...)[0m [1;32m 32[0m [1;32m 33[0m [38;5;124;03m"""[39;00m [1;32m 35[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m---> 37[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v1 [1;32m 38[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v2 [1;32m 39[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m [38;5;28;01mimport[39;00m forward_compatibility_horizon File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/compat/v1/__init__.py:47[0m, in [0;36m<module>[0;34m[0m [1;32m 45[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m layers [1;32m 46[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m linalg [0;32m---> 47[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m lite [1;32m 48[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m logging [1;32m 49[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m lookup File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/lite/__init__.py:9[0m, in [0;36m<module>[0;34m[0m [1;32m 6[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [1;32m 8[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m constants [0;32m----> 9[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m experimental [1;32m 10[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m [38;5;28;01mimport[39;00m Interpreter [1;32m 11[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m [38;5;28;01mimport[39;00m OpHint File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/lite/experimental/__init__.py:8[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Public API for tf.lite.experimental namespace.[39;00m [1;32m 4[0m [38;5;124;03m"""[39;00m [1;32m 6[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m----> 8[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m authoring [1;32m 9[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01manalyzer[39;00m [38;5;28;01mimport[39;00m ModelAnalyzer [38;5;28;01mas[39;00m Analyzer [1;32m 10[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m [38;5;28;01mimport[39;00m OpResolverType File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/lite/experimental/authoring/__init__.py:8[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Public API for tf.lite.experimental.authoring namespace.[39;00m [1;32m 4[0m [38;5;124;03m"""[39;00m [1;32m 6[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m----> 8[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mauthoring[39;00m[38;5;21;01m.[39;00m[38;5;21;01mauthoring[39;00m [38;5;28;01mimport[39;00m compatible File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/lite/python/authoring/authoring.py:43[0m, in [0;36m<module>[0;34m[0m [1;32m 39[0m [38;5;28;01mimport[39;00m [38;5;21;01mfunctools[39;00m [1;32m 42[0m [38;5;66;03m# pylint: disable=g-import-not-at-top[39;00m [0;32m---> 43[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m convert [1;32m 44[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m lite [1;32m 45[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mmetrics[39;00m [38;5;28;01mimport[39;00m converter_error_data_pb2 File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/lite/python/convert.py:29[0m, in [0;36m<module>[0;34m[0m [1;32m 26[0m [38;5;28;01mimport[39;00m [38;5;21;01msix[39;00m [1;32m 28[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m lite_constants [0;32m---> 29[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m util [1;32m 30[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m wrap_toco [1;32m 31[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mconvert_phase[39;00m [38;5;28;01mimport[39;00m Component File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/lite/python/util.py:51[0m, in [0;36m<module>[0;34m[0m [1;32m 47[0m [38;5;66;03m# Jax functions used by TFLite[39;00m [1;32m 48[0m [38;5;66;03m# pylint: disable=g-import-not-at-top[39;00m [1;32m 49[0m [38;5;66;03m# pylint: disable=unused-import[39;00m [1;32m 50[0m [38;5;28;01mtry[39;00m: [0;32m---> 51[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m [38;5;28;01mimport[39;00m xla_computation [38;5;28;01mas[39;00m _xla_computation [1;32m 52[0m [38;5;28;01mexcept[39;00m [38;5;167;01mImportError[39;00m: [1;32m 53[0m _xla_computation [38;5;241m=[39m [38;5;28;01mNone[39;00m File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/__init__.py:116[0m, in [0;36m<module>[0;34m[0m [1;32m 40[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mconfig[39;00m [38;5;28;01mimport[39;00m ( [1;32m 41[0m config [38;5;28;01mas[39;00m config, [1;32m 42[0m enable_checks [38;5;28;01mas[39;00m enable_checks, [0;32m (...)[0m [1;32m 51[0m numpy_rank_promotion [38;5;28;01mas[39;00m numpy_rank_promotion, [1;32m 52[0m ) [1;32m 53[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mapi[39;00m [38;5;28;01mimport[39;00m ( [1;32m 54[0m ad, [38;5;66;03m# TODO(phawkins): update users to avoid this.[39;00m [1;32m 55[0m checkpoint [38;5;28;01mas[39;00m checkpoint, [0;32m (...)[0m [1;32m 114[0m xla_computation [38;5;28;01mas[39;00m xla_computation, [1;32m 115[0m ) [0;32m--> 116[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01mexperimental[39;00m[38;5;21;01m.[39;00m[38;5;21;01mmaps[39;00m [38;5;28;01mimport[39;00m soft_pmap [38;5;28;01mas[39;00m soft_pmap [1;32m 117[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01mversion[39;00m [38;5;28;01mimport[39;00m __version__ [38;5;28;01mas[39;00m __version__ [1;32m 119[0m [38;5;66;03m# These submodules are separate because they are in an import cycle with[39;00m [1;32m 120[0m [38;5;66;03m# jax and rely on the names imported above.[39;00m File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/experimental/maps.py:26[0m, in [0;36m<module>[0;34m[0m [1;32m 23[0m [38;5;28;01mfrom[39;00m [38;5;21;01mfunctools[39;00m [38;5;28;01mimport[39;00m wraps, partial, partialmethod [1;32m 24[0m [38;5;28;01mfrom[39;00m [38;5;21;01menum[39;00m [38;5;28;01mimport[39;00m Enum [0;32m---> 26[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m numpy [38;5;28;01mas[39;00m jnp [1;32m 27[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m core [1;32m 28[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m linear_util [38;5;28;01mas[39;00m lu File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/numpy/__init__.py:19[0m, in [0;36m<module>[0;34m[0m [1;32m 1[0m [38;5;66;03m# Copyright 2018 Google LLC[39;00m [1;32m 2[0m [38;5;66;03m#[39;00m [1;32m 3[0m [38;5;66;03m# Licensed under the Apache License, Version 2.0 (the "License");[39;00m [0;32m (...)[0m [1;32m 17[0m [1;32m 18[0m [38;5;66;03m# flake8: noqa: F401[39;00m [0;32m---> 19[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m fft [38;5;28;01mas[39;00m fft [1;32m 20[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m linalg [38;5;28;01mas[39;00m linalg [1;32m 22[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01minterpreters[39;00m[38;5;21;01m.[39;00m[38;5;21;01mxla[39;00m [38;5;28;01mimport[39;00m DeviceArray [38;5;28;01mas[39;00m DeviceArray File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/numpy/fft.py:17[0m, in [0;36m<module>[0;34m[0m [1;32m 1[0m [38;5;66;03m# Copyright 2020 Google LLC[39;00m [1;32m 2[0m [38;5;66;03m#[39;00m [1;32m 3[0m [38;5;66;03m# Licensed under the Apache License, Version 2.0 (the "License");[39;00m [0;32m (...)[0m [1;32m 14[0m [1;32m 15[0m [38;5;66;03m# flake8: noqa: F401[39;00m [0;32m---> 17[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mnumpy[39;00m[38;5;21;01m.[39;00m[38;5;21;01mfft[39;00m [38;5;28;01mimport[39;00m ( [1;32m 18[0m ifft [38;5;28;01mas[39;00m ifft, [1;32m 19[0m ifft2 [38;5;28;01mas[39;00m ifft2, [1;32m 20[0m ifftn [38;5;28;01mas[39;00m ifftn, [1;32m 21[0m ifftshift [38;5;28;01mas[39;00m ifftshift, [1;32m 22[0m ihfft [38;5;28;01mas[39;00m ihfft, [1;32m 23[0m irfft [38;5;28;01mas[39;00m irfft, [1;32m 24[0m irfft2 [38;5;28;01mas[39;00m irfft2, [1;32m 25[0m irfftn [38;5;28;01mas[39;00m irfftn, [1;32m 26[0m fft [38;5;28;01mas[39;00m fft, [1;32m 27[0m fft2 [38;5;28;01mas[39;00m fft2, [1;32m 28[0m fftfreq [38;5;28;01mas[39;00m fftfreq, [1;32m 29[0m fftn [38;5;28;01mas[39;00m fftn, [1;32m 30[0m fftshift [38;5;28;01mas[39;00m fftshift, [1;32m 31[0m hfft [38;5;28;01mas[39;00m hfft, [1;32m 32[0m rfft [38;5;28;01mas[39;00m rfft, [1;32m 33[0m rfft2 [38;5;28;01mas[39;00m rfft2, [1;32m 34[0m rfftfreq [38;5;28;01mas[39;00m rfftfreq, [1;32m 35[0m rfftn [38;5;28;01mas[39;00m rfftn, [1;32m 36[0m ) [1;32m 38[0m [38;5;66;03m# Module initialization is encapsulated in a function to avoid accidental[39;00m [1;32m 39[0m [38;5;66;03m# namespace pollution.[39;00m [1;32m 40[0m _NOT_IMPLEMENTED [38;5;241m=[39m [] File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/_src/numpy/fft.py:19[0m, in [0;36m<module>[0;34m[0m [1;32m 16[0m [38;5;28;01mimport[39;00m [38;5;21;01moperator[39;00m [1;32m 17[0m [38;5;28;01mimport[39;00m [38;5;21;01mnumpy[39;00m [38;5;28;01mas[39;00m [38;5;21;01mnp[39;00m [0;32m---> 19[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m [38;5;28;01mimport[39;00m lax [1;32m 20[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlib[39;00m [38;5;28;01mimport[39;00m xla_client [1;32m 21[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mutil[39;00m [38;5;28;01mimport[39;00m safe_zip File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/lax/__init__.py:332[0m, in [0;36m<module>[0;34m[0m [1;32m 299[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m [38;5;28;01mimport[39;00m (_reduce_sum, _reduce_max, _reduce_min, _reduce_or, [1;32m 300[0m _reduce_and, _reduce_window_sum, _reduce_window_max, [1;32m 301[0m _reduce_window_min, _reduce_window_prod, [0;32m (...)[0m [1;32m 306[0m _upcast_fp16_for_computation, _broadcasting_shape_rule, [1;32m 307[0m _eye, _tri, _delta, _ones, _zeros, _dilate_shape) [1;32m 308[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcontrol_flow[39;00m [38;5;28;01mimport[39;00m ( [1;32m 309[0m associative_scan [38;5;28;01mas[39;00m associative_scan, [1;32m 310[0m cond [38;5;28;01mas[39;00m cond, [0;32m (...)[0m [1;32m 330[0m while_p [38;5;28;01mas[39;00m while_p, [1;32m 331[0m ) [0;32m--> 332[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mfft[39;00m [38;5;28;01mimport[39;00m ( [1;32m 333[0m fft [38;5;28;01mas[39;00m fft, [1;32m 334[0m fft_p [38;5;28;01mas[39;00m fft_p, [1;32m 335[0m ) [1;32m 336[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mparallel[39;00m [38;5;28;01mimport[39;00m ( [1;32m 337[0m all_gather [38;5;28;01mas[39;00m all_gather, [1;32m 338[0m all_to_all [38;5;28;01mas[39;00m all_to_all, [0;32m (...)[0m [1;32m 355[0m xeinsum [38;5;28;01mas[39;00m xeinsum, [1;32m 356[0m ) [1;32m 357[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mother[39;00m [38;5;28;01mimport[39;00m ( [1;32m 358[0m conv_general_dilated_patches [38;5;28;01mas[39;00m conv_general_dilated_patches [1;32m 359[0m ) File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/_src/lax/fft.py:145[0m, in [0;36m<module>[0;34m[0m [1;32m 143[0m batching[38;5;241m.[39mprimitive_batchers[fft_p] [38;5;241m=[39m fft_batching_rule [1;32m 144[0m [38;5;28;01mif[39;00m pocketfft: [0;32m--> 145[0m xla[38;5;241m.[39mbackend_specific_translations[[38;5;124m'[39m[38;5;124mcpu[39m[38;5;124m'[39m][fft_p] [38;5;241m=[39m [43mpocketfft[49m[38;5;241;43m.[39;49m[43mpocketfft[49m [0;31mAttributeError[0m: module 'jaxlib.pocketfft' has no attribute 'pocketfft'
In [2]:
def gaussian_points(dim=2, n_points=1000, mean_vector=np.array([0, 0]),
sample_variance=1):
"""
Very simple custom function to generate gaussian distributed point clusters
with variable dimension, number of points, means in each direction
(must match dim) and sample variance.
Inputs:
dim (int)
n_points (int)
mean_vector (np.array) (where index 0 is x, index 1 is y etc.)
sample_variance (float)
Returns:
data (np.array): with dimensions (dim x n_points)
"""
mean_matrix = np.zeros(dim) + mean_vector
covariance_matrix = np.eye(dim) * sample_variance
data = np.random.multivariate_normal(mean_matrix, covariance_matrix,
n_points)
return data
def generate_simple_clustering_dataset(dim=2, n_points=1000, plotting=True,
return_data=True):
"""
Toy model to illustrate k-means clustering
"""
data1 = gaussian_points(mean_vector=np.array([5, 5]))
data2 = gaussian_points()
data3 = gaussian_points(mean_vector=np.array([1, 4.5]))
data4 = gaussian_points(mean_vector=np.array([5, 1]))
data = np.concatenate((data1, data2, data3, data4), axis=0)
if plotting:
fig, ax = plt.subplots()
ax.scatter(data[:, 0], data[:, 1], alpha=0.2)
ax.set_title('Toy Model Dataset')
plt.show()
if return_data:
return data
data = generate_simple_clustering_dataset()In [3]:
n_samples, dimensions = data.shape
n_clusters = 4
# we randomly initialize our centroids
np.random.seed(2021)
centroids = data[np.random.choice(n_samples, n_clusters, replace=False), :]
distances = np.zeros((n_samples, n_clusters))
# first we need to calculate the distance to each centroid from our data
for k in range(n_clusters):
for n in range(n_samples):
dist = 0
for d in range(dimensions):
dist += np.abs(data[n, d] - centroids[k, d])**2
distances[n, k] = dist
# we initialize an array to keep track of to which cluster each point belongs
# the way we set it up here the index tracks which point and the value which
# cluster the point belongs to
cluster_labels = np.zeros(n_samples, dtype='int')
# next we loop through our samples and for every point assign it to the cluster
# to which it has the smallest distance to
for n in range(n_samples):
# tracking variables (all of this is basically just an argmin)
smallest = 1e10
smallest_row_index = 1e10
for k in range(n_clusters):
if distances[n, k] < smallest:
smallest = distances[n, k]
smallest_row_index = k
cluster_labels[n] = smallest_row_indexIn [4]:
fig = plt.figure()
ax = fig.add_subplot()
unique_cluster_labels = np.unique(cluster_labels)
for i in unique_cluster_labels:
ax.scatter(data[cluster_labels == i, 0],
data[cluster_labels == i, 1],
label = i,
alpha = 0.2)
ax.scatter(centroids[:, 0], centroids[:, 1], c='black')
ax.set_title("First Grouping of Points to Centroids")
plt.show()In [5]:
max_iterations = 100
tolerance = 1e-8
for iteration in range(max_iterations):
prev_centroids = centroids.copy()
for k in range(n_clusters):
# this array will be used to update our centroid positions
vector_mean = np.zeros(dimensions)
mean_divisor = 0
for n in range(n_samples):
if cluster_labels[n] == k:
vector_mean += data[n, :]
mean_divisor += 1
# update according to the k means
centroids[k, :] = vector_mean / mean_divisor
# we find the dissimilarity
for k in range(n_clusters):
for n in range(n_samples):
dist = 0
for d in range(dimensions):
dist += np.abs(data[n, d] - centroids[k, d])**2
distances[n, k] = dist
# assign each point
for n in range(n_samples):
smallest = 1e10
smallest_row_index = 1e10
for k in range(n_clusters):
if distances[n, k] < smallest:
smallest = distances[n, k]
smallest_row_index = k
cluster_labels[n] = smallest_row_index
# convergence criteria
centroid_difference = np.sum(np.abs(centroids - prev_centroids))
if centroid_difference < tolerance:
print(f'Converged at iteration {iteration}')
break
elif iteration == max_iterations:
print(f'Did not converge in {max_iterations} iterations')In [6]:
fig = plt.figure()
ax = fig.add_subplot()
unique_cluster_labels = np.unique(cluster_labels)
for i in unique_cluster_labels:
ax.scatter(data[cluster_labels == i, 0],
data[cluster_labels == i, 1],
label = i,
alpha = 0.2)
ax.scatter(centroids[:, 0], centroids[:, 1], c='black')
ax.set_title("Final Result of K-means Clustering")
plt.show()In [7]:
def naive_kmeans(data, n_clusters=4, max_iterations=100, tolerance=1e-8):
start_time = time.time()
n_samples, dimensions = data.shape
n_clusters = 4
#np.random.seed(2021)
centroids = data[np.random.choice(n_samples, n_clusters, replace=False), :]
distances = np.zeros((n_samples, n_clusters))
for k in range(n_clusters):
for n in range(n_samples):
dist = 0
for d in range(dimensions):
dist += np.abs(data[n, d] - centroids[k, d])**2
distances[n, k] = dist
cluster_labels = np.zeros(n_samples, dtype='int')
for n in range(n_samples):
smallest = 1e10
smallest_row_index = 1e10
for k in range(n_clusters):
if distances[n, k] < smallest:
smallest = distances[n, k]
smallest_row_index = k
cluster_labels[n] = smallest_row_index
for iteration in range(max_iterations):
prev_centroids = centroids.copy()
for k in range(n_clusters):
vector_mean = np.zeros(dimensions)
mean_divisor = 0
for n in range(n_samples):
if cluster_labels[n] == k:
vector_mean += data[n, :]
mean_divisor += 1
centroids[k, :] = vector_mean / mean_divisor
for k in range(n_clusters):
for n in range(n_samples):
dist = 0
for d in range(dimensions):
dist += np.abs(data[n, d] - centroids[k, d])**2
distances[n, k] = dist
for n in range(n_samples):
smallest = 1e10
smallest_row_index = 1e10
for k in range(n_clusters):
if distances[n, k] < smallest:
smallest = distances[n, k]
smallest_row_index = k
cluster_labels[n] = smallest_row_index
centroid_difference = np.sum(np.abs(centroids - prev_centroids))
if centroid_difference < tolerance:
print(f'Converged at iteration {iteration}')
print(f'Runtime: {time.time() - start_time} seconds')
return cluster_labels, centroids
print(f'Did not converge in {max_iterations} iterations')
print(f'Runtime: {time.time() - start_time} seconds')
return cluster_labels, centroids