Second-Me/lpm_kernel/api/domains/trainprocess/training_params_manager.py
JimmyZQX f5bb0dad59
Data Filtering with Gemma (#396)
* Add code for data filtering llm judge

* Ignore log file created on root (mainly for synthetic_data_generation.log)

* Fix metadata API compatibility issues by commenting out metadata tags in LLM API calls

- Commented out metadata.tags parameters in all LLM API calls across the codebase
- This fixes compatibility issues with custom LLM providers that don't support metadata
- Affects shades generation, topics generation, wiki generation, bio QA, and question generation
- Preserves the original code structure for future re-enabling if needed

* feat: add data filtering pipeline with Ollama integration

- Add MergedDataJudge class for intelligent data filtering using Ollama Gemma
- Integrate automatic Ollama CLI installation into project setup process
- Add DATA_FILTERING step to training pipeline with concurrent processing
- Include testing for MergedDataJudge in its local main() function
- Add Ollama dependency to pyproject.toml

* feat: add automatic Ollama model cleanup after data filtering

* Add logging for outputting data filtering parameters

* fix: adjust error handling for MergedDataJudge:
- Keep original merged.json unchanged when any error occurs
- Exit filtering process immediately on errors instead of continuing with defaults
- Ensure training pipeline continues safely even if data filtering fails

* Add frontend for data filtering pipeline

* resolve data filtering quality_level error by commenting out problematic fields, change TrainProcessService back to original class definition

* fix: quote unquoted shade icons to prevent JSON parsing errors

* Fixed wiki_res.json missing due to no database connection at wiki/base.py module import

* Added scoring reasoning as part of the merged data

* fix: filter ANSI escape sequences from Ollama logs in data filtering step

* fix: Add data filtering steps to cloud training to resolve KeyError

- Added 'Data Filtering' step to cloud training progress holder
- Added data filtering step execution in cloud training service
- Added data filtering parameters to cloud training routes
- Updated frontend to send data filtering parameters
- Fixed missing except clause in cloud training service

This resolves the KeyError: 'data_filtering' when switching from cloud to local training.
2025-08-15 11:19:12 +08:00

109 lines
No EOL
3.9 KiB
Python

"""
Training parameters management module.
This module provides functions for managing and accessing training parameters.
"""
import logging
import json
import os
# Configure logger
logger = logging.getLogger(__name__)
class TrainingParamsManager:
"""
Training parameters manager class.
"""
# Default training parameters
_default_training_params = {
"model_name": "Qwen3-0.6B",
"learning_rate": 1e-4,
"number_of_epochs": 3,
"concurrency_threads": 2,
"data_synthesis_mode": "high",
"use_cuda": False, # Default to using CUDA when available
"is_cot": True,
"language": "chinese",
"data_filtering_model": "gemma:2b",
"data_filtering_workers": 5,
"data_filtering_keep_ratio": 0.8
}
_params_file_path = None
@classmethod
def _get_params_file_path(cls):
"""
Get the training parameters file path
"""
if cls._params_file_path is None:
# Set the parameters file path
progress_dir = os.path.join(os.getcwd(), "data", "progress")
if not os.path.exists(progress_dir):
os.makedirs(progress_dir)
cls._params_file_path = os.path.join(progress_dir, "training_params.json")
return cls._params_file_path
@classmethod
def update_training_params(cls, params, use_previous_params=True):
"""
Update the latest training parameters and save to file
Args:
params: Dictionary containing training parameters
use_previous_params: Whether to use previous training parameters as base
"""
# First try to load existing parameters
current_params = cls.get_latest_training_params() if use_previous_params else cls._default_training_params.copy()
# Update parameters
for key, value in params.items():
if key in cls._default_training_params:
current_params[key] = value
logger.debug(f"Updated training parameter {key} to {value}")
else:
logger.warning(f"Ignoring unknown parameter: {key}")
# Save to file
params_file = cls._get_params_file_path()
try:
with open(params_file, 'w', encoding='utf-8') as f:
json.dump(current_params, f, indent=2)
logger.info(f"Training parameters saved to {params_file}")
except Exception as e:
logger.error(f"Failed to save training parameters to file: {str(e)}", exc_info=True)
@classmethod
def get_latest_training_params(cls):
"""
Get the latest training parameters from file
Returns:
dict: Dictionary containing the latest training parameters
"""
params_file = cls._get_params_file_path()
# If file exists, read from file
if os.path.exists(params_file):
try:
with open(params_file, 'r', encoding='utf-8') as f:
params = json.load(f)
# Replace null values with default values
default_params = cls._default_training_params.copy()
for key, value in default_params.items():
if key not in params or params[key] is None:
params[key] = value
logger.debug(f"Loaded training parameters from {params_file}")
return params
except Exception as e:
logger.error(f"Failed to load training parameters from file: {str(e)}", exc_info=True)
# If reading fails, return default parameters
return cls._default_training_params.copy()
else:
# If file does not exist, return default parameters
return cls._default_training_params.copy()