summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rwxr-xr-xgpt_adventure.py23
1 files changed, 15 insertions, 8 deletions
diff --git a/gpt_adventure.py b/gpt_adventure.py
index 49138ff..dd86adf 100755
--- a/gpt_adventure.py
+++ b/gpt_adventure.py
@@ -54,12 +54,11 @@ class State:
def rollback_history(self) -> None:
self.history = self.history[:-1]
- def prompt(self) -> str:
- return self.prologue + "\n".join(self.history) + f"\n{self.character}: "
+ def continuation_prompt(self) -> str:
+ return self.prologue + "\n".join(self.history)
- def trim_history(self):
- while self.history and self.model.token_length(self.prompt()) > 1024:
- self.history = self.history[1:]
+ def prompt(self) -> str:
+ return self.continuation_prompt() + f"\n{self.character}: "
def print_context(self) -> None:
for elem in self.history[-10:]:
@@ -68,6 +67,11 @@ class State:
else:
print(elem.removeprefix(f"{self.character}: "))
+ def generate(self, prompt: str) -> str:
+ while self.history and self.model.token_length(self.prompt()) > 1024:
+ self.history = self.history[1:]
+ return self.model.generate(prompt)
+
def autosave(self) -> None:
if self.filename:
self.save(self.filename)
@@ -127,6 +131,7 @@ def main() -> None:
if msg == "/help":
commands = (
"/help",
+ "/more",
"/retry",
"/edit",
"/undo",
@@ -135,6 +140,10 @@ def main() -> None:
print("Commands:")
for cmd in commands:
print(f"\t{cmd}")
+ elif msg == "/more":
+ response = state.generate(state.prompt())
+ print(response)
+ state.history[-1] += response
elif msg == "/retry":
state.rollback_history()
clear()
@@ -162,9 +171,7 @@ def main() -> None:
else:
state.add_history("You", msg)
- state.trim_history()
-
- response = state.model.generate(state.prompt())
+ response = state.generate(state.prompt())
if not response:
response = "haha I'm not sure what you mean"
print(response)