USER
Das ist der Java-Code für einen Chatbot mit GPT und neuronalen Netz, das lernen soll auf meine Fragen zu antworten. Der Bot antwortet immer mit demselben Wort, doch nach jedem training ändert sich das Wort wieder. Vielleicht wäre es eine Lösung, eine Temperatur vor oder nach dem Softmax einzubauen: package com.mycompany.simplechatbot;
import java.io.BufferedReader;
import java.io.FileReader;
import java.io.IOException;
import java.util.*;
import java.util.regex.Matcher;
import java.util.regex.Pattern;
import java.util.stream.Collectors;
/**
* Ein einfacher Chatbot mit Wortvektoren, einem neuronalen Netz und einem einfachen Attention-Mechanismus.
*/
public class SimpleChatbot {
/**
* Repräsentiert ein Wort und seinen zugehörigen Vektor.
*/
static class WordVector {
private String word;
private double[] vector;
public WordVector(String word, double[] vector) {
this.word = word;
this.vector = vector;
}
public String getWord() {
return word;
}
public double[] getVector() {
return vector;
}
}
// Eine Liste deutscher Stopwörter (angepasst)
static final Set<String> stopwords = new HashSet<>(Arrays.asList(
"a", "an", "the"
));
/**
* Lädt und verwaltet Wortvektoren aus einer Datei.
*/
static class WordVectorModel {
private Map<String, WordVector> wordVectors;
private Map<String, Integer> wordToIndex;
private Map<Integer, String> indexToWord;
private int vectorSize = -1;
public WordVectorModel(String filePath) throws IOException {
wordVectors = new HashMap<>();
wordToIndex = new HashMap<>();
indexToWord = new HashMap<>();
loadWordVectors(filePath);
}
private void loadWordVectors(String filePath) throws IOException {
String line;
int index = 0;
try (BufferedReader reader = new BufferedReader(new FileReader(filePath))) {
while ((line = reader.readLine()) != null) {
String[] parts = line.split(" ");
if (parts.length < 2) continue; // Sicherstellen, dass mindestens Wort und eine Dimension vorhanden sind
String word = parts[0];
double[] vector = new double[parts.length - 1];
for (int i = 0; i < parts.length - 1; i++) {
try {
vector[i] = Double.parseDouble(parts[i + 1]);
} catch (NumberFormatException e) {
System.out.println("Fehler beim Parsen der Zahl für Wort: " + word + ". Überspringe dieses Wort.");
vector = null;
break;
}
}
if(vector == null) continue;
if (!stopwords.contains(word)) {
// Überprüfung der Vektorlänge
if(vectorSize == -1) {
vectorSize = vector.length;
} else if(vector.length != vectorSize) {
System.out.println("Vektorgröße inkonsistent für Wort: " + word + ". Erwartet: " + vectorSize + ", Gefunden: " + vector.length);
continue; // Überspringe Vektoren mit inkonsistenter Größe
}
wordVectors.put(word.toLowerCase(), new WordVector(word, vector));
wordToIndex.put(word.toLowerCase(), index);
indexToWord.put(index, word.toLowerCase());
index++;
}
}
System.out.println("Vektoren erfolgreich aus der Datei geladen: " + filePath);
System.out.println("Anzahl Wörter im Model: " + wordVectors.size());
System.out.println("Dimension der Vektoren: " + vectorSize);
// Optional: Ausgabe einiger Vektoren zur Überprüfung
int count = 0;
for (Map.Entry<String, WordVector> entry : wordVectors.entrySet()) {
String word = entry.getKey();
double[] vector = entry.getValue().getVector();
System.out.print("Word: " + word + ", Vector: [");
for (int i = 0; i < Math.min(vector.length, 5); i++) { // Begrenze die Ausgabe für Lesbarkeit
System.out.print(vector[i]);
if (i < vector.length - 1 && i < 4) { // Anpassung der Bedingung zur Begrenzung auf 5 Elemente
System.out.print(", ");
}
}
System.out.println("...]");
count++;
if (count >= 1) break; // Entferne dies, um mehr Vektoren zu sehen
}
} catch (IOException e) {
System.out.println("Fehler beim Laden der Vektoren.");
e.printStackTrace();
throw e; // Weiterwerfen der Ausnahme nach der Protokollierung
}
}
public WordVector getWordVector(String word) {
return wordVectors.get(word.toLowerCase());
}
public int getVectorSize() {
return vectorSize;
}
public Collection<WordVector> getAllWordVectors() {
return wordVectors.values();
}
public Map<String, Integer> getWordToIndexMap() {
return wordToIndex;
}
public Map<Integer, String> getIndexToWordMap() {
return indexToWord;
}
public int getVocabSize() {
return wordVectors.size();
}
}
/**
* Stellt Hilfsfunktionen für Vektoroperationen bereit.
*/
static class VectorUtils {
public static double[] add(double[] a, double[] b) {
double[] result = new double[a.length];
for(int i=0; i<a.length; i++) {
result[i] = a[i] + b[i];
}
return result;
}
public static double[] subtract(double[] a, double[] b) {
double[] result = new double[a.length];
for(int i=0; i<a.length; i++) {
result[i] = a[i] - b[i];
}
return result;
}
public static double[] multiply(double scalar, double[] a) {
double[] result = new double[a.length];
for(int i=0; i<a.length; i++) {
result[i] = scalar * a[i];
}
return result;
}
public static double dotProduct(double[] a, double[] b) {
double sum = 0.0;
for(int i=0; i<a.length; i++) {
sum += a[i] * b[i];
}
return sum;
}
public static double norm(double[] a) {
return Math.sqrt(dotProduct(a, a));
}
public static double cosineSimilarity(double[] a, double[] b) {
double dot = dotProduct(a, b);
double normA = norm(a);
double normB = norm(b);
if(normA == 0 || normB == 0) return 0.0;
return dot / (normA * normB);
}
public static double[] average(List<double[]> vectors) {
if(vectors == null || vectors.isEmpty()) return null;
int size = vectors.get(0).length;
double[] avg = new double[size];
for(double[] vec : vectors) {
for(int i=0; i<size; i++) {
avg[i] += vec[i];
}
}
for(int i=0; i<size; i++) {
avg[i] /= vectors.size();
}
return avg;
}
public static double[] normalize(double[] a) {
double n = norm(a);
if(n == 0) return a;
return multiply(1.0 / n, a);
}
}
/**
* Ein einfaches neuronales Netzwerk mit einer versteckten Schicht.
*/
static class NeuralNetwork {
private int inputSize;
private int hiddenSize;
private int outputSize;
private double[][] weightsInputHidden;
private double[][] weightsHiddenOutput;
private double learningRate = 0.001; // Reduzierte Lernrate für stabiles Training
private Random rand = new Random();
public NeuralNetwork(int inputSize, int hiddenSize, int outputSize) {
this.inputSize = inputSize;
this.hiddenSize = hiddenSize;
this.outputSize = outputSize;
weightsInputHidden = new double[inputSize][hiddenSize];
weightsHiddenOutput = new double[hiddenSize][outputSize];
initializeWeights(weightsInputHidden);
initializeWeights(weightsHiddenOutput);
}
public int getHiddenSize() {
return hiddenSize;
}
private void initializeWeights(double[][] weights) {
for(int i=0; i<weights.length; i++) {
for(int j=0; j<weights[i].length; j++) {
weights[i][j] = rand.nextGaussian() * 0.01;
}
}
}
/**
* Führt einen Forward-Pass durch das Netzwerk durch und gibt die aktivierten Output-Werte zurück.
*
* @param input Der Eingabevektor.
* @return Der aktivierte Ausgabevektor nach Softmax.
*/
public double[] forward(double[] input) {
// Berechne Hidden Layer Aktivierungen
double[] hiddenOutput = new double[hiddenSize];
for(int j=0; j<hiddenSize; j++) {
double sum = 0.0;
for(int i=0; i<inputSize; i++) {
sum += input[i] * weightsInputHidden[i][j];
}
hiddenOutput[j] = relu(sum);
}
// Berechne Output Layer (Logits)
double[] logits = new double[outputSize];
for(int k=0; k<outputSize; k++) {
for(int j=0; j<hiddenSize; j++) {
logits[k] += hiddenOutput[j] * weightsHiddenOutput[j][k];
}
}
// Wende Softmax an, um Wahrscheinlichkeiten zu erhalten
double[] probabilities = softmax(logits);
return probabilities;
}
/**
* Aktivierungsfunktion ReLU.
*
* @param x Eingabewert.
* @return Aktivierungswert.
*/
private double relu(double x) {
return Math.max(0, x);
}
/**
* Trainiert das Netzwerk mit der Kreuzentropie-Verlustfunktion und Backpropagation.
*
* @param input Der Eingabevektor.
* @param target Der Zielindex des Zielwortes.
*/
public void train(double[] input, int target) {
// Forward pass
double[] hidden = new double[hiddenSize];
for(int j=0; j<hiddenSize; j++) {
double sum = 0.0;
for(int i=0; i<inputSize; i++) {
sum += input[i] * weightsInputHidden[i][j];
}
hidden[j] = relu(sum);
}
// Berechne Output Layer (Logits)
double[] logits = new double[outputSize];
for(int k=0; k<outputSize; k++) {
for(int j=0; j<hiddenSize; j++) {
logits[k] += hidden[j] * weightsHiddenOutput[j][k];
}
}
// Softmax-Ausgabe
double[] probabilities = softmax(logits);
// Ziel-Probabilities (One-Hot)
double[] targetProbabilities = new double[outputSize];
targetProbabilities[target] = 1.0;
// Kreuzentropie-Verlust und Gradienten berechnen
double[] outputErrors = new double[outputSize];
for(int k=0; k<outputSize; k++) {
outputErrors[k] = probabilities[k] - targetProbabilities[k];
}
// Backpropagate errors to hidden layer
double[] hiddenErrors = new double[hiddenSize];
for(int j=0; j<hiddenSize; j++) {
double error = 0.0;
for(int k=0; k<outputSize; k++) {
error += outputErrors[k] * weightsHiddenOutput[j][k];
}
// ReLU-Derivative
hiddenErrors[j] = hidden[j] > 0 ? error : 0.0;
}
// Update weightsHiddenOutput
for(int j=0; j<hiddenSize; j++) {
for(int k=0; k<outputSize; k++) {
weightsHiddenOutput[j][k] -= learningRate * outputErrors[k] * hidden[j];
}
}
// Update weightsInputHidden
for(int i=0; i<inputSize; i++) {
for(int j=0; j<hiddenSize; j++) {
weightsInputHidden[i][j] -= learningRate * hiddenErrors[j] * input[i];
}
}
}
/**
* Wendet die Softmax-Funktion auf ein Array von Logits an.
*
* @param logits Die Eingabewerte (Logits).
* @return Die normalisierten Wahrscheinlichkeiten.
*/
private double[] softmax(double[] logits) {
double max = Double.NEGATIVE_INFINITY;
for(double logit : logits) {
if(logit > max) max = logit;
}
double sum = 0.0;
double[] expScores = new double[logits.length];
for(int i=0; i<logits.length; i++) {
expScores[i] = Math.exp(logits[i] - max);
sum += expScores[i];
}
double[] softmax = new double[logits.length];
for(int i=0; i<logits.length; i++) {
softmax[i] = expScores[i] / sum;
}
return softmax;
}
}
/**
* Lädt die Trainingsdaten und speichert Eingabe-Antwort-Paare.
*/
public static class TrainingData {
private List<String[]> inputPairs;
private List<String[]> outputPairs;
private Map<String, Integer> wordFrequencies;
public TrainingData(String trainingFilePath) throws IOException {
inputPairs = new ArrayList<>();
outputPairs = new ArrayList<>();
wordFrequencies = new HashMap<>();
loadTrainingData(trainingFilePath);
}
public Map<String, Integer> getWordFrequencies() {
return wordFrequencies;
}
private void loadTrainingData(String filePath) throws IOException {
Pattern pattern = Pattern.compile("\\b[a-zA-ZäöüÄÖÜß0-9']+\\b");
try (BufferedReader reader = new BufferedReader(new FileReader(filePath))) {
String line;
while ((line = reader.readLine()) != null) {
String[] parts = line.split("\t");
if (parts.length != 2) continue;
String[] inputWords = extractFilteredWords(parts[0], pattern);
String[] outputWords = extractFilteredWords(parts[1], pattern);
if (inputWords.length > 0 && outputWords.length > 0) { // Optionale Überprüfung
inputPairs.add(inputWords);
outputPairs.add(outputWords);
// Frequenzzählung der Zielwörter
for(String word : outputWords) {
wordFrequencies.put(word, wordFrequencies.getOrDefault(word, 0) + 1);
}
}
}
}
System.out.println("Geladene Trainingsbeispiele: " + inputPairs.size());
System.out.println("Input: ");
inputPairs.forEach(input -> System.out.println(String.join(" ", input)));
System.out.println("Output: ");
outputPairs.forEach(output -> System.out.println(String.join(" ", output)));
}
/**
* Extrahiert Wörter aus einem Text, wendet das Regex-Muster an, konvertiert sie zu Kleinbuchstaben
* und filtert dabei Stopwörter heraus.
*
* @param text Der Eingabetext.
* @param pattern Das Regex-Muster zum Extrahieren der Wörter.
* @return Ein Array von gefilterten Wörtern.
*/
private String[] extractFilteredWords(String text, Pattern pattern) {
Matcher matcher = pattern.matcher(text.toLowerCase());
List<String> words = new ArrayList<>();
while (matcher.find()) {
String word = matcher.group();
if (!stopwords.contains(word)) {
words.add(word);
}
}
return words.toArray(new String[0]);
}
public List<String[]> getInputPairs() {
return inputPairs;
}
public List<String[]> getOutputPairs() {
return outputPairs;
}
}
/**
* Integriert alle Komponenten und ermöglicht das Training und die Interaktion.
*/
static class Chatbot {
private WordVectorModel model;
private NeuralNetwork neuralNet;
private Scanner scanner;
private Map<String, Double> samplingProbabilities;
// Trainingsdaten
private TrainingData trainingData;
private final int EPOCHS = 100000; // Reduzierte Epochen für Debugging
// Für Einfachheit definieren wir eine feste Antwortlänge
private final int MAX_RESPONSE_LENGTH = 5;
public Chatbot(String vectorFilePath, String trainingFilePath) throws IOException {
model = new WordVectorModel(vectorFilePath);
int inputSize = model.getVectorSize();
int outputSize = model.getVocabSize(); // Setze outputSize auf die Vokabulargröße
// Einfaches neuronales Netz: Eingabe -> Hidden (133) -> Ausgabe (Vokabulargröße)
neuralNet = new NeuralNetwork(inputSize, 134, outputSize);
scanner = new Scanner(System.in);
trainingData = new TrainingData(trainingFilePath);
computeSamplingProbabilities();
}
/*
* Berechnet die Sampling-Gewichte für jedes Wort, sodass jedes Wort
* im Output gleich oft vorkommt.
*/
private void computeSamplingProbabilities() {
Map<String, Integer> frequencies = trainingData.getWordFrequencies();
samplingProbabilities = new HashMap<>();
// Summe der inversen Frequenzen für Normalisierung
double inverseSum = 0.0;
// Berechne die Inversen der Frequenzen und addiere sie zur Normalisierung
for (Map.Entry<String, Integer> entry : frequencies.entrySet()) {
int frequency = entry.getValue();
double inverseFreq = 1.0 / frequency;
samplingProbabilities.put(entry.getKey(), inverseFreq);
inverseSum += inverseFreq;
}
// Normalisiere die Gewichte, sodass ihre Summe 1 ergibt (damit es Wahrscheinlichkeiten sind)
for (Map.Entry<String, Double> entry : samplingProbabilities.entrySet()) {
samplingProbabilities.put(entry.getKey(), entry.getValue() / inverseSum);
}
}
private void trainNetwork() {
List<String[]> inputs = trainingData.getInputPairs();
List<String[]> outputs = trainingData.getOutputPairs();
Map<String, Integer> wordToIndex = model.getWordToIndexMap();
Random random = new Random();
for(int epoch = 0; epoch < EPOCHS; epoch++) {
double totalLoss = 0.0;
for(int i = 0; i < inputs.size(); i++) {
String[] inputWords = inputs.get(i);
String[] outputWords = outputs.get(i);
// Verarbeitung der Eingabe
List<WordVector> inputVectorsList = new ArrayList<>();
for(String word : inputWords) {
WordVector wv = model.getWordVector(word);
if(wv != null) inputVectorsList.add(wv);
}
if(inputVectorsList.isEmpty()) continue;
double[] contextVector = VectorUtils.average(
inputVectorsList.stream().map(WordVector::getVector).collect(Collectors.toList())
);
// Iteration über jedes Wort in der Ausgabe
for(int k = 0; k < outputWords.length && k < MAX_RESPONSE_LENGTH; k++) {
String targetWord = outputWords[k];
Integer targetIndex = wordToIndex.get(targetWord);
if(targetIndex == null) continue;
// Bestimme, ob das Wort trainiert werden soll basierend auf Sampling-Wahrscheinlichkeit
double prob = samplingProbabilities.getOrDefault(targetWord, 1.0);
if(random.nextDouble() > prob) {
// Wort wird übersprungen
// Update Kontextvektor mit dem Vektor des Zielwortes
WordVector targetWordVector = model.getWordVector(targetWord);
if(targetWordVector == null) continue;
contextVector = VectorUtils.add(contextVector, targetWordVector.getVector());
contextVector = VectorUtils.normalize(contextVector); // Optional: Normalisierung
continue;
}
// Trainiere das neuronale Netzwerk mit contextVector und targetIndex
neuralNet.train(contextVector, targetIndex);
// Forward-Pass zur Verlustberechnung
double[] outputProbabilities = neuralNet.forward(contextVector);
double loss = -Math.log(outputProbabilities[targetIndex] + 1e-9); // Kreuzentropie-Verlust
totalLoss += loss;
// Update Kontextvektor mit dem Vektor des Zielwortes
WordVector targetWordVector = model.getWordVector(targetWord);
if(targetWordVector == null) continue;
contextVector = VectorUtils.add(contextVector, targetWordVector.getVector());
contextVector = VectorUtils.normalize(contextVector); // Optional: Normalisierung
}
}
// Ausgabe des Verlusts alle 100 Epochen
if(epoch % 100 == 0) { // Anpassung der Frequenz für Verlustaussage
double averageLoss = totalLoss / (inputs.size() * MAX_RESPONSE_LENGTH);
System.out.println("Epoch " + epoch + " - Durchschnittlicher Verlust: " + averageLoss);
}
// Optional: Break if loss is sufficiently low
// if(totalLoss / (inputs.size() * MAX_RESPONSE_LENGTH) < desiredThreshold) break;
}
}
/**
* Startet das Training und die Benutzerinteraktion.
*/
public void start() {
System.out.println("Starte Training...");
trainNetwork();
System.out.println("Training abgeschlossen.");
System.out.println("Willkommen beim erweiterten Chatbot mit einfachem Attention! Tippe 'exit' zum Beenden.");
while (true) {
System.out.print("Du: ");
String input = scanner.nextLine().trim().toLowerCase();
if ("exit".equalsIgnoreCase(input)) {
System.out.println("Chatbot beendet. Auf Wiedersehen!");
break;
}
String response = generateResponse(input);
System.out.println("Bot: " + response);
}
scanner.close();
}
/**
* Generiert eine Antwort basierend auf der Benutzereingabe.
*
* @param input Die Benutzereingabe.
* @return Die generierte Antwort.
*/
private String generateResponse(String input) {
String[] words = input.split("\\s+");
List<WordVector> vectors = new ArrayList<>();
for(String word : words) {
WordVector wv = model.getWordVector(word);
if(wv != null) {
vectors.add(wv);
}
}
if(vectors.isEmpty()) {
return "Entschuldigung, ich verstehe dich nicht.";
}
// Sammle die Vektoren
double[][] inputVectors = vectors.stream().map(WordVector::getVector).toArray(double[][]::new);
// Initialer Kontextvektor (z.B. Durchschnitt der Eingabevektoren)
double[] contextVector = VectorUtils.average(Arrays.asList(inputVectors));
contextVector = VectorUtils.normalize(contextVector); // Normalisierung zur Stabilität
StringBuilder responseBuilder = new StringBuilder();
Map<Integer, String> indexToWord = model.getIndexToWordMap();
for(int i=0; i<MAX_RESPONSE_LENGTH; i++) {
// Berechne Aufmerksamkeitsscores basierend auf der Kosinusähnlichkeit
double[] attentionWeights = computeAttentionWeights(inputVectors, contextVector);
// Berechne gewichtete Summe der Eingabevektoren
double[] weightedInput = new double[contextVector.length];
for(int j=0; j<inputVectors.length; j++) {
double[] weightedWord = VectorUtils.multiply(attentionWeights[j], inputVectors[j]);
weightedInput = VectorUtils.add(weightedInput, weightedWord);
}
// Forward-Pass durch das neuronale Netzwerk
double[] outputProbabilities = neuralNet.forward(weightedInput);
// Finde das Wort mit der höchsten Wahrscheinlichkeit
int predictedIndex = argMax(outputProbabilities);
String nextWord = indexToWord.get(predictedIndex);
if(nextWord == null) {
break;
}
responseBuilder.append(nextWord).append(" ");
// Update Kontextvektor durch Akkumulation
WordVector nextWordVector = model.getWordVector(nextWord);
if(nextWordVector != null) {
contextVector = VectorUtils.add(contextVector, nextWordVector.getVector());
contextVector = VectorUtils.normalize(contextVector); // Optional: Normalisierung
}
}
return responseBuilder.toString().trim();
}
/**
* Berechnet die Aufmerksamkeitsgewichte basierend auf der Kosinusähnlichkeit zwischen
* dem Kontextvektor und jedem Eingabewortvektor.
*
* @param inputVectors Die Eingabewortvektoren.
* @param contextVector Der Kontextvektor.
* @return Die normalisierten Aufmerksamkeitsscores.
*/
private double[] computeAttentionWeights(double[][] inputVectors, double[] contextVector) {
int n = inputVectors.length;
double[] scores = new double[n];
for(int i=0; i<n; i++) {
scores[i] = VectorUtils.cosineSimilarity(inputVectors[i], contextVector);
}
// Softmax-Funktion zur Normalisierung der Scores
double[] softmaxScores = softmax(scores);
return softmaxScores;
}
/**
* Wendet die Softmax-Funktion auf ein Array von Scores an.
*
* @param scores Die Eingabescores.
* @return Die normalisierten Softmax-Scores.
*/
private double[] softmax(double[] scores) {
double max = Double.NEGATIVE_INFINITY;
for(double s : scores) {
if(s > max) max = s;
}
double sum = 0.0;
double[] expScores = new double[scores.length];
for(int i=0; i<scores.length; i++) {
expScores[i] = Math.exp(scores[i] - max);
sum += expScores[i];
}
double[] softmax = new double[scores.length];
for(int i=0; i<scores.length; i++) {
softmax[i] = expScores[i] / sum;
}
return softmax;
}
/**
* Findet das Index des höchsten Wertes im Array.
*
* @param array Das Eingabe-Array.
* @return Der Index des größten Elements.
*/
private int argMax(double[] array) {
int maxIndex = 0;
double max = array[0];
for(int i=1; i<array.length; i++) {
if(array[i] > max) {
max = array[i];
maxIndex = i;
}
}
return maxIndex;
}
}
public static void main(String[] args) {
String trainingFilePath = "/home/benito/Schreibtisch/train.txt";
String vectorFilePath = "/home/benito/Schreibtisch/skynetModel.txt";
try {
Chatbot chatbot = new Chatbot(vectorFilePath, trainingFilePath);
chatbot.start();
} catch(IOException e) {
System.err.println("Fehler beim Laden der Dateien: " + e.getMessage());
}
}
}ASSISTANT
Es scheint, dass Ihr Chatbot derzeit immer dasselbe Wort als Antwort generiert, was auf die Verwendung von `argMax` in der `generateResponse`-Methode zurückzuführen ist. Dieser Ansatz wählt immer das Wort mit der höchsten Wahrscheinlichkeit aus, was zu deterministischen und wenig variierenden Antworten führt. Um die Antworten des Chatbots abwechslungsreicher zu gestalten, können Sie zwei Hauptansätze verfolgen:
1. **Einführen einer Temperaturparameter in die Softmax-Funktion**: Dies ermöglicht es, die Wahrscheinlichkeitsverteilung der Ausgaben zu steuern, wodurch die Modelle mehr oder weniger Entropie erzeugen können.
2. **Stochastic Sampling statt deterministischen Auswahl**: Anstatt immer das Wort mit der höchsten Wahrscheinlichkeit (`argMax`) auszuwählen, können Sie basierend auf den Wahrscheinlichkeiten nach der Softmax-Funktion zufällig ein Wort auswählen.
Im Folgenden werde ich erläutern, wie Sie diese Änderungen in Ihrem bestehenden Code implementieren können.
### 1. Einführung eines Temperaturparameters in die Softmax-Funktion
Die **Temperatur** ist ein Parameter, der die Schärfe oder Flachheit der resultierenden Wahrscheinlichkeitsverteilung steuert. Eine höhere Temperatur (z.B. >1) führt zu einer flacheren Verteilung, wodurch weniger wahrscheinlich auch niedrig bewertete Optionen ausgewählt werden können. Eine niedrigere Temperatur (z.B. <1) macht die Verteilung schärfer und verstärkt die Unterschiede zwischen den Wahrscheinlichkeiten der einzelnen Klassen.
#### a. Modifikation der Softmax-Funktion
Fügen Sie einen Temperaturparameter zu Ihrer Softmax-Funktion hinzu. Dies kann in der `NeuralNetwork`-Klasse erfolgen.
```java
/**
* Wendet die Softmax-Funktion auf ein Array von Logits an, unter Berücksichtigung der Temperatur.
*
* @param logits Die Eingabewerte (Logits).
* @param temperature Die Temperatur zur Steuerung der Entropie.
* @return Die normalisierten Wahrscheinlichkeiten.
*/
private double[] softmax(double[] logits, double temperature) {
double max = Double.NEGATIVE_INFINITY;
for (double logit : logits) {
if (logit > max) max = logit;
}
double sum = 0.0;
double[] expScores = new double[logits.length];
for (int i = 0; i < logits.length; i++) {
expScores[i] = Math.exp((logits[i] - max) / temperature);
sum += expScores[i];
}
double[] softmax = new double[logits.length];
for (int i = 0; i < logits.length; i++) {
softmax[i] = expScores[i] / sum;
}
return softmax;
}
```
#### b. Anpassen des Forward-Passes
Ändern Sie die `forward`-Methode, um die Temperatur zu berücksichtigen. Sie können die Temperatur als Parameter übergeben oder einen Standardwert verwenden.
```java
/**
* Führt einen Forward-Pass durch das Netzwerk durch und gibt die aktivierten Output-Werte zurück.
*
* @param input Der Eingabevektor.
* @param temperature Die Temperatur zur Steuerung der Softmax-Verteilung.
* @return Der aktivierte Ausgabevektor nach Softmax.
*/
public double[] forward(double[] input, double temperature) {
// Berechne Hidden Layer Aktivierungen
double[] hiddenOutput = new double[hiddenSize];
for (int j = 0; j < hiddenSize; j++) {
double sum = 0.0;
for (int i = 0; i < inputSize; i++) {
sum += input[i] * weightsInputHidden[i][j];
}
hiddenOutput[j] = relu(sum);
}
// Berechne Output Layer (Logits)
double[] logits = new double[outputSize];
for (int k = 0; k < outputSize; k++) {
for (int j = 0; j < hiddenSize; j++) {
logits[k] += hiddenOutput[j] * weightsHiddenOutput[j][k];
}
}
// Wende Softmax mit Temperatur an
double[] probabilities = softmax(logits, temperature);
return probabilities;
}
```
### 2. Implementierung von Stochastic Sampling
Statt `argMax` zu verwenden, können Sie basierend auf den Wahrscheinlichkeiten nach der Softmax-Funktion zufällig ein Wort auswählen. Dies erhöht die Varianz der Antworten und vermeidet wiederholte gleiche Antworten.
#### a. Hinzufügen einer Sampling-Methode
Fügen Sie eine Methode hinzu, die basierend auf den Wahrscheinlichkeiten ein Wort index auswählt.
```java
/**
* Wählt einen Index basierend auf den Wahrscheinlichkeiten mittels stochastic sampling aus.
*
* @param probabilities Das Array der Wahrscheinlichkeiten.
* @param rand Ein Random-Objekt.
* @return Der ausgewählte Index.
*/
private int sampleFromDistribution(double[] probabilities, Random rand) {
double p = rand.nextDouble();
double cumulative = 0.0;
for (int i = 0; i < probabilities.length; i++) {
cumulative += probabilities[i];
if (p < cumulative) {
return i;
}
}
return probabilities.length - 1; // Rückfall, falls numerische Instabilität
}
```
#### b. Anpassen der `generateResponse`-Methode
Nutzen Sie die modifizierte `forward`-Methode mit Temperatur und die Sampling-Methode, um das nächste Wort auszuwählen.
```java
/**
* Generiert eine Antwort basierend auf der Benutzereingabe.
*
* @param input Die Benutzereingabe.
* @return Die generierte Antwort.
*/
private String generateResponse(String input) {
String[] words = input.split("\\s+");
List<WordVector> vectors = new ArrayList<>();
for (String word : words) {
WordVector wv = model.getWordVector(word);
if (wv != null) {
vectors.add(wv);
}
}
if (vectors.isEmpty()) {
return "Entschuldigung, ich verstehe dich nicht.";
}
// Sammle die Vektoren
double[][] inputVectors = vectors.stream().map(WordVector::getVector).toArray(double[][]::new);
// Initialer Kontextvektor (z.B. Durchschnitt der Eingabevektoren)
double[] contextVector = VectorUtils.average(Arrays.asList(inputVectors));
contextVector = VectorUtils.normalize(contextVector); // Normalisierung zur Stabilität
StringBuilder responseBuilder = new StringBuilder();
Map<Integer, String> indexToWord = model.getIndexToWordMap();
Random rand = new Random(); // Initialisiere Random-Objekt hier
// Definieren Sie die Temperatur (experimentieren Sie mit verschiedenen Werten, z.B. 0.7, 1.0, 1.5)
double temperature = 1.0;
for (int i = 0; i < MAX_RESPONSE_LENGTH; i++) {
// Berechne Aufmerksamkeitsscores basierend auf der Kosinusähnlichkeit
double[] attentionWeights = computeAttentionWeights(inputVectors, contextVector);
// Berechne gewichtete Summe der Eingabevektoren
double[] weightedInput = new double[contextVector.length];
for (int j = 0; j < inputVectors.length; j++) {
double[] weightedWord = VectorUtils.multiply(attentionWeights[j], inputVectors[j]);
weightedInput = VectorUtils.add(weightedInput, weightedWord);
}
// Forward-Pass durch das neuronale Netzwerk
double[] outputProbabilities = neuralNet.forward(weightedInput, temperature);
// Wählen Sie ein Wort basierend auf den Wahrscheinlichkeiten
int predictedIndex = neuralNet.sampleFromDistribution(outputProbabilities, rand);
String nextWord = indexToWord.get(predictedIndex);
if (nextWord == null) {
break;
}
responseBuilder.append(nextWord).append(" ");
// Update Kontextvektor durch Akkumulation
WordVector nextWordVector = model.getWordVector(nextWord);
if (nextWordVector != null) {
contextVector = VectorUtils.add(contextVector, nextWordVector.getVector());
contextVector = VectorUtils.normalize(contextVector); // Optional: Normalisierung
}
}
return responseBuilder.toString().trim();
}
```
### 3. Zusätzliche Empfehlungen
Neben der Einführung von Temperatur und stochastic sampling gibt es weitere Ansätze, die die Qualität und Vielfalt der Antworten Ihres Chatbots verbessern können:
- **Abweichung von `MAX_RESPONSE_LENGTH`**: Anstatt eine feste maximale Antwortlänge festzulegen, könnten Sie eine when-basierte Methode implementieren, die das Ende der Antwort erkennt, z.B. durch ein spezielles Token oder die Wahrscheinlichkeit weiterer Wörter unter einem bestimmten Schwellenwert.
- **Erhöhung der Modellkapazität**: Ihr aktuelles neuronales Netzwerk ist relativ einfach. Erwägen Sie die Verwendung von tieferen oder komplexeren Architekturen, um komplexere Muster in den Daten zu erfassen.
- **Verbesserung der Trainingsdaten**: Stellen Sie sicher, dass Ihre Trainingsdaten vielfältig und repräsentativ für die gewünschten Antworten sind. Eine größere und vielfältigere Datenbasis kann die Generalisierungsfähigkeit des Modells verbessern.
- **Validierung des Trainingsprozesses**: Überprüfen Sie regelmäßig den Trainingsfortschritt und stellen Sie sicher, dass das Modell tatsächlich lernt und sich nicht über- oder unteranpasst. Dies kann durch die Überwachung von Trainings- und Validierungsverlusten erfolgen.
- **Implementierung von Regularisierungstechniken**: Methoden wie Dropout oder L2-Regularisierung können helfen, Überanpassung zu verhindern und die Generalisierungsfähigkeit des Modells zu verbessern.
### Beispielcode für Temperaturanpassung und Sampling
Hier ist eine zusammengefasste Version der Änderungen, die auf Ihrem Code basieren:
```java
// In der NeuralNetwork-Klasse
public double[] forward(double[] input, double temperature) {
// ... bestehender Forward-Pass-Code ...
// Wende Softmax mit Temperatur an
double[] probabilities = softmax(logits, temperature);
return probabilities;
}
private double[] softmax(double[] logits, double temperature) {
double max = Double.NEGATIVE_INFINITY;
for (double logit : logits) {
if (logit > max) max = logit;
}
double sum = 0.0;
double[] expScores = new double[logits.length];
for (int i = 0; i < logits.length; i++) {
expScores[i] = Math.exp((logits[i] - max) / temperature);
sum += expScores[i];
}
double[] softmax = new double[logits.length];
for (int i = 0; i < logits.length; i++) {
softmax[i] = expScores[i] / sum;
}
return softmax;
}
/**
* Stichproben Methode basierend auf den Wahrscheinlichkeiten.
*/
public int sampleFromDistribution(double[] probabilities, Random rand) {
double p = rand.nextDouble();
double cumulative = 0.0;
for (int i = 0; i < probabilities.length; i++) {
cumulative += probabilities[i];
if (p < cumulative) {
return i;
}
}
return probabilities.length - 1; // Fallback
}
```
```java
// In der Chatbot-Klasse
private String generateResponse(String input) {
// ... bestehender Code ...
double temperature = 1.0; // Passen Sie die Temperatur an
for (int i = 0; i < MAX_RESPONSE_LENGTH; i++) {
// ... bestehender Code ...
double[] outputProbabilities = neuralNet.forward(weightedInput, temperature);
// Wählen Sie ein Wort basierend auf den Wahrscheinlichkeiten
int predictedIndex = neuralNet.sampleFromDistribution(outputProbabilities, rand);
String nextWord = indexToWord.get(predictedIndex);
if (nextWord == null) {
break;
}
responseBuilder.append(nextWord).append(" ");
// ... bestehender Code ...
}
return responseBuilder.toString().trim();
}
```
### Fazit
Durch die Einführung eines Temperaturparameters und das Verwenden von stochastic sampling anstelle von `argMax` können Sie die Vielfalt und Natürlichkeit der Antworten Ihres Chatbots erheblich verbessern. Experimentieren Sie mit verschiedenen Temperaturwerten, um ein gutes Gleichgewicht zwischen Konsistenz und Vielfalt zu finden. Zusätzlich sollten Sie die Trainingsdaten und die Modellarchitektur weiter optimieren, um die Gesamtleistung Ihres Chatbots zu steigern.
Falls Sie weitere Fragen haben oder zusätzliche Unterstützung benötigen, stehe ich gerne zur Verfügung!