{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":8250323,"sourceType":"datasetVersion","datasetId":4895216},{"sourceId":27858,"sourceType":"modelInstanceVersion","modelInstanceId":22049},{"sourceId":40740,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":34287}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Published on May 21, 2024. By Marília Prata, mpwolke","metadata":{}},{"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\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 read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-05-21T20:03:33.465797Z","iopub.execute_input":"2024-05-21T20:03:33.466178Z","iopub.status.idle":"2024-05-21T20:04:52.557610Z","shell.execute_reply.started":"2024-05-21T20:03:33.466147Z","shell.execute_reply":"2024-05-21T20:04:52.556537Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#GPU must be ON!","metadata":{}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:05:49.837253Z","iopub.execute_input":"2024-05-21T20:05:49.837659Z","iopub.status.idle":"2024-05-21T20:05:50.988999Z","shell.execute_reply.started":"2024-05-21T20:05:49.837628Z","shell.execute_reply":"2024-05-21T20:05:50.987762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#Install","metadata":{}},{"cell_type":"code","source":"!pip3 install -U -q git+https://github.com/james77777778/keras-nlp.git@int8-gemma\n!pip3 install -U -q keras # git+https://github.com/keras-team/keras.git","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-05-21T20:05:57.431163Z","iopub.execute_input":"2024-05-21T20:05:57.431992Z","iopub.status.idle":"2024-05-21T20:06:51.846866Z","shell.execute_reply.started":"2024-05-21T20:05:57.431950Z","shell.execute_reply":"2024-05-21T20:06:51.845624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#Import","metadata":{}},{"cell_type":"code","source":"import os\nos.environ[\"KERAS_BACKEND\"] = \"jax\"\n\nimport json\nimport keras\nimport keras_nlp\n\nkeras.config.set_dtype_policy(\"bfloat16\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-05-21T20:06:58.592715Z","iopub.execute_input":"2024-05-21T20:06:58.593163Z","iopub.status.idle":"2024-05-21T20:07:14.440506Z","shell.execute_reply.started":"2024-05-21T20:06:58.593128Z","shell.execute_reply":"2024-05-21T20:07:14.439517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import jax\n\njax.devices()","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:07:20.020192Z","iopub.execute_input":"2024-05-21T20:07:20.021210Z","iopub.status.idle":"2024-05-21T20:07:20.693261Z","shell.execute_reply.started":"2024-05-21T20:07:20.021174Z","shell.execute_reply":"2024-05-21T20:07:20.692272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_DIR = \"/kaggle/input/gemma/keras/gemma_1.1_instruct_7b_en_int8/2\"\nPRESET = \"gemma_1.1_instruct_7b_en\"","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:07:25.202756Z","iopub.execute_input":"2024-05-21T20:07:25.203636Z","iopub.status.idle":"2024-05-21T20:07:25.208315Z","shell.execute_reply.started":"2024-05-21T20:07:25.203591Z","shell.execute_reply":"2024-05-21T20:07:25.207170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#Distributed Model","metadata":{}},{"cell_type":"code","source":"# Create a device mesh with (1, 2) shape so that the weights are sharded across\n# all 2 GPUs\ndevice_mesh = keras.distribution.DeviceMesh(\n    (1, 2),\n    [\"batch\", \"model\"],\n    devices=keras.distribution.list_devices()\n)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:07:30.402969Z","iopub.execute_input":"2024-05-21T20:07:30.403659Z","iopub.status.idle":"2024-05-21T20:07:30.410483Z","shell.execute_reply.started":"2024-05-21T20:07:30.403628Z","shell.execute_reply":"2024-05-21T20:07:30.409284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_dim = \"model\"\n\nlayout_map = keras.distribution.LayoutMap(device_mesh)\n\n# Weights that match 'token_embedding/embeddings' will be sharded on 8 TPUs\nlayout_map[\"token_embedding/embeddings$\"] = (None, model_dim)\n# Regex to match against the query, key and value matrices in the decoder\n# attention layers\nlayout_map[\"decoder_block_.*/attention/(query|key|value)/kernel$\"] = (\n    None,\n    model_dim,\n    None,\n)\nlayout_map[\"decoder_block_.*/attention_output/kernel$\"] = (\n    None,\n    None,\n    model_dim,\n)\nlayout_map[\"decoder_block_.*/ffw_gating/kernel$\"] = (model_dim, None)\nlayout_map[\"decoder_block_.*/ffw_linear/kernel$\"] = (None, model_dim)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:07:36.364176Z","iopub.execute_input":"2024-05-21T20:07:36.365038Z","iopub.status.idle":"2024-05-21T20:07:36.371189Z","shell.execute_reply.started":"2024-05-21T20:07:36.365008Z","shell.execute_reply":"2024-05-21T20:07:36.370160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_parallel = keras.distribution.ModelParallel(\n    device_mesh, layout_map, batch_dim_name=\"batch\")\n\nkeras.distribution.set_distribution(model_parallel)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:07:43.130060Z","iopub.execute_input":"2024-05-21T20:07:43.130768Z","iopub.status.idle":"2024-05-21T20:07:43.136053Z","shell.execute_reply.started":"2024-05-21T20:07:43.130736Z","shell.execute_reply":"2024-05-21T20:07:43.134967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df= pd.read_json('../input/gemma-1-1-instruct-7b-en-int8/gemma_1.1_instruct_7b_en_int8.json', lines=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:07:54.513665Z","iopub.execute_input":"2024-05-21T20:07:54.514158Z","iopub.status.idle":"2024-05-21T20:07:54.550699Z","shell.execute_reply.started":"2024-05-21T20:07:54.514118Z","shell.execute_reply":"2024-05-21T20:07:54.549501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#GemmaCausalLM","metadata":{}},{"cell_type":"code","source":"# Load the int8 quantized model\nwith open(f\"{MODEL_DIR}/{PRESET}_int8.json\", \"r\") as f:\n    config = json.loads(f.read())\nmodel: keras_nlp.models.GemmaCausalLM = (\n    keras.saving.deserialize_keras_object(config)\n)\nmodel.load_weights(f\"{MODEL_DIR}/{PRESET}_int8.weights.h5\")\nmodel.preprocessor.tokenizer.load_preset_assets(PRESET)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:08:00.484028Z","iopub.execute_input":"2024-05-21T20:08:00.484670Z","iopub.status.idle":"2024-05-21T20:10:03.264735Z","shell.execute_reply.started":"2024-05-21T20:08:00.484642Z","shell.execute_reply":"2024-05-21T20:10:03.263859Z"},"_kg_hide-output":false,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(model.generate(\"What are the severity levels of Spinal Canal Stenosis?\", max_length=1024))","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:10:41.724278Z","iopub.execute_input":"2024-05-21T20:10:41.724668Z","iopub.status.idle":"2024-05-21T20:12:04.891714Z","shell.execute_reply.started":"2024-05-21T20:10:41.724638Z","shell.execute_reply":"2024-05-21T20:12:04.890515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(model.generate(\"Which is the percent of Moderate level in right_subarticular_stenosis_l1_l2?\", max_length=1024))","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:12:11.649411Z","iopub.execute_input":"2024-05-21T20:12:11.650356Z","iopub.status.idle":"2024-05-21T20:12:14.676743Z","shell.execute_reply.started":"2024-05-21T20:12:11.650323Z","shell.execute_reply":"2024-05-21T20:12:14.675699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(model.generate(\"What are the core conditions of Lumbar Spine degeneration?\", max_length=1024))","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:12:21.427686Z","iopub.execute_input":"2024-05-21T20:12:21.428619Z","iopub.status.idle":"2024-05-21T20:12:29.330537Z","shell.execute_reply.started":"2024-05-21T20:12:21.428588Z","shell.execute_reply":"2024-05-21T20:12:29.329452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#Check","metadata":{}},{"cell_type":"code","source":"decoder_block_1 = model.backbone.get_layer('decoder_block_1')\nprint(type(decoder_block_1))\nfor variable in decoder_block_1.weights:\n    print(f'{variable.path:<58}  {str(variable.shape):<16}  {str(variable.value.sharding.spec)}')","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:12:39.828567Z","iopub.execute_input":"2024-05-21T20:12:39.829336Z","iopub.status.idle":"2024-05-21T20:12:39.835528Z","shell.execute_reply.started":"2024-05-21T20:12:39.829306Z","shell.execute_reply":"2024-05-21T20:12:39.834518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv', nrows=20)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:12:53.365713Z","iopub.execute_input":"2024-05-21T20:12:53.366453Z","iopub.status.idle":"2024-05-21T20:12:53.387033Z","shell.execute_reply.started":"2024-05-21T20:12:53.366416Z","shell.execute_reply":"2024-05-21T20:12:53.386025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re\n\n#Juan Merino https://www.kaggle.com/code/juanmerinobermejo/aimo-starter-notebook-gemma-7b/notebook\n\nresponses = []\n\nfor i in train['condition']:\n    prompt = (f\"Which are the core conditions that affects lumbar degeneration? PROBLEM: {i}\")\n    response = model.generate(prompt,max_length=850) #Original was gemma_lm.generate\n    print(response)\n    responses.append(response)\n\ntrain['gemma_1.1_instruct_7b_answer'] = responses\n\ndef extract_integer(text):\n    match = re.search(r'The answer is: (\\d+)', text)\n    if match:\n        return int(match.group(1))\n    else:\n        return None\n\ntrain['gemma_1.1_instruct_7b_answer_integer'] = train['gemma_1.1_instruct_7b_answer'].apply(extract_integer)\ntrain['gemma_1.1_instruct_7b_answer'] = train['gemma_1.1_instruct_7b_answer_integer']\ntrain = train.drop('gemma_1.1_instruct_7b_answer_integer', axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T20:13:00.787743Z","iopub.execute_input":"2024-05-21T20:13:00.788254Z","iopub.status.idle":"2024-05-21T20:18:34.592919Z","shell.execute_reply.started":"2024-05-21T20:13:00.788218Z","shell.execute_reply":"2024-05-21T20:18:34.591793Z"},"_kg_hide-output":true,"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#Embracing answers Gemma 1.1_instruct_7b. Unhide to read them all.\n\nAll that in less than 17 minutes. Well-done","metadata":{}},{"cell_type":"markdown","source":"#Acknowledgements:\n\nDatasets/Models:\n\nAwsaf https://www.kaggle.com/datasets/awsaf49/gemma-1-1-instruct-7b-en-int8\n\nAwsaf https://www.kaggle.com/models/awsaf49/gemma/Keras/gemma_1.1_instruct_7b_en_int8/2\n\nAwsaf https://www.kaggle.com/code/awsaf49/gemma-1-1-7b-int8-load\n\nJuan Merino https://www.kaggle.com/code/juanmerinobermejo/aimo-starter-notebook-gemma-7b/notebook\n\nmpwolke https://www.kaggle.com/code/mpwolke/why-is-420-so-funny-gemma-1-1-instruct-7b-en","metadata":{}}]}