circuitkit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- circuitkit/__init__.py +128 -0
- circuitkit/__main__.py +9 -0
- circuitkit/analysis/__init__.py +19 -0
- circuitkit/analysis/cross_method_jaccard.py +116 -0
- circuitkit/analysis/metrics.py +54 -0
- circuitkit/analysis/scores.py +44 -0
- circuitkit/api.py +2682 -0
- circuitkit/applications/__init__.py +70 -0
- circuitkit/applications/arch_registry.py +315 -0
- circuitkit/applications/arch_utils.py +302 -0
- circuitkit/applications/common_utils/__init__.py +15 -0
- circuitkit/applications/common_utils/_covariance.py +223 -0
- circuitkit/applications/common_utils/_device.py +33 -0
- circuitkit/applications/common_utils/_metrics.py +429 -0
- circuitkit/applications/common_utils/_tokenization.py +510 -0
- circuitkit/applications/common_utils/benchmark_analysis.py +400 -0
- circuitkit/applications/common_utils/cure_clue.py +338 -0
- circuitkit/applications/common_utils/hallucination_detection.py +497 -0
- circuitkit/applications/common_utils/linear_probe.py +294 -0
- circuitkit/applications/editing/__init__.py +30 -0
- circuitkit/applications/editing/cake.py +237 -0
- circuitkit/applications/editing/circuit_guided_editing.py +487 -0
- circuitkit/applications/editing/fine_tune_editing.py +327 -0
- circuitkit/applications/editing/knowledge_editing.py +598 -0
- circuitkit/applications/editing/knowledge_editing_enhanced.py +863 -0
- circuitkit/applications/editing/mcircke.py +263 -0
- circuitkit/applications/editing/memit_wrapper.py +947 -0
- circuitkit/applications/editing/rome_wrapper.py +770 -0
- circuitkit/applications/finetuning/__init__.py +17 -0
- circuitkit/applications/finetuning/benchmark_peft.py +451 -0
- circuitkit/applications/finetuning/circuit_tuning.py +377 -0
- circuitkit/applications/finetuning/healing_metrics.py +304 -0
- circuitkit/applications/finetuning/peft_methods.py +563 -0
- circuitkit/applications/finetuning/soft_healing.py +717 -0
- circuitkit/applications/pruning/__init__.py +13 -0
- circuitkit/applications/pruning/eval_utils.py +347 -0
- circuitkit/applications/pruning/examples/__init__.py +0 -0
- circuitkit/applications/pruning/examples/prune.py +787 -0
- circuitkit/applications/pruning/examples/prune_llama.py +661 -0
- circuitkit/applications/pruning/examples/prune_qwen.py +591 -0
- circuitkit/applications/pruning/finetune_utils.py +477 -0
- circuitkit/applications/pruning/importance.py +97 -0
- circuitkit/applications/pruning/neuron_pruner.py +36 -0
- circuitkit/applications/pruning/node_pruner.py +186 -0
- circuitkit/applications/pruning/pruner.py +529 -0
- circuitkit/applications/pruning/score_extractor.py +541 -0
- circuitkit/applications/pruning/selectors/__init__.py +0 -0
- circuitkit/applications/pruning/selectors/multi_granular_selector.py +119 -0
- circuitkit/applications/pruning/selectors/taylor_selector.py +112 -0
- circuitkit/applications/pruning/weight_pruner.py +652 -0
- circuitkit/applications/quantization/__init__.py +19 -0
- circuitkit/applications/quantization/examples/__init__.py +0 -0
- circuitkit/applications/quantization/examples/quantize_llama.py +745 -0
- circuitkit/applications/quantization/examples/quantize_qwen.py +726 -0
- circuitkit/applications/quantization/llmcompressor_quantize.py +388 -0
- circuitkit/applications/quantization/quant_utils.py +825 -0
- circuitkit/applications/quantization/score_extractor.py +465 -0
- circuitkit/applications/quantization/selectors/__init__.py +0 -0
- circuitkit/applications/quantization/selectors/awq_selector.py +109 -0
- circuitkit/applications/quantization/selectors/tacq_selector.py +145 -0
- circuitkit/applications/selective_finetuning/__init__.py +0 -0
- circuitkit/applications/selective_finetuning/examples/__init__.py +0 -0
- circuitkit/applications/selective_finetuning/examples/finetune_llama.py +668 -0
- circuitkit/applications/selective_finetuning/examples/finetune_qwen.py +654 -0
- circuitkit/applications/selective_finetuning/finetune_utils.py +643 -0
- circuitkit/applications/selective_finetuning/score_loader.py +585 -0
- circuitkit/applications/selective_finetuning/selector.py +616 -0
- circuitkit/applications/steering/__init__.py +31 -0
- circuitkit/applications/steering/steering.py +791 -0
- circuitkit/applications/steering/steering_enhanced.py +556 -0
- circuitkit/applications/steering/weight_steering.py +407 -0
- circuitkit/artifacts/__init__.py +25 -0
- circuitkit/artifacts/circuit_artifact.py +559 -0
- circuitkit/artifacts/converters.py +405 -0
- circuitkit/artifacts/scores.py +195 -0
- circuitkit/backends/__init__.py +96 -0
- circuitkit/backends/acdc/__init__.py +0 -0
- circuitkit/backends/acdc/artifact_export.py +140 -0
- circuitkit/backends/acdc/data.py +183 -0
- circuitkit/backends/acdc/model_utils/__init__.py +0 -0
- circuitkit/backends/acdc/model_utils/micro_model_utils.py +143 -0
- circuitkit/backends/acdc/model_utils/transformer_lens_utils.py +232 -0
- circuitkit/backends/acdc/prune.py +123 -0
- circuitkit/backends/acdc/prune_algos/ACDC.py +136 -0
- circuitkit/backends/acdc/prune_algos/__init__.py +0 -0
- circuitkit/backends/acdc/prune_algos/mask_gradient.py +129 -0
- circuitkit/backends/acdc/prune_algos/prune_algos.py +33 -0
- circuitkit/backends/acdc/tasks/__init__.py +3 -0
- circuitkit/backends/acdc/tasks/docstring_prompts.py +837 -0
- circuitkit/backends/acdc/tasks/docstring_utils.py +88 -0
- circuitkit/backends/acdc/tasks/induction_utils.py +122 -0
- circuitkit/backends/acdc/tasks/ioi_dataset.py +156 -0
- circuitkit/backends/acdc/tasks/ioi_utils.py +96 -0
- circuitkit/backends/acdc/types.py +248 -0
- circuitkit/backends/acdc/utils/__init__.py +0 -0
- circuitkit/backends/acdc/utils/ablation_activations.py +160 -0
- circuitkit/backends/acdc/utils/custom_tqdm.py +14 -0
- circuitkit/backends/acdc/utils/graph_utils.py +497 -0
- circuitkit/backends/acdc/utils/misc.py +68 -0
- circuitkit/backends/acdc/utils/patch_wrapper.py +106 -0
- circuitkit/backends/acdc/utils/patchable_model.py +143 -0
- circuitkit/backends/acdc/utils/task_utils.py +31 -0
- circuitkit/backends/acdc/utils/tensor_ops.py +118 -0
- circuitkit/backends/acdc/visualize.py +253 -0
- circuitkit/backends/cdt/__init__.py +30 -0
- circuitkit/backends/cdt/adapter.py +245 -0
- circuitkit/backends/cdt/propagation.py +357 -0
- circuitkit/backends/cdt/pyfunctions/__init__.py +14 -0
- circuitkit/backends/cdt/pyfunctions/cdt_ablations.py +215 -0
- circuitkit/backends/cdt/pyfunctions/cdt_basic.py +235 -0
- circuitkit/backends/cdt/pyfunctions/cdt_core.py +418 -0
- circuitkit/backends/cdt/pyfunctions/cdt_from_source_nodes.py +237 -0
- circuitkit/backends/cdt/pyfunctions/cdt_source_to_target.py +685 -0
- circuitkit/backends/cdt/pyfunctions/faithfulness_ablations.py +251 -0
- circuitkit/backends/cdt/pyfunctions/general.py +314 -0
- circuitkit/backends/cdt/pyfunctions/ioi_dataset.py +958 -0
- circuitkit/backends/cdt/pyfunctions/local_importance.py +809 -0
- circuitkit/backends/cdt/pyfunctions/pathology.py +460 -0
- circuitkit/backends/cdt/pyfunctions/toy_model.py +190 -0
- circuitkit/backends/cdt/pyfunctions/wrappers.py +159 -0
- circuitkit/backends/eap/__init__.py +2 -0
- circuitkit/backends/eap/artifact_export.py +137 -0
- circuitkit/backends/eap/attribute.py +784 -0
- circuitkit/backends/eap/attribute_node.py +1795 -0
- circuitkit/backends/eap/circuit_kit_adapter.py +121 -0
- circuitkit/backends/eap/eap_utils.py +582 -0
- circuitkit/backends/eap/evaluate.py +762 -0
- circuitkit/backends/eap/graph.py +1569 -0
- circuitkit/backends/eap/metrics.py +793 -0
- circuitkit/backends/eap/py.typed +0 -0
- circuitkit/backends/eap/visualization.py +101 -0
- circuitkit/backends/ibcircuit/__init__.py +0 -0
- circuitkit/backends/ibcircuit/artifact_export.py +127 -0
- circuitkit/backends/ibcircuit/ib_noise.py +216 -0
- circuitkit/backends/ibcircuit/ib_utils.py +194 -0
- circuitkit/backends/ibcircuit/model_wrapper.py +537 -0
- circuitkit/backends/ibcircuit/trainer.py +603 -0
- circuitkit/benchmarks/__init__.py +47 -0
- circuitkit/benchmarks/baselines/__init__.py +20 -0
- circuitkit/benchmarks/baselines/gptq.py +200 -0
- circuitkit/benchmarks/baselines/magnitude.py +204 -0
- circuitkit/benchmarks/baselines/random.py +142 -0
- circuitkit/benchmarks/baselines/sparsegpt.py +236 -0
- circuitkit/benchmarks/baselines/wanda.py +263 -0
- circuitkit/benchmarks/benchmark.py +764 -0
- circuitkit/benchmarks/reporting.py +639 -0
- circuitkit/circuit.py +390 -0
- circuitkit/cli/__init__.py +1 -0
- circuitkit/cli/config.py +74 -0
- circuitkit/cli/debug.py +279 -0
- circuitkit/cli/main.py +2208 -0
- circuitkit/cli/utils.py +351 -0
- circuitkit/corruption/__init__.py +61 -0
- circuitkit/corruption/base.py +132 -0
- circuitkit/corruption/color_swap.py +170 -0
- circuitkit/corruption/distractor.py +319 -0
- circuitkit/corruption/distractor_variation.py +313 -0
- circuitkit/corruption/effectiveness.py +297 -0
- circuitkit/corruption/entity_swap.py +288 -0
- circuitkit/corruption/negation.py +364 -0
- circuitkit/corruption/paraphrase.py +390 -0
- circuitkit/corruption/pipeline.py +333 -0
- circuitkit/corruption/position_shift.py +106 -0
- circuitkit/corruption/role_swap.py +381 -0
- circuitkit/corruption/token_swap.py +257 -0
- circuitkit/corruption/validators.py +570 -0
- circuitkit/corruption/voice_swap.py +393 -0
- circuitkit/data/__init__.py +8 -0
- circuitkit/data/adapters/__init__.py +20 -0
- circuitkit/data/adapters/base.py +123 -0
- circuitkit/data/adapters/code.py +106 -0
- circuitkit/data/adapters/conversational.py +167 -0
- circuitkit/data/adapters/forget_retain.py +152 -0
- circuitkit/data/adapters/instruction.py +125 -0
- circuitkit/data/adapters/math.py +133 -0
- circuitkit/data/adapters/mcq.py +194 -0
- circuitkit/data/adapters/pairwise.py +182 -0
- circuitkit/data/adapters/safety_prompt.py +258 -0
- circuitkit/data/auto_detect.py +187 -0
- circuitkit/data/clean_only.py +124 -0
- circuitkit/data/corruption/__init__.py +38 -0
- circuitkit/data/corruption/base.py +239 -0
- circuitkit/data/corruption/benign_rewrite.py +127 -0
- circuitkit/data/corruption/code_syntax_corrupt.py +98 -0
- circuitkit/data/corruption/entity_swap.py +113 -0
- circuitkit/data/corruption/final_answer_swap.py +231 -0
- circuitkit/data/corruption/instruction_swap.py +149 -0
- circuitkit/data/corruption/llm_counterfactual.py +161 -0
- circuitkit/data/corruption/logical_negation.py +96 -0
- circuitkit/data/corruption/math_step_corrupt.py +88 -0
- circuitkit/data/corruption/mcq_choice_swap.py +106 -0
- circuitkit/data/corruption/operand_swap.py +103 -0
- circuitkit/data/corruption/profession_swap.py +125 -0
- circuitkit/data/corruption/resample.py +75 -0
- circuitkit/data/corruption/template.py +195 -0
- circuitkit/data/corruption/template_utils.py +328 -0
- circuitkit/data/corruption/token_swap.py +90 -0
- circuitkit/data/dataset_schema.py +169 -0
- circuitkit/data/eap_dataset.py +98 -0
- circuitkit/data/invariance_groups/__init__.py +33 -0
- circuitkit/data/invariance_groups/builder.py +323 -0
- circuitkit/data/invariance_groups/schema.py +274 -0
- circuitkit/data/normalized.py +259 -0
- circuitkit/data/normalized_task.py +594 -0
- circuitkit/data/task_data/__init__.py +11 -0
- circuitkit/data/task_data/core/TLACDCCorrespondence.py +263 -0
- circuitkit/data/task_data/core/TLACDCEdge.py +113 -0
- circuitkit/data/task_data/core/TLACDCExperiment.py +1052 -0
- circuitkit/data/task_data/core/TLACDCInterpNode.py +96 -0
- circuitkit/data/task_data/core/__init__.py +12 -0
- circuitkit/data/task_data/core/acdc_utils.py +614 -0
- circuitkit/data/task_data/generation/__init__.py +11 -0
- circuitkit/data/task_data/generation/cache.py +267 -0
- circuitkit/data/task_data/generation/manager.py +562 -0
- circuitkit/data/task_data/generation/utils.py +323 -0
- circuitkit/data/task_data/storage/__init__.py +18 -0
- circuitkit/data/task_data/storage/greaterthan/greaterthan_32_ffd33106.json +23 -0
- circuitkit/data/task_data/storage/ioi/ioi_16_8c879ddb.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_32_a432ca4a.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_500_1f7e7324.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_64_3bba747e.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_64_f4a164db.json +43 -0
- circuitkit/data/task_data/storage/ioi/ioi_8_e87df42e.json +43 -0
- circuitkit/data/task_data/tasks/__init__.py +10 -0
- circuitkit/data/task_data/tasks/binary_align/generate_binary_align.py +1167 -0
- circuitkit/data/task_data/tasks/binary_align/jailbreak_binary.csv +335 -0
- circuitkit/data/task_data/tasks/binary_align/safe_binary.csv +335 -0
- circuitkit/data/task_data/tasks/capital_country/__init__.py +5 -0
- circuitkit/data/task_data/tasks/capital_country/utils.py +395 -0
- circuitkit/data/task_data/tasks/docstring/__init__.py +5 -0
- circuitkit/data/task_data/tasks/docstring/prompts.py +1175 -0
- circuitkit/data/task_data/tasks/docstring/utils.py +282 -0
- circuitkit/data/task_data/tasks/double_io/__init__.py +0 -0
- circuitkit/data/task_data/tasks/double_io/double_io_dataset.py +485 -0
- circuitkit/data/task_data/tasks/gender_bias/__init__.py +5 -0
- circuitkit/data/task_data/tasks/gender_bias/utils.py +396 -0
- circuitkit/data/task_data/tasks/gender_bias/utils2.py +143 -0
- circuitkit/data/task_data/tasks/greaterthan/__init__.py +5 -0
- circuitkit/data/task_data/tasks/greaterthan/utils.py +534 -0
- circuitkit/data/task_data/tasks/hypernymy/__init__.py +5 -0
- circuitkit/data/task_data/tasks/hypernymy/utils.py +326 -0
- circuitkit/data/task_data/tasks/induction/__init__.py +5 -0
- circuitkit/data/task_data/tasks/induction/utils.py +222 -0
- circuitkit/data/task_data/tasks/ioi/__init__.py +8 -0
- circuitkit/data/task_data/tasks/ioi/ioi_dataset.py +962 -0
- circuitkit/data/task_data/tasks/ioi/utils.py +656 -0
- circuitkit/data/task_data/tasks/sva/__init__.py +5 -0
- circuitkit/data/task_data/tasks/sva/utils.py +132 -0
- circuitkit/data/task_data/tasks/wmdp/wmdp_utils.py +296 -0
- circuitkit/data/template.py +392 -0
- circuitkit/data/wikitext_calibration.py +164 -0
- circuitkit/data/worthiness.py +746 -0
- circuitkit/evaluation/__init__.py +83 -0
- circuitkit/evaluation/checkpoint_benchmark.py +857 -0
- circuitkit/evaluation/evaluate.py +929 -0
- circuitkit/evaluation/full.py +556 -0
- circuitkit/evaluation/hf_checkpoint.py +1219 -0
- circuitkit/evaluation/intervention_faithfulness.py +198 -0
- circuitkit/evaluation/lm_eval_simple.py +223 -0
- circuitkit/evaluation/lm_harness.py +681 -0
- circuitkit/evaluation/master_grid.py +307 -0
- circuitkit/evaluation/mmlu_eval.py +208 -0
- circuitkit/evaluation/pillars/__init__.py +28 -0
- circuitkit/evaluation/pillars/ablation.py +404 -0
- circuitkit/evaluation/pillars/baselines.py +902 -0
- circuitkit/evaluation/pillars/causal_patching.py +371 -0
- circuitkit/evaluation/pillars/generalization.py +623 -0
- circuitkit/evaluation/pillars/intervention_reliability.py +325 -0
- circuitkit/evaluation/pillars/robustness.py +854 -0
- circuitkit/evaluation/pillars/stability.py +571 -0
- circuitkit/evaluation/report.py +318 -0
- circuitkit/evaluation/reports/__init__.py +20 -0
- circuitkit/evaluation/reports/aggregator.py +533 -0
- circuitkit/evaluation/reports/robustness_report.py +331 -0
- circuitkit/evaluation/reports/stability_report.py +298 -0
- circuitkit/evaluation/stability_discovery.py +439 -0
- circuitkit/evaluation/transfer.py +510 -0
- circuitkit/evaluation/transfer_analysis.py +315 -0
- circuitkit/evaluation/transfer_visualizer.py +327 -0
- circuitkit/evaluation/weight_based_eval.py +227 -0
- circuitkit/pipeline.py +1000 -0
- circuitkit/quick.py +1157 -0
- circuitkit/selection/__init__.py +54 -0
- circuitkit/selection/cdt_selector.py +64 -0
- circuitkit/selection/eap_gp_selector.py +67 -0
- circuitkit/selection/eap_selector.py +83 -0
- circuitkit/selection/gptq_selector.py +164 -0
- circuitkit/selection/ibcircuit_selector.py +120 -0
- circuitkit/selection/magnitude_selector.py +28 -0
- circuitkit/selection/random_selector.py +16 -0
- circuitkit/selection/relp_selector.py +66 -0
- circuitkit/selection/wanda_selector.py +174 -0
- circuitkit/tasks/__init__.py +40 -0
- circuitkit/tasks/_algorithm_families.py +106 -0
- circuitkit/tasks/_chat.py +262 -0
- circuitkit/tasks/auto_schema.py +526 -0
- circuitkit/tasks/bootstrap.py +93 -0
- circuitkit/tasks/builtins/__init__.py +44 -0
- circuitkit/tasks/builtins/boolq.py +501 -0
- circuitkit/tasks/builtins/capital_country.py +261 -0
- circuitkit/tasks/builtins/double_io.py +380 -0
- circuitkit/tasks/builtins/gender_bias.py +284 -0
- circuitkit/tasks/builtins/glue.py +683 -0
- circuitkit/tasks/builtins/greater_than.py +471 -0
- circuitkit/tasks/builtins/gsm8k.py +563 -0
- circuitkit/tasks/builtins/hypernymy.py +262 -0
- circuitkit/tasks/builtins/ifeval.py +116 -0
- circuitkit/tasks/builtins/ioi.py +447 -0
- circuitkit/tasks/builtins/ioi_acdc.py +323 -0
- circuitkit/tasks/builtins/ioi_legacy.py +473 -0
- circuitkit/tasks/builtins/mmlu.py +1578 -0
- circuitkit/tasks/builtins/sva.py +257 -0
- circuitkit/tasks/builtins/truthfulqa.py +520 -0
- circuitkit/tasks/builtins/winogrande.py +647 -0
- circuitkit/tasks/builtins/winogrande_mc.py +484 -0
- circuitkit/tasks/builtins/wmdp.py +1120 -0
- circuitkit/tasks/generic.py +1405 -0
- circuitkit/tasks/hf_factory.py +480 -0
- circuitkit/tasks/inspect.py +95 -0
- circuitkit/tasks/registry.py +72 -0
- circuitkit/tasks/safety_datasets.py +139 -0
- circuitkit/tasks/specs.py +246 -0
- circuitkit/tasks/type_specs/__init__.py +30 -0
- circuitkit/tasks/type_specs/classification_spec.py +49 -0
- circuitkit/tasks/type_specs/generation_spec.py +47 -0
- circuitkit/tasks/type_specs/mcq_spec.py +49 -0
- circuitkit/tasks/type_specs/qa_spec.py +129 -0
- circuitkit/tasks/type_specs/summarization_spec.py +45 -0
- circuitkit/tasks/type_specs/translation_spec.py +45 -0
- circuitkit/tasks/validator.py +481 -0
- circuitkit/tasks/yaml_loader.py +381 -0
- circuitkit/tooling/__init__.py +7 -0
- circuitkit/tooling/validate_environment.py +61 -0
- circuitkit/utils/__init__.py +0 -0
- circuitkit/utils/artifacts.py +39 -0
- circuitkit/utils/async_processing.py +371 -0
- circuitkit/utils/bootstrap.py +194 -0
- circuitkit/utils/config.py +330 -0
- circuitkit/utils/corruption_validation.py +369 -0
- circuitkit/utils/dataset_cache.py +303 -0
- circuitkit/utils/debug.py +345 -0
- circuitkit/utils/debugging.py +354 -0
- circuitkit/utils/device.py +51 -0
- circuitkit/utils/distributed.py +531 -0
- circuitkit/utils/exceptions.py +355 -0
- circuitkit/utils/logging.py +382 -0
- circuitkit/utils/memory.py +191 -0
- circuitkit/utils/optimization.py +316 -0
- circuitkit/utils/profiling.py +414 -0
- circuitkit/utils/token_utils.py +159 -0
- circuitkit/visualize/__init__.py +79 -0
- circuitkit/visualize/comparison.py +499 -0
- circuitkit/visualize/d3_template.py +1064 -0
- circuitkit/visualize/editor.py +385 -0
- circuitkit/visualize/feature_saliency.py +408 -0
- circuitkit/visualize/gallery.py +392 -0
- circuitkit/visualize/graph_viz.py +884 -0
- circuitkit/visualize/jupyter_suite.py +223 -0
- circuitkit/visualize/plotter.py +246 -0
- circuitkit/visualize/saliency.py +402 -0
- circuitkit/visualize/streamlit_app.py +471 -0
- circuitkit/visualize/theme.py +331 -0
- circuitkit-0.1.0.dist-info/METADATA +192 -0
- circuitkit-0.1.0.dist-info/RECORD +368 -0
- circuitkit-0.1.0.dist-info/WHEEL +5 -0
- circuitkit-0.1.0.dist-info/entry_points.txt +2 -0
- circuitkit-0.1.0.dist-info/licenses/LICENSE.md +73 -0
- circuitkit-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,962 @@
|
|
|
1
|
+
"""
|
|
2
|
+
This is taken from mlab2 repo; arthur/induction branch
|
|
3
|
+
|
|
4
|
+
It is a very slightly edited version of https://github.com/redwoodresearch/Easy-Transformer/blob/main/easy_transformer/ioi_dataset.py
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import copy
|
|
8
|
+
import random
|
|
9
|
+
import re
|
|
10
|
+
import warnings
|
|
11
|
+
from typing import List, Union
|
|
12
|
+
|
|
13
|
+
import numpy as np
|
|
14
|
+
import torch
|
|
15
|
+
|
|
16
|
+
from .....utils.logging import get_logger
|
|
17
|
+
|
|
18
|
+
logger = get_logger("data.task_data.ioi_dataset")
|
|
19
|
+
|
|
20
|
+
NAMES = [
|
|
21
|
+
"Michael",
|
|
22
|
+
"Christopher",
|
|
23
|
+
"Jessica",
|
|
24
|
+
"Matthew",
|
|
25
|
+
"Ashley",
|
|
26
|
+
"Jennifer",
|
|
27
|
+
"Joshua",
|
|
28
|
+
"Amanda",
|
|
29
|
+
"Daniel",
|
|
30
|
+
"David",
|
|
31
|
+
"James",
|
|
32
|
+
"Robert",
|
|
33
|
+
"John",
|
|
34
|
+
"Joseph",
|
|
35
|
+
"Andrew",
|
|
36
|
+
"Ryan",
|
|
37
|
+
"Brandon",
|
|
38
|
+
"Jason",
|
|
39
|
+
"Justin",
|
|
40
|
+
"Sarah",
|
|
41
|
+
"William",
|
|
42
|
+
"Jonathan",
|
|
43
|
+
"Stephanie",
|
|
44
|
+
"Brian",
|
|
45
|
+
"Nicole",
|
|
46
|
+
"Nicholas",
|
|
47
|
+
"Anthony",
|
|
48
|
+
"Heather",
|
|
49
|
+
"Eric",
|
|
50
|
+
"Elizabeth",
|
|
51
|
+
"Adam",
|
|
52
|
+
"Megan",
|
|
53
|
+
"Melissa",
|
|
54
|
+
"Kevin",
|
|
55
|
+
"Steven",
|
|
56
|
+
"Thomas",
|
|
57
|
+
"Timothy",
|
|
58
|
+
"Christina",
|
|
59
|
+
"Kyle",
|
|
60
|
+
"Rachel",
|
|
61
|
+
"Laura",
|
|
62
|
+
"Lauren",
|
|
63
|
+
"Amber",
|
|
64
|
+
"Brittany",
|
|
65
|
+
"Danielle",
|
|
66
|
+
"Richard",
|
|
67
|
+
"Kimberly",
|
|
68
|
+
"Jeffrey",
|
|
69
|
+
"Amy",
|
|
70
|
+
"Crystal",
|
|
71
|
+
"Michelle",
|
|
72
|
+
"Tiffany",
|
|
73
|
+
"Jeremy",
|
|
74
|
+
"Benjamin",
|
|
75
|
+
"Mark",
|
|
76
|
+
"Emily",
|
|
77
|
+
"Aaron",
|
|
78
|
+
"Charles",
|
|
79
|
+
"Rebecca",
|
|
80
|
+
"Jacob",
|
|
81
|
+
"Stephen",
|
|
82
|
+
"Patrick",
|
|
83
|
+
"Sean",
|
|
84
|
+
"Erin",
|
|
85
|
+
"Jamie",
|
|
86
|
+
"Kelly",
|
|
87
|
+
"Samantha",
|
|
88
|
+
"Nathan",
|
|
89
|
+
"Sara",
|
|
90
|
+
"Dustin",
|
|
91
|
+
"Paul",
|
|
92
|
+
"Angela",
|
|
93
|
+
"Tyler",
|
|
94
|
+
"Scott",
|
|
95
|
+
"Katherine",
|
|
96
|
+
"Andrea",
|
|
97
|
+
"Gregory",
|
|
98
|
+
"Erica",
|
|
99
|
+
"Mary",
|
|
100
|
+
"Travis",
|
|
101
|
+
"Lisa",
|
|
102
|
+
"Kenneth",
|
|
103
|
+
"Bryan",
|
|
104
|
+
"Lindsey",
|
|
105
|
+
"Kristen",
|
|
106
|
+
"Jose",
|
|
107
|
+
"Alexander",
|
|
108
|
+
"Jesse",
|
|
109
|
+
"Katie",
|
|
110
|
+
"Lindsay",
|
|
111
|
+
"Shannon",
|
|
112
|
+
"Vanessa",
|
|
113
|
+
"Courtney",
|
|
114
|
+
"Christine",
|
|
115
|
+
"Alicia",
|
|
116
|
+
"Cody",
|
|
117
|
+
"Allison",
|
|
118
|
+
"Bradley",
|
|
119
|
+
"Samuel",
|
|
120
|
+
]
|
|
121
|
+
|
|
122
|
+
ABC_TEMPLATES = [
|
|
123
|
+
"Then, [A], [B] and [C] went to the [PLACE]. [B] and [C] gave a [OBJECT] to [A]",
|
|
124
|
+
"Afterwards [A], [B] and [C] went to the [PLACE]. [B] and [C] gave a [OBJECT] to [A]",
|
|
125
|
+
"When [A], [B] and [C] arrived at the [PLACE], [B] and [C] gave a [OBJECT] to [A]",
|
|
126
|
+
"Friends [A], [B] and [C] went to the [PLACE]. [B] and [C] gave a [OBJECT] to [A]",
|
|
127
|
+
]
|
|
128
|
+
|
|
129
|
+
BAC_TEMPLATES = [
|
|
130
|
+
template.replace("[B]", "[A]", 1).replace("[A]", "[B]", 1) for template in ABC_TEMPLATES
|
|
131
|
+
]
|
|
132
|
+
|
|
133
|
+
BABA_TEMPLATES = [
|
|
134
|
+
"Then, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
135
|
+
"Then, [B] and [A] had a lot of fun at the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
136
|
+
"Then, [B] and [A] were working at the [PLACE]. [B] decided to give a [OBJECT] to [A]",
|
|
137
|
+
"Then, [B] and [A] were thinking about going to the [PLACE]. [B] wanted to give a [OBJECT] to [A]",
|
|
138
|
+
"Then, [B] and [A] had a long argument, and afterwards [B] said to [A]",
|
|
139
|
+
"After [B] and [A] went to the [PLACE], [B] gave a [OBJECT] to [A]",
|
|
140
|
+
"When [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give it to [A]",
|
|
141
|
+
"When [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give the [OBJECT] to [A]",
|
|
142
|
+
"While [B] and [A] were working at the [PLACE], [B] gave a [OBJECT] to [A]",
|
|
143
|
+
"While [B] and [A] were commuting to the [PLACE], [B] gave a [OBJECT] to [A]",
|
|
144
|
+
"After the lunch, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
145
|
+
"Afterwards, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
146
|
+
"Then, [B] and [A] had a long argument. Afterwards [B] said to [A]",
|
|
147
|
+
"The [PLACE] [B] and [A] went to had a [OBJECT]. [B] gave it to [A]",
|
|
148
|
+
"Friends [B] and [A] found a [OBJECT] at the [PLACE]. [B] gave it to [A]",
|
|
149
|
+
]
|
|
150
|
+
|
|
151
|
+
BABA_LONG_TEMPLATES = [
|
|
152
|
+
"Then in the morning, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
153
|
+
"Then in the morning, [B] and [A] had a lot of fun at the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
154
|
+
"Then in the morning, [B] and [A] were working at the [PLACE]. [B] decided to give a [OBJECT] to [A]",
|
|
155
|
+
"Then in the morning, [B] and [A] were thinking about going to the [PLACE]. [B] wanted to give a [OBJECT] to [A]",
|
|
156
|
+
"Then in the morning, [B] and [A] had a long argument, and afterwards [B] said to [A]",
|
|
157
|
+
"After taking a long break [B] and [A] went to the [PLACE], [B] gave a [OBJECT] to [A]",
|
|
158
|
+
"When soon afterwards [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give it to [A]",
|
|
159
|
+
"When soon afterwards [B] and [A] got a [OBJECT] at the [PLACE], [B] decided to give the [OBJECT] to [A]",
|
|
160
|
+
"While spending time together [B] and [A] were working at the [PLACE], [B] gave a [OBJECT] to [A]",
|
|
161
|
+
"While spending time together [B] and [A] were commuting to the [PLACE], [B] gave a [OBJECT] to [A]",
|
|
162
|
+
"After the lunch in the afternoon, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
163
|
+
"Afterwards, while spending time together [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
164
|
+
"Then in the morning afterwards, [B] and [A] had a long argument. Afterwards [B] said to [A]",
|
|
165
|
+
"The local big [PLACE] [B] and [A] went to had a [OBJECT]. [B] gave it to [A]",
|
|
166
|
+
"Friends separated at birth [B] and [A] found a [OBJECT] at the [PLACE]. [B] gave it to [A]",
|
|
167
|
+
]
|
|
168
|
+
|
|
169
|
+
BABA_LATE_IOS = [
|
|
170
|
+
"Then, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
171
|
+
"Then, [B] and [A] had a lot of fun at the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
172
|
+
"Then, [B] and [A] were working at the [PLACE]. [B] decided to give a [OBJECT] to [A]",
|
|
173
|
+
"Then, [B] and [A] were thinking about going to the [PLACE]. [B] wanted to give a [OBJECT] to [A]",
|
|
174
|
+
"Then, [B] and [A] had a long argument and after that [B] said to [A]",
|
|
175
|
+
"After the lunch, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
176
|
+
"Afterwards, [B] and [A] went to the [PLACE]. [B] gave a [OBJECT] to [A]",
|
|
177
|
+
"Then, [B] and [A] had a long argument. Afterwards [B] said to [A]",
|
|
178
|
+
]
|
|
179
|
+
|
|
180
|
+
BABA_EARLY_IOS = [
|
|
181
|
+
"Then [B] and [A] went to the [PLACE], and [B] gave a [OBJECT] to [A]",
|
|
182
|
+
"Then [B] and [A] had a lot of fun at the [PLACE], and [B] gave a [OBJECT] to [A]",
|
|
183
|
+
"Then [B] and [A] were working at the [PLACE], and [B] decided to give a [OBJECT] to [A]",
|
|
184
|
+
"Then [B] and [A] were thinking about going to the [PLACE], and [B] wanted to give a [OBJECT] to [A]",
|
|
185
|
+
"Then [B] and [A] had a long argument, and after that [B] said to [A]",
|
|
186
|
+
"After the lunch [B] and [A] went to the [PLACE], and [B] gave a [OBJECT] to [A]",
|
|
187
|
+
"Afterwards [B] and [A] went to the [PLACE], and [B] gave a [OBJECT] to [A]",
|
|
188
|
+
"Then [B] and [A] had a long argument, and afterwards [B] said to [A]",
|
|
189
|
+
]
|
|
190
|
+
|
|
191
|
+
TEMPLATES_VARIED_MIDDLE = [
|
|
192
|
+
"",
|
|
193
|
+
]
|
|
194
|
+
|
|
195
|
+
# no end of texts, GPT-2 small wasn't trained this way (ask Arthur)
|
|
196
|
+
# warnings.warn("Adding end of text prefixes!")
|
|
197
|
+
# for TEMPLATES in [BABA_TEMPLATES, BABA_EARLY_IOS, BABA_LATE_IOS]:
|
|
198
|
+
# for i in range(len(TEMPLATES)):
|
|
199
|
+
# TEMPLATES[i] = "<|endoftext|>" + TEMPLATES[i]
|
|
200
|
+
|
|
201
|
+
ABBA_TEMPLATES = BABA_TEMPLATES[:]
|
|
202
|
+
ABBA_LATE_IOS = BABA_LATE_IOS[:]
|
|
203
|
+
ABBA_EARLY_IOS = BABA_EARLY_IOS[:]
|
|
204
|
+
|
|
205
|
+
for TEMPLATES in [ABBA_TEMPLATES, ABBA_LATE_IOS, ABBA_EARLY_IOS]:
|
|
206
|
+
for i in range(len(TEMPLATES)):
|
|
207
|
+
first_clause = True
|
|
208
|
+
for j in range(1, len(TEMPLATES[i]) - 1):
|
|
209
|
+
if TEMPLATES[i][j - 1 : j + 2] == "[B]" and first_clause:
|
|
210
|
+
TEMPLATES[i] = TEMPLATES[i][:j] + "A" + TEMPLATES[i][j + 1 :]
|
|
211
|
+
elif TEMPLATES[i][j - 1 : j + 2] == "[A]" and first_clause:
|
|
212
|
+
first_clause = False
|
|
213
|
+
TEMPLATES[i] = TEMPLATES[i][:j] + "B" + TEMPLATES[i][j + 1 :]
|
|
214
|
+
|
|
215
|
+
VERBS = [" tried", " said", " decided", " wanted", " gave"]
|
|
216
|
+
PLACES = [
|
|
217
|
+
"store",
|
|
218
|
+
"garden",
|
|
219
|
+
"restaurant",
|
|
220
|
+
"school",
|
|
221
|
+
"hospital",
|
|
222
|
+
"office",
|
|
223
|
+
"house",
|
|
224
|
+
"station",
|
|
225
|
+
]
|
|
226
|
+
OBJECTS = [
|
|
227
|
+
"ring",
|
|
228
|
+
"kiss",
|
|
229
|
+
"bone",
|
|
230
|
+
"basketball",
|
|
231
|
+
"computer",
|
|
232
|
+
"necklace",
|
|
233
|
+
"drink",
|
|
234
|
+
"snack",
|
|
235
|
+
]
|
|
236
|
+
|
|
237
|
+
ANIMALS = [
|
|
238
|
+
"dog",
|
|
239
|
+
"cat",
|
|
240
|
+
"snake",
|
|
241
|
+
"elephant",
|
|
242
|
+
"beetle",
|
|
243
|
+
"hippo",
|
|
244
|
+
"giraffe",
|
|
245
|
+
"tiger",
|
|
246
|
+
"husky",
|
|
247
|
+
"lion",
|
|
248
|
+
"panther",
|
|
249
|
+
"whale",
|
|
250
|
+
"dolphin",
|
|
251
|
+
"beaver",
|
|
252
|
+
"rabbit",
|
|
253
|
+
"fox",
|
|
254
|
+
"lamb",
|
|
255
|
+
"ferret",
|
|
256
|
+
]
|
|
257
|
+
|
|
258
|
+
def multiple_replace(dict, text):
|
|
259
|
+
# from: https://stackoverflow.com/questions/15175142/how-can-i-do-multiple-substitutions-using-regex
|
|
260
|
+
# Create a regular expression from the dictionary keys
|
|
261
|
+
regex = re.compile("(%s)" % "|".join(map(re.escape, dict.keys())))
|
|
262
|
+
|
|
263
|
+
# For each match, look-up corresponding value in dictionary
|
|
264
|
+
return regex.sub(lambda mo: dict[mo.string[mo.start() : mo.end()]], text)
|
|
265
|
+
|
|
266
|
+
def iter_sample_fast(iterable, samplesize, seed):
|
|
267
|
+
random.seed(seed)
|
|
268
|
+
results = []
|
|
269
|
+
# Fill in the first samplesize elements:
|
|
270
|
+
try:
|
|
271
|
+
for _ in range(samplesize):
|
|
272
|
+
results.append(next(iterable))
|
|
273
|
+
except StopIteration:
|
|
274
|
+
raise ValueError("Sample larger than population.")
|
|
275
|
+
random.shuffle(results) # Randomize their positions
|
|
276
|
+
|
|
277
|
+
return results
|
|
278
|
+
|
|
279
|
+
NOUNS_DICT = NOUNS_DICT = {"[PLACE]": PLACES, "[OBJECT]": OBJECTS}
|
|
280
|
+
|
|
281
|
+
def gen_prompt_uniform(
|
|
282
|
+
templates,
|
|
283
|
+
names,
|
|
284
|
+
nouns_dict,
|
|
285
|
+
N,
|
|
286
|
+
symmetric,
|
|
287
|
+
prefixes=None,
|
|
288
|
+
abc=False,
|
|
289
|
+
seed=None,
|
|
290
|
+
):
|
|
291
|
+
assert seed is not None
|
|
292
|
+
random.seed(seed)
|
|
293
|
+
|
|
294
|
+
nb_gen = 0
|
|
295
|
+
ioi_prompts = []
|
|
296
|
+
while nb_gen < N:
|
|
297
|
+
temp = random.choice(templates)
|
|
298
|
+
temp_id = templates.index(temp)
|
|
299
|
+
name_1 = ""
|
|
300
|
+
name_2 = ""
|
|
301
|
+
name_3 = ""
|
|
302
|
+
while len(set([name_1, name_2, name_3])) < 3:
|
|
303
|
+
name_1 = random.choice(names)
|
|
304
|
+
name_2 = random.choice(names)
|
|
305
|
+
name_3 = random.choice(names)
|
|
306
|
+
|
|
307
|
+
nouns = {}
|
|
308
|
+
ioi_prompt = {}
|
|
309
|
+
for k in nouns_dict:
|
|
310
|
+
nouns[k] = random.choice(nouns_dict[k])
|
|
311
|
+
ioi_prompt[k] = nouns[k]
|
|
312
|
+
prompt = temp
|
|
313
|
+
for k in nouns_dict:
|
|
314
|
+
prompt = prompt.replace(k, nouns[k])
|
|
315
|
+
|
|
316
|
+
if prefixes is not None:
|
|
317
|
+
L = random.randint(30, 40)
|
|
318
|
+
pref = ".".join(random.choice(prefixes).split(".")[:L])
|
|
319
|
+
pref += "<|endoftext|>"
|
|
320
|
+
else:
|
|
321
|
+
pref = ""
|
|
322
|
+
|
|
323
|
+
prompt1 = prompt.replace("[A]", name_1)
|
|
324
|
+
prompt1 = prompt1.replace("[B]", name_2)
|
|
325
|
+
if abc:
|
|
326
|
+
prompt1 = prompt1.replace("[C]", name_3)
|
|
327
|
+
prompt1 = pref + prompt1
|
|
328
|
+
ioi_prompt["text"] = prompt1
|
|
329
|
+
ioi_prompt["IO"] = name_1
|
|
330
|
+
ioi_prompt["S"] = name_2
|
|
331
|
+
ioi_prompt["TEMPLATE_IDX"] = temp_id
|
|
332
|
+
ioi_prompts.append(ioi_prompt)
|
|
333
|
+
if abc:
|
|
334
|
+
ioi_prompts[-1]["C"] = name_3
|
|
335
|
+
|
|
336
|
+
nb_gen += 1
|
|
337
|
+
|
|
338
|
+
if symmetric and nb_gen < N:
|
|
339
|
+
prompt2 = prompt.replace("[A]", name_2)
|
|
340
|
+
prompt2 = prompt2.replace("[B]", name_1)
|
|
341
|
+
prompt2 = pref + prompt2
|
|
342
|
+
ioi_prompts.append(
|
|
343
|
+
{"text": prompt2, "IO": name_2, "S": name_1, "TEMPLATE_IDX": temp_id}
|
|
344
|
+
)
|
|
345
|
+
nb_gen += 1
|
|
346
|
+
return ioi_prompts
|
|
347
|
+
|
|
348
|
+
def gen_flipped_prompts( # noqa: C901 - complex function, refactor out of scope for lint pass
|
|
349
|
+
prompts, names, flip=("S2", "IO"), seed=None
|
|
350
|
+
):
|
|
351
|
+
"""_summary_
|
|
352
|
+
|
|
353
|
+
Args:
|
|
354
|
+
prompts (List[D]): _description_
|
|
355
|
+
flip (tuple, optional): First element is the string to be replaced, Second is what to replace with. Defaults to ("S2", "IO").
|
|
356
|
+
|
|
357
|
+
Returns:
|
|
358
|
+
_type_: _description_
|
|
359
|
+
"""
|
|
360
|
+
|
|
361
|
+
assert seed is not None
|
|
362
|
+
np.random.seed(seed)
|
|
363
|
+
|
|
364
|
+
flipped_prompts = []
|
|
365
|
+
|
|
366
|
+
for prompt in prompts:
|
|
367
|
+
t = prompt["text"].split(" ")
|
|
368
|
+
prompt = prompt.copy()
|
|
369
|
+
if flip[0] == "S2":
|
|
370
|
+
if flip[1] == "IO":
|
|
371
|
+
t[len(t) - t[::-1].index(prompt["S"]) - 1] = prompt["IO"]
|
|
372
|
+
temp = prompt["IO"]
|
|
373
|
+
prompt["IO"] = prompt["S"]
|
|
374
|
+
prompt["S"] = temp
|
|
375
|
+
elif flip[1] == "RAND":
|
|
376
|
+
rand_name = names[np.random.randint(len(names))]
|
|
377
|
+
while rand_name == prompt["IO"] or rand_name == prompt["S"]:
|
|
378
|
+
rand_name = names[np.random.randint(len(names))]
|
|
379
|
+
t[len(t) - t[::-1].index(prompt["S"]) - 1] = rand_name
|
|
380
|
+
else:
|
|
381
|
+
raise ValueError("Invalid flip[1] value")
|
|
382
|
+
|
|
383
|
+
elif flip[0] == "IO":
|
|
384
|
+
if flip[1] == "RAND":
|
|
385
|
+
rand_name = names[np.random.randint(len(names))]
|
|
386
|
+
while rand_name == prompt["IO"] or rand_name == prompt["S"]:
|
|
387
|
+
rand_name = names[np.random.randint(len(names))]
|
|
388
|
+
|
|
389
|
+
t[t.index(prompt["IO"])] = rand_name
|
|
390
|
+
t[t.index(prompt["IO"])] = rand_name
|
|
391
|
+
prompt["IO"] = rand_name
|
|
392
|
+
elif flip[1] == "ANIMAL":
|
|
393
|
+
rand_animal = ANIMALS[np.random.randint(len(ANIMALS))]
|
|
394
|
+
t[t.index(prompt["IO"])] = rand_animal
|
|
395
|
+
prompt["IO"] = rand_animal
|
|
396
|
+
elif flip[1] == "S1":
|
|
397
|
+
io_index = t.index(prompt["IO"])
|
|
398
|
+
s1_index = t.index(prompt["S"])
|
|
399
|
+
io = t[io_index]
|
|
400
|
+
s1 = t[s1_index]
|
|
401
|
+
t[io_index] = s1
|
|
402
|
+
t[s1_index] = io
|
|
403
|
+
else:
|
|
404
|
+
raise ValueError("Invalid flip[1] value")
|
|
405
|
+
|
|
406
|
+
elif flip[0] in ["S", "S1"]:
|
|
407
|
+
if flip[1] == "ANIMAL":
|
|
408
|
+
new_s = ANIMALS[np.random.randint(len(ANIMALS))]
|
|
409
|
+
if flip[1] == "RAND":
|
|
410
|
+
new_s = names[np.random.randint(len(names))]
|
|
411
|
+
t[t.index(prompt["S"])] = new_s
|
|
412
|
+
if flip[0] == "S": # literally just change the first S if this is S1
|
|
413
|
+
t[len(t) - t[::-1].index(prompt["S"]) - 1] = new_s
|
|
414
|
+
prompt["S"] = new_s
|
|
415
|
+
elif flip[0] == "END":
|
|
416
|
+
if flip[1] == "S":
|
|
417
|
+
t[len(t) - t[::-1].index(prompt["IO"]) - 1] = prompt["S"]
|
|
418
|
+
elif flip[0] == "PUNC":
|
|
419
|
+
n = []
|
|
420
|
+
|
|
421
|
+
# separate the punctuation from the words
|
|
422
|
+
for i, word in enumerate(t):
|
|
423
|
+
if "." in word:
|
|
424
|
+
n.append(word[:-1])
|
|
425
|
+
n.append(".")
|
|
426
|
+
elif "," in word:
|
|
427
|
+
n.append(word[:-1])
|
|
428
|
+
n.append(",")
|
|
429
|
+
else:
|
|
430
|
+
n.append(word)
|
|
431
|
+
|
|
432
|
+
# remove punctuation, important that you check for period first
|
|
433
|
+
if flip[1] == "NONE":
|
|
434
|
+
if "." in n:
|
|
435
|
+
n[n.index(".")] = ""
|
|
436
|
+
elif "," in n:
|
|
437
|
+
n[len(n) - n[::-1].index(",") - 1] = ""
|
|
438
|
+
|
|
439
|
+
# remove empty strings
|
|
440
|
+
while "" in n:
|
|
441
|
+
n.remove("")
|
|
442
|
+
|
|
443
|
+
# add punctuation back to the word before it
|
|
444
|
+
while "," in n:
|
|
445
|
+
n[n.index(",") - 1] += ","
|
|
446
|
+
n.remove(",")
|
|
447
|
+
|
|
448
|
+
while "." in n:
|
|
449
|
+
n[n.index(".") - 1] += "."
|
|
450
|
+
n.remove(".")
|
|
451
|
+
|
|
452
|
+
t = n
|
|
453
|
+
|
|
454
|
+
elif flip[0] == "C2":
|
|
455
|
+
if flip[1] == "A":
|
|
456
|
+
t[len(t) - t[::-1].index(prompt["C"]) - 1] = prompt["A"]
|
|
457
|
+
elif flip[0] == "S+1":
|
|
458
|
+
if t[t.index(prompt["S"]) + 1] == "and":
|
|
459
|
+
t[t.index(prompt["S"]) + 1] = [
|
|
460
|
+
"with one friend named",
|
|
461
|
+
"accompanied by",
|
|
462
|
+
][np.random.randint(2)]
|
|
463
|
+
else:
|
|
464
|
+
t[t.index(prompt["S"]) + 1] = (
|
|
465
|
+
t[t.index(prompt["S"])] + ", after a great day, " + t[t.index(prompt["S"]) + 1]
|
|
466
|
+
)
|
|
467
|
+
del t[t.index(prompt["S"])]
|
|
468
|
+
else:
|
|
469
|
+
raise ValueError(f"Invalid flipper {flip[0]}")
|
|
470
|
+
|
|
471
|
+
if "IO" in prompt:
|
|
472
|
+
prompt["text"] = " ".join(t)
|
|
473
|
+
flipped_prompts.append(prompt)
|
|
474
|
+
else:
|
|
475
|
+
flipped_prompts.append(
|
|
476
|
+
{
|
|
477
|
+
"A": prompt["A"],
|
|
478
|
+
"B": prompt["B"],
|
|
479
|
+
"C": prompt["C"],
|
|
480
|
+
"text": " ".join(t),
|
|
481
|
+
}
|
|
482
|
+
)
|
|
483
|
+
|
|
484
|
+
return flipped_prompts
|
|
485
|
+
|
|
486
|
+
# *Tok Idxs Methods
|
|
487
|
+
|
|
488
|
+
def get_name_idxs(prompts, tokenizer, idx_types=["IO", "S", "S2"], prepend_bos=False):
|
|
489
|
+
name_idx_dict = dict((idx_type, []) for idx_type in idx_types)
|
|
490
|
+
double_s2 = False
|
|
491
|
+
for prompt in prompts:
|
|
492
|
+
t = prompt["text"].split(" ")
|
|
493
|
+
toks = tokenizer.tokenize(" ".join(t[:-1]))
|
|
494
|
+
for idx_type in idx_types:
|
|
495
|
+
if "2" in idx_type:
|
|
496
|
+
idx = (
|
|
497
|
+
len(toks)
|
|
498
|
+
- toks[::-1].index(tokenizer.tokenize(" " + prompt[idx_type[:-1]])[0])
|
|
499
|
+
- 1
|
|
500
|
+
)
|
|
501
|
+
else:
|
|
502
|
+
idx = toks.index(tokenizer.tokenize(" " + prompt[idx_type])[0])
|
|
503
|
+
name_idx_dict[idx_type].append(idx)
|
|
504
|
+
if "S" in idx_types and "S2" in idx_types:
|
|
505
|
+
if name_idx_dict["S"][-1] == name_idx_dict["S2"][-1]:
|
|
506
|
+
double_s2 = True
|
|
507
|
+
if double_s2:
|
|
508
|
+
warnings.warn("S2 index has been computed as the same for S and S2")
|
|
509
|
+
|
|
510
|
+
return [int(prepend_bos) + torch.tensor(name_idx_dict[idx_type]) for idx_type in idx_types]
|
|
511
|
+
|
|
512
|
+
def get_word_idxs(prompts, word_list, tokenizer):
|
|
513
|
+
"""Get the index of the words in word_list in the prompts. Exactly one of the word_list word has to be present in each prompt"""
|
|
514
|
+
idxs = []
|
|
515
|
+
tokenized_words = [tokenizer.decode(tokenizer(word)["input_ids"][0]) for word in word_list]
|
|
516
|
+
for pr_idx, prompt in enumerate(prompts):
|
|
517
|
+
toks = [
|
|
518
|
+
tokenizer.decode(t)
|
|
519
|
+
for t in tokenizer(prompt["text"], return_tensors="pt", padding=True)["input_ids"][0]
|
|
520
|
+
]
|
|
521
|
+
idx = None
|
|
522
|
+
for i, w_tok in enumerate(tokenized_words):
|
|
523
|
+
if word_list[i] in prompt["text"]:
|
|
524
|
+
try:
|
|
525
|
+
idx = toks.index(w_tok)
|
|
526
|
+
if toks.count(w_tok) > 1:
|
|
527
|
+
idx = len(toks) - toks[::-1].index(w_tok) - 1
|
|
528
|
+
except Exception:
|
|
529
|
+
idx = toks.index(w_tok)
|
|
530
|
+
# raise ValueError(toks, w_tok, prompt["text"])
|
|
531
|
+
if idx is None:
|
|
532
|
+
raise ValueError(f"Word {word_list} and {i} not found {prompt}")
|
|
533
|
+
idxs.append(idx)
|
|
534
|
+
return torch.tensor(idxs)
|
|
535
|
+
|
|
536
|
+
def get_end_idxs(prompts, tokenizer, name_tok_len=1, prepend_bos=False, toks=None):
|
|
537
|
+
|
|
538
|
+
# toks = torch.Tensor(tokenizer([prompt["text"] for prompt in prompts], padding=True).input_ids).type(torch.int)
|
|
539
|
+
relevant_idx = int(prepend_bos)
|
|
540
|
+
# if the sentence begins with an end token
|
|
541
|
+
# AND the model pads at the end with the same end token,
|
|
542
|
+
# then we need make special arrangements
|
|
543
|
+
|
|
544
|
+
pad_token_id = tokenizer.pad_token_id
|
|
545
|
+
|
|
546
|
+
end_idxs_raw = []
|
|
547
|
+
for i in range(toks.shape[0]):
|
|
548
|
+
if pad_token_id not in toks[i][1:]:
|
|
549
|
+
end_idxs_raw.append(toks.shape[1])
|
|
550
|
+
continue
|
|
551
|
+
nonzers = (toks[i] == pad_token_id).nonzero()
|
|
552
|
+
try:
|
|
553
|
+
nonzers = nonzers[relevant_idx]
|
|
554
|
+
except Exception:
|
|
555
|
+
logger.error(toks[i])
|
|
556
|
+
logger.error(nonzers)
|
|
557
|
+
logger.error(relevant_idx)
|
|
558
|
+
logger.error(i)
|
|
559
|
+
raise ValueError("Something went wrong")
|
|
560
|
+
nonzers = nonzers[0]
|
|
561
|
+
nonzers = nonzers.item()
|
|
562
|
+
end_idxs_raw.append(nonzers)
|
|
563
|
+
end_idxs = torch.tensor(end_idxs_raw)
|
|
564
|
+
end_idxs = end_idxs - 1 - name_tok_len
|
|
565
|
+
|
|
566
|
+
for i in range(toks.shape[0]):
|
|
567
|
+
assert toks[i][end_idxs[i] + 1] != 0 and (
|
|
568
|
+
toks.shape[1] == end_idxs[i] + 2 or toks[i][end_idxs[i] + 2] == pad_token_id
|
|
569
|
+
), (
|
|
570
|
+
toks[i],
|
|
571
|
+
end_idxs[i],
|
|
572
|
+
toks[i].shape,
|
|
573
|
+
"the END idxs aren't properly formatted",
|
|
574
|
+
)
|
|
575
|
+
|
|
576
|
+
return end_idxs
|
|
577
|
+
|
|
578
|
+
ALL_SEM = [
|
|
579
|
+
"S",
|
|
580
|
+
"IO",
|
|
581
|
+
"S2",
|
|
582
|
+
"end",
|
|
583
|
+
"S+1",
|
|
584
|
+
"and",
|
|
585
|
+
] # , "verb", "starts", "S-1", "punct"] # Kevin's antic averages
|
|
586
|
+
|
|
587
|
+
def get_idx_dict(ioi_prompts, tokenizer, prepend_bos=False, toks=None):
|
|
588
|
+
(
|
|
589
|
+
IO_idxs,
|
|
590
|
+
S_idxs,
|
|
591
|
+
S2_idxs,
|
|
592
|
+
) = get_name_idxs(
|
|
593
|
+
ioi_prompts,
|
|
594
|
+
tokenizer,
|
|
595
|
+
idx_types=["IO", "S", "S2"],
|
|
596
|
+
prepend_bos=prepend_bos,
|
|
597
|
+
)
|
|
598
|
+
|
|
599
|
+
end_idxs = get_end_idxs(
|
|
600
|
+
ioi_prompts,
|
|
601
|
+
tokenizer,
|
|
602
|
+
name_tok_len=1,
|
|
603
|
+
prepend_bos=prepend_bos,
|
|
604
|
+
toks=toks,
|
|
605
|
+
)
|
|
606
|
+
|
|
607
|
+
punct_idxs = get_word_idxs(ioi_prompts, [",", "."], tokenizer)
|
|
608
|
+
|
|
609
|
+
return {
|
|
610
|
+
"IO": IO_idxs,
|
|
611
|
+
"IO-1": IO_idxs - 1,
|
|
612
|
+
"IO+1": IO_idxs + 1,
|
|
613
|
+
"S": S_idxs,
|
|
614
|
+
"S-1": S_idxs - 1,
|
|
615
|
+
"S+1": S_idxs + 1,
|
|
616
|
+
"S2": S2_idxs,
|
|
617
|
+
"end": end_idxs,
|
|
618
|
+
"starts": torch.zeros_like(end_idxs),
|
|
619
|
+
"punct": punct_idxs,
|
|
620
|
+
}
|
|
621
|
+
|
|
622
|
+
# Some functions for experiments on Pointer Arithmetic
|
|
623
|
+
|
|
624
|
+
PREFIXES = [
|
|
625
|
+
" Afterwards,",
|
|
626
|
+
" Two friends met at a bar. Then,",
|
|
627
|
+
" After a long day,",
|
|
628
|
+
" After a long day,",
|
|
629
|
+
" Then,",
|
|
630
|
+
" Then,",
|
|
631
|
+
]
|
|
632
|
+
|
|
633
|
+
def flip_prefixes(ioi_prompts):
|
|
634
|
+
ioi_prompts = copy.deepcopy(ioi_prompts)
|
|
635
|
+
for prompt in ioi_prompts:
|
|
636
|
+
if prompt["text"].startswith("The "):
|
|
637
|
+
prompt["text"] = "After the lunch, the" + prompt["text"][4:]
|
|
638
|
+
else:
|
|
639
|
+
io_idx = prompt["text"].index(prompt["IO"])
|
|
640
|
+
s_idx = prompt["text"].index(prompt["S"])
|
|
641
|
+
first_idx = min(io_idx, s_idx)
|
|
642
|
+
prompt["text"] = random.choice(PREFIXES) + " " + prompt["text"][first_idx:]
|
|
643
|
+
|
|
644
|
+
return ioi_prompts
|
|
645
|
+
|
|
646
|
+
def flip_names(ioi_prompts):
|
|
647
|
+
ioi_prompts = copy.deepcopy(ioi_prompts)
|
|
648
|
+
for prompt in ioi_prompts:
|
|
649
|
+
punct_idx = max(
|
|
650
|
+
[i for i, x in enumerate(list(prompt["text"])) if x in [",", "."]]
|
|
651
|
+
) # only flip name in the first clause
|
|
652
|
+
io = prompt["IO"]
|
|
653
|
+
s = prompt["S"]
|
|
654
|
+
prompt["text"] = (
|
|
655
|
+
prompt["text"][:punct_idx]
|
|
656
|
+
.replace(io, "#")
|
|
657
|
+
.replace(s, "@")
|
|
658
|
+
.replace("#", s)
|
|
659
|
+
.replace("@", io)
|
|
660
|
+
) + prompt["text"][punct_idx:]
|
|
661
|
+
|
|
662
|
+
return ioi_prompts
|
|
663
|
+
|
|
664
|
+
class IOIDataset:
|
|
665
|
+
def __init__(
|
|
666
|
+
self,
|
|
667
|
+
prompt_type: Union[str, List[str]], # if list, then it will be a list of templates
|
|
668
|
+
N=500,
|
|
669
|
+
model=None, # Required model parameter for TokenIDGenerator
|
|
670
|
+
prompts=None,
|
|
671
|
+
symmetric=False,
|
|
672
|
+
prefixes=None,
|
|
673
|
+
nb_templates=None,
|
|
674
|
+
ioi_prompts_for_word_idxs=None,
|
|
675
|
+
prepend_bos=False,
|
|
676
|
+
manual_word_idx=None,
|
|
677
|
+
seed=None,
|
|
678
|
+
):
|
|
679
|
+
"""
|
|
680
|
+
ioi_prompts_for_word_idxs:
|
|
681
|
+
if you want to use a different set of prompts to get the word indices, you can pass it here
|
|
682
|
+
(example use case: making a ABCA dataset)
|
|
683
|
+
"""
|
|
684
|
+
|
|
685
|
+
assert seed is not None
|
|
686
|
+
random.seed(seed)
|
|
687
|
+
|
|
688
|
+
if not (
|
|
689
|
+
N == 1
|
|
690
|
+
or prepend_bos is False
|
|
691
|
+
or tokenizer.bos_token_id # noqa: F821 - pre-existing vendored bug
|
|
692
|
+
== tokenizer.eos_token_id # noqa: F821 - pre-existing vendored bug
|
|
693
|
+
):
|
|
694
|
+
warnings.warn("Probably word_idx will be calculated incorrectly due to this formatting")
|
|
695
|
+
assert not (symmetric and prompt_type == "ABC")
|
|
696
|
+
assert (prompts is not None) or (not symmetric) or (N % 2 == 0), f"{symmetric} {N}"
|
|
697
|
+
assert nb_templates is None or (nb_templates % 2 == 0 or prompt_type != "mixed")
|
|
698
|
+
self.prompt_type = prompt_type
|
|
699
|
+
|
|
700
|
+
if nb_templates is None:
|
|
701
|
+
nb_templates = len(BABA_TEMPLATES)
|
|
702
|
+
|
|
703
|
+
if prompt_type == "ABBA":
|
|
704
|
+
self.templates = ABBA_TEMPLATES[:nb_templates].copy()
|
|
705
|
+
elif prompt_type == "BABA":
|
|
706
|
+
self.templates = BABA_TEMPLATES[:nb_templates].copy()
|
|
707
|
+
elif prompt_type == "mixed":
|
|
708
|
+
self.templates = (
|
|
709
|
+
BABA_TEMPLATES[: nb_templates // 2].copy()
|
|
710
|
+
+ ABBA_TEMPLATES[: nb_templates // 2].copy()
|
|
711
|
+
)
|
|
712
|
+
random.shuffle(self.templates)
|
|
713
|
+
elif prompt_type == "ABC":
|
|
714
|
+
self.templates = ABC_TEMPLATES[:nb_templates].copy()
|
|
715
|
+
elif prompt_type == "BAC":
|
|
716
|
+
self.templates = BAC_TEMPLATES[:nb_templates].copy()
|
|
717
|
+
elif prompt_type == "ABC mixed":
|
|
718
|
+
self.templates = (
|
|
719
|
+
ABC_TEMPLATES[: nb_templates // 2].copy()
|
|
720
|
+
+ BAC_TEMPLATES[: nb_templates // 2].copy()
|
|
721
|
+
)
|
|
722
|
+
random.shuffle(self.templates)
|
|
723
|
+
elif isinstance(prompt_type, list):
|
|
724
|
+
self.templates = prompt_type
|
|
725
|
+
else:
|
|
726
|
+
raise ValueError(prompt_type)
|
|
727
|
+
|
|
728
|
+
if model is None:
|
|
729
|
+
raise ValueError(
|
|
730
|
+
"Model is required for IOIDataset. "
|
|
731
|
+
"No default model to ensure model compatibility."
|
|
732
|
+
)
|
|
733
|
+
self.tokenizer = model.tokenizer
|
|
734
|
+
self.model = model
|
|
735
|
+
|
|
736
|
+
self.prefixes = prefixes
|
|
737
|
+
self.prompt_type = prompt_type
|
|
738
|
+
if prompts is None:
|
|
739
|
+
self.ioi_prompts = gen_prompt_uniform( # a list of dict of the form {"text": "Alice and Bob bla bla. Bob gave bla to Alice", "IO": "Alice", "S": "Bob"}
|
|
740
|
+
self.templates,
|
|
741
|
+
NAMES,
|
|
742
|
+
nouns_dict={"[PLACE]": PLACES, "[OBJECT]": OBJECTS},
|
|
743
|
+
N=N,
|
|
744
|
+
symmetric=symmetric,
|
|
745
|
+
prefixes=self.prefixes,
|
|
746
|
+
abc=(prompt_type in ["ABC", "ABC mixed", "BAC"]),
|
|
747
|
+
seed=(seed + 987654321) % 123456789,
|
|
748
|
+
)
|
|
749
|
+
else:
|
|
750
|
+
assert N == len(prompts), f"{N} and {len(prompts)}"
|
|
751
|
+
self.ioi_prompts = prompts
|
|
752
|
+
|
|
753
|
+
all_ids = [prompt["TEMPLATE_IDX"] for prompt in self.ioi_prompts]
|
|
754
|
+
all_ids_ar = np.array(all_ids)
|
|
755
|
+
self.groups = []
|
|
756
|
+
for id in list(set(all_ids)):
|
|
757
|
+
self.groups.append(np.where(all_ids_ar == id)[0])
|
|
758
|
+
|
|
759
|
+
small_groups = []
|
|
760
|
+
for group in self.groups:
|
|
761
|
+
if len(group) < 5:
|
|
762
|
+
small_groups.append(len(group))
|
|
763
|
+
if len(small_groups) > 0:
|
|
764
|
+
warnings.warn(f"Some groups have less than 5 prompts, they have lengths {small_groups}")
|
|
765
|
+
|
|
766
|
+
self.sentences = [
|
|
767
|
+
prompt["text"] for prompt in self.ioi_prompts
|
|
768
|
+
] # a list of strings. Renamed as this should NOT be forward passed
|
|
769
|
+
|
|
770
|
+
self.templates_by_prompt = [] # for each prompt if it's ABBA or BABA
|
|
771
|
+
for i in range(N):
|
|
772
|
+
if self.sentences[i].index(self.ioi_prompts[i]["IO"]) < self.sentences[i].index(
|
|
773
|
+
self.ioi_prompts[i]["S"]
|
|
774
|
+
):
|
|
775
|
+
self.templates_by_prompt.append("ABBA")
|
|
776
|
+
else:
|
|
777
|
+
self.templates_by_prompt.append("BABA")
|
|
778
|
+
|
|
779
|
+
texts = [
|
|
780
|
+
(self.tokenizer.bos_token if prepend_bos else "") + prompt["text"]
|
|
781
|
+
for prompt in self.ioi_prompts
|
|
782
|
+
]
|
|
783
|
+
self.toks = torch.Tensor(self.tokenizer(texts, padding=True).input_ids).type(torch.int)
|
|
784
|
+
|
|
785
|
+
if ioi_prompts_for_word_idxs is None:
|
|
786
|
+
ioi_prompts_for_word_idxs = self.ioi_prompts
|
|
787
|
+
self.word_idx = get_idx_dict(
|
|
788
|
+
ioi_prompts_for_word_idxs,
|
|
789
|
+
self.tokenizer,
|
|
790
|
+
prepend_bos=prepend_bos,
|
|
791
|
+
toks=self.toks,
|
|
792
|
+
)
|
|
793
|
+
self.prepend_bos = prepend_bos
|
|
794
|
+
if manual_word_idx is not None:
|
|
795
|
+
self.word_idx = manual_word_idx
|
|
796
|
+
|
|
797
|
+
self.sem_tok_idx = {
|
|
798
|
+
k: v for k, v in self.word_idx.items() if k in ALL_SEM
|
|
799
|
+
} # the semantic indices that kevin uses
|
|
800
|
+
self.N = N
|
|
801
|
+
self.max_len = max(
|
|
802
|
+
[len(self.tokenizer(prompt["text"]).input_ids) for prompt in self.ioi_prompts]
|
|
803
|
+
)
|
|
804
|
+
|
|
805
|
+
# Use TokenIDGenerator for consistent token ID generation
|
|
806
|
+
from circuitkit.utils.token_utils import TokenIDGenerator
|
|
807
|
+
|
|
808
|
+
# Create a mock model object if only tokenizer is provided
|
|
809
|
+
if model is None:
|
|
810
|
+
# Create a minimal model-like object for TokenIDGenerator
|
|
811
|
+
class MockModel:
|
|
812
|
+
def __init__(self, tokenizer):
|
|
813
|
+
self.tokenizer = tokenizer
|
|
814
|
+
self.cfg = type("Config", (), {"model_name": "unknown"})()
|
|
815
|
+
|
|
816
|
+
mock_model = MockModel(self.tokenizer)
|
|
817
|
+
token_gen = TokenIDGenerator(mock_model)
|
|
818
|
+
else:
|
|
819
|
+
token_gen = TokenIDGenerator(model)
|
|
820
|
+
|
|
821
|
+
self.io_tokenIDs = token_gen.get_token_ids_batch([" " + p["IO"] for p in self.ioi_prompts])
|
|
822
|
+
self.s_tokenIDs = token_gen.get_token_ids_batch([" " + p["S"] for p in self.ioi_prompts])
|
|
823
|
+
|
|
824
|
+
self.tokenized_prompts = []
|
|
825
|
+
|
|
826
|
+
for i in range(self.N):
|
|
827
|
+
self.tokenized_prompts.append(
|
|
828
|
+
"|".join([self.tokenizer.decode(tok) for tok in self.toks[i]])
|
|
829
|
+
)
|
|
830
|
+
|
|
831
|
+
@classmethod
|
|
832
|
+
def construct_from_ioi_prompts_metadata(cls, templates, ioi_prompts_data, **kwargs):
|
|
833
|
+
"""
|
|
834
|
+
Given a list of dictionaries (ioi_prompts_data)
|
|
835
|
+
{
|
|
836
|
+
"S": "Bob",
|
|
837
|
+
"IO": "Alice",
|
|
838
|
+
"TEMPLATE_IDX": 0
|
|
839
|
+
}
|
|
840
|
+
|
|
841
|
+
create and IOIDataset from these
|
|
842
|
+
"""
|
|
843
|
+
|
|
844
|
+
prompts = []
|
|
845
|
+
for metadata in ioi_prompts_data:
|
|
846
|
+
cur_template = templates[metadata["TEMPLATE_IDX"]]
|
|
847
|
+
prompts.append(metadata)
|
|
848
|
+
prompts[-1]["text"] = (
|
|
849
|
+
cur_template.replace("[A]", metadata["IO"])
|
|
850
|
+
.replace("[B]", metadata["S"])
|
|
851
|
+
.replace("[PLACE]", metadata["[PLACE]"])
|
|
852
|
+
.replace("[OBJECT]", metadata["[OBJECT]"])
|
|
853
|
+
)
|
|
854
|
+
# prompts[-1]["[PLACE]"] = metadata["[PLACE]"]
|
|
855
|
+
# prompts[-1]["[OBJECT]"] = metadata["[OBJECT]"]
|
|
856
|
+
return IOIDataset(prompt_type=templates, prompts=prompts, **kwargs)
|
|
857
|
+
|
|
858
|
+
def gen_flipped_prompts(self, flip, seed=None):
|
|
859
|
+
"""
|
|
860
|
+
Return a IOIDataset where the name to flip has been replaced by a random name.
|
|
861
|
+
"""
|
|
862
|
+
|
|
863
|
+
assert seed is not None
|
|
864
|
+
|
|
865
|
+
assert isinstance(flip, tuple) or flip in [
|
|
866
|
+
"prefix",
|
|
867
|
+
], f"{flip} is not a tuple. Probably change to ('IO', 'RAND') or equivalent?"
|
|
868
|
+
|
|
869
|
+
if flip == "prefix":
|
|
870
|
+
flipped_prompts = flip_prefixes(self.ioi_prompts)
|
|
871
|
+
else:
|
|
872
|
+
if flip in [("IO", "S1"), ("S", "IO")]:
|
|
873
|
+
flipped_prompts = gen_flipped_prompts(
|
|
874
|
+
self.ioi_prompts,
|
|
875
|
+
None,
|
|
876
|
+
flip,
|
|
877
|
+
seed=(seed + 12345) % 9876,
|
|
878
|
+
)
|
|
879
|
+
elif flip == ("S2", "IO"):
|
|
880
|
+
flipped_prompts = gen_flipped_prompts(
|
|
881
|
+
self.ioi_prompts,
|
|
882
|
+
None,
|
|
883
|
+
flip,
|
|
884
|
+
seed=(seed + 12345) % 6543,
|
|
885
|
+
)
|
|
886
|
+
|
|
887
|
+
else:
|
|
888
|
+
assert flip[1] == "RAND" and flip[0] in [
|
|
889
|
+
"S",
|
|
890
|
+
"RAND",
|
|
891
|
+
"S2",
|
|
892
|
+
"IO",
|
|
893
|
+
"S1",
|
|
894
|
+
"S+1",
|
|
895
|
+
], flip
|
|
896
|
+
flipped_prompts = gen_flipped_prompts(
|
|
897
|
+
self.ioi_prompts, NAMES, flip, seed=(seed + 345467) % 5432
|
|
898
|
+
)
|
|
899
|
+
|
|
900
|
+
flipped_ioi_dataset = IOIDataset(
|
|
901
|
+
prompt_type=self.prompt_type,
|
|
902
|
+
N=self.N,
|
|
903
|
+
model=self.model,
|
|
904
|
+
prompts=flipped_prompts,
|
|
905
|
+
prefixes=self.prefixes,
|
|
906
|
+
ioi_prompts_for_word_idxs=flipped_prompts if flip[0] == "RAND" else None,
|
|
907
|
+
prepend_bos=self.prepend_bos,
|
|
908
|
+
manual_word_idx=self.word_idx,
|
|
909
|
+
seed=(seed + 23456) % 963,
|
|
910
|
+
)
|
|
911
|
+
return flipped_ioi_dataset
|
|
912
|
+
|
|
913
|
+
def copy(self):
|
|
914
|
+
copy_ioi_dataset = IOIDataset(
|
|
915
|
+
prompt_type=self.prompt_type,
|
|
916
|
+
N=self.N,
|
|
917
|
+
model=self.model,
|
|
918
|
+
prompts=self.ioi_prompts.copy(),
|
|
919
|
+
prefixes=self.prefixes.copy() if self.prefixes is not None else self.prefixes,
|
|
920
|
+
ioi_prompts_for_word_idxs=self.ioi_prompts.copy(),
|
|
921
|
+
)
|
|
922
|
+
return copy_ioi_dataset
|
|
923
|
+
|
|
924
|
+
def __getitem__(self, key):
|
|
925
|
+
sliced_prompts = self.ioi_prompts[key]
|
|
926
|
+
sliced_dataset = IOIDataset(
|
|
927
|
+
prompt_type=self.prompt_type,
|
|
928
|
+
N=len(sliced_prompts),
|
|
929
|
+
model=self.model,
|
|
930
|
+
prompts=sliced_prompts,
|
|
931
|
+
prefixes=self.prefixes,
|
|
932
|
+
prepend_bos=self.prepend_bos,
|
|
933
|
+
)
|
|
934
|
+
return sliced_dataset
|
|
935
|
+
|
|
936
|
+
def __setitem__(self, key, value):
|
|
937
|
+
raise TypeError(
|
|
938
|
+
"IOIDataset is immutable. To create a modified dataset, "
|
|
939
|
+
"construct a new IOIDataset() with the desired parameters."
|
|
940
|
+
)
|
|
941
|
+
|
|
942
|
+
def __delitem__(self, key):
|
|
943
|
+
raise TypeError(
|
|
944
|
+
"IOIDataset is immutable. To create a modified dataset, "
|
|
945
|
+
"construct a new IOIDataset() with the desired parameters."
|
|
946
|
+
)
|
|
947
|
+
|
|
948
|
+
def __len__(self):
|
|
949
|
+
return self.N
|
|
950
|
+
|
|
951
|
+
def tokenized_prompts(self):
|
|
952
|
+
return self.toks
|
|
953
|
+
|
|
954
|
+
# tests that the templates work as intended
|
|
955
|
+
# assert len(BABA_EARLY_IOS) == len(BABA_LATE_IOS), (len(BABA_EARLY_IOS), len(BABA_LATE_IOS))
|
|
956
|
+
# for i in range(len(BABA_EARLY_IOS)):
|
|
957
|
+
# d1 = IOIDataset(N=1, prompt_type=BABA_EARLY_IOS[i:i+1])
|
|
958
|
+
# d2 = IOIDataset(N=1, prompt_type=BABA_LATE_IOS[i:i+1])
|
|
959
|
+
# for tok in ["IO", "S"]: # occur one earlier and one later
|
|
960
|
+
# assert d1.word_idx[tok] + 1 == d2.word_idx[tok], (d1.word_idx[tok], d2.word_idx[tok])
|
|
961
|
+
# for tok in ["S2"]:
|
|
962
|
+
# assert d1.word_idx[tok] == d2.word_idx[tok], (d1.word_idx[tok], d2.word_idx[tok])
|