Tutorial
Leveraging LLMs to boost the performance of classic ML models
Enhance ML model performance by reducing overfitting risks, preserving data quality, and expanding the capabilities of classic ML applicationsIn classic machine learning (ML), data quality and diversity are essential for building effective, generalizable models. Limited or imbalanced data sets can lead to overfitting, underperformance on minority classes, and biased predictions. While synthetic data generation techniques like oversampling (or upsampling) help to a degree, they often result in redundant or noisy samples, which can reduce data quality and increase the risk of overfitting.
Large language models (LLMs), such as OpenAI’s GPT models or IBM's Granite models, offer a new approach by generating unique, contextually rich data that maintains diversity and quality. Unlike traditional synthetic methods, LLMs can produce high-quality, varied data that better reflects real-world patterns, making them powerful tools to augment ML training data effectively.
This tutorial explores how you can use LLMs to generate pre-processed, high-quality data for classic ML models. Through code examples and practical guidance, you'll learn how LLM-driven data generation can enhance ML model performance by reducing overfitting risks, preserving data quality, and expanding the capabilities of classic ML applications.
Let’s start by first understanding the limitations of the existing synthetic data generation approach.
Limitations of traditional synthetic data generation for classic ML
Traditional synthetic data generation methods, such as oversampling techniques like SMOTE (Synthetic Minority Over-sampling Technique), have been widely used to address data scarcity and class imbalances in classical machine learning. While these techniques help expand data sets, they come with significant limitations that can impact model performance and generalizability:
Data duplicity - Oversampling, which duplicates existing data from underrepresented classes, often results in repetitive data that lacks informational diversity. This redundancy fails to provide the model with new patterns to learn from, limiting its ability to generalize effectively. Consequently, models trained on such duplicated data may show misleadingly high performance on training data but fail to adapt well to new, unseen samples.
Quality and relevance - Traditional synthetic data generation, particularly SMOTE, creates new samples by interpolating between existing data points. While this adds some variation, the generated samples often lack the complex, nuanced patterns seen in real-world data. This gap in data realism affects the model’s ability to generalize, as it may learn to recognize overly simplified or interpolated patterns that don’t accurately reflect real-world scenarios. Consequently, the model’s performance on diverse, real-world data can suffer.
Overfitting - Synthetic data with limited diversity can inadvertently increase the risk of overfitting, especially when the data lacks genuine variation. Models trained on such repetitive or overly simplistic data can become highly tuned to the specific patterns in the training set, reducing their capacity to handle novel examples. This overfitting issue often makes synthetic data an inadequate solution for achieving a balanced and realistic data set.
These limitations highlight the need for more advanced data augmentation methods that can generate realistic, high-quality data without duplicity, and with rich contextual relevance, which is an area where large language models (LLMs) can provide a powerful solution.
Let’s understand how that can be done with using langchain, watsonx.ai and Watson Studio.
Prerequisites
- Python 3.9+
- An IBM Cloud account.
Steps
Step 1. Set up your environment
First, you need to set up a watson.ai project to be able to access the watsonx LLMs for data generation.
- Log in to watsonx.ai using your IBM Cloud account
- Create a watsonx.ai project.
- Create a Watson Machine Learning service instance (choose the Lite plan, which is a free instance).
- Generate an API Key in WML.
- Associate the WML service to the project you created in watsonx.ai.
Now, you are finally ready to rock!
Simply import the necessary libraries and use the watsonX API key.
import watson_nlp
import pandas as pd
from tqdm import tqdm
from langchain_ibm import WatsonxLLM
import numpy as np
import multiprocessing
import time
import os
import re
from watson_nlp.blocks.classification.transformer import Transformer
from watson_core.data_model.streams.resolver import DataStreamResolver
os.environ["WATSONX_APIKEY"] = 'API-KEY'
Step 2. Train a model on the generated data set
We define a function to train a model on the generated data set using a pre-trained model from Watson NLP. The function train_model sets batch size, epochs, and training data, and saves the trained model.
def train_model(training_data_file, model_name, target_column):
batch_size = 64
epochs = 50
data_stream_resolver = DataStreamResolver(target_stream_type=list, expected_keys={'TEXT': str, target_column: str})
train_stream = data_stream_resolver.as_data_stream(training_data_file)
pretrained_model_resource = watson_nlp.load('pretrained-model_slate.153m.distilled_many_transformer_multilingual_uncased')
model = Transformer.train(train_stream, pretrained_model_resource, num_train_epochs=epochs, train_batch_size=batch_size, verbose=2, learning_rate=5e-5)
model.save(model_name)
return model
Step 3. Clean the data
Create a utility function, clean_text, to remove unwanted characters and spaces from the text data, making it easier for the model to process the input text.
def clean_text(text):
text = re.sub(r'[^a-zA-Z0-9]', ' ', text)
text = re.sub(r'\s+', ' ', text)
return text
Step 4. Generate the data
The function generate_response interacts with the watsonx LLM by sending prompts, which outputs paraphrased text samples based on user complaints, mimicking realistic data without duplicity.
def generate_response(prompt):
parameters = {
"decoding_method": "sample",
"min_new_tokens": 1,
"max_new_tokens": 4096,
"stop_sequences": [],
"repetition_penalty": 1,
'temperature': 0
}
watsonx_llm = WatsonxLLM(
model_id="meta-llama/llama-3-1-70b-instruct",
url="https://us-south.ml.cloud.ibm.com",
project_id='PROJECT-ID',
params=parameters
)
response = watsonx_llm.invoke(prompt)
return response
Step 5. Oversampling data using LLMs
The oversample_data function structures a prompt that generates multiple variations of text data to create unique samples for oversampling.
def oversample_data(text):
prompt = f'''
<|begin_of_text|><|start_header_id|>system<|end_header_id|>
You are an expert english paraphraser. Generate new sentences by adhering to the INSTRUCTIONS shared below.
ABOUT TEXT:
1. The text comprises of 8 complaints each represented by ">"
INSTRUCTIONS:
1. Respond by generating 8 new complaints from the user complaints mentioned in the text at all cost!
2. Generate 8 new complaints by paraphrasing or by changing the way it is written in the TEXT section
3. Don't include any irrelevant information. Generate only relevant text.
4. Don't include any header or footer in your response at any cost!
5. Each new complaint should be represented by ">" by all means!
6. Maintain the same format/style of the text that has been written in the TEXT section
7. The new complaints shouldn't be similar to the original complaints at all - I need new ones!
Expected Input:
>Complaint 1
>Complaint 2
>Complaint 3
>Complaint 4
>Complaint 5
>Complaint 6
>Complaint 7
>Complaint 8
Expected Output:
>New Complaint 1
>New Complaint 2
>New Complaint 3
>New Complaint 4
>New Complaint 5
>New Complaint 6
>New Complaint 7
>New Complaint 8
TEXT:
{text}
<|eot_id|><|start_header_id|>user<|end_header_id|>
'''
response = generate_response(prompt)
return response
Step 6. Generate new samples
The function generate_new_samples uses the oversample_data function to create new samples, specifically for underrepresented classes in the data set.
def generate_new_samples(df, target_column, class_name):
final_list = []
original_df = df[df[target_column]==class_name]
new_df = df[df[target_column]==class_name].sample(frac = 1).reset_index()
new_complaints = []
for i in range(0, len(new_df), 8):
try:
text = '>'
text += '\n>'.join(new_df['TEXT'].iloc[i:i+8].to_list())
response = oversample_data(text)
idx = response.index('>')
response = response[idx:]
response = response.split('>')
new_complaints.extend([resp for resp in response if resp not in ('\n', ' ', '')])
except Exception as e:
print('Printing Error: ', e)
pass
for complaint in new_complaints:
if 'New Complaint' not in complaint:
final_list.append(complaint)
new_class_list = [class_name]*len(final_list)
text_list = new_df.TEXT.to_list()
class_list = new_df[target_column].to_list()
text_list += final_list
class_list += new_class_list
new_df = pd.DataFrame()
new_df['TEXT'] = text_list
new_df[target_column] = class_list
new_df = pd.concat([original_df, new_df])
return new_df
Step 7. Accelerate data generation
The worker function and multiprocessing code apply parallel processing to accelerate data generation across multiple classes.
def worker(class_name):
return generate_new_samples(train_set,target_column, class_name)
# Input files
files = ['trainset.csv', 'testset.csv']
target_column = 'YOUR_TARGET_COLUMN'
# Loading and cleaning data
df = pd.DataFrame()
for file in files:
wslib.download_file(file)
df_train = pd.read_csv(files[0])
train_set = df_train[df_train[target_column]!=' '].reset_index(drop=True)
train_set.TEXT = [text.lower() for text in train_set.TEXT]
list_of_cleaned_text = []
for text in tqdm(train_set.TEXT):
cleaned_text = clean_text(text)
list_of_cleaned_text.append(cleaned_text)
train_set.TEXT = list_of_cleaned_text
class_distribution = dict(train_set[target_column].value_counts())
median_value = np.median(np.array(list(class_distribution.values())))
classes_less_than_median = [class_name for class_name in class_distribution if class_distribution[class_name]<median_value]
start = time.time()
with multiprocessing.Pool(processes=20) as pool:
with tqdm(total=len(classes_less_than_median), desc="Processing classes") as pbar:
results = pool.imap_unordered(worker, classes_less_than_median)
oversampled_df = pd.concat(results, ignore_index=True)
pbar.update(1)
end = time.time()
print('Total Time taken in mins: ', round((end-start)/60, 2))
Step 8. Save and train the model with the new data set
This code saves the oversampled data set, trains the model, and then evaluates it on a test data set by calculating Top-1, Top-3, and Top-5 accuracies.
oversampled_df = oversampled_df.reset_index(drop=True)
oversampled_df = pd.concat([train_set[~train_set[target_column].isin(classes_less_than_median)], oversampled_df]).reset_index(drop=True)
oversampled_df = oversampled_df[['TEXT', target_column]]
oversampled_df.to_csv('oversampled_data.csv', index=False)
wslib.upload_file('oversampled_data.csv', overwrite=True)
model_name = 'model_name'
model = train_model('oversampled_data.csv', model_name, target_column)
# Test and Evaluate
df_test = pd.read_csv(files[1])
results = []
top_5_sum = 0
top_3_sum = 0
top_1_sum = 0
for i in tqdm(range(len(df_test.index))):
result = model.run(df_test.iloc[i]["TEXT"]).to_dict()
result_dict = {
"Top1": result["classes"][0]["class_name"],
"Top2": result["classes"][1]["class_name"],
"Top3": result["classes"][2]["class_name"],
"Top4": result["classes"][3]["class_name"],
"Top5": result["classes"][4]["class_name"]
}
results.append(result_dict)
if result["classes"][0]["class_name"] == df_test.iloc[i][target_column]:
top_1_sum += 1
top_3_sum += 1
top_5_sum += 1
elif result["classes"][1]["class_name"] == df_test.iloc[i][target_column]:
top_3_sum += 1
top_5_sum += 1
elif result["classes"][2]["class_name"] == df_test.iloc[i][target_column]:
top_3_sum += 1
top_5_sum += 1
elif result["classes"][3]["class_name"] == df_test.iloc[i][target_column]:
top_5_sum += 1
elif result["classes"][4]["class_name"] == df_test.iloc[i][target_column]:
top_5_sum += 1
print("Top1 Accuracy: {0}".format(top_1_sum/len(df_test.index)))
print("Top3 Accuracy: {0}".format(top_3_sum/len(df_test.index)))
print("Top5 Accuracy: {0}".format(top_5_sum/len(df_test.index)))
Addressing challenges and practical considerations
While LLMs can significantly enhance data generation for classic ML models, implementing them effectively requires addressing several key challenges:
Data privacy: When generating synthetic data, it’s essential to ensure that no sensitive or personal information is inadvertently produced. Using generalized, anonymized input data and validating outputs can help prevent any privacy breaches. Additionally, carefully monitoring and controlling the prompt inputs will mitigate the risk of any sensitive information making its way into the generated text.
Bias mitigation: LLMs can sometimes amplify biases present in training data or prompt phrasing. To avoid biased data generation, it’s crucial to design prompts that encourage neutrality and diversity and to validate outputs rigorously. This may involve additional data preprocessing or prompt-engineering strategies that encourage a wide range of perspectives within generated data, fostering more balanced model training.
Cost and computational resources: Generating and training on large data sets with LLMs can be computationally demanding, especially for small or resource-limited teams. Careful planning, such as optimizing prompt length and setting efficient generation parameters, helps reduce computational overhead. Balancing generation quality with the resources required can ensure efficient, cost-effective use of LLMs without sacrificing performance.
Summary and next steps
Incorporating LLM-based data generation methods into classic ML pipelines provides a powerful alternative to traditional synthetic data approaches, enhancing model robustness, accuracy, and generalizability. By leveraging LLMs to generate varied, high-quality training samples, classic ML models can overcome data limitations and avoid common issues such as overfitting or insufficient data diversity.
The synergy between LLMs and classic ML techniques represents a promising area for experimentation and innovation. By fine-tuning LLM prompts and exploring varied data generation methods, practitioners can continue to improve training workflows, ultimately driving better outcomes across diverse applications. As LLMs evolve, their potential as data augmentation tools only grows, opening up new possibilities for classical ML model performance in increasingly complex and dynamic environments.