add negative_nq + visualize no_context base model

This commit is contained in:
51616 2025-05-20 14:37:56 +00:00
parent 19ca0f99a5
commit e692df98e9
8 changed files with 283 additions and 46 deletions

View file

@ -116,51 +116,134 @@ def get_run_data(run_path):
return grouped_data
def get_generated_text_data(run_path):
def get_generated_text_data(run_path: str) -> dict:
"""
Loads generated text data from all *_generated_text.jsonl files.
Loads generated text data from all *_generated_text.jsonl files from the current run's output.
It specifically excludes any files ending with _no_context_generated_text.jsonl,
as those are handled by get_base_model_generated_text.
Also recursively searches for generated text files in subfolders.
Args:
run_path: Path to the directory containing the .jsonl files.
run_path: Path to the directory containing the .jsonl files for the current run.
Returns:
A dictionary containing the generated text data for each split found.
A dictionary containing the generated text data for each primary split found.
"""
generated_data = {}
files = sorted(glob(f"{run_path}/*_generated_text.jsonl"))
print(files)
# Find all *_generated_text.jsonl files in the current directory
# Find all *_generated_text.jsonl files in the current directory, excluding _no_context_ variants
files = sorted(
[
f
for f in glob(os.path.join(run_path, "*_generated_text.jsonl"))
if "_no_context_generated_text.jsonl" not in os.path.basename(f)
]
)
print(f"Processing generated text files from {run_path}: {files}")
for filename in files:
# Extract split name from filename (remove _generated_text.jsonl)
split = filename.split("/")[-1].split("_generated_text.jsonl")[0]
split = os.path.basename(filename).replace("_generated_text.jsonl", "")
try:
with open(filename) as f:
lines = f.readlines()
lines = f.readlines(100_000)
data = [json.loads(line) for line in lines]
generated_data[split] = data
except FileNotFoundError:
generated_data[split] = None
except json.JSONDecodeError as e:
print(f"Error decoding JSON from {filename}: {e}")
generated_data[split] = None
# Recursively search for *_generated_text.jsonl files in subdirectories
for subdir in [d for d in glob(os.path.join(run_path, "*")) if os.path.isdir(d)]:
subdir_name = os.path.basename(subdir)
subdir_files = sorted(glob(f"{subdir}/*_generated_text.jsonl"))
for subdir_path in [
d for d in glob(os.path.join(run_path, "*")) if os.path.isdir(d)
]:
subdir_name = os.path.basename(subdir_path)
subdir_files = sorted(
[
f
for f in glob(os.path.join(subdir_path, "*_generated_text.jsonl"))
if "_no_context_generated_text.jsonl" not in os.path.basename(f)
]
)
print(
f"Processing generated text files from subdirectory {subdir_path}: {subdir_files}"
)
for filename in subdir_files:
# Extract split name from filename and include subfolder name
base_split = filename.split("/")[-1].split("_generated_text.jsonl")[0]
base_split = os.path.basename(filename).replace("_generated_text.jsonl", "")
# Create a unique split key including the subdirectory name
split = f"{subdir_name}_{base_split}"
try:
with open(filename) as f:
lines = f.readlines()
lines = f.readlines(100_000)
data = [json.loads(line) for line in lines]
generated_data[split] = data
except FileNotFoundError:
generated_data[split] = None
except json.JSONDecodeError as e:
print(f"Error decoding JSON from {filename}: {e}")
generated_data[split] = None
return generated_data
def get_base_model_generated_text(model_name, generated_data):
"""
Loads generated text data from the base model for comparison.
Also loads "no context" base model data if available.
Args:
model_name: The name of the base model
generated_data: Dictionary of generated data from the fine-tuned model
Returns:
Tuple: (base_model_data, base_model_no_context_data)
Dictionaries containing base model generated texts for matching splits
"""
if not model_name:
return {}, {}
base_model_data = {}
base_model_no_context_data = {}
# # Create normalized model name for directory lookup
# normalized_model_name = model_name.replace("/", "_")
for split in generated_data:
# Skip if no data for this split
if not generated_data[split]:
base_model_data[split] = None
base_model_no_context_data[split] = None
continue
# Construct path to the base model's output for this split
base_model_path = f"eval_results/{model_name}/{split}_generated_text.jsonl"
base_model_no_context_path = (
f"eval_results/{model_name}/{split}_no_context_generated_text.jsonl"
)
try:
with open(base_model_path) as f:
lines = f.readlines(100_000)
data = [json.loads(line) for line in lines]
base_model_data[split] = data
except FileNotFoundError:
print(f"Base model output not found for {split} at {base_model_path}")
# Store None to indicate we tried but didn't find matching data
base_model_data[split] = None
try:
with open(base_model_no_context_path) as f:
lines = f.readlines(100_000)
data = [json.loads(line) for line in lines]
base_model_no_context_data[split] = data
except FileNotFoundError:
print(
f"Base model (no context) output not found for {split} at {base_model_no_context_path}"
)
base_model_no_context_data[split] = None
return base_model_data, base_model_no_context_data
def get_available_checkpoints(run):
"""
Finds all checkpoint directories in a run folder.
@ -292,6 +375,31 @@ def visualize(run):
except FileNotFoundError:
model_name = None
# If config.yaml doesn't have the model name, try args.yaml
if not model_name:
args_path = os.path.join(logdir, "args.yaml")
try:
with open(args_path) as f:
yaml.add_constructor("!", lambda loader, node: None)
yaml.add_multi_constructor(
"tag:yaml.org,2002:python/object",
lambda loader, suffix, node: None,
Loader=yaml.SafeLoader,
)
args = yaml.safe_load(f)
model_name = args.get("model_name_or_path")
except FileNotFoundError:
model_name = None
# Get base model outputs if we have the model name and modulated outputs
base_model_data = {}
base_model_no_context_data = {}
print(generated_data.keys())
if model_name and generated_data:
base_model_data, base_model_no_context_data = get_base_model_generated_text(
model_name, generated_data
)
return render_template(
"visualize.html",
run=run,
@ -299,6 +407,8 @@ def visualize(run):
selected_eval_folder=selected_eval_folder,
data=data,
generated_data=generated_data,
base_model_data=base_model_data,
base_model_no_context_data=base_model_no_context_data,
model_name=model_name,
checkpoints=checkpoint_names,
)

View file

@ -452,6 +452,43 @@
background-color: #bbdefb;
cursor: not-allowed;
}
.model-comparison {
display: flex;
flex-wrap: wrap;
gap: 20px;
margin-bottom: 15px;
}
.model-output {
flex: 1;
min-width: 300px;
border: 1px solid #e0e0e0;
border-radius: 5px;
padding: 10px;
background-color: #f9f9f9;
}
.model-output p {
margin: 0;
}
.model-output span {
display: block;
padding: 10px;
background-color: white;
border: 1px solid #eee;
border-radius: 4px;
min-height: 80px;
max-height: 150px;
overflow-y: auto;
white-space: pre-wrap;
margin-top: 5px;
}
.highlight-diff {
background-color: #fff8e1;
}
</style>
{% endblock %}
@ -536,7 +573,22 @@
<strong>Context:</strong> <span id="{{ split }}-context"></span>
</p>
<p><strong>Input:</strong> <span id="{{ split }}-input"></span></p>
<p><strong>Generated:</strong> <span id="{{ split }}-generated"></span></p>
<!-- Display both models with enhanced styling -->
<div class="model-comparison">
<div class="model-output">
<p><strong>Modulated Model:</strong> <span id="{{ split }}-generated"></span></p>
</div>
<div class="model-output">
<p><strong>Base Model ({{ model_name }}):</strong> <span id="{{ split }}-base-generated">Not
available</span></p>
</div>
<div class="model-output">
<p><strong>Base Model (No Context - {{ model_name }}):</strong> <span
id="{{ split }}-base-no-context-generated">Not available</span></p>
</div>
</div>
<p><strong>Label:</strong> <span id="{{ split }}-label"></span></p>
</div>
<div class="generated-text-controls">
@ -632,10 +684,13 @@
<script>
var generatedData = {{ generated_data | tojson }};
var baseModelData = {{ base_model_data | tojson }};
var baseModelNoContextData = {{ base_model_no_context_data | tojson }};
function updateText(split) {
var index = parseInt(document.getElementById(split + '-index').value) - 1;
var data = generatedData[split][index];
// Check if 'context' exists in the data and update it
if ('context' in data) {
document.getElementById(split + '-context').textContent = data.context;
@ -643,9 +698,26 @@
} else {
document.getElementById(split + '-context-container').style.display = 'none';
}
document.getElementById(split + '-input').textContent = data.input;
document.getElementById(split + '-generated').textContent = data.generated;
document.getElementById(split + '-label').textContent = data.label;
// Update base model output if available
var baseModelOutput = document.getElementById(split + '-base-generated');
if (baseModelData && baseModelData[split] && baseModelData[split].length > index) {
baseModelOutput.textContent = baseModelData[split][index].generated;
} else {
baseModelOutput.textContent = "Not available";
}
// Update base model (no context) output if available
var baseModelNoContextOutput = document.getElementById(split + '-base-no-context-generated');
if (baseModelNoContextData && baseModelNoContextData[split] && baseModelNoContextData[split].length > index) {
baseModelNoContextOutput.textContent = baseModelNoContextData[split][index].generated;
} else {
baseModelNoContextOutput.textContent = "Not available";
}
}
function next(split) {