Skip to content

Commit c1fdce8

Browse files
committed
pylint fixes
1 parent df82415 commit c1fdce8

20 files changed

Lines changed: 4991 additions & 1893 deletions

docs/guides/llm_calculator.ipynb

Lines changed: 3108 additions & 1 deletion
Large diffs are not rendered by default.

src/MaxText/examples/demo_decoding.ipynb

Lines changed: 17 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -165,7 +165,7 @@
165165
"metadata": {},
166166
"outputs": [],
167167
"source": [
168-
"jax.distributed.initialize() # distributed.initialize should only be called once.\n",
168+
"jax.distributed.initialize() # distributed.initialize should only be called once.\n",
169169
"jax.devices()"
170170
]
171171
},
@@ -189,7 +189,7 @@
189189
"RUN_NAME = datetime.datetime.now().strftime(\"%Y-%m-%d-%H-%M-%S\")\n",
190190
"MODEL_CHECKPOINT_PATH = f\"/tmp/checkpoints/{MODEL_NAME}/{RUN_NAME}/unscanned\"\n",
191191
"\n",
192-
"HF_TOKEN = userdata.get('HF_TOKEN')\n",
192+
"HF_TOKEN = userdata.get(\"HF_TOKEN\")\n",
193193
"login(token=HF_TOKEN)\n",
194194
"max_logging.log(\"Authenticated with Hugging Face successfully!\")"
195195
]
@@ -211,13 +211,13 @@
211211
"source": [
212212
"%%capture\n",
213213
"argv = [\n",
214-
" \"\", # This is a placeholder, it's not actually used by the script's logic\n",
214+
" \"\", # This is a placeholder, it's not actually used by the script's logic\n",
215215
" f\"{MAXTEXT_PKG_DIR}/configs/base.yml\",\n",
216216
" f\"model_name={MODEL_NAME}\",\n",
217217
" f\"base_output_directory={MODEL_CHECKPOINT_PATH}\",\n",
218218
" f\"hf_access_token={HF_TOKEN}\",\n",
219219
" \"use_multimodal=false\",\n",
220-
" \"scan_layers=false\"\n",
220+
" \"scan_layers=false\",\n",
221221
"]\n",
222222
"\n",
223223
"to_maxtext.main(argv)"
@@ -250,7 +250,7 @@
250250
"source": [
251251
"%%capture\n",
252252
"config = pyconfig.initialize(\n",
253-
" [\"\", f\"{MAXTEXT_PKG_DIR}/configs/base.yml\"], \n",
253+
" [\"\", f\"{MAXTEXT_PKG_DIR}/configs/base.yml\"],\n",
254254
" per_device_batch_size=1.0,\n",
255255
" run_name=\"test\",\n",
256256
" max_target_length=4,\n",
@@ -339,7 +339,7 @@
339339
"\n",
340340
"# Pad input_ids to max_target_length\n",
341341
"padded_ids = np.zeros(config.max_target_length, dtype=np.int32)\n",
342-
"padded_ids[:len(input_ids)] = input_ids\n",
342+
"padded_ids[: len(input_ids)] = input_ids\n",
343343
"ids = np.asarray(padded_ids, dtype=np.int32)\n",
344344
"\n",
345345
"s = (config.global_batch_size_to_train_on, config.max_target_length)\n",
@@ -349,7 +349,9 @@
349349
")\n",
350350
"\n",
351351
"ids = np.stack([ids for _ in range(config.global_batch_size_to_train_on)])\n",
352-
"max_logging.log(f\"input_ids={input_ids}, \\n\\nids={ids}, \\n\\ndecoder_segment_ids = {decoder_segment_ids}, \\n\\ndecoder_positions= {decoder_positions}\")"
352+
"max_logging.log(\n",
353+
" f\"input_ids={input_ids}, \\n\\nids={ids}, \\n\\ndecoder_segment_ids = {decoder_segment_ids}, \\n\\ndecoder_positions= {decoder_positions}\"\n",
354+
")"
353355
]
354356
},
355357
{
@@ -370,12 +372,12 @@
370372
"outputs": [],
371373
"source": [
372374
"full_train_logits = model.apply(\n",
373-
" state.params,\n",
374-
" ids,\n",
375-
" decoder_positions,\n",
376-
" decoder_segment_ids,\n",
377-
" enable_dropout=False,\n",
378-
" rngs={\"aqt\": init_rng},\n",
375+
" state.params,\n",
376+
" ids,\n",
377+
" decoder_positions,\n",
378+
" decoder_segment_ids,\n",
379+
" enable_dropout=False,\n",
380+
" rngs={\"aqt\": init_rng},\n",
379381
")\n",
380382
"full_train_logits = jax.experimental.multihost_utils.process_allgather(full_train_logits)\n",
381383
"max_logging.log(f\"{full_train_logits[0, 0, :]=}\")"
@@ -397,17 +399,15 @@
397399
"outputs": [],
398400
"source": [
399401
"selected_logits = jax.lax.dynamic_slice(\n",
400-
" full_train_logits,\n",
401-
" (0, 0, full_train_logits.shape[2]-2, 0),\n",
402-
" (1, 1, 1, full_train_logits.shape[3])\n",
402+
" full_train_logits, (0, 0, full_train_logits.shape[2] - 2, 0), (1, 1, 1, full_train_logits.shape[3])\n",
403403
")\n",
404404
"\n",
405405
"# Consider the greedily sampled token\n",
406406
"init_rng, new_rng = jax.random.split(init_rng)\n",
407407
"first_generated_token = inference_utils.sampling(\n",
408408
" selected_logits,\n",
409409
" new_rng,\n",
410-
" config.decode_sampling_strategy, #\"greedy\"\n",
410+
" config.decode_sampling_strategy, # \"greedy\"\n",
411411
")\n",
412412
"output = tokenizer.decode([first_generated_token.item()])\n",
413413
"max_logging.log(f\"Next predicted token is `{output}` for the input prompt: `{config.prompt}`.\")"

0 commit comments

Comments
 (0)