{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install \"tensorflow>=2\"\n!pip install \"tensorflow_hub>=0.7\"\n!pip install bert-for-tf2\n!pip install sentencepiece","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import tensorflow as tf\nimport tensorflow_hub as hub\nprint(\"TF version: \", tf.__version__)\nprint(\"Hub version: \", hub.__version__)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import tensorflow_hub as hub\nimport tensorflow as tf\nimport bert\nFullTokenizer = bert.bert_tokenization.FullTokenizer\nfrom tensorflow.keras.models import Model       # Keras is the new high level API for TensorFlow\nimport math\nimport numpy as np","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"max_seq_length = 128  # Your choice here.\ninput_word_ids = tf.keras.layers.Input(shape=(max_seq_length,), dtype=tf.int32,\n                                       name=\"input_word_ids\")\ninput_mask = tf.keras.layers.Input(shape=(max_seq_length,), dtype=tf.int32,\n                                   name=\"input_mask\")\nsegment_ids = tf.keras.layers.Input(shape=(max_seq_length,), dtype=tf.int32,\n                                    name=\"segment_ids\")\nbert_layer = hub.KerasLayer(\"https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/1\",\n                            trainable=True)\npooled_output, sequence_output = bert_layer([input_word_ids, input_mask, segment_ids])\n\nmodel = Model(inputs=[input_word_ids, input_mask, segment_ids], outputs=[pooled_output, sequence_output])\n\nvocab_file = bert_layer.resolved_object.vocab_file.asset_path.numpy()\ndo_lower_case = bert_layer.resolved_object.do_lower_case.numpy()\ntokenizer = FullTokenizer(vocab_file, do_lower_case)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# See BERT paper: https://arxiv.org/pdf/1810.04805.pdf\n# And BERT implementation convert_single_example() at https://github.com/google-research/bert/blob/master/run_classifier.py\n\ndef get_masks(tokens, max_seq_length):\n    \"\"\"Mask for padding\"\"\"\n    if len(tokens)>max_seq_length:\n        raise IndexError(\"Token length more than max seq length!\")\n    return [1]*len(tokens) + [0] * (max_seq_length - len(tokens))\n\n\ndef get_segments(tokens, max_seq_length):\n    \"\"\"Segments: 0 for the first sequence, 1 for the second\"\"\"\n    if len(tokens)>max_seq_length:\n        raise IndexError(\"Token length more than max seq length!\")\n    segments = []\n    current_segment_id = 0\n    for token in tokens:\n        segments.append(current_segment_id)\n        if token == \"[SEP]\":\n            current_segment_id = 1\n    return segments + [0] * (max_seq_length - len(tokens))\n\n\ndef get_ids(tokens, tokenizer, max_seq_length):\n    \"\"\"Token ids from Tokenizer vocab\"\"\"\n    token_ids = tokenizer.convert_tokens_to_ids(tokens)\n    input_ids = token_ids + [0] * (max_seq_length-len(token_ids))\n    return input_ids","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def tokenize_sentence(sentence):\n    stokens = tokenizer.tokenize(sentence)\n    stokens = [\"[CLS]\"] + stokens + [\"[SEP]\"]\n    \n    input_ids = get_ids(stokens, tokenizer, max_seq_length)\n    input_masks = get_masks(stokens, max_seq_length)\n    input_segments = get_segments(stokens, max_seq_length)\n    \n    return input_ids, input_masks, input_segments\n\ndef compare_sentences(sentence_1, sentence_2, distance_metric):\n    input_ids_1, input_masks_1, input_segments_1 = tokenize_sentence(sentence_1)\n    input_ids_2, input_masks_2, input_segments_2 = tokenize_sentence(sentence_2)\n    \n    pool_embs_1, all_embs_1 = model.predict([[input_ids_1],[input_masks_1],[input_segments_1]])\n    pool_embs_2, all_embs_2 = model.predict([[input_ids_2],[input_masks_2],[input_segments_2]])\n    \n    return distance_metric(pool_embs_1[0], pool_embs_2[0])\n    \ndef square_rooted(x):\n    return math.sqrt(sum([a*a for a in x]))\n\ndef cosine_similarity(x,y):\n    numerator = sum(a*b for a,b in zip(x,y))\n    denominator = square_rooted(x)*square_rooted(y)\n    return numerator/float(denominator)\n\ndef dummy_metric(x,y):\n    return 42","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"s1 = 'How are you doing?'\n# s2 = '''Right after Ricky Gervais talks about how the Hollywood Foreign Press is racist and doesn't include people of color the cameraman zooms out to show just how few people of color were invited to this event.'''\ns2 = 'How are we feeling?'\ncompare_sentences(s1, s2, cosine_similarity)","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":1}