{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# Generatiivsed võrgud\n", "\n", "Korduvad närvivõrgud (RNN-id) ja nende väratiga rakutüübid, nagu Long Short Term Memory Cells (LSTM-id) ja Gated Recurrent Units (GRU-d), pakuvad mehhanismi keele modelleerimiseks, st nad suudavad õppida sõnade järjestust ja ennustada järgmise sõna järjestuses. See võimaldab meil kasutada RNN-e **generatiivseteks ülesanneteks**, nagu tavaline tekstigeneratsioon, masintõlge ja isegi pildiallkirjade loomine.\n", "\n", "RNN arhitektuuris, mida käsitlesime eelmises osas, genereeris iga RNN üksus järgmise varjatud oleku väljundina. Kuid me saame lisada igale korduva üksusele veel ühe väljundi, mis võimaldab meil luua **järjestuse** (mis on sama pikk kui algne järjestus). Lisaks saame kasutada RNN üksusi, mis ei võta igal sammul sisendit, vaid kasutavad ainult algset olekuvektorit ja seejärel genereerivad väljundite järjestuse.\n", "\n", "Selles märkmikus keskendume lihtsatele generatiivsetele mudelitele, mis aitavad meil teksti genereerida. Lihtsuse huvides loome **tähemärgi tasemel võrgu**, mis genereerib teksti täht-tähelt. Treeningu ajal peame võtma mõne tekstikorpuse ja jagama selle tähemärkide järjestusteks.\n" ] }, { "cell_type": "code", "execution_count": 1, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Loading dataset...\n", "Building vocab...\n" ] } ], "source": [ "import torch\n", "import torchtext\n", "import numpy as np\n", "from torchnlp import *\n", "train_dataset,test_dataset,classes,vocab = load_dataset()" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Tegelaskujude sõnavara loomine\n", "\n", "Tegelaskujude tasemel generatiivse võrgu loomiseks peame teksti jagama üksikuteks tähtedeks, mitte sõnadeks. Seda saab teha, määratledes teistsuguse tokeniseerija:\n" ] }, { "cell_type": "code", "execution_count": 2, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Vocabulary size = 82\n", "Encoding of 'a' is 1\n", "Character with code 13 is c\n" ] } ], "source": [ "def char_tokenizer(words):\n", " return list(words) #[word for word in words]\n", "\n", "counter = collections.Counter()\n", "for (label, line) in train_dataset:\n", " counter.update(char_tokenizer(line))\n", "vocab = torchtext.vocab.vocab(counter)\n", "\n", "vocab_size = len(vocab)\n", "print(f\"Vocabulary size = {vocab_size}\")\n", "print(f\"Encoding of 'a' is {vocab.get_stoi()['a']}\")\n", "print(f\"Character with code 13 is {vocab.get_itos()[13]}\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "Vaatame näidet, kuidas saame oma andmekogumist teksti kodeerida:\n" ] }, { "cell_type": "code", "execution_count": 3, "metadata": {}, "outputs": [ { "data": { "text/plain": [ "tensor([ 0, 1, 2, 2, 3, 4, 5, 6, 3, 7, 8, 1, 9, 10, 3, 11, 2, 1,\n", " 12, 3, 7, 1, 13, 14, 3, 15, 16, 5, 17, 3, 5, 18, 8, 3, 7, 2,\n", " 1, 13, 14, 3, 19, 20, 8, 21, 5, 8, 9, 10, 22, 3, 20, 8, 21, 5,\n", " 8, 9, 10, 3, 23, 3, 4, 18, 17, 9, 5, 23, 10, 8, 2, 2, 8, 9,\n", " 10, 24, 3, 0, 1, 2, 2, 3, 4, 5, 9, 8, 8, 5, 25, 10, 3, 26,\n", " 12, 27, 16, 26, 2, 27, 16, 28, 29, 30, 1, 16, 26, 3, 17, 31, 3, 21,\n", " 2, 5, 9, 1, 23, 13, 32, 16, 27, 13, 10, 24, 3, 1, 9, 8, 3, 10,\n", " 8, 8, 27, 16, 28, 3, 28, 9, 8, 8, 16, 3, 1, 28, 1, 27, 16, 6])" ] }, "execution_count": 3, "metadata": {}, "output_type": "execute_result" } ], "source": [ "def enc(x):\n", " return torch.LongTensor(encode(x,voc=vocab,tokenizer=char_tokenizer))\n", "\n", "enc(train_dataset[0][1])" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Generatiivse RNN-i treenimine\n", "\n", "RNN-i treenimiseks teksti genereerimiseks toimime järgmiselt. Igal sammul võtame `nchars` pikkuse tähemärkide jada ja palume võrgul genereerida iga sisendtähemärgi jaoks järgmine väljundtähemärk:\n", "\n", "![Pilt, mis näitab näidet RNN-i genereerimisest sõnaga 'HELLO'.](../../../../../translated_images/et/rnn-generate.56c54afb52f9781d.webp)\n", "\n", "Olenevalt konkreetsest stsenaariumist võime soovida lisada ka erimärke, näiteks *järjestuse lõpp* ``. Meie puhul tahame lihtsalt treenida võrku lõputu teksti genereerimiseks, seega määrame iga jada suuruseks `nchars` tokenit. Järelikult koosneb iga treeningnäide `nchars` sisendist ja `nchars` väljundist (mis on sisendjada, nihutatud ühe sümboli võrra vasakule). Minipartii koosneb mitmest sellisest jadast.\n", "\n", "Minipartiide genereerimiseks võtame iga uudisteksti pikkusega `l` ja loome sellest kõik võimalikud sisend-väljund kombinatsioonid (neid kombinatsioone on `l-nchars`). Need moodustavad ühe minipartii ja minipartiide suurus on igal treenimissammul erinev.\n" ] }, { "cell_type": "code", "execution_count": 4, "metadata": {}, "outputs": [ { "data": { "text/plain": [ "(tensor([[ 0, 1, 2, ..., 28, 29, 30],\n", " [ 1, 2, 2, ..., 29, 30, 1],\n", " [ 2, 2, 3, ..., 30, 1, 16],\n", " ...,\n", " [20, 8, 21, ..., 1, 28, 1],\n", " [ 8, 21, 5, ..., 28, 1, 27],\n", " [21, 5, 8, ..., 1, 27, 16]]),\n", " tensor([[ 1, 2, 2, ..., 29, 30, 1],\n", " [ 2, 2, 3, ..., 30, 1, 16],\n", " [ 2, 3, 4, ..., 1, 16, 26],\n", " ...,\n", " [ 8, 21, 5, ..., 28, 1, 27],\n", " [21, 5, 8, ..., 1, 27, 16],\n", " [ 5, 8, 9, ..., 27, 16, 6]]))" ] }, "execution_count": 4, "metadata": {}, "output_type": "execute_result" } ], "source": [ "nchars = 100\n", "\n", "def get_batch(s,nchars=nchars):\n", " ins = torch.zeros(len(s)-nchars,nchars,dtype=torch.long,device=device)\n", " outs = torch.zeros(len(s)-nchars,nchars,dtype=torch.long,device=device)\n", " for i in range(len(s)-nchars):\n", " ins[i] = enc(s[i:i+nchars])\n", " outs[i] = enc(s[i+1:i+nchars+1])\n", " return ins,outs\n", "\n", "get_batch(train_dataset[0][1])" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "Nüüd määratleme generaatori võrgu. See võib põhineda mis tahes korduvrakul, mida arutasime eelmises osas (lihtne, LSTM või GRU). Meie näites kasutame LSTM-i.\n", "\n", "Kuna võrk võtab sisendiks tähemärke ja sõnavara suurus on üsna väike, ei ole meil vaja sisendkihtide jaoks embedding-kihte; ühekuumkoodiga (one-hot-encoded) sisend võib otse minna LSTM-rakku. Kuid kuna me edastame tähemärkide numbreid sisendina, peame need enne LSTM-i edastamist ühekuumkoodiga kodeerima. Seda tehakse `forward` läbimise ajal, kutsudes `one_hot` funktsiooni. Väljundkooder oleks lineaarne kiht, mis teisendab peidetud oleku ühekuumkoodiga väljundiks.\n" ] }, { "cell_type": "code", "execution_count": 5, "metadata": {}, "outputs": [], "source": [ "class LSTMGenerator(torch.nn.Module):\n", " def __init__(self, vocab_size, hidden_dim):\n", " super().__init__()\n", " self.rnn = torch.nn.LSTM(vocab_size,hidden_dim,batch_first=True)\n", " self.fc = torch.nn.Linear(hidden_dim, vocab_size)\n", "\n", " def forward(self, x, s=None):\n", " x = torch.nn.functional.one_hot(x,vocab_size).to(torch.float32)\n", " x,s = self.rnn(x,s)\n", " return self.fc(x),s" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "Treeningu ajal soovime olla võimelised genereeritud teksti näidiseid võtma. Selleks määratleme funktsiooni `generate`, mis loob väljundstringi pikkusega `size`, alustades algstringist `start`.\n", "\n", "Selle tööpõhimõte on järgmine. Kõigepealt edastame kogu algstringi läbi võrgu ja võtame väljundoleku `s` ning järgmise ennustatud tähemärgi `out`. Kuna `out` on ühekujulise kodeeringuga, võtame `argmax`, et saada tähemärgi indeks `nc` sõnavaras, ning kasutame `itos`, et leida tegelik tähemärk ja lisada see tulemuseks olevate tähemärkide loendisse `chars`. Seda ühe tähemärgi genereerimise protsessi korratakse `size` korda, et luua vajalik arv tähemärke.\n" ] }, { "cell_type": "code", "execution_count": 8, "metadata": {}, "outputs": [], "source": [ "def generate(net,size=100,start='today '):\n", " chars = list(start)\n", " out, s = net(enc(chars).view(1,-1).to(device))\n", " for i in range(size):\n", " nc = torch.argmax(out[0][-1])\n", " chars.append(vocab.get_itos()[nc])\n", " out, s = net(nc.view(1,-1),s)\n", " return ''.join(chars)" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "Nüüd alustame treenimist! Treeningtsükkel on peaaegu sama nagu kõigis meie varasemates näidetes, kuid täpsuse asemel prindime iga 1000 epohhi järel genereeritud teksti näidise.\n", "\n", "Erilist tähelepanu tuleb pöörata sellele, kuidas me arvutame kaotust. Kaotuse arvutamiseks vajame üheselt kodeeritud väljundit `out` ja oodatud teksti `text_out`, mis on tähemärkide indeksite loend. Õnneks ootab funktsioon `cross_entropy` esimeseks argumendiks normaliseerimata võrgu väljundit ja teiseks klassi numbrit, mis on täpselt see, mis meil olemas on. Lisaks teeb see automaatse keskmistamise minibatch'i suuruse järgi.\n", "\n", "Samuti piirame treenimist `samples_to_train` näidiste arvuga, et mitte liiga kaua oodata. Soovitame teil katsetada ja proovida pikemat treenimist, võimalusel mitme epohhi jooksul (sellisel juhul peaksite selle koodi ümber looma uue tsükli).\n" ] }, { "cell_type": "code", "execution_count": 9, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Current loss = 4.398899078369141\n", "today sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr s\n", "Current loss = 2.161320447921753\n", "today and to the tor to to the tor to to the tor to to the tor to to the tor to to the tor to to the tor t\n", "Current loss = 1.6722588539123535\n", "today and the court to the could to the could to the could to the could to the could to the could to the c\n", "Current loss = 2.423795223236084\n", "today and a second to the conternation of the conternation of the conternation of the conternation of the \n", "Current loss = 1.702607274055481\n", "today and the company to the company to the company to the company to the company to the company to the co\n", "Current loss = 1.692358136177063\n", "today and the company to the company to the company to the company to the company to the company to the co\n", "Current loss = 1.9722288846969604\n", "today and the control the control the control the control the control the control the control the control \n", "Current loss = 1.8705692291259766\n", "today and the second to the second to the second to the second to the second to the second to the second t\n", "Current loss = 1.7626899480819702\n", "today and a security and a security and a security and a security and a security and a security and a secu\n", "Current loss = 1.5574463605880737\n", "today and the company and the company and the company and the company and the company and the company and \n", "Current loss = 1.5620026588439941\n", "today and the be that the be the be that the be the be that the be the be that the be the be that the be t\n" ] } ], "source": [ "net = LSTMGenerator(vocab_size,64).to(device)\n", "\n", "samples_to_train = 10000\n", "optimizer = torch.optim.Adam(net.parameters(),0.01)\n", "loss_fn = torch.nn.CrossEntropyLoss()\n", "net.train()\n", "for i,x in enumerate(train_dataset):\n", " # x[0] is class label, x[1] is text\n", " if len(x[1])-nchars<10:\n", " continue\n", " samples_to_train-=1\n", " if not samples_to_train: break\n", " text_in, text_out = get_batch(x[1])\n", " optimizer.zero_grad()\n", " out,s = net(text_in)\n", " loss = torch.nn.functional.cross_entropy(out.view(-1,vocab_size),text_out.flatten()) #cross_entropy(out,labels)\n", " loss.backward()\n", " optimizer.step()\n", " if i%1000==0:\n", " print(f\"Current loss = {loss.item()}\")\n", " print(generate(net))" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "See näide genereerib juba üsna head teksti, kuid seda saab mitmel viisil veelgi paremaks muuta:\n", "* **Parem minibatchide genereerimine**. Viis, kuidas me treeningu jaoks andmeid ette valmistasime, oli ühe minibatchi genereerimine ühest proovist. See pole ideaalne, kuna minibatchid on kõik erineva suurusega ja mõnda neist ei saa isegi genereerida, kuna tekst on väiksem kui `nchars`. Lisaks ei koorma väikesed minibatchid GPU-d piisavalt. Targem oleks võtta üks suur tekstilõik kõigist proovidest, seejärel genereerida kõik sisend-väljund paarid, segada need ja luua võrdse suurusega minibatchid.\n", "* **Mitmekihiline LSTM**. Tasub proovida 2 või 3 kihti LSTM-rakke. Nagu mainisime eelmises osas, eraldab iga LSTM kiht tekstist teatud mustreid, ja tähemärgi tasemel generaatori puhul võime eeldada, et madalam LSTM tase vastutab silpide eraldamise eest, kõrgemad tasemed aga sõnade ja sõnakombinatsioonide eest. Seda saab lihtsalt rakendada, andes LSTM konstruktorile kihtide arvu parameetri.\n", "* Võid samuti katsetada **GRU üksustega** ja vaadata, millised annavad paremaid tulemusi, ning **erinevate varjatud kihtide suurustega**. Liiga suur varjatud kiht võib viia üleõppimisele (näiteks õpib võrk täpse teksti ära), ja väiksem suurus ei pruugi anda häid tulemusi.\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Pehme teksti genereerimine ja temperatuur\n", "\n", "Eelmises `generate` definitsioonis valisime alati järgmise tähemärgina selle, millel oli suurim tõenäosus. See viis sageli selleni, et tekst \"kordas\" samu tähemärkide järjestusi uuesti ja uuesti, nagu selles näites:\n", "```\n", "today of the second the company and a second the company ...\n", "```\n", "\n", "Kui aga vaatame järgmise tähemärgi tõenäosusjaotust, võib juhtuda, et mõne kõrgeima tõenäosuse vahe ei ole suur, näiteks ühel tähemärgil võib olla tõenäosus 0.2 ja teisel 0.19 jne. Näiteks, kui otsime järgmist tähemärki järjestuses '*play*', võib järgmine tähemärk sama hästi olla kas tühik või **e** (nagu sõnas *player*).\n", "\n", "See viib meid järelduseni, et alati ei ole \"õiglane\" valida tähemärki, millel on suurem tõenäosus, sest teise kõrgeima valimine võib siiski viia tähendusliku tekstini. Mõistlikum on **valida juhuslikult** tähemärke tõenäosusjaotusest, mille annab võrgu väljund.\n", "\n", "Seda juhuslikku valimist saab teha `multinomial` funktsiooni abil, mis rakendab nn **multinomiaalset jaotust**. Funktsioon, mis rakendab seda **pehmet** teksti genereerimist, on defineeritud allpool:\n" ] }, { "cell_type": "code", "execution_count": 10, "metadata": { "scrolled": true }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "--- Temperature = 0.3\n", "Today and a company and complete an all the land the restrational the as a security and has provers the pay to and a report and the computer in the stand has filities and working the law the stations for a company and with the company and the final the first company and refight of the state and and workin\n", "\n", "--- Temperature = 0.8\n", "Today he oniis its first to Aus bomblaties the marmation a to manan boogot that pirate assaid a relaid their that goverfin the the Cappets Ecrotional Assonia Cition targets it annight the w scyments Blamity #39;s TVeer Diercheg Reserals fran envyuil that of ster said access what succers of Dour-provelith\n", "\n", "--- Temperature = 1.0\n", "Today holy they a 11 will meda a toket subsuaties, engins for Chanos, they's has stainger past to opening orital his thempting new Nattona was al innerforder advan-than #36;s night year his religuled talitatian what the but with Wednesday to Justment will wemen of Mark CCC Camp as Timed Nae wome a leaders\n", "\n", "--- Temperature = 1.3\n", "Today gpone 2.5 fech atcusion poor cocles toparsdorM.cht Line Pamage put 43 his calt lowed to the book, that has authh-the silia rruch ailing to'ory andhes beutirsimi- Aefffive heading offil an auf eacklets is charged evis, Gunymy oy) Mony has it after-sloythyor loveId out filme, the Natabl -Najuntaxiggs \n", "\n", "--- Temperature = 1.8\n", "Today plary, P.slan chly\\401 mardregationly #39;t 8.1Mide) closes ,filtcon alfly playin roven!\\grea.-QFBEP: Iss onfarchQ/itilia CCf Zivesigntwasta orce.-Peul-aw.uicrin of fuglinfsut aftaningwo, MIEX awayew Aice Woiduar Corvagiugge oppo esig ThusBratourid canthly-RyI.co lagitems\\eexciaishes.conBabntusmor I\n", "\n" ] } ], "source": [ "def generate_soft(net,size=100,start='today ',temperature=1.0):\n", " chars = list(start)\n", " out, s = net(enc(chars).view(1,-1).to(device))\n", " for i in range(size):\n", " #nc = torch.argmax(out[0][-1])\n", " out_dist = out[0][-1].div(temperature).exp()\n", " nc = torch.multinomial(out_dist,1)[0]\n", " chars.append(vocab.get_itos()[nc])\n", " out, s = net(nc.view(1,-1),s)\n", " return ''.join(chars)\n", " \n", "for i in [0.3,0.8,1.0,1.3,1.8]:\n", " print(f\"--- Temperature = {i}\\n{generate_soft(net,size=300,start='Today ',temperature=i)}\\n\")" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "Oleme lisanud veel ühe parameetri, mida nimetatakse **temperatuuriks**, ja seda kasutatakse näitamaks, kui tugevalt peaksime järgima kõrgeimat tõenäosust. Kui temperatuur on 1.0, teeme õiglast multinomiaalset valimit, ja kui temperatuur läheneb lõpmatusele - muutuvad kõik tõenäosused võrdseks ning me valime järgmise tähemärgi juhuslikult. Allolevas näites näeme, et tekst muutub mõttetuks, kui temperatuuri liiga palju suurendame, ja meenutab \"tsüklilist\" rangelt genereeritud teksti, kui see läheneb 0-le.\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "\n---\n\n**Lahtiütlus**: \nSee dokument on tõlgitud AI tõlketeenuse [Co-op Translator](https://github.com/Azure/co-op-translator) abil. Kuigi püüame tagada täpsust, palume arvestada, et automaatsed tõlked võivad sisaldada vigu või ebatäpsusi. Algne dokument selle algses keeles tuleks pidada autoriteetseks allikaks. Olulise teabe puhul soovitame kasutada professionaalset inimtõlget. Me ei vastuta selle tõlke kasutamisest tulenevate arusaamatuste või valesti tõlgenduste eest.\n" ] } ], "metadata": { "interpreter": { "hash": "16af2a8bbb083ea23e5e41c7f5787656b2ce26968575d8763f2c4b17f9cd711f" }, "kernelspec": { "display_name": "Python 3.8.12 ('py38')", "language": "python", "name": "python3" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.8.12" }, "coopTranslator": { "original_hash": "7673cd150d96c74c6d6011460094efb4", "translation_date": "2025-10-11T12:51:09+00:00", "source_file": "lessons/5-NLP/17-GenerativeNetworks/GenerativePyTorch.ipynb", "language_code": "et" } }, "nbformat": 4, "nbformat_minor": 4 }